feat(kvcache): add nvfp4 kv quantization - #408
Conversation
… fp8)
One (token, kv head) row of K and of V becomes head_dim e4m3 codes plus ONE
fp32 symmetric scale, in a code buffer with exactly the geometry of the 16-bit
KV buffer -- only the element type changes. That halves the bytes per cached
token (the scale sidecar costs 4/head_dim of it back, ~3% at head_dim 128), and
it is what lets Qwen3.8-Flash-Next serve a 1M-token context on this card.
Codes are kept in a plain uint8 buffer on EVERY architecture, and the fp8e4nv
type never appears in a kernel signature. Both ways of choosing that per target
failed on real hardware and are recorded here so nobody reopens them: the
compile-time fp8-native probe (e4m3_compat.e4m3_native_cx) answers the question
independently from the host that allocated the buffer and disagreed with it on
sm_100, and branching on a pointer's element type is NOT statically pruned --
triton still type-checked the dead arm, whose int mask fill is illegal against
an fp8 pointer ("cannot cast int32 to fp8e4nv", raised at CUDA graph capture).
What remains is the software encode/decode that already runs wherever the fp8
type is unavailable and is bit-exact per e4m3_compat's header, so the cache
holds the same bytes and produces the same numbers on every card (docs/cli.md).
- server/args.py, engine/config.py: --kv-cache-dtype {auto,bf16,fp8}, refused at
startup for the pools and backends that cannot apply the row scales
(attention/__init__.py: BackendInfo.supports_fp8_kv) rather than ignored.
- kernel/triton/kv_quant.py: fused quantize+scatter -- one launch under CUDA
graph capture, where the slot ids arrive as a device tensor.
- kvcache: unit_bytes() counts codes plus the scale sidecar, so ft ctl stats and
cache --kv N follow the smaller footprint, and rebuild reallocates the scale
buffers alongside the codes (mha, hybrid-SWA and QSA pools).
- kvcache/base.py: pool.dtype is the COMPUTE dtype -- what store_kv receives and
what a backend sizes its scratch with -- while pool.store_dtype is what the
buffer holds. Reporting codes as dtype handed e4m3 to QSA's 16-bit indexer and
died compiling qsa_mqa_paged; the contract is now asserted at backend init and
in the kernel wrapper. QSA's block-selection keys stay 16-bit: only the
selected K/V rows are read back as codes.
Tested on: sm_100, 148 SMs, Linux; 524,480 fp8 KV tokens = 6.47 GiB,
Qwen3.8-Flash-Next with: ft serve --kv-cache-dtype fp8 -> 1M-token context.
Covered by tests/kernels/test_kv_fp8.py, tests/kernels/test_qsa_fp8.py,
tests/kernels/test_triton_attention.py, tests/kernels/test_e4m3_compat.py,
tests/kvcache/test_mha_pool_fp8.py, tests/kvcache/test_qsa_pool_fp8.py and
tests/engine/test_kv_quant_config.py (CUDA-gated; not run on the Windows
development box, which has neither triton nor pytest installed).
Not included here, on purpose: unifying the two fp8-native probes (triton's
cache-key walk rejects a constexpr function that defers to a host one, so
warn_if_probes_disagree() reports the disagreement instead), and a hardware
decode fast path on sm_89+ (that needs a constexpr flag threaded from the host
plus the matching AOT variants, since testing the dtype does not prune).
`quantize_kv_to_cache` passed `k.stride(0)` as the only source pitch and
`_kv_quant_scatter_kernel` used it for both tensors:
src = t * stride_xs + h * D + d
xk = tl.load(k_src + src, ...)
xv = tl.load(v_src + src, ...)
The guard above it checks only the inner stride (`k.stride(1) == 1 and
v.stride(1) == 1`), never `k.stride(0) == v.stride(0)`, so the kernel carries an
undocumented contract: K and V must share one row pitch.
When they do not, V is read at K's pitch. In the failing test K is a view of the
qkv slice (pitch 1152) while V is materialised by `.clamp()` (pitch 384), so with
8 tokens of 3072 elements:
token 0 reads 0 correct by coincidence
token 1-2 reads 1152, 2304 in range, WRONG rows
token 3-7 reads 3456 .. 8064 past the initialised data
2684 of 3072 codes wrong, all in V, K byte-perfect. Deterministic addressing;
only the contents of the uninitialised tail vary with allocator history, which is
why the mismatch count drifts (2684 / 2663 / 2676 across runs) while the
mismatching positions do not -- the in-range half is exactly 764 every time.
Found with `compute-sanitizer --tool initcheck` (TRITON_DISABLE_LINE_INFO=0),
which named `kv_quant.py:110`. `memcheck` reports 0 errors because PyTorch's
caching allocator rounds allocations up and the bad read stays inside the pooled
segment; `racecheck` reports 0 hazards because it is not a race.
Fix: pass `v.stride(0)` as its own kernel argument and load each tensor with its
own pitch.
Verified on RTX 4090 (sm_89): tests/kernels/test_kv_fp8.py 2 failed -> 1 failed,
the flip being test_codes_match_the_reference_quantizer_and_reconstruction_is_close;
five consecutive standalone runs give K 0/3072 and V 0/3072 with got.sort() ==
exp.sort(); compute-sanitizer initcheck reports 0 errors on the patched build.
Independently confirmed on RTX 5090 D (sm_120) by @Kaempferia: same single flip,
same multiset property, 5 runs clean.
Note for reviewers: the test's SECOND assertion (dequantised error <= 0.08)
passes at 0.035 while V is wrong, so a reconstruction-level check does not catch
this class. Only the exact-code assertion does.
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
fix(kernels): give the V tensor its own row pitch in the fp8 KV store
`test_store_kv_writes_the_slot_the_attend_kernel_will_read` reaches row 256 -- its comment asks for the 255/256 page boundary -- against a four-page, 256-slot pool, so `codes[out_loc]` gathers one row past the view. Five pages is what that row list needs. The gather raises a device-side assert, and the CUDA context does not recover from it, so in a single-process run everything scheduled afterwards is reported as failing as well -- not because those tests stop working, but because there is no longer a context to run them in. One file per process does not show that. Assisted-by: Claude Opus 5
`test_encoder_inverts_the_grid_through_the_scale_one_path` moves `rows` to the device and leaves the V tensor beside it on the host, so the kernel is handed a CPU pointer. Assisted-by: Claude Opus 5
…its ids `test_layer_ids_remap_applies_to_scales_too` backs `layer_ids=(1, 3)` on a pool built with `num_layers=LAYERS`, and LAYERS is 3, so id 3 is one past the end and the constructor raises before the assertion it is there to make. The helper now takes the depth. Assisted-by: Claude Opus 5
…the backend now reads `k_scale` / `v_scale` arrived with the fp8 store, and the backend reads them on every path, including a 16-bit pool -- which answers None. The two hand-rolled `FakeKVCache` classes in this file do not inherit the base pool, so they were left without them and raise `AttributeError` instead. Assisted-by: Claude Opus 5
…e tile
The per-(token, kv_head) dequant scale is constant down each dot's reduction
dim, so it never has to touch K or V:
scores[m,n] = (sum_d q[m,d] * k[d,n]) * s_k[n]
p @ (diag(s_v) @ v) = (p * s_v[None,:]) @ v
Scaling the BLOCK_M x BLOCK_N result instead of the BLOCK_D x BLOCK_N K tile
and the BLOCK_N x BLOCK_DV V tile is head_dim/BLOCK_M fewer multiplies. p
itself stays unscaled, since l_i accumulates it as the softmax denominator and
knows nothing about V's quantization.
That also removes the only reason the tile was widened to fp32.
kv_load_e4m3_tile_f32 builds an fp16 bit pattern and widens purely so a
* 256.0 can put the value back on the true e4m3 scale, and 2^8 is a power of
two, so once the scale rides the dot output it folds into that scale exactly.
kv_load_e4m3_tile_scaled16 stops before the widen and leaves the fold to the
caller, keeping the tile 16-bit through the whole loop.
Accuracy improves rather than degrades. The general scale used to multiply
before the narrow to the compute dtype, so the product rounded; now the tile
reaches the dot exactly (the code's own 3 mantissa bits, |x| <= 1.75) and the
scale is applied in fp32 afterwards. Worst-case absolute error in
test_extend_paged_attention_decodes_fp8_scales drops 0.281 -> 0.0996 on sm_86.
(That test still exceeds its 2e-2 tolerance on this card both before and after
-- it fails on 3e5bbdd unpatched too, so it is not introduced here.)
The loader's bit placement is also the same number in 4 ops instead of 7: for
v = 128s + r, ((v & 0x80) << 8) | ((v & 0x7F) << 7) and (v + (v & 0x80)) << 7
are both (256s + r) << 7. Verified identical on all 256 codes, NaN patterns
included, by the new test in tests/kernels/test_e4m3_compat.py.
RTX 3070 (sm_86, 8GB, driver 610.57.04), i7-11700KF, Qwen3.6-35B-A3B-NVFP4,
--moe-backend hybrid --kv-cache-dtype fp8 --max-seq-len-override 180000
--memory-ratio 0.9 --max-running-requests 1 --max-prefill-length 1024,
2 reps, median, 127 output tokens, unique nonce per request:
ctx TTFT 3e5bbdd -> here decode 3e5bbdd -> here
33k 34.13 -> 33.95 s 47.04 -> 48.05 t/s
65k 88.79 -> 85.03 s 38.99 -> 41.61 t/s
100k 168.27 -> 160.42 s 33.59 -> 37.13 t/s
Most of the prefill win needs the tile-sizing fix in the next commit; this one
is mainly a decode gain on its own.
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
_select_extend_tile budgets shared memory as
(BLOCK_M + 2 * BLOCK_N) * BLOCK_D * 2
which charges K and V at 2 bytes/element whatever the cache actually holds. The
q tile is always 2 bytes/element, but K and V follow the cache, so a 1-byte fp8
cache is billed for twice the shared memory it uses and falls through to a
smaller tile than it has room for. On an RTX 3070 (sm_86, 99KB opt-in) at
head_dim 256 that is BLOCK_N 32 where 64 fits.
Take the element size as a parameter and bill K/V at it:
(BLOCK_M * 2 + 2 * BLOCK_N * kv_bytes) * BLOCK_D
kv_bytes=2 is algebraically the previous expression, so every 16-bit cache
keeps the tile it had; the existing parametrisation in
test_select_extend_tile_is_shared_memory_aware still passes unchanged. The
head_dim <= 256 ladder gains a 64x64 rung between 128x64 and 64x32, which only
an fp8 cache can reach on a consumer card.
The budget stays a conservative proxy rather than an exact model. On this card
it correctly rejects both tiles that fail to launch (128x64 and 64x128, which
raise OutOfResources: shared memory, Required: 114688, Hardware limit: 101376)
and correctly accepts the two the ladder uses. It also rejects 128x32 and
32x128, which do launch -- but those are not on the ladder, and rejecting a
tile that would have worked only costs a smaller tile, never a failure.
Same setup as the previous commit, measured on top of it:
ctx TTFT before -> after decode before -> after
33k 33.95 -> 28.18 s 48.05 -> 47.94 t/s
65k 85.03 -> 64.36 s 41.61 -> 40.92 t/s
100k 160.42 -> 112.78 s 37.13 -> 36.26 t/s
Decode is untouched by this commit (it only moves the extend/prefill tile); the
small differences there are run-to-run noise at 2 reps.
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
test: keep the fp8 tests runnable in a single pytest process
…p4-kv-quantization
…p4-kv-quantization
Port of upstream PR FlashML-org#408's codec delta: packed E2M1 nibbles (low first), one E4M3 block scale per 16 values, one FP32 row scale, quantization against the STORED (rounded) block scale; MHAKVCache packs head_dim//2 bytes + block scales + row scales and stores through quantize_nvfp4_to_cache, with k/v_block_scale accessors and block-scale-aware unit_bytes. Memory contract: D=128 row = 64+8+4 = 76 B, K/V pair 152 B, and the planner's spec_kv_bytes_per_token equals the pool's allocation; head_dim % 16 is rejected. Tests (ported + extended): E2M1 grid/ties/boundaries oracle, scatter vs independent reference across dims and dtypes, CUDA-graph replay, budget/rebuild/layer-map, 76 B contract, divisibility. GPU on the RTX 3090 (sm_86): 304 passed, prod restored. Attention/SWA/backend readers for nvfp4 land with 5y2.15; the CLI still cannot select nvfp4 (engine aliases) until that task.
…didate Rebuilt: the previous eval branch was cut from fb5f3b4 and carried the lm_head slice that is now reverted, so its warmup crash was mine, not FlashML-org#408's. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_0173pf9k9fSVtwbm3f898HDt
|
Ran this on sm_89 (2 × RTX 6000 Ada, TP=2,
Accuracy is free. 294/300 both ways, same questions, greedy. On this model 4-bit KV costs nothing measurable, which is the result I was least confident about going in and the one most worth having. The capacity claim is almost exactly right. Before running it I sized the win from your 76 B/row figure and predicted 6.04 GiB/rank freed and 88.1% residency. Measured: 5.25 GiB and 88.2%. The throughput win is zero, and I think that is architectural on Ada. Our decode is PCIe-bound on the offload backend, so a 71.7% → 88.2% residency jump cuts expert fetch traffic ~58% and should have shown up. It does not, because nvfp4 dequantises on every attention read and sm_89 has no native FP4 math — the Triton unpack-and-scale in the hot path costs about what the residency saves. On Blackwell, where the dequant is cheaper or native, I would expect your gain to survive; this is not an argument against the PR, it is a note that the benefit is arch-dependent and a reader on Ada should size it as capacity, not speed. That framing is still valuable to us: 5.25 GiB/rank at flat tok/s and identical accuracy is budget we can spend on a larger KV pool or more concurrent requests. Worth saying explicitly in the PR text, because "3.4× smaller KV" reads as a throughput claim and on at least one architecture it is not one. Two things for the PR itself. It merges onto a current qwen4_exp deploy line with exactly one conflict, a comment in And a caution from my own mistake, in case it saves someone else the run: my first attempt reported that #408 crashed at 🤖 Generated with Claude Code |
Port of upstream PR FlashML-org#408's reader half (attention/triton.py, kernel/triton/attention.py, hybrid_swa_pool, qsa attend/qsa_sparse/qsa_pool, args + engine aliases/gates): packed E2M1 codes are decoded with their E4M3 block scales and fp32 row scales in the paged, split-k decode, extend, extend-split and QSA kernels, while index/GDN/PLE tiers stay 16-bit. BackendInfo.supports_nvfp4_kv is set only for triton and qsa_sparse; MLA/DSA stay rejected (backlog 5y2.18) and head_dim%16 is enforced config-time and in the factory. GPU on the box: focused 77 passed; wider 931 passed with 5 failures that A/B-revert to the pre-change tree (2 MTP test-design + 3 skeleton/driver-state issues, not this change).
Summary
Add opt-in
--kv-cache-dtype nvfp4KV-cache storage for paged MHA/GQA, hybrid-SWA, and QSA sparse attention.K/V use packed E2M1 values, one E4M3 scale per 16 values, and one FP32 row scale per token/KV head. QSA keeps its index, pending-ring, and scratch tiers in BF16; only paged K/V rows are quantized.
At head_dim 128, each K/V row uses 76 bytes versus 256 bytes for BF16. The implementation includes pool budgeting/rebuild, CUDA Graph-safe KV writes, Triton restore paths, configuration validation, and a microbenchmark.
Validation
Hardware: RTX 5090 (32607 MiB), driver 591.86, WSL2/Docker, PyTorch 2.11.0+cu130, Triton 3.6.0.