Skip to content

feat(qwen4_exp): support ModelOpt QAD checkpoints (NVFP4 PLE tables, MXFP8 dense) - #594

Open
skavans wants to merge 2 commits into
FlashML-org:mainfrom
skavans:feat/qad-ckpt-support
Open

skavans wants to merge 2 commits into
FlashML-org:mainfrom
skavans:feat/qad-ckpt-support

Conversation

@skavans

@skavans skavans commented Oct 3, 2026 •

Copy link
Copy Markdown

What

Makes local-inference-lab/Qwen3.8-Flash-Next-NVFP4 (a ModelOpt QAD requant of
Qwen3.8-Flash-Next) load and serve out of the box. On main the loader dies in two places:

  1. the n-gram PLE table shards are NVFP4-packed (U8, two e2m1 codes per byte + an fp8-e4m3
    block scale per 16 elements + a global scalar under .weight_scale_2) and load_ple_table
    only accepts F8_E4M3;
  2. ModelOpt exports MXFP8 dense block scales as <proj>.weight_scale while the scheme loads
    them as weight_scale_inv, so _rename skipped them and the module KeyErrors mid load.

Two commits, one per defect:

  • load and gather NVFP4-packed PLE tables - the pinned host table becomes a packed uint8
    bank plus a pinned fp8 block-scale bank; the gather kernel gains an NVFP4 branch (LUT unpack,
    block scale x global scale, int64 row offsets at 320M rows). Reuses the canonical _e2m1_lut
    from kernel/triton/nvfp4_dequant.py and the e4m3_compat helpers; no new dequant math.
    Packed rows are served by the pinned PLE backend only; the disk backend raises
    NotImplementedError with a reason instead of misreading rows.
  • read modelopt MXFP8 dense scales stored under .weight_scale - the ModelOpt dialect's
    storage table carries the truth (MXFP8 -> weight_scale) and _rename renames through it
    instead of blanket-skipping scale suffixes, so the scales fuse and validate on the existing
    block-scale path. Covers the split GDN layout with in_proj_b|a quantized.

Tests

CPU-only unit tests (no GPU needed), green on this tree: pytest tests/models/qwen4_exp -m "not slow".

  • test_ple_nvfp4.py: synthetic packed checkpoint (real safetensors files in tmp_path) loaded
    through load_ple_table; the pinned gather is bitwise-checked against a torch reference that
    mirrors the checkpoint formula (e2m1 code x block scale x global scale, fp32, stored bf16);
    out-of-range ids store zeros; prefetch path included.
  • MXFP8 fixture parametrizes the reader-emits-keys test and a dedicated fusion test asserts the
    scales come out under weight_scale_inv fused across qkv/z/ba parts.

Scope: what is model-specific and what is not

  • The storage-table fix (modelopt.py) is engine-wide: every checkpoint loading through the
    ModelOpt dialect with an MXFP8 scheme now resolves its block scales correctly, not just this
    checkpoint.
  • The _rename wiring is local to qwen4_exp because that dialect has its own weight reader
    that blanket-skipped scale suffixes; other dialect readers already consult the storage table.
  • The NVFP4 PLE path is qwen4_exp-specific for the simple reason that only this family ships an
    n-gram PLE table; the unpack itself reuses the shared E2M1 LUT and e4m3_compat helpers.

Evaluated on

  • Hardware: RTX 5090 32 GB, 192 GiB DDR5 dual-channel, Ryzen 9850X3D, Linux, CUDA 13.2.
  • Command:
    ft serve --model local-inference-lab/Qwen3.8-Flash-Next-NVFP4 --max-seq-len-override 262144 --memory-ratio 0.94 --kv-reserve-tokens 262144 --ple-backend pinned
    (env PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True; 0.94 is tuned for 32 GB cards, lower it elsewhere).
  • Baseline for every row: the official RadixArk/Qwen3.8-Flash-Next-NVFP4 checkpoint measured
    by us on the same machine and the same pinned engine stack
    ; the QAD side additionally carries
    exactly the two commits of this PR - without them the checkpoint does not load at all.

Same-day A/B, one box, fixed seeds, both sides at --memory-ratio 0.93; the 0.94 in the
command above is the battle configuration we also run this checkpoint on (the bench rows below
are directly comparable):

metric RadixArk NVFP4 (official) this QAD checkpoint
disk size 126 GiB 99 GiB
ARC-400 (5-shot) 97.5% (390/400) in 370 s 97.2% (389/400) in 297 s
153-task diverse suite, total 103/153 102/153
- GPQA-Main (24, medium effort) 14/24, suite 891 s 15/24, suite 691 s (-22%)
- MMLU-Pro (60, 10-option, none) 39/60 38/60
- IFEval strict (44, none) 25/44 24/44
- HumanEval (25, sandbox, medium) 25/25 25/25

Decode, from passive engine-side bs=1 step windows in the logs (long agent sessions, 130-190k
context; same log category both columns, different days - treat as indicative, the bench rows
above are the controlled comparison):

decode median RadixArk NVFP4 this QAD checkpoint
130-150k tokens 64.2 t/s 75.8 t/s (+18%)
>150k tokens 70.2 t/s 76.9 t/s (+10%)

Read of the whole table: quality is parity within noise (the two benchmark rows differ by 1 task
in each direction, and the four differing ARC items split 2/2), while the QAD checkpoint is
27 GiB smaller and faster - MXFP8 dense + NVFP4 PLE put fewer bytes on a bandwidth-bound decode
path, and the bench wall clocks above are the same-day proof that speed did not cost accuracy.

Notes

  • The modelopt.py one-liner changes MXFP8 scale storage for the ModelOpt dialect as a whole.
    It matches the exporter's naming in this checkpoint family; FP8_BLOCK keeps
    weight_scale_inv. If another ModelOpt MXFP8 export stores the scales under
    weight_scale_inv, the rename falls through and the missing key surfaces as it did before
    this PR - no silent wrong-scale load.
  • reasoning_effort/tool-calling contract is unchanged (same architecture, same chat template).

Eva (port) added 2 commits October 3, 2026 19:17
…t_scale

ModelOpt exports MXFP8 block scales as '<proj>.weight_scale' while the scheme
loads them as 'weight_scale_inv': the reader skipped the suffix and the module
KeyError'd mid WeightLoad. The dialect's storage table now carries the truth
(MXFP8 -> .weight_scale) and _rename consults it instead of blanket-skipping,
so the scale fuses and validates on the existing block-scale path. Covers the
split GDN layout with b|a quantized (local-inference-lab/Qwen3.8-Flash-Next-NVFP4).
KarrAcaRn pushed a commit to KarrAcaRn/FreeToken-ByAI that referenced this pull request Oct 3, 2026
…QAD checkpoints (NVFP4 PLE tables, MXFP8 dense)

Conflict: main had split load_ple_table into the validating scan_ple_table plus a host-RAM
admission check. The NVFP4 table (packed rows, per-shard block scales, global weight_scale_2)
now goes through that scan with the same checks, and load_ple_table fills both banks.

Assisted-by: Claude Opus 5.5
KarrAcaRn pushed a commit to KarrAcaRn/FreeToken-ByAI that referenced this pull request Oct 3, 2026
…fig's rows and index the oob lookup per row

The table scan now checks the row count against the n-gram config, and the cuda gather test assigned 8 ids into a 16-wide row.

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