Skip to content

Score a DSA sparse indexer in FP4 - #2216

Merged
valarLip merged 5 commits into
mainfrom
perf-sparse-indexer-fp4
Sep 15, 2026
Merged

valarLip merged 5 commits into
mainfrom
perf-sparse-indexer-fp4

Conversation

@XiaobingSuper

@XiaobingSuper XiaobingSuper commented Sep 13, 2026 •

Copy link
Copy Markdown
Contributor

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 and 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.

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 of MFMA_M, and gfx950 -- so any model of that shape gets it, and a mismatch warns once and falls back to FP8. There is no model_type test 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.py is unchanged (0 differing lines); deepseek_v4_attn.py differs in 2 of 157 definitions, both the merge. Constants live once in sparse_indexer_fp4.py and v4_kernels re-exports FP4_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.py with fixed-length random prompts, prefix caching off, TP4, no MTP:

ctx bs FP8 FP4 delta
4096 1 11.157 11.177 +0.020 ms (+0.18%)
4096 8 12.512 12.610 +0.098 ms (+0.78%)
131072 8 13.445 12.995 -0.450 ms
131072 16 16.817 16.108 -0.709 ms
262144 8 14.503 13.569 -0.934 ms

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_topk short-circuit and the prefill indexer genuinely runs, TP4, no MTP:

flexible-extract strict-match
FP4 0.9666 +/- 0.0049 0.9674 +/- 0.0049
FP8 0.9682 +/- 0.0048 0.9682 +/- 0.0048

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_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 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_prefill performs 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 column j is flat KV index j, which is the column space cu_seqlen_ks/cu_seqlen_ke and 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 over pages, zeros in the tail.

Under 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. 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 -- DeepseekV4AttentionMetadataBuilder is MLA and is not an AiterMLAMetadataBuilder.

Accuracy under DCP

GSM8K 20-shot over 1323 prompts, every one >= 2048 tokens (min 2634 / p50 3382 / max 4295) so the index_topk short-circuit never fires, TP4 + dcp=4, identical harness on both legs, no MTP:

exact_match wall
FP4 + DCP4 0.9636 / 0.9644 +/- 0.0052 1915 s
FP8 + DCP4 0.9682 / 0.9689 +/- 0.0048 1972 s

FP4 lands 0.45 points below FP8, about 0.65 sigma on the combined error, and fp4-fallback-warnings-final is 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_prefill and is withdrawn: it scored keys against neighbouring rows' exponents. It looked healthy because k_norm runs 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_cache from the widened fused writer and conv_out from the Triton prefill converter whose kernel gained a SEQ_LOCAL constexpr, 6.8 MB total, all equal between main and this branch over identical input bytes. The pool's entry_bytes, field groups, view names and shapes, transfer roles, the Indexer's parameters, buffers, modules and forward_impl source 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:

leg prompt 0 (in=23) prompt 1 (in=3666) prompt 2 (in=5627)
base cf707d07 cb05891e ac80dd99
base1 7e9d9470 dca46aee f7f79539
base2 35dac085 4000ebc0 2a8a4ca0

Three 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_mqa aliases weights when indexer_fp4 is false; the new total_kv argument is numpy host arithmetic whose value is consumed only past an _indexer_fp4 guard; both schedule publishers return on that guard; the FP4 scoring branch is gated on q_scale, and its FP8 arm is a verbatim copy of the original call; and dcp_local_logits_width is the identity -(-a//b) == (a+b-1)//b, checked over 19,500,000 pairs with no mismatch. Everything else it adds sits behind if self._indexer_fp4.

Notes

PCP is refused under FP4 because its candidate exchange is the one reader of the fp32 weights the FP4 writer does not produce. KV transfer is refused because the region map cannot describe the separate e8m0 scale plane. next_n > 1 is 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 them model.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.

Copilot AI lite review requested due to automatic review settings September 13, 2026 12:49

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot was unable to review this pull request because the user who requested the review has reached their quota limit.

@github-actions

Copy link
Copy Markdown
Contributor

🏷️ CI Guide

Runs automatically on every eligible PR before approval:

  • ✅ Pre Checkin: Black, Ruff, catalog schema validation, non-GPU unit tests

Heavy model tests:

  • ✅ Run after the PR is approved and Pre Checkin passes
  • ✅ Run immediately when an approval review is submitted
  • ✅ Can be requested before approval with labels
Label Tests
ci:full Run all heavy PR model tests: native ATOM, vLLM, and SGLang
ci:atom Run native ATOM model accuracy tests
ci:vllm Run ATOM vLLM OOT model accuracy tests
ci:sglang Run ATOM SGLang model accuracy tests

Heavy jobs are skipped when the PR is not approved and no matching ci:* label is present.
Add labels via the sidebar or gh pr edit 2216 --add-label <label>

@XiaobingSuper
XiaobingSuper force-pushed the perf-sparse-indexer-fp4 branch from f6555c1 to 2611718 Compare September 13, 2026 14:25
Copilot AI review requested due to automatic review settings September 13, 2026 14:25

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot was unable to review this pull request because the user who requested the review has reached their quota limit.

@XiaobingSuper
XiaobingSuper force-pushed the perf-sparse-indexer-fp4 branch from 2611718 to 99704a7 Compare September 14, 2026 02:08
Copilot AI review requested due to automatic review settings September 14, 2026 02:08

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot was unable to review this pull request because the user who requested the review has reached their quota limit.

Copilot AI review requested due to automatic review settings September 14, 2026 13:05

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot was unable to review this pull request because the user who requested the review is ineligible. To be eligible to request a review, you need a paid Copilot license, or your organization must enable Copilot code review.

@gbyu-amd
gbyu-amd force-pushed the perf-sparse-indexer-fp4 branch from f575520 to eb8044f Compare September 14, 2026 13:14
Copilot AI review requested due to automatic review settings September 14, 2026 13:14

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot was unable to review this pull request because the user who requested the review is ineligible. To be eligible to request a review, you need a paid Copilot license, or your organization must enable Copilot code review.

@valarLip

Copy link
Copy Markdown
Collaborator

Reviewed at eb8044fa5.

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

eagle_proposer.py:482:

# 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)

builder is self.runner.attn_metadata_builder (line 442) — whatever backend the target uses. The surrounding code branches on target_uses_mla at 443/456/488, and the dcp_token_block_tables write fourteen lines above is if ... is not None guarded. This call is neither.

The method is defined exactly once in the tree:

$ grep -rn "def _publish_indexer_fp4_decode_schedule" atom/
atom/model_ops/attentions/aiter_mla.py:2231

and there is no base-class definition (_publish_indexer_fp4_decode_schedule does not appear in backends.py). The non-MLA builders do not inherit from AiterMLAMetadataBuilder:

AiterAttentionMetadataBuilder(CommonAttentionBuilder)
GDNAttentionMetadataBuilder(GDNStateMixin, AiterAttentionMetadataBuilder)
TritonMHAMetadataBuilder(AiterAttentionMetadataBuilder)
DeepseekV4AttentionMetadataBuilder(CommonAttentionBuilder)

So Llama-3 + EAGLE3, Qwen3-Next + MTP, MiniMax-M3 eagle3 and gpt-oss all raise AttributeError: '...MetadataBuilder' object has no attribute '_publish_indexer_fp4_decode_schedule' at the tail of draft step 0.

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 hasattr/is not None treatment the neighbouring line already uses.

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

_dcp_stage_indexer_fp4_prefill (deepseek_v2.py:1485-1499) uses the plain in-block row for the scale plane on both the read and the write:

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 row

But the scale plane's trailing axis is swizzled. The writer indexer_fp4_kv_scale_offset (aiter csrc/kernels/cache_kernels.cu:1454-1466) computes sflat = (pos_in_block % 16) * tiles_per_block + pos_in_block / 16, which at kv_block_size=64 is Q(r) = (r % 16) * 4 + r // 16; the reader loads the same swizzled byte, while the data offset uses the plain pos_in_block.

This PR's own oracle encodes exactly that asymmetry — tests/test_sparse_indexer_fp4.py:239-240:

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 r_src yields the exponent belonging to logical position Q⁻¹(r_src); writing it at physical row r_dst makes it the exponent for Q⁻¹(r_dst). The composition is the identity only when r_src == r_dst — which under DCP interleave is exactly what does not hold (the default DCPConfig.interleave_size is 1). Worked example at W=4, S=1: global token u=1 ends up reading the scale of global token 64, out of rank 0's physical block rather than rank 1's.

Trigger: dcp_world_size > 1 + --index-cache-dtype fp4 + any prefill above index_topk. No fault, no shape error.

It is also likely invisible in this PR's GSM8K DCP leg, because k_norm makes per-block e8m0 exponents nearly uniform — which is precisely what makes it dangerous rather than reassuring. There is no test: the only FP4 prefill GPU test is non-DCP and never calls this function.

The fix must permute both the source read and the destination write; correcting one is not enough.

3. [verified] There is no -inf sentinel any more, which converts every schedule mismatch below from a detectable NaN into plausible garbage

ATOM allocates the FP4 logits buffer with torch.empty at both sites:

deepseek_v2.py:1542   logits = torch.empty(q_fp4.shape[0], width, dtype=torch.float32, ...)
deepseek_v2.py:1956   logits = torch.empty([num_rows, max_model_len], dtype=torch.float32, ...)

and aiter fills -inf only when it builds the schedule itself:

# aiter pa_mqa_logits_fp4.py:708-721, pa_mqa_logits_fp4_prefill.py:900-918
schedule_internal = cta_info is None
...
elif schedule_internal:
    out.fill_(-inf)

All three ATOM call sites (prefill :1542, decode :1956, DCP decode dcp_ops.py:1000) supply a precomputed schedule, so the fill never runs. The deleted comment in deepseek_v4_attn.py stated the failure mode as "logits stay at the -inf/NaN pre-fill → wrong top-k"; that sentinel no longer exists on this path.

Any row or chunk the schedule drops therefore reads as recycled allocator memory, and top_k_per_row_* ranks it as a real score. At minimum a debug-mode fill, or an assertion that n_ctas covers q.shape[0].


4. [reported] The DCP decode schedule re-derives lengths from index 0, so it reads the wrong rows in two callers

aiter_mla.py:2253 reads forward_vars["dcp_local_context_lens"].gpu[: bs * next_n] — always from 0 — instead of the per-ubatch/per-step slice the caller already holds.

  • TBO: build_ubatch_metadata (2949-2955) deliberately offsets this very buffer by request_start = ubatch_idx * requests_per_ubatch and stores the result on attn.dcp_local_context_lens (2974), then calls the publisher at 3007 with ubatch=ubatch_idx+1. Ubatch 1's CTA schedule is built from ubatch 0's sequence lengths.
  • MTP draft: at the eagle call site the buffer still holds the VERIFY step's layout — one row per query token over scheduled_tokens = bs * max_seqlen_q (2397-2411). Reading [:bs] with next_n=1 takes sequence 0's K tokens plus sequence 1's first token, and so on. The per-request layout is written later, by prepare_mtp_decode (1090-1094), whose own comment says the verify copy is stale.

Either way the scorer covers the wrong KV column ranges for part of the batch. The non-DCP branch (2261) is correct because it reads attn_metadata.context_lens[:bs]; read attn_metadata.dcp_local_context_lens symmetrically.

5. [reported] Draft steps 1..K-1 score with a schedule built at step 0

_enter_decode_metadata, and with it the publish at eagle_proposer.py:482, runs only at i == 0 (eagle_proposer.py:731-738). The very next statements advance the lengths — attn_metadata.context_lens[:running_bs] += 1 / prepare_mtp_decode(update_context_lens=True) at 754-784 — and prepare_mtp_decode re-derives dcp_local_context_lens (aiter_mla.py:1086-1094) without republishing.

compute_varctx_schedule distributes CTAs over ceil(ctx / FP4_MQA_BLOCK_K=256) chunks, so with --num-speculative-tokens > 1, whenever a sequence's context crosses a 256 boundary during the draft loop the trailing chunk gets no CTA and its logits columns are never written — and per #3 those columns are allocator garbage, not -inf. Data-dependent, roughly 1-in-256 rows per step.

DeepSeek-V4 rebuilds the schedule on every build() (deepseek_v4_attn.py:_refresh_fp4_windows). This PR's accuracy legs were all run without MTP.

6. [reported] TBO prefill + FP4 dies on a missing attribute

split_attn_metadata (atom/utils/tbo/ubatch_splitting.py:476-505) constructs a fresh AttentionMetaData from declared dataclass fields. indexer_fp4_cta_info / _n_ctas / _local_starts / _local_ends / _max_seq_len and dcp_indexer_fp4_local_slots / _block_tables are set dynamically by the publisher and are not declared in atom/utils/forward_context.py:561.

The MLA override at aiter_mla.py:3012 carefully re-slices cu_seqlen_ks, cu_seqlen_ke, sparse_cu_seqlens_q, batch_id_per_q_token and sparse_kv_indptr — and adds nothing for the FP4 twins. With --enable-tbo + --index-cache-dtype fp4 on a DSA model, the first sparse prefill ubatch raises AttributeError: 'AttentionMetaData' object has no attribute 'indexer_fp4_local_starts' at deepseek_v2.py:1783.

Carrying them by reference would not fix it either: build_ubatch_prefill_metadata rebases batch_id_per_q_token by - req_start, while cta_info encodes un-rebased absolute row ids, and _prefill_mqa_logits_fp4 takes the whole_batch fast path for any ubatch that fits one chunk.

The decode side was handled (3007), which makes the prefill omission look like an oversight. assert_fp4_indexer_supported rejects only fused_writer=False and PCP; TBO is unguarded.

7. [reported] The enable predicate reads head geometry only, so pooled/hybrid indexers turn FP4 on and then meet consumers that have no FP4 spelling

sparse_indexer_fp4_enabled checks index_topk, index_head_dim == 128, index_n_heads % 16 and the chip — nothing about the consumer, despite being described as "structural, never per-model".

  • Glm5NextIndexer (glm5_next.py:555) inherits Indexer.__init__ verbatim, so GLM-5.3-Flash-class configs (index_head_dim=128, index_n_heads=32) set _indexer_fp4 = True; its forward then routes to torch.ops.aiter.sparse_attn_indexer_kpool (glm5_next.py:714), which has no FP4 branch.
  • KimiAiterMLAGDNMetadataBuilder overrides build_kv_cache_tensor (kimi_mla_gdn_attn.py:320-369) and unconditionally does index_cache.view(rows, 1, aligned_index_dim), never setting kv_cache_scale — so the two-plane FP4 pool is allocated and then either raises a bare RuntimeError: shape ... is invalid for input of size ..., or silently binds E2M1 bytes to an FP8 reader when the numels happen to match.

Related: Indexer.__init__ computes _indexer_fp4 before use_qk_rope_cache_fusion, so a NoPE indexer (rope_dim 0 ≠ head_dim//2) passes the structural gate and then hard-fails in assert_fp4_indexer_supported with a message telling the user to unset ATOM_DISABLE_DS_INDEXER_QK_ROPE_CACHE_FUSION — a knob that cannot help.

The predicate should require index_kpool == 1 (and the fused writer) rather than discovering this three layers down.

8. [reported] Both FP4 scorer call sites pass kv_block_size=runner_block_size, with no assert kv_block_size == 64

aiter's contract is kv_cache_scale: u8 [num_blocks, k_tiles, 4, kv_block_size], i.e. kv_block_size is indexer rows per paged block, which fp4_index_block_shapes pins to FP4_KV_BLOCK_SIZE = 64. The pool is built from self._index_rows_per_block() (aiter_mla.py:1192-1198), which the base builder returns as block_size but kimi_mla_gdn_attn.py:165 overrides to block_size // index_kpool.

So --block-size 128 --index-kpool 2 gives index_rows_per_block = 64 — the pool builds with no raise — while the scorer is told kv_block_size = 128 at deepseek_v2.py:1975 and :1849. Page/row misaddressing, silently wrong logits. deepseek_v4.py:2178 carries an explicit assert kv_block_size == 64; the new MLA path has no equivalent.

9. [reported] fp4_decode_parallel_units is not monotonic in next_n, and the truncation that follows is silent

next_n * max(-(-512 // next_n), max_bs): with max_bs = 128, units(4) = 512 but units(3) = 513; with max_bs <= 74, units(8) = 512 but units(7) = 518.

The buffer is allocated once at fp4_decode_parallel_units(self.max_bs, max_seqlen_qo) (aiter_mla.py:546). If any later publish is reached with a next_n that is not max_seqlen_qo, [:parallel_units] (aiter_mla.py:2245) silently yields fewer rows while line 2273 still publishes indexer_fp4_n_ctas = parallel_units, so the kernel reads cta_info rows past the allocation.

The comment "Monotonic in next_n, which is what lets one buffer serve every speculation width" is the load-bearing claim and it is false. assert parallel_units <= self._indexer_fp4_cta_info[ubatch].shape[0] costs nothing. (Same file, line 41: "max(floor, rows) never lifts off the floor" is contradicted by this PR's own test at max_bs=8192.)

10. [reported] fp4_index_block_shapes hard-raises where the rest of the module degrades, and names the wrong knob

raise ValueError(f"The FP4 sparse indexer requires --block-size {FP4_KV_BLOCK_SIZE}, got {rows}")

rows is self._index_rows_per_block(), which is block_size // index_kpool on kimi_mla_gdn_attn.py:165. This PR's own test asserts the distinction — # The indexer rows per block are constrained, the KV block size is not, followed by _pool(index_fp4=True, block_size=32, index_rows_per_block=64) passing.

So a kpool user with --block-size 128 --index-kpool 2 is told to set a flag that will not move the checked quantity, and a plain user with --block-size 128 gets a hard crash deep inside MlaKvPool.__init__ at allocate_kv_cache time — after Indexer.__init__ already committed _indexer_fp4 = True and warmup traced the FP4 branch — where the same module's predicate degrades gracefully with a named reason for every other unsupported geometry.

Move the rows == 64 test into sparse_indexer_fp4_enabled, or at minimum name the real constraint.

11. [reported] _build_dcp_indexer_fp4_prefill_meta raises mid-serving on a schedule the scheduler may legally produce

pages = ceil(total_kv / block) counts the whole co-scheduled prefill batch, while cols = block_tables.shape[1] is ceil(max_model_len / block) — one sequence's allowance. Three co-scheduled 40k-token prefills under a 128k model exceed it, which is a legal schedule.

The engine then dies mid-serving with "lower max_num_seqs or max_num_batched_tokens" — knobs that bound request count and new tokens, not the sum of co-scheduled context lengths, so they do not close the hole. Either size the staged table once at startup from max_num_seqs * max_model_len / block, or drop it: it is the identity permutation and exists only because the kernel insists on a paged indirection this path does not need.

12. [reported] Under FP4 the registered custom op returns an uninitialized torch.empty

weights = torch.empty(weights.shape, device=weights.device, dtype=torch.float32) escapes via all three return weights — :1727 (the sub-threshold prefill early return, which fires before any FP4 scoring happens at all), :1953 (DCP decode) and :2024.

Today the only consumer is deepseek_v2.py:2960 (hidden_states_or_q_c = idx_ret) under pcp_is_enabled(), which assert_fp4_indexer_supported refuses. But that refusal lives in Indexer.__init__, three files from the escape, while sparse_attn_indexer_fake (:2062) promises a real tensor to torch.compile and atom/plugin/rtpllm/.../rtp_sparse_mla_backend.py:1869 returns the op's value straight to its caller. Relax the PCP refusal, or have a plugin consume the return, and NaNs flow into the attention query with no allocation error.

torch.zeros costs one small memset per layer and deletes the class of bug. At minimum the comment should name Indexer.__init__'s assert as the thing that must not be relaxed.

13. [reported] The KV-transfer refusal is broader than the V4 policy for the same feature

aiter_mla.py:1334 get_kv_transfer_tensors raises NotImplementedError whenever runner.config.kv_transfer_config is truthy. The same FP4 feature on the V4 path narrows the identical refusal to _uses_pd_staging(transfer_config) (deepseek_v4_attn.py:1843-1850) and explicitly states "Standalone LMCache offload can carry both FP4 indexer pools".

So MLA-sparse + --index-cache-dtype fp4 + an LMCache offload connector crashes at startup on a combination the V4 path deliberately allows — two divergent policies for one feature. The return None that follows also silently disables PD rather than reporting it.

Latent companion: MlaKvPool.region_tensors() now yields index_scale.layer_N roles, and the dispatch below keys on role.startswith("index."), so the moment this refusal is relaxed the scale plane is published as an mla.index_scale.layer_N region with a unit_bytes a connector would act on.

14. [reported] The FP4 verdict is declared three times with no cross-check, and the chip gate is spelled three ways

Indexer.__init__ computes self._indexer_fp4 (deepseek_v2.py:2307), the builder computes it again (aiter_mla.py:396), and the builder then unconditionally overwrites the module's value at KV-cache bind time (aiter_mla.py:1301) — while the module's own docstring names this exact hazard ("a layer that guessed wrong bakes in the other branch") and nothing asserts they agree.

If the builder says True where the Indexer said False, assert_fp4_indexer_supported never ran, so a PCP or ATOM_DISABLE_DS_INDEXER_QK_ROPE_CACHE_FUSION=1 run proceeds into the FP4 pool layout instead of raising.

Separately, sparse_indexer_fp4_enabled (:71, "The one place that decides") requires gfx == "gfx950" while v4_kernels.fp4_indexer_enabled (:105, "Single source of truth for the predicate") requires gfx != "gfx942", and config.py:2143 spells the chip test a third time. Any future chip silently enables one FP4 indexer and disables the other, and a reader who hits either docstring stops looking for the other.

assert module.indexer._indexer_fp4 == self._indexer_fp4 at the bind site costs nothing.

15. [reported] The DCP staging path moves as many bytes as the FP8 path it replaces

Per layer, per DCP prefill forward, _dcp_stage_indexer_fp4_prefill does: torch.arange(total_kv) plus four int64 div/mod kernels (1484, 1493-1494) that depend only on total_kv/block_size/the published slots; two new_zeros memsets of the whole staged key set that the next line fully overwrites; and an index_select temporary.

That is four full-key-set writes at ~68 B/key against the FP8 gather's two at ~132 B/key — roughly 1.50 GB versus 1.45 GB per forward at total_kv=262144 × 21 layers. The dtype halving is eaten by the staging round-trip, and it is unbudgeted: sparse_attn_indexer_fake still models only the FP8 gather, so KV sizing under-reserves for it.

Same shape at deepseek_v2.py:1533-1540, where _prefill_mqa_logits_fp4 rebuilds an identical per-chunk schedule 21× instead of once, and at aiter_mla.py:1650, where staged_tables is re-allocated per step to hold a constant arange(cols).

Cheapest fixes: publish page/row (and the composed destination index) once in the builder; new_empty plus zeroing only the sub-64-row tail; hoist the identity table into __init__.


Docs and tests

No docs/ file is touched. Per the repo's Docs-match-code checklist, a new module (sparse_indexer_fp4.py, 8 public symbols) needs a row in docs/model_ops_guide.md's table, and docs/configuration_guide.md:43/366 plus docs/scheduling_kv_cache_guide.md:27 now disagree with the code on what index_cache_dtype=fp4 reaches, on the hard --block-size 64 requirement, and on the PCP / KV-transfer refusals.

The new test file does import cleanly without aiter or triton, so there is no CI collection-abort risk. But only 6 of its 10 tests can run in CI — the sole coverage of the DCP staged page table is behind importorskip("atom.model_ops.attentions.aiter_mla") — and there is no test at all for _dcp_stage_indexer_fp4_prefill, which is where #2 lives.

@gbyu-amd
gbyu-amd marked this pull request as draft September 14, 2026 14:09
@gbyu-amd
gbyu-amd force-pushed the perf-sparse-indexer-fp4 branch from 45ad00e to b31edac Compare September 14, 2026 15:22
@gbyu-amd
gbyu-amd marked this pull request as ready for review September 14, 2026 15:22
@gbyu-amd
gbyu-amd force-pushed the perf-sparse-indexer-fp4 branch from b31edac to 3c32026 Compare September 14, 2026 15:32
Copilot AI review requested due to automatic review settings September 14, 2026 15:32

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot was unable to review this pull request because the user who requested the review is ineligible. To be eligible to request a review, you need a paid Copilot license, or your organization must enable Copilot code review.

@ROCm ROCm deleted a comment from gbyu-amd Sep 15, 2026
@XiaobingSuper

Copy link
Copy Markdown
Contributor Author

Accuracy (Kimi-K2.7-Code-MXFP4) on this PR is pre-existing breakage on main, not introduced here.

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 exact_match is ever produced:

Memory access fault by GPU node-4 ... Reason: Unknown.
[t=20s] GPU fault detected (5 signals) - exiting 2

The identical signature appears on this PR's base commit c74e8ce in two independent main runs — push and scheduled — and on four earlier main commits (c93b75d, 02f3fa0, ae51a34, 82f0b24).

This PR's paths are not implicated: the engine log for the failing job shows index_cache_dtype: None, no sparse_indexer_fp4 activation, and no shape/view error.

@XiaobingSuper

XiaobingSuper commented Sep 15, 2026 •

Copy link
Copy Markdown
Contributor Author

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

#1. Correct, and my "the row count is why this is unguarded" comment was defending the wrong thing. Worth adding to your analysis: gating on target_uses_mla, the idiom fourteen lines up, would not have been enough either — DeepseekV4AttentionMetadataBuilder is an MLA model whose builder is a CommonAttentionBuilder sibling, not an AiterMLAMetadataBuilder, so DeepSeek-V4 + MTP would still have crashed, and that is in the CI matrix. Fixed with a base-class no-op, which is total over the backends a draft can reach. The test asserts each non-MLA builder resolves to the base implementation, that the MLA one overrides, and that the base is inert (the draft shares the target's metadata object).

Accuracy (gpt-oss-120b), Accuracy (Qwen3-Next-80B-A3B-Thinking) and Accuracy (DeepSeek-V4-Pro) now pass on the matrix, which exercises this more directly than the unit test does.

#2. Confirmed, including your reading of the writer — sflat = (pos_in_block % 16) * tiles_per_block + pos_in_block / 16 generalises to (r % _MFMA_M) * (block_size // _MFMA_M) + r // _MFMA_M, which is now fp4_index_scale_rows, the single place the two planes' row axes are related. Both the read and the write bend through it, and the oracle imports it rather than restating the constant. It also takes block_size and rejects a mismatch itself, rather than relying on a raise two functions away.

There is now a test that calls _dcp_stage_indexer_fp4_prefill. I checked it covers the bug rather than just the function: reverting only the two bends makes it fail, restoring them passes.

Accuracy

You were right that the leg could not certify this path. Three legs, same harness, same 1323 prompts (all above index_topk, min 2634 / p50 3382 / max 4295), dcp=4, --gpu-memory-utilization unchanged, zero FP4 fallback warnings on every leg:

leg flexible strict block_bytes / blocks wall
FP8 + DCP4 0.9682 ± 0.0048 0.9689 ± 0.0048 3068928 / 45969 1972 s
FP4 + DCP4, sample 1 0.9636 ± 0.0052 0.9644 ± 0.0051 2966784 / 47607 1915 s
FP4 + DCP4, sample 2 0.9682 ± 0.0048 0.9689 ± 0.0048 2966784 / 47607 1883 s

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 #2 exponent bug live, scored 0.9704 — the highest of any leg. k_norm precedes the quantizer, so per-block e8m0 exponents are nearly uniform and a key wearing its neighbour's exponent still dequantizes close to correct. End-to-end accuracy is not a detector for layout bugs of this class, which is what the writer-vs-scorer cross-checks are for.

Reported items that held

#4 — both callers, and the draft one is an ordering bug rather than a slicing one: _enter_decode_metadata publishes before prepare_mtp_decode refreshes the buffer, and that function's own comment says the verify copy is stale. Now reads get_published_dcp_local_context_lens, the accessor the scorer already uses, so a ubatch gets the slice it publishes; the global buffer stays the fallback and every non-ubatch path is unchanged.

#5 — confirmed. Published once per step, immediately after prepare_mtp_decode, which is also the draft half of the wrong-rows bug above.

#3 — true but narrower than stated, and the narrowing changes the fix. top_k_per_row_decode bounds its scan by local_ctx, so columns past the context are never read and the uninitialised tail is inert. The real exposure is the case you pair it with: rows under-covered inside the window, where [:parallel_units] truncates silently while indexer_fp4_n_ctas still reports the full count. That is now an assert. A per-step -inf fill over [rows, max_model_len] is the cost this path exists to avoid; happy to add a debug-gated fill if you would still rather have the belt.

#6, #7 — confirmed, including that carrying the fields by reference would not fix #6. Prefill micro-batching now raises where PCP already did; --enable-tbo-decode is handled and stays allowed. The predicate rejects index_kpool > 1. Accuracy (GLM-5.3-Flash-kpool-16shot) passes on the matrix, which is the configuration #7 was about.

#9 — the claim is false (f(3) = 513 > f(4) = 512) but the buffer is safe: the property carrying the weight is the weaker f(n) >= f(1), and the only widths ever requested are 1 and max_seqlen_qo. The docstring now states the true property, with a test pinning both it and your counterexample.

#11 — you were right and my first answer was wrong. I rejected this as reserving for a batch that cannot co-exist; 16 x 131072 is 2.1M tokens against this configuration's 3.05M cache slots, so the scheduler can absolutely build it. One correction to the example: block_tables is a fixed global width, so the threshold is a combined context above max_model_len — 3 x 40k is 1920 pages and does not trip it, 4 x 40k or 2 x 70k does. The staged table is now sized at max_bs * block_table_cols and the raise is gone. Cost is about 2 MB and one extra scorer variant for the DCP path, which beats dying mid-serving on a legal schedule.

#12 — torch.zeros. A self-review pass found the same class one call earlier: weights_mqa was torch.empty_like and is read as weights_mqa[:num_padded_tokens], which includes padded rows. aiter's fused writer skips rows whose slot is negative unless compute_all_q_rope is set, and neither the FP4 nor the FP8 branch passes it — so the padded tail was uninitialised on every fused run, not only under FP4. Now zeros_like. The effect was benign (a row's weight only ever enters its own logits row, and padded rows are discarded downstream) but it was an uninitialised read on a shipped path. The FP8 branch's weights_out has the same shape and predates this work; I have left it alone.

#14 — worse than a missing assertion. That line was an assignment: the builder overwrote the verdict the Indexer had computed. It cannot fix a disagreement, because Indexer.__init__ has already built k_cache from its own answer, so overwriting the flag only makes the flag describe an object it no longer matches and defers the split to a graph/eager dtype mismatch — exactly the drift the predicate's own docstring warns about. It is now a comparison that names both values and the layer, and it was fixed up into the commit that introduced the overwrite, so no commit on the branch carries it.

Docs — row added for sparse_indexer_fp4.py, and the block-size-64 requirement plus the gfx950 / geometry / PCP / KV-transfer conditions written into both config tables.

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 ATOM_DISABLE_DS_INDEXER_QK_ROPE_CACHE_FUSION was advice that could not help.

Items I do not think should change here

#8, #10 — the reachable case was --block-size 128 --index-kpool 2, which #7 closes: a pooled indexer can no longer be FP4, so index_rows_per_block == block_size on every FP4 run, fp4_index_block_shapes already rejects anything but 64, and the message now names the quantity it checks. The assert would be unreachable.

#13 — a real inconsistency, and I think the narrower _uses_pd_staging policy is right, but it changes behaviour on a connector path with no coverage here. It should be its own change rather than riding along with a blocker fix.

#15 — the arithmetic looks right and the cheap parts (hoisting the identity table, publishing page/row once) are worth doing. They carry no correctness content, and folding a performance pass into this diff would make the bugs above harder to review.

The three chip gates — real, and worse for a reader than the count suggests, since sparse_indexer_fp4_enabled says "The one place that decides" while fp4_indexer_enabled says "Single source of truth for the predicate". Collapsing them spans config.py and the V4 path, neither of which this PR otherwise touches. Same reasoning as #13.

Tests

tests/test_sparse_indexer_fp4.py is 18 passing on gfx950, of which 11 need no GPU and run in CI. The four that do not are the writer-vs-scorer cross-checks for decode, DCP decode and prefill; they are the only coverage that the layout the fused FP4 writer produces is the layout each scorer reads, which is where #2 lived, so I have kept them despite CI not reaching them. Their green is ours to report rather than the matrix's.

@ROCm ROCm deleted a comment from gbyu-amd Sep 15, 2026
@ROCm ROCm deleted a comment from gbyu-amd Sep 15, 2026
@gbyu-amd
gbyu-amd force-pushed the perf-sparse-indexer-fp4 branch from 3c32026 to c1b699d Compare September 15, 2026 02:03
Copilot AI review requested due to automatic review settings September 15, 2026 02:03

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot was unable to review this pull request because the user who requested the review has reached their quota limit.

@ROCm ROCm deleted a comment from gbyu-amd Sep 15, 2026
Copilot AI review requested due to automatic review settings September 15, 2026 02:18
@gbyu-amd
gbyu-amd force-pushed the perf-sparse-indexer-fp4 branch from c1b699d to 8ee2745 Compare September 15, 2026 02:18

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot was unable to review this pull request because the user who requested the review has reached their quota limit.

XiaobingSuper and others added 5 commits September 15, 2026 03:46
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>
@gbyu-amd
gbyu-amd force-pushed the perf-sparse-indexer-fp4 branch from 8ee2745 to ef4dfec Compare September 15, 2026 03:47
Copilot AI review requested due to automatic review settings September 15, 2026 03:47

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot was unable to review this pull request because the user who requested the review has reached their quota limit.

@valarLip
valarLip merged commit 3dc0cea into main Sep 15, 2026
56 of 61 checks passed
@valarLip
valarLip deleted the perf-sparse-indexer-fp4 branch September 15, 2026 09:18
ZhiweiYan-96 added a commit that referenced this pull request Sep 15, 2026
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.
sajandhy pushed a commit to sajandhy/ATOM that referenced this pull request Sep 17, 2026
* 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>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants