Skip to content

feat(quant): opt-in fp8-block QAT for the lm_head - #636

Open
zhangnju wants to merge 1 commit into
FlashML-org:mainfrom
zhangnju:lmhead-fp8-qat
Open

zhangnju wants to merge 1 commit into
FlashML-org:mainfrom
zhangnju:lmhead-fp8-qat

Conversation

@zhangnju

@zhangnju zhangnju commented Oct 8, 2026 •

Copy link
Copy Markdown

Summary

This change adds an opt-in path that quantizes the bf16 lm_head to fp8 block-scale at load time. The lm_head is a large [vocab, hidden] matrix, so decoding one token reads a lot of bytes from it; storing it in fp8 halves that traffic. The feature reuses the existing fp8-block triton GEMV end-to-end and adds no new kernel. It is disabled by default and is turned on with FREETOKEN_FP8_LM_HEAD=1, so the default behavior and logit quality are unchanged.

How it works

The checkpoint still ships a bf16 lm_head. When the flag is set, the loader assigns it a new quantization scheme whose weight is quantized to fp8 with a bf16 128-block scale after the weight is loaded, in the linear method's finalize step. From then on it is served by the same fp8-block GEMV used for every other fp8-block linear.

Changes

  • layers/quantization/scheme.py: adds QuantKind.FP8_BLOCK_QAT and fp8_block_qat_scheme. The scheme declares only a weight role, because the checkpoint carries no scale for it.
  • kernel/triton/fp8_block_linear.py: adds per_block_quant_fp8, which quantizes a bf16 weight to fp8 plus a bf16 128-block scale. It is the inverse of the existing dequant_block_fp8.
  • layers/quantization/linear/fp8_block.py: adds Fp8BlockQuantizeAtLoadLinearMethod and its kernel. The kernel's finalize block-quantizes the bf16 weight to fp8 after load; the method otherwise reuses the standard fp8-block GEMV.
  • layers/quantization/configs/fp8.py: when the flag is set, scheme_for_name returns the BLOCK_QAT scheme for lm_head before the default exclusion that would have kept it in bf16.
  • models/qwen3_5_moe/weight.py: two loader validations are relaxed for FP8_BLOCK_QAT, which ships a bf16 weight with no scale companion, so the bf16 tensor is accepted instead of being rejected against an fp8 scheme.

Validation

  • per_block_quant_fp8 and dequant_block_fp8 round-trip on CPU with a relative error of about 2.25%, which matches fp8 block granularity.
  • On an R9700 (gfx1201) serving Qwen3.6-35B-A3B-FP8 at TP=2, toggling the flag (with the native MoE path disabled on both runs, to isolate the lm_head): bf16 lm_head gives 86.9 tok/s and fp8 lm_head gives 88.9 tok/s, a 2.3% decode improvement, with byte-identical greedy output (the fp8 path reuses the same triton kernel, so greedy decoding is unchanged).

Notes

  • The feature is off by default, so it does not affect the standard path.
  • This touches only the fp8-block linear path; it is independent of the MoE kernels and of the tensor-parallel sharding work.

Test plan

  • CPU: per_block_quant_fp8 / dequant_block_fp8 round-trip, rel-err ~2.25%
  • ROCm (R9700, gfx1201): Qwen3.6-35B-A3B-FP8 TP=2, flag on vs off -> +2.3% decode, greedy byte-identical
  • CUDA sanity on any fp8-block model: flag off is a no-op; flag on exercises the finalize path

…t load)

Quantize the bf16 lm_head to fp8 block-scale at load so decode reads fewer
bytes from the big [vocab, hidden] matrix. Reuses the existing fp8-block
triton GEMV end-to-end; no new kernel. Opt-in via FREETOKEN_FP8_LM_HEAD=1
(default OFF -> zero change to logit quality).

- scheme.py: QuantKind.FP8_BLOCK_QAT + fp8_block_qat_scheme (weight role only,
  no scale in the checkpoint).
- kernel/triton/fp8_block_linear.py: per_block_quant_fp8 (bf16 -> fp8 + bf16
  128-block scale), the inverse of dequant_block_fp8.
- linear/fp8_block.py: Fp8BlockQuantizeAtLoadLinearMethod + its kernel's
  finalize() block-quantizes the bf16 weight post-load.
- configs/fp8.py: flag-gated scheme_for_name returns BLOCK_QAT for lm_head
  before the default not_convert exclusion.
- qwen3_5_moe/weight.py: the two loader validations exempt FP8_BLOCK_QAT (the
  checkpoint ships a bf16 weight with no scale companion).
KarrAcaRn pushed a commit to KarrAcaRn/FreeToken-ByAI that referenced this pull request Oct 10, 2026
KarrAcaRn pushed a commit to KarrAcaRn/FreeToken-ByAI that referenced this pull request Oct 10, 2026
…kernel skip the fp8_block-only scheme check

FREETOKEN_FP8_LM_HEAD=1 asserted while building the lm_head: the kernel
inherited unusable_reason from the float-scale fp8-block kernel, which
calls fp8_block_size on a scheme that is FP8_BLOCK_QAT.

Assisted-by: Claude Opus 5.5

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