Skip to content

fix(turbomind): load per-channel FP8 compressed-tensors checkpoints (Ornith) - #5018

Open
bltcn wants to merge 5 commits into
InternLM:mainfrom
bltcn:pr/ornith-fp8-per-channel
Open

bltcn wants to merge 5 commits into
InternLM:mainfrom
bltcn:pr/ornith-fp8-per-channel

Conversation

@bltcn

@bltcn bltcn commented Oct 3, 2026

Copy link
Copy Markdown
Contributor

Summary

TurboMind cannot load per-channel (channel-quantized) FP8 compressed-tensors checkpoints such as Ornith-1.5-35B-A3B-FP8:

  • quant_method: compressed-tensors, format: float-quantized
  • 8-bit float weights (F8_E4M3), strategy: channel, group_size: null
  • scales named *.weight_scale (BF16, shape [N, 1]), group key config_group_0

The converter hard-asserted pack-quantized + 4-bit int for compressed-tensors, so these checkpoints were rejected.

Changes (python-only; C++ already supports it)

Mainline already has the FP8 compute path for pre-SM89 GPUs (#4871) and groupwise 128x1 e4m3 kernels (#4943). This PR only closes the python-side gap:

  • converter.py: accept float-quantized compressed-tensors with 8-bit float weights (per-channel -> FP8Format(block_out=1), 128x128 blocked -> block_out=128); allow explicit --model-format fp8 for such checkpoints; group-key fallback config_group_0/group_0
  • weight_format.py: FP8Format accepts .weight_scale suffix and per-channel [N, 1] scales; new post_process hook expands per-channel scales into the kernel's K-grouped layout [K//128, N] (identity for blocked, so no behavior change for existing fp8 checkpoints)
  • builders/linear.py: _build_linear invokes the format's post_process
  • tests/turbomind/linear/test_fp8_per_channel.py: new unit tests

Validation (2x RTX 2080Ti, sm75, CUDA 12.8)

  • unit tests: 22/22 pass
  • real Ornith-1.5-35B-A3B-FP8 e2e: converter parse -> per-channel scale expansion -> e4m3 groupwise GEMM plan -> serve + chat all pass
  • full eval suite vs 9/7 baseline (same model): accuracy within sampling noise on all 8 datasets (gsm8k 0.870 vs 0.883, mmlu 0.691 vs 0.707, arc_c 0.931 vs 0.942, ...); longctx needle recall 1.0 up to 128k; agent c8 throughput 168 tok/s (vs 138 baseline)

bltcn added 5 commits October 2, 2026 11:46
…P8 (Ornith)

Some compressed-tensors checkpoints (e.g. Ornith-1.5) store FP8 weights in
the float-quantized format with a per-output-channel .weight_scale
(BF16, shape [N, 1]) instead of the blocked .weight_scale_inv
([K//128, N//128]). The converter rejected them outright
(pack-quantized-only assert) and the FP8 path did not recognize the
.weight_scale key.

- converter: accept float-quantized 8-bit float compressed-tensors groups
  (group_0 / config_group_0), map to FP8Format(block_out=1) for
  per-channel and block_out=128 for blocked
- FP8Format: recognize .weight_scale, accept per-channel [N, 1] scales,
  add post_process hook that expands per-channel scales to the K-grouped
  kernel layout [K//128, N] (identical rows; weight bytes unchanged)
- WeightFormat: identity post_process hook on the base class

The C++ side already supports per-channel FP8 (groupwise 128x1 e4m3,
InternLM#4943); this only closes the python loading gap.
…_process

normalize() transposes 2-D tensors to TM layout before post_process runs, so
the per-channel weight_scale [N, 1] reaches the hook as [1, N]. Match the
transposed shape when expanding to [K//128, N]; guard also rejects any
other layout instead of crashing on expand.
…ors float-quantized checkpoints

Ornith-style checkpoints carry quant_method=compressed-tensors with
float-quantized FP8 weights. The production startup script passes an
explicit --model-format fp8, which previously hit the strict
'user input == quant_method' assert. fp8 and compressed-tensors float-quantized
are the same underlying format, so accept the override; the float-quantized
branch still validates 8-bit float weights.
CI lint/unit_test fixes for the per-channel FP8 PR:
- converter.py: wrap the _build_quantized_formats signature to respect the
  120-char line-length limit (ruff E501).
- test_fp8_per_channel.py: sort imports (ruff I001) and fix the two test
  cases that call FP8Format.post_process / dequant directly to pass the
  normalize()-transposed per-channel scale [1, N] (not the raw [N, 1]);
  this matches the production path where normalize() transposes 2-D
  tensors to TM layout before post_process runs.

Verified in a real-CUDA container (lmdeploy 0.18.0, cu128): all 6 tests
pass; regression A/B (patched vs pristine) on the existing linear tests is
identical (1 pre-existing skip-path failure in both), so the patch breaks
nothing.
CI lint's docformatter hook (v1.7.7, --wrap-descriptions 120) requires a
blank line after the summary sentence of the base-class post_process
docstring I added. Split it into summary + body so pre-commit passes.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant