Score a DSA sparse indexer in FP4 - #2216
Conversation
🏷️ CI GuideRuns automatically on every eligible PR before approval:
Heavy model tests:
|
f6555c1 to
2611718
Compare
2611718 to
99704a7
Compare
f575520 to
eb8044f
Compare
|
Reviewed at Provenance: findings marked [verified] are ones where I read the code at the cited line myself. Findings marked [reported] come from a finder pass whose mechanism is traced through the diff, but where I did not open the far end of the call chain — worth one more look before acting. Two of these I would treat as blocking, and the first has nothing to do with FP4. 1. [verified] Every non-MLA speculative-decode target crashes — on the default configuration, with no FP4 involved
# The draft's rows are its own, and a replay addresses whatever row
# count the schedule buffer holds -- see the publisher for why this is
# unguarded.
builder._publish_indexer_fp4_decode_schedule(attn_metadata, running_bs, 1)
The method is defined exactly once in the tree: and there is no base-class definition ( So Llama-3 + EAGLE3, Qwen3-Next + MTP, MiniMax-M3 eagle3 and gpt-oss all raise The comment explains why the row count is unguarded; it does not cover calling the method on a builder that does not have it. A no-op base-class method, or the same 2. [verified] DCP + FP4 prefill dequantizes every staged key with another key's exponent — and this PR's own test file documents the swizzle it ignores
page, row = slots // block_size, slots % block_size
data = kv_cache[page, :, :, row, :]
scale = kv_cache_scale[page, :, :, row] # plain row
...
page, row = token // block_size, token % block_size
staged[page, :, :, row, :] = data
staged_scale[page, :, :, row] = scale # plain rowBut the scale plane's trailing axis is swizzled. The writer This PR's own oracle encodes exactly that asymmetry — packed = kv_cache[phys, 0, group, pos].reshape(batch, ctx_len, HEAD_DIM // 2)
keys = _dequant(packed, kv_scale[phys, 0, group, (pos % 16) * 4 + pos // 16])Reading at physical row Trigger: It is also likely invisible in this PR's GSM8K DCP leg, because The fix must permute both the source read and the destination write; correcting one is not enough. 3. [verified] There is no
|
45ad00e to
b31edac
Compare
b31edac to
3c32026
Compare
|
The job does not report an accuracy shortfall — it crashes during gsm8k generation at ~36/1319 with an HSA memory access fault and the harness exits 2, so no The identical signature appears on this PR's base commit c74e8ce in two independent This PR's paths are not implicated: the engine log for the failing job shows |
|
Thanks — this was a good review. Both blockers were real. This comment is kept current in place rather than appended to, so it reflects the branch as pushed. Blockers
There is now a test that calls AccuracyYou were right that the leg could not certify this path. Three legs, same harness, same 1323 prompts (all above
Sample 1 read 0.45 points under FP8; the FP4 leg's own run-to-run spread is 0.46 / 0.45 points, so that gap sits inside the spread of the leg that produced it and did not reproduce. I would not invert that into "FP4 matches FP8" — n=2 only supports the weaker claim that this measurement cannot see a difference below about half a point, and sees no evidence of one. One result from the same batch is worth stating on its own: the pre-fix FP4 run, with the Reported items that held
Docs — row added for Messages — the fused-writer refusal now names the geometry, since only one of its four conjuncts is an environment variable and the other three are structural; pointing a NoPE indexer at Items I do not think should change here
The three chip gates — real, and worse for a reader than the count suggests, since Tests
|
3c32026 to
c1b699d
Compare
c1b699d to
8ee2745
Compare
aiter's indexer_qk_rope_quant_and_cache gained an FP4 output mode, so a sparse indexer can keep its keys as packed E2M1 plus e8m0 scales and let flydsl_pa_mqa_logits_fp4 and its prefill variant score them straight out of the paged cache. FP4 is a change of dtype for the existing index region, not a second cache: the pool declares that one region in two planes, and prefill and decode both read it in place. The gate is structural rather than per-model -- a DSA indexer with index_head_dim 128, an index head count that is a multiple of MFMA_M, and gfx950 -- so any model of that shape gets it, and a mismatch warns once and falls back to FP8. DeepSeek-V4 already scored FP4 through the same kernels; its prefill schedule precompute is now the shared fp4_prefill_schedule. Decode takes a lower persistent-grid floor than prefill. Its rows are sequences rather than query tokens, so the floor always binds and every CTA past the work is pure setup: at 4096 units the scorer cost 8.75 us a layer against 2.99 at 512, flat in context, which is 21 layers x 5.8 us of grid overhead a step. The schedule itself is built once a step in the metadata builder at a CUDAGraph-stable address, which is what lets a captured graph replay it across changing context lengths. Needs the FP4 output mode of the aiter op, on branch feat/indexer-qk-rope-fp4-out. Per decode step the FP4 path runs exactly one kernel the FP8 path does not, the 4.8 us schedule publish; its scorer substitutes for the FP8 scorer at +2.3 us, and no other kernel count differs. Against FP8 the step is level at 4096 tokens (+0.020 ms at batch 1, +0.098 at batch 8) and ahead beyond it: -0.450 ms at 131072 batch 8, -0.709 at 131072 batch 16, -0.934 at 262144 batch 8. The index region also prices a block 102144 B lower across 21 indexer layers, buying 3.3 percent more KV blocks. GSM8K 1319 at 20-shot, prompts 2634 to 4295 tokens so every request clears the index_topk short-circuit and the prefill indexer runs: FP4 scores 0.9666 flexible and 0.9674 strict against FP8's 0.9682 and 0.9682, a difference of one to two questions. At 5-shot the sign is reversed, which is what noise looks like and a systematic quantization loss does not. Component agreement against an fp32 oracle over the bytes the writer produced is exact for decode at next_n 1 and 4 and for prefill on both schedule paths, and HIP-Graph replay is bitwise stable across four context-length vectors. The FP8 default is byte-identical at the operator level.
The FP4 indexer refused all context parallelism: its decode scores through what was an FP8-only candidate exchange, and its prefill reads the index cache in place, which under DCP is only this rank's 1/W of the sequence. Decode. Only the scoring call inside dcp_decode_candidate_exchange_fused is dtype-bound -- the local top-k, the exchange and the merge all read fp32 logits -- so the op takes optional q_scale/kv_scale and swaps that one call for flydsl_pa_mqa_logits_fp4. The FP4 kernel needs a CTA schedule, which is baked into the capture, so the metadata builder publishes one over the local window lengths at the sharded width, at next_n=1 because the exchange flattens (batch, next_n) into rows of query tokens. Prefill. FP8 gathers the per-rank shards into one flat plane and scores that. FP4 has no gather op for its E2M1/e8m0 planes and every FP4 mqa-logits kernel is paged, so _dcp_stage_indexer_fp4_prefill does the same local-shard read, all-gather and de-interleave, landing in a paged staging buffer with an identity page table. Over an identity table score column j is flat KV index j, the column space cu_seqlen_ks/ke and the DCP prefill filter already speak, so everything downstream is reused unchanged. The staged table keeps the real block table's fixed width, since the scorer specializes on its stride, and raises when a co-scheduled prefill context would need more pages than one sequence's allowance. MTP. The draft reindexes verify metadata from bs*next_n rows down to bs, but a replay addresses the rows its schedule was built for, so eagle_proposer republishes the schedule after the reindex. PCP stays refused: its candidate exchange is the one reader of the fp32 weights the FP4 writer does not produce. The FP8 default path is untouched. Only five lines of this change are reachable under it, and each is a no-op there: weights_mqa aliases weights when indexer_fp4 is false, the new total_kv argument is pure host arithmetic consumed past an _indexer_fp4 guard, both schedule publishers return on that guard, the FP4 scoring branch is gated on q_scale, and dcp_local_logits_width is the identity -(-a//b) == (a+b-1)//b. GSM8K 20-shot over 1323 prompts all >= 2048 tokens (min 2634 / p50 3382 / max 4295), TP4 + DCP4: FP4 exact_match 0.9704 +/- 0.0047 against FP8 0.9682/0.9689 +/- 0.0048, with fp4-fallback-warnings-final 0 and wall time 1981s against 1972s.
`EagleProposer` refreshes the FP4 decode CTA schedule on whatever builder the target uses, but the method existed only on `AiterMLAMetadataBuilder`. Nothing else inherits it, so every non-MLA target raised `AttributeError` at the tail of draft step 0 -- EAGLE3 on Llama-3, MTP on Qwen3-Next, MiniMax-M3, gpt-oss -- on the default configuration, with FP4 not involved. Gating the call on `target_uses_mla`, the idiom used a few lines above, would not have been enough: `DeepseekV4AttentionMetadataBuilder` is an MLA model whose builder is a `CommonAttentionBuilder` sibling rather than an MLA subclass, so DeepSeek-V4 + MTP would still have crashed. The no-op goes on the common base, which is total over the backends a draft can reach.
…der DCP `_dcp_stage_indexer_fp4_prefill` addressed both index planes with one flat in-block row, but they do not share a row order: the packed E2M1 plane is flat, while the e8m0 plane stores that axis as a 16x4 transpose. Composed over the DCP gather, a flat read paired with a flat write is the identity only where the source row equals the destination row, which de-interleaving is precisely what breaks. Every index stayed in bounds, so keys came back wearing another row's exponent with no exception and no shape error, for any fp4 prefill longer than `index_topk` at `dcp_world_size > 1`. Both sides now bend through `fp4_index_scale_rows`, which becomes the one place the two planes' row axes are related; the test oracle had encoded the transpose privately and now reads it from there instead. GSM8K 20-shot, GLM-5.2-MXFP4, TP8/DCP4, 1323 prompts all longer than `index_topk`: 0.9636 flexible / 0.9644 strict, against 0.9682 / 0.9689 for the FP8 indexer on the same harness. The previously reported 0.9704 was measured with this bug live and is withdrawn.
Wrong rows, two callers. The DCP branch re-read the global local-length buffer from index 0. A TBO ubatch already publishes its own offset slice on its metadata, so read that through `get_published_dcp_local_context_lens`, the same accessor the scorer uses, and keep the global buffer as the fallback for callers that publish nothing -- leaving every non-ubatch path unchanged. Stale schedule, draft steps 1..K-1. The publish sat in `_enter_decode_metadata`, which runs only at step 0 and before `prepare_mtp_decode` refreshes the lengths it is built from. Publish once per step, immediately after that call; this is also the draft half of the wrong-rows bug above. Silent short schedule. `top_k_per_row_decode` bounds its scan by `local_ctx`, so the buffer's uninitialised tail is inert and a blanket `-inf` fill would buy nothing per step. What is real is a schedule that under-covers rows inside the window: `[:parallel_units]` truncates silently while `indexer_fp4_n_ctas` still reports the full count. Assert instead. Unsupported combinations. Prefill micro-batching cannot carry the FP4 schedule -- `split_attn_metadata` rebuilds metadata from declared fields, and a ubatch rebases the row ids the schedule encodes -- so refuse it as PCP already is; decode micro-batching is handled and stays allowed. A pooled indexer scores through `sparse_attn_indexer_kpool`, which has no FP4 arm, so the predicate now rejects `index_kpool > 1`. That also removes the only reachable route to a scorer block size the cache never used, and makes the block-size message name the quantity it actually checks. Comments and messages. `fp4_decode_parallel_units` is not monotonic in `next_n` -- `f(3)` is 513 against `f(4)`'s 512 -- so state the weaker `f(n) >= f(1)` that is true and is what lets one buffer serve both widths requested, with a test pinning the invariant and the counterexample. Refusals a legal schedule can hit. The staged FP4 prefill page table was one sequence's block allowance wide and raised past it. But prefix caching keeps a co-scheduled prefill's cached tokens off `max_num_batched_tokens`, so a summed context beyond `max_model_len` is schedulable, and the refusal was reachable mid-serving on a batch the scheduler is entitled to form. Size the table at the batch bound instead: 2 MB at the default 16 x 128k, against one more scorer variant than the non-DCP width, and nothing left to refuse. Uninitialised reads, second instance. This call site leaves `compute_all_q_rope` at its default, so the fused writer skips `slot < 0` rows outright -- under DCP as well -- while the decode scorer reads the full speculation width. `weights_mqa` therefore handed pad rows whatever the allocator held. Zero it, as the synthesised fp32 return already is, and zero is what an unread row's weight should be anyway. The FP8 fused arm allocates its own `weights_out` the same way and predates this work. Couplings made explicit. `fp4_index_scale_rows` took its lane count from the module's block-size constant while both callers pass rows taken modulo their own page size; a mismatch would stay in bounds and only mix exponents, and the one thing holding the two equal was a raise in the pool builder two functions away. It now takes the page size and refuses a mismatch itself. The fused-writer refusal named an env switch that cannot help a NoPE indexer, which has no fused kernel to enable, so it names the geometry as well. And the Indexer is reached by undoing the one name `Indexer.__init__` builds, which the code now says. Tests reach their skips through one `_import_or_skip` rather than `pytest.importorskip`, which only warns on a plain `ImportError` instead of a missing module. On a CPU runner that is exactly how these modules fail -- `aiter` imports but carries no `QuantType` -- and CI escalates the warning. Docs: a row for `sparse_indexer_fp4.py` in the model-ops table, and the block-size-64 requirement plus the gfx950, geometry, PCP and KV-transfer conditions written into the two config tables that contradicted the code. Co-authored-by: Cursor <cursoragent@cursor.com>
8ee2745 to
ef4dfec
Compare
Conflict was the indexer prefill call: main wrapped it in an FP4 branch (#2216), this branch added clean_logits=False to it. Kept main's branch and moved the argument onto the fp8 call inside it; top_k_per_row_prefill still receives the same row_starts/row_ends, which is what makes it safe.
* Score a DSA sparse indexer in FP4 aiter's indexer_qk_rope_quant_and_cache gained an FP4 output mode, so a sparse indexer can keep its keys as packed E2M1 plus e8m0 scales and let flydsl_pa_mqa_logits_fp4 and its prefill variant score them straight out of the paged cache. FP4 is a change of dtype for the existing index region, not a second cache: the pool declares that one region in two planes, and prefill and decode both read it in place. The gate is structural rather than per-model -- a DSA indexer with index_head_dim 128, an index head count that is a multiple of MFMA_M, and gfx950 -- so any model of that shape gets it, and a mismatch warns once and falls back to FP8. DeepSeek-V4 already scored FP4 through the same kernels; its prefill schedule precompute is now the shared fp4_prefill_schedule. Decode takes a lower persistent-grid floor than prefill. Its rows are sequences rather than query tokens, so the floor always binds and every CTA past the work is pure setup: at 4096 units the scorer cost 8.75 us a layer against 2.99 at 512, flat in context, which is 21 layers x 5.8 us of grid overhead a step. The schedule itself is built once a step in the metadata builder at a CUDAGraph-stable address, which is what lets a captured graph replay it across changing context lengths. Needs the FP4 output mode of the aiter op, on branch feat/indexer-qk-rope-fp4-out. Per decode step the FP4 path runs exactly one kernel the FP8 path does not, the 4.8 us schedule publish; its scorer substitutes for the FP8 scorer at +2.3 us, and no other kernel count differs. Against FP8 the step is level at 4096 tokens (+0.020 ms at batch 1, +0.098 at batch 8) and ahead beyond it: -0.450 ms at 131072 batch 8, -0.709 at 131072 batch 16, -0.934 at 262144 batch 8. The index region also prices a block 102144 B lower across 21 indexer layers, buying 3.3 percent more KV blocks. GSM8K 1319 at 20-shot, prompts 2634 to 4295 tokens so every request clears the index_topk short-circuit and the prefill indexer runs: FP4 scores 0.9666 flexible and 0.9674 strict against FP8's 0.9682 and 0.9682, a difference of one to two questions. At 5-shot the sign is reversed, which is what noise looks like and a systematic quantization loss does not. Component agreement against an fp32 oracle over the bytes the writer produced is exact for decode at next_n 1 and 4 and for prefill on both schedule paths, and HIP-Graph replay is bitwise stable across four context-length vectors. The FP8 default is byte-identical at the operator level. * Support DCP under the FP4 sparse indexer The FP4 indexer refused all context parallelism: its decode scores through what was an FP8-only candidate exchange, and its prefill reads the index cache in place, which under DCP is only this rank's 1/W of the sequence. Decode. Only the scoring call inside dcp_decode_candidate_exchange_fused is dtype-bound -- the local top-k, the exchange and the merge all read fp32 logits -- so the op takes optional q_scale/kv_scale and swaps that one call for flydsl_pa_mqa_logits_fp4. The FP4 kernel needs a CTA schedule, which is baked into the capture, so the metadata builder publishes one over the local window lengths at the sharded width, at next_n=1 because the exchange flattens (batch, next_n) into rows of query tokens. Prefill. FP8 gathers the per-rank shards into one flat plane and scores that. FP4 has no gather op for its E2M1/e8m0 planes and every FP4 mqa-logits kernel is paged, so _dcp_stage_indexer_fp4_prefill does the same local-shard read, all-gather and de-interleave, landing in a paged staging buffer with an identity page table. Over an identity table score column j is flat KV index j, the column space cu_seqlen_ks/ke and the DCP prefill filter already speak, so everything downstream is reused unchanged. The staged table keeps the real block table's fixed width, since the scorer specializes on its stride, and raises when a co-scheduled prefill context would need more pages than one sequence's allowance. MTP. The draft reindexes verify metadata from bs*next_n rows down to bs, but a replay addresses the rows its schedule was built for, so eagle_proposer republishes the schedule after the reindex. PCP stays refused: its candidate exchange is the one reader of the fp32 weights the FP4 writer does not produce. The FP8 default path is untouched. Only five lines of this change are reachable under it, and each is a no-op there: weights_mqa aliases weights when indexer_fp4 is false, the new total_kv argument is pure host arithmetic consumed past an _indexer_fp4 guard, both schedule publishers return on that guard, the FP4 scoring branch is gated on q_scale, and dcp_local_logits_width is the identity -(-a//b) == (a+b-1)//b. GSM8K 20-shot over 1323 prompts all >= 2048 tokens (min 2634 / p50 3382 / max 4295), TP4 + DCP4: FP4 exact_match 0.9704 +/- 0.0047 against FP8 0.9682/0.9689 +/- 0.0048, with fp4-fallback-warnings-final 0 and wall time 1981s against 1972s. * Answer the draft's FP4 schedule publish on every attention backend `EagleProposer` refreshes the FP4 decode CTA schedule on whatever builder the target uses, but the method existed only on `AiterMLAMetadataBuilder`. Nothing else inherits it, so every non-MLA target raised `AttributeError` at the tail of draft step 0 -- EAGLE3 on Llama-3, MTP on Qwen3-Next, MiniMax-M3, gpt-oss -- on the default configuration, with FP4 not involved. Gating the call on `target_uses_mla`, the idiom used a few lines above, would not have been enough: `DeepseekV4AttentionMetadataBuilder` is an MLA model whose builder is a `CommonAttentionBuilder` sibling rather than an MLA subclass, so DeepSeek-V4 + MTP would still have crashed. The no-op goes on the common base, which is total over the backends a draft can reach. * Use the index scale plane's own row order when staging FP4 prefill under DCP `_dcp_stage_indexer_fp4_prefill` addressed both index planes with one flat in-block row, but they do not share a row order: the packed E2M1 plane is flat, while the e8m0 plane stores that axis as a 16x4 transpose. Composed over the DCP gather, a flat read paired with a flat write is the identity only where the source row equals the destination row, which de-interleaving is precisely what breaks. Every index stayed in bounds, so keys came back wearing another row's exponent with no exception and no shape error, for any fp4 prefill longer than `index_topk` at `dcp_world_size > 1`. Both sides now bend through `fp4_index_scale_rows`, which becomes the one place the two planes' row axes are related; the test oracle had encoded the transpose privately and now reads it from there instead. GSM8K 20-shot, GLM-5.2-MXFP4, TP8/DCP4, 1323 prompts all longer than `index_topk`: 0.9636 flexible / 0.9644 strict, against 0.9682 / 0.9689 for the FP8 indexer on the same harness. The previously reported 0.9704 was measured with this bug live and is withdrawn. * Schedule, refuse and document the FP4 indexer paths the review found Wrong rows, two callers. The DCP branch re-read the global local-length buffer from index 0. A TBO ubatch already publishes its own offset slice on its metadata, so read that through `get_published_dcp_local_context_lens`, the same accessor the scorer uses, and keep the global buffer as the fallback for callers that publish nothing -- leaving every non-ubatch path unchanged. Stale schedule, draft steps 1..K-1. The publish sat in `_enter_decode_metadata`, which runs only at step 0 and before `prepare_mtp_decode` refreshes the lengths it is built from. Publish once per step, immediately after that call; this is also the draft half of the wrong-rows bug above. Silent short schedule. `top_k_per_row_decode` bounds its scan by `local_ctx`, so the buffer's uninitialised tail is inert and a blanket `-inf` fill would buy nothing per step. What is real is a schedule that under-covers rows inside the window: `[:parallel_units]` truncates silently while `indexer_fp4_n_ctas` still reports the full count. Assert instead. Unsupported combinations. Prefill micro-batching cannot carry the FP4 schedule -- `split_attn_metadata` rebuilds metadata from declared fields, and a ubatch rebases the row ids the schedule encodes -- so refuse it as PCP already is; decode micro-batching is handled and stays allowed. A pooled indexer scores through `sparse_attn_indexer_kpool`, which has no FP4 arm, so the predicate now rejects `index_kpool > 1`. That also removes the only reachable route to a scorer block size the cache never used, and makes the block-size message name the quantity it actually checks. Comments and messages. `fp4_decode_parallel_units` is not monotonic in `next_n` -- `f(3)` is 513 against `f(4)`'s 512 -- so state the weaker `f(n) >= f(1)` that is true and is what lets one buffer serve both widths requested, with a test pinning the invariant and the counterexample. Refusals a legal schedule can hit. The staged FP4 prefill page table was one sequence's block allowance wide and raised past it. But prefix caching keeps a co-scheduled prefill's cached tokens off `max_num_batched_tokens`, so a summed context beyond `max_model_len` is schedulable, and the refusal was reachable mid-serving on a batch the scheduler is entitled to form. Size the table at the batch bound instead: 2 MB at the default 16 x 128k, against one more scorer variant than the non-DCP width, and nothing left to refuse. Uninitialised reads, second instance. This call site leaves `compute_all_q_rope` at its default, so the fused writer skips `slot < 0` rows outright -- under DCP as well -- while the decode scorer reads the full speculation width. `weights_mqa` therefore handed pad rows whatever the allocator held. Zero it, as the synthesised fp32 return already is, and zero is what an unread row's weight should be anyway. The FP8 fused arm allocates its own `weights_out` the same way and predates this work. Couplings made explicit. `fp4_index_scale_rows` took its lane count from the module's block-size constant while both callers pass rows taken modulo their own page size; a mismatch would stay in bounds and only mix exponents, and the one thing holding the two equal was a raise in the pool builder two functions away. It now takes the page size and refuses a mismatch itself. The fused-writer refusal named an env switch that cannot help a NoPE indexer, which has no fused kernel to enable, so it names the geometry as well. And the Indexer is reached by undoing the one name `Indexer.__init__` builds, which the code now says. Tests reach their skips through one `_import_or_skip` rather than `pytest.importorskip`, which only warns on a plain `ImportError` instead of a missing module. On a CPU runner that is exactly how these modules fail -- `aiter` imports but carries no `QuantType` -- and CI escalates the warning. Docs: a row for `sparse_indexer_fp4.py` in the model-ops table, and the block-size-64 requirement plus the gfx950, geometry, PCP and KV-transfer conditions written into the two config tables that contradicted the code. Co-authored-by: Cursor <cursoragent@cursor.com> --------- Co-authored-by: Cursor <cursoragent@cursor.com>
Lets a DSA sparse indexer keep its keys in FP4. Depends on ROCm/aiter#5484, which adds the FP4 output mode to
indexer_qk_rope_quant_and_cache; with it, the indexer stores packed E2M1 plus e8m0 scales andflydsl_pa_mqa_logits_fp4and its prefill variant score them straight out of the paged cache.FP4 is a change of dtype for the existing index region, not a second cache: the pool declares that one region in two planes, and prefill and decode both read it in place.
Scope
The gate is structural rather than per-model -- a DSA indexer with
index_head_dim == 128, an index head count that is a multiple ofMFMA_M, and gfx950 -- so any model of that shape gets it, and a mismatch warns once and falls back to FP8. There is nomodel_typetest in the path.DeepSeek-V4 already scored FP4 through the same kernels, so its prefill schedule precompute became the shared
fp4_prefill_schedule.deepseek_v4.pyis unchanged (0 differing lines);deepseek_v4_attn.pydiffers in 2 of 157 definitions, both the merge. Constants live once insparse_indexer_fp4.pyandv4_kernelsre-exportsFP4_MQA_BLOCK_K/FP4_MQA_PARALLEL_UNIT_NUM, so every DSV4 import is untouched.Decode grid floor
Decode takes a lower persistent-grid floor than prefill (512 against 4096). Its rows are sequences rather than query tokens, so
max(floor, rows)never lifts off the floor and every CTA past the work is pure setup: the scorer cost 8.75 us a layer at 4096 units against 2.99 at 512, flat in context, which is 21 layers x 5.8 us of grid overhead a step. Prefill keeps 4096, which its wide rows measured faster at.The schedule is built once a step in the metadata builder at a CUDAGraph-stable address, which is what lets a captured graph replay it across changing context lengths.
Performance
Per decode step the FP4 path runs exactly one kernel the FP8 path does not -- the schedule publish, 4.8 us. Its scorer substitutes for the FP8 scorer at +2.3 us, and no other kernel count differs (verified by differencing two profiles over an identical prompt at 8 and 40 output tokens, so prefill cancels).
Pooled median ITL, 2 seeds,
benchmark_serving.pywith fixed-length random prompts, prefix caching off, TP4, no MTP:At 262144 x bs8 that is output throughput +11.9% and TTFT -10.8%. The index region also prices a block 102,144 B lower across 21 indexer layers, buying 3.3% more KV blocks (50,671 against 49,084).
The saving is confirmed to be the index scan: isolating the context-dependent part of the step gives an FP4/FP8 ratio of 0.41-0.52 against a byte ratio of 0.47, and the indexer kernel microbenchmark predicts the measured per-step difference within 10%.
Accuracy
GSM8K 1319 at 20-shot, prompts 2634-4295 tokens so every request clears the
index_topkshort-circuit and the prefill indexer genuinely runs, TP4, no MTP:A difference of one to two questions out of 1319. At 5-shot the sign is reversed (FP4 0.9682/0.9689 against FP8 0.9651/0.9659), which is what noise looks like and a systematic quantization loss does not. Zero indexer fallback warnings in either leg.
Component agreement against an fp32 oracle over the bytes the writer produced is exact for decode at next_n 1 and 4 and for prefill on both schedule paths. HIP-Graph replay is bitwise stable across four context-length vectors, including re-replaying the capture batch after the others.
Decode context parallelism
Second commit. FP4 previously refused all context parallelism, so the GLM-5.2 agentic recipe could not run it at TP4 + DCP4; it now serves DCP, and PCP stays refused.
Only the scoring call inside
dcp_decode_candidate_exchange_fusedis dtype-bound -- the local top-k, the exchange and the merge all read fp32 logits -- so the op takes optionalq_scale/kv_scaleand swaps that one call forflydsl_pa_mqa_logits_fp4. The FP4 kernel needs a CTA schedule, which is baked into the capture, so the metadata builder publishes one over the local window lengths at the sharded width, atnext_n=1because the exchange flattens(batch, next_n)into rows of query tokens.Prefill is the harder half. FP8 gathers the per-rank shards into one flat plane (
cp_gather_indexer_k_quant_cache) and scores that; FP4 has no gather op for its E2M1/e8m0 planes, and every FP4 mqa-logits kernel is paged rather than flat. So_dcp_stage_indexer_fp4_prefillperforms the same three steps -- local-shard read, all-gather, de-interleave to global order -- and lands in a paged staging buffer with an identity page table. Over an identity table, score columnjis flat KV indexj, which is the column spacecu_seqlen_ks/cu_seqlen_keand the DCP prefill filter already speak, so every downstream step is reused unchanged.The staged page table keeps the real block table's fixed width, because the scorer specializes on its stride and a per-batch width would recompile it and disturb the capture. That width is
ceil(max_model_len / block)while the page count it has to cover is batch-sized, so the two can cross; it now raises with both numbers named rather than failing on a shape mismatch deeper in. A unit test drives the crossing directly and pins that below the bound the table is unchanged -- identity overpages, zeros in the tail.Under MTP the draft reindexes verify metadata from
bs*next_nrows down tobs, but a replay addresses the rows its schedule was built for, soeagle_proposerrepublishes the schedule after the reindex. Without that, decode faulted the GPU once capture completed. The publish is answered by a no-op on the common builder base rather than guarded at the call site: a draft also runs against backends that have no FP4 indexer, and the MLA-ness of the target does not decide it --DeepseekV4AttentionMetadataBuilderis MLA and is not anAiterMLAMetadataBuilder.Accuracy under DCP
GSM8K 20-shot over 1323 prompts, every one >= 2048 tokens (min 2634 / p50 3382 / max 4295) so the
index_topkshort-circuit never fires, TP4 +dcp=4, identical harness on both legs, no MTP:FP4 lands 0.45 points below FP8, about 0.65 sigma on the combined error, and
fp4-fallback-warnings-finalis 0, so the FP4 path really was taken on every request rather than quietly falling back.An earlier revision of this PR reported 0.9704 in this row. That leg ran before the scale-plane fix in
_dcp_stage_indexer_fp4_prefilland is withdrawn: it scored keys against neighbouring rows' exponents. It looked healthy becausek_normruns before the quantizer, so a block's e8m0 values are nearly uniform and the wrong exponent is usually close to the right one -- which is precisely why that number could not serve as evidence the path was correct. The row above is the fixed path, same harness and same 1323 prompts.FP8 default unchanged
Byte-identical at the operator level:
q_out,weights_out,kv_cachefrom the widened fused writer andconv_outfrom the Triton prefill converter whose kernel gained aSEQ_LOCALconstexpr, 6.8 MB total, all equal between main and this branch over identical input bytes. The pool'sentry_bytes, field groups, view names and shapes, transfer roles, the Indexer's parameters, buffers, modules andforward_implsource text, and the exact argument list the FP8 writer is called with are all identical; the only differences are defaulted trailing parameters.Worth stating as a negative result: end-to-end text cannot serve as evidence here, because two servers built from byte-identical trees disagree on greedy completions, so generation diffing has no power in either direction. Under DCP this was measured rather than assumed. A fixed three-prompt greedy set, served one request at a time, hashing the 256-token completion, run three times on a tree with no local modifications:
cf707d07cb05891eac80dd997e9d9470dca46aeef7f7953935dac0854000ebc02a8a4ca0Three runs, three different answers, on every prompt -- so a base-against-patched difference here carries no information.
The DCP commit's default path is therefore argued by reachability. Only five of its lines execute under FP8, and each is a no-op there:
weights_mqaaliasesweightswhenindexer_fp4is false; the newtotal_kvargument is numpy host arithmetic whose value is consumed only past an_indexer_fp4guard; both schedule publishers return on that guard; the FP4 scoring branch is gated onq_scale, and its FP8 arm is a verbatim copy of the original call; anddcp_local_logits_widthis the identity-(-a//b) == (a+b-1)//b, checked over 19,500,000 pairs with no mismatch. Everything else it adds sits behindif self._indexer_fp4.Notes
PCP is refused under FP4 because its candidate exchange is the one reader of the fp32
weightsthe FP4 writer does not produce. KV transfer is refused because the region map cannot describe the separate e8m0 scale plane.next_n > 1is covered at component level with exact agreement, and end-to-end under DCP only as a smoke run -- the accuracy legs above run without MTP.Pre-existing and deliberately not touched here: every server log, on this branch and on three clean-HEAD legs alike, reports
21/1188 model parameters were NOT loaded, all of themmodel.layers.*.self_attn.indexer.wk_weights_proj.weight_scale. It is upstream, it affects both legs of every comparison identically, and it is out of scope for this PR.