Skip to content

test: add block-level DeepSeek-V4 attention test (real DeepseekV4Attention) - #1723

Open
carlushuang wants to merge 4 commits into
mainfrom
carhuang/dsv4-attn-block-test
Open

carlushuang wants to merge 4 commits into
mainfrom
carhuang/dsv4-attn-block-test

Conversation

@carlushuang

@carlushuang carlushuang commented Jul 29, 2026 •

Copy link
Copy Markdown
Collaborator

DSV4 sibling of #1402's test_attention_block_gptoss.py. Drives the real DeepseekV4Attention end to end — wqkv_a -> q_norm -> wq_b -> qk_norm+RoPE -> SWA write -> Compressor/Indexer -> sparse paged attention -> inverse RoPE -> grouped output LoRA -> wo_b — against a torch reference, for all three V4 layer types (compress_ratio 0=Dense, 4=CSA, 128=HCA) in both prefill and decode.

python3 tests/block/test_attention_block_dsv4.py
python3 tests/block/test_attention_block_dsv4.py --layer csa --seqlen 2560

Why it drives the real metadata builder

The GPT-OSS test hand-fills a 5-field AttentionMetaData. AttentionMetaData_DSV4 carries compress plans, per-seq state slots, batch_id_per_token and three ragged paged index sets, so hand-filling it would mean reimplementing most of DeepseekV4AttentionMetadataBuilder (~2700 lines).

Instead the harness duck-types the slice of ModelRunner the builder touches and drives the real builder: allocate_kv_cache_tensors -> allocate_per_req_cache -> build_kv_cache_tensor (per sub-module: Attention / Indexer / Compressor) -> prepare_prefill / prepare_decode. So the test also covers production metadata construction. No ModelRunner, no checkpoint — a synthetic 3-layer config (one layer of each ratio) built from the real config.json via DeepseekV4Args.from_hf_config.

Three things the harness must get right, each of which fails confusingly:

  • cudagraph_mode must be present and None, else forward() takes the PIECEWISE custom-op split instead of forward_impl.
  • args.quant_config must come from make_v4_quant_config. Without it q_norm's fused_quant path is inactive and qr, qr_scale = self.q_norm(...) fails to unpack.
  • model_runner.max_per_req_cache_slots must be non-zero, or _build_paged_prefill_meta hits its warmup guard and returns early, leaving kv_indices_prefix_* as None for the prefill kernel to dereference.

What the reference proves, and what it doesn't

The attention core is recomputed independently in torch; the projections, RoPE and output LoRA are the real module, so a mismatch localises to the sparse attention rather than to a projection. Same split as the GPT-OSS test.

Layer Reference
Dense (0) Independent. check_dense_window also verifies the builder's gather list against the analytic sliding window.
CSA (4) Gathers the indices the module selected — validates gather + attention math + inverse RoPE + output LoRA, but takes the indexer's choice of top-k as given.
HCA (128) Same, via the compress-plan indices.

Reproducing the DSA indexer scoring in torch is deliberately out of scope; it belongs with the aiter-side port.

Each config prints what it exercised (per-token KV-set sizes, whether the CSA top-k truncated), because CSA only truncates once ctx//4 > index_topk — a short-context run silently covers only the easy case.

Pass criterion

Zero mismatching elements, not the err < 0.05 idiom. checkAllclose returns the fraction of mismatching elements, so that idiom silently accepts a 5%-wrong tensor — verified: a x1.05 perturbation of the reference passes under it and fails here.

Validation

gfx1250 (MI400-class, single GPU), bf16 KV, whole matrix in ~14 s:

  dense window check: 0 token(s) with wrong extend length, prefix_total=0
  prefill dense ratio=  0 [kv/token max=128 (prefix=0, extend=128)]:   max_abs=0.01562 mismatch=0/2097152 -> PASS
  decode  dense ratio=  0 [kv/token max=128 (prefix=128, extend=0)]:   max_abs=0.00317 mismatch=0/8192    -> PASS
  prefill csa   ratio=  4 [kv/token max=192 (prefix=64, extend=128)]:  max_abs=0.01660 mismatch=0/2097152 -> PASS
  decode  csa   ratio=  4 [kv/token max=203 (prefix=203, extend=0)]:   max_abs=0.00208 mismatch=0/8192    -> PASS
  prefill hca   ratio=128 [kv/token max=130 (prefix=2, extend=128)]:   max_abs=0.01953 mismatch=0/2097152 -> PASS
  decode  hca   ratio=128 [kv/token max=130 (prefix=130, extend=0)]:   max_abs=0.00293 mismatch=0/8192    -> PASS

Plus batch 1 / 2 / 4, and a long-context CSA run where the top-k actually truncates (ctx 2560 -> 640 committed vs index_topk 512), which also passes. Worst error throughout is bf16 rounding.

--kv-cache-dtype fp8 skips on gfx1250 with an explanation rather than a wall of template errors: kv_fp8 sets fp8_2buff=True, routing qk_norm_rope into aiter's fused_qk_norm_rope_group_quant, whose JIT module needs CK-tile — and CK-tile's arch list (ck_tile/core/config.hpp) has no gfx1250, so it cannot build there at all. This is the same reason DSV4-Flash serving needs bf16 KV on gfx1250.

Validation environment

Rebased onto current main (merge 6cf4ac12), and the test was then run against this branch's own tree inside the gfx1250 bring-up container, not only against the older pinned commit.

Running against main surfaced three harness assumptions that my earlier signature-only API check had missed — all fixed in b45036b4, all stub-config/geometry issues in the test rather than problems with the module under test:

Drift Effect Fix
DeepseekV4AttentionMetadataBuilder.block_size moved 128 -> 256 CommonAttentionBuilder asserts model_runner.block_size % builder.block_size == 0; the harness hardcoded 128 --block-size now defaults to the builder's own class attribute, so it tracks that constant (resolves to 128 on the pinned tree, 256 on main -- both verified)
config.index_cache_dtype is new (fp4 indexer, #1709) the indexer reads it unguarded via get_current_atom_config() stub mirrors the real Config.__post_init__ default of kv_cache_dtype
config.dspark is new (DSpark refactor, #1700) prepare_decode reads config.dspark.ragged unguarded stub uses the real DSparkConfig when present, else an all-defaults stand-in

On this branch's tree (i.e. current main): Dense and HCA pass in both prefill and decode — 4 of the 6 configs.

CSA cannot be executed in that container: main's ATOM calls aiter::rope_rotate_activation() with 9 arguments and the container's pinned aiter fork declares 8. That is an ATOM<->aiter version coupling in the bring-up image, not something the test can work around — the image is the only environment where DSV4 currently runs on gfx1250, and its aiter is a fork carrying gfx1250 tuning rather than upstream. CSA needs a box with an aiter new enough for current main.

All 6 configs (including CSA, and the long-context run where the top-k truncates) pass on the pinned tree d13bfb0c, with numbers unchanged before and after the merge.

So: the harness is confirmed compatible with main, 4/6 configs are confirmed passing on main, and the remaining 2 are blocked by an environment version skew rather than by anything in this change. Worth a CI run on a box with matching aiter before merging.

What this test actually drives

The kernel chain the block exercises, for one CSA layer of a decode step. Captured on gfx1250 from a real serving run (capture-trace skill, DSV4-Flash TP1, ISL 8192 / OSL 1024 / conc 16, bf16 KV) — these are server numbers, not this test's, included so a reviewer can see what the block covers and how each stage is implemented.

The layer runs on two concurrent GPU streams, so it is laid out as two columns rather than one ordered list — the three side-stream kernels genuinely overlap the main stream and are not sequenced with it. t is µs relative to the first kernel of the layer.

t start–end impl main stream (0) side stream (22) purpose
0.0–4.0 HIP aiter::mhc_fused_post_pre_gemm_sqrsum mHC post/pre + sq-sum
6.0–10.4 HIP aiter::mhc_pre_big_fuse_rmsnorm attn input RMSNorm
12.7–17.9 Triton _gemm_a16w16…BLOCK_M_16 mHC head GEMM
13.9–19.0 Triton _gemm_a16w16…BLOCK_M_16 mHC head GEMM (parallel)
20.3–22.7 HIP aiter::dynamic_per_group_scaled_quant act -> fp8 per-group
20.8–26.2 FlyDSL fused_compress_attn_w32_D128…Q_ue8m0 compressor, indexer side
25.2–31.4 FlyDSL fused_compress_attn_w32_D512… compressor, main KV
28.1–30.8 Triton _update_compressor_states_kernel compressor ring state
33.1–35.8 Triton _gemm_a8w8_blockscale_preshuffle wqkv_a
37.7–40.6 Triton _update_compressor_states_kernel state update (main)
42.4–44.3 Triton _gemm_a8w8_blockscale_reduce split-K reduce for wqkv_a
46.1–48.1 HIP aiter::add_rmsnorm_quant q_norm + fp8 quant
50.2–54.9 Triton _gemm_a8w8_blockscale_preshuffle wq_b
56.9–60.3 FlyDSL qk_norm_rope…kvw_paged…flydsl qk-norm + RoPE + paged KV write
62.2–65.2 Triton _gemm_a8w8_blockscale_preshuffle indexer q proj
67.5–70.3 HIP aiter::rope_hadamard_rotate_activation_quant indexer q RoPE+Hadamard+quant
72.1–75.8 Triton _gemm_a16w16 indexer weights_proj
77.8–79.1 Triton _scale_indexer_weights_kernel indexer weight scale
81.1–88.8 Triton (Gluon) _gluon_deepgemm_fp8_paged_mqa_logits_preshuffle indexer scoring
90.8–100.1 HIP aiter::ob::radix_topk_one_block_kernel top-k select (hipcub)
101.9–104.6 Triton _csa_translate_pack_kernel top-k -> paged offsets
106.5–117.9 Triton _pa_decode_sparse…KV_SPLITS_16 sparse attention core
120.4–128.1 Triton _pa_decode_sparse_reduce… split-KV reduce
130.5–132.2 Triton _inverse_rope_gptj_kernel inverse RoPE
133.9–140.4 Triton _batched_gemm_bf16…BLOCK_M_16 wo_a grouped LoRA
142.5–144.5 Triton _batched_gemm_bf16_reduce split-K reduce
146.3–148.9 HIP aiter::dynamic_per_group_scaled_quant quant for wo_b
151.2–156.0 Triton _gemm_a8w8_blockscale_preshuffle wo_b
157.9–160.0 Triton _gemm_a8w8_blockscale_reduce split-K reduce
161.8 (next layer begins)

Implementation mix — 122.0 µs of kernel time in a ~162 µs layer:

impl kernels µs share source
Triton 19 79.7 65% aiter ops/triton/_triton_kernels/ + ATOM model_ops/v4_kernels/
HIP C++ 7 27.4 22% aiter csrc/kernels/*.cu
FlyDSL 3 14.9 12% aiter ops/flydsl/kernels/
asm (.co) 0 0 — —
opus 0 0 — —

Notes for reviewers:

  • No asm and no opus kernels on this path on gfx1250 — ENABLE_CK=0 plus the ATOM_USE_TRITON_* recipe routes everything to Triton/FlyDSL. On gfx950 the same path would pull in hand-written .co assembly, so the test covers a materially different kernel set per arch.
  • Only the indexer-side compressor is offloaded. fused_compress_attn_D128 and one state update run on stream 22; the main-KV D512 compressor and the second state update stay on stream 0. So maybe_compressors_async hides part of the compressor cost, not all of it.
  • swa_write is fused into qk-norm+RoPE in decode (note kvw_paged in the kernel name). The standalone swa_write kernel only runs in prefill — which is why the test's stage accounting shows it under prefill only.
  • ~40 µs of the ~162 µs layer is inter-kernel gap (122 µs of kernel time). At bs=16 decode this path is launch-bound as much as compute-bound.
  • Dense layers drop the indexer block (indexer q proj through csa_translate_pack); HCA keeps the compressor but not the indexer. That is the per-layer-type difference the three test configs cover.
  • The run uses 18 GPU streams in total, but stream 0 carries 88% of all GPU time (11002 of 12527 ms); streams 20/21/22 are ~450 ms each and the rest are negligible.

…ntion)

DSV4 sibling of tests/block/test_attention_block_gptoss.py (#1402). Drives the
real DeepseekV4Attention end to end -- wqkv_a -> q_norm -> wq_b -> qk_norm+RoPE
-> SWA write -> Compressor/Indexer -> sparse paged attention -> inverse RoPE ->
grouped output LoRA -> wo_b -- against a torch reference, for all three V4 layer
types (compress_ratio 0=Dense, 4=CSA, 128=HCA) in prefill and decode.

Unlike the GPT-OSS test, the metadata is not hand-built. AttentionMetaData_DSV4
carries compress plans, per-seq state slots, batch_id_per_token and three ragged
paged index sets, so filling it by hand would mean reimplementing most of
DeepseekV4AttentionMetadataBuilder. Instead the harness duck-types the slice of
ModelRunner the builder needs and drives the real builder
(allocate_kv_cache_tensors -> allocate_per_req_cache -> build_kv_cache_tensor ->
prepare_prefill / prepare_decode), so the test also covers production metadata
construction.

Three things the harness has to get right, each of which fails confusingly:
  - cudagraph_mode must be present and None, else forward() takes the PIECEWISE
    custom-op split instead of forward_impl.
  - args.quant_config must come from make_v4_quant_config; without it q_norm's
    fused_quant path is inactive and `qr, qr_scale = self.q_norm(...)` fails to
    unpack.
  - model_runner.max_per_req_cache_slots must be non-zero, or
    _build_paged_prefill_meta hits its warmup guard and returns early, leaving
    kv_indices_prefix_* as None for the prefill kernel to dereference.

Reference: the attention core is recomputed independently in torch; the
projections, RoPE and output LoRA are the real module, so a mismatch localises
to the sparse attention. Dense additionally gets an independent check of the
builder's gather list against the analytic sliding window. CSA/HCA gather the
indices the module selected, so they validate the gather + attention math but
take the indexer's choice of top-k as given (reproducing DSA scoring in torch is
out of scope here). Each config prints the per-token KV-set sizes and whether
the CSA top-k truncated, since a short-context run only covers the easy case.

The pass criterion is zero mismatching elements, not aiter checkAllclose's
`err < 0.05` -- that returns the FRACTION of mismatching elements, so the usual
idiom silently accepts a 5%-wrong tensor (verified: a x1.05 perturbation passes
under it and fails here).

Validated on gfx1250 (MI400-class, 1 GPU), bf16 KV, all 6 configs, plus
batch 1/2/4 and a long-context CSA run where the top-k actually truncates
(ctx 2560 -> 640 committed vs index_topk 512). Worst error is bf16 rounding:
max_abs 0.0195 on prefill, 0.0039 on decode; zero mismatches at atol/rtol 3e-2.

--kv-cache-dtype fp8 skips on gfx1250 with an explanation: kv_fp8 sets
fp8_2buff=True, routing qk_norm_rope into aiter's fused_qk_norm_rope_group_quant,
whose JIT module needs CK-tile -- and CK-tile's arch list has no gfx1250, so it
cannot build there at all.
@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 1723 --add-label <label>

…_dtype, dspark)

Merging main and actually running against it surfaced three harness
assumptions that a signature-only API check had missed. All three are
stub-config/geometry issues in the test, not problems with the module
under test.

  - DeepseekV4AttentionMetadataBuilder.block_size moved 128 -> 256, and
    CommonAttentionBuilder asserts model_runner.block_size is a multiple
    of it. --block-size now defaults to the builder's own class attribute
    instead of a hardcoded 128, so the harness tracks that constant
    wherever it goes (resolves to 128 on the older pinned tree, 256 on
    main -- both verified).

  - config.index_cache_dtype is new (fp4 indexer, #1709) and the indexer
    reads it unguarded via get_current_atom_config(). The stub now mirrors
    the real Config.__post_init__ default of kv_cache_dtype.

  - config.dspark is new (DSpark drafter refactor, #1700) and
    prepare_decode reads config.dspark.ragged unguarded. The stub uses the
    real DSparkConfig when the tree has one and falls back to an
    all-defaults stand-in otherwise.

Re-validated after the merge: all 6 configs (3 layer types x prefill/decode)
still pass on the pinned gfx1250 tree, unchanged numbers.
@carlushuang

Copy link
Copy Markdown
Collaborator Author

Merged current main (6cf4ac12) into the branch — clean, no conflicts, the test file itself untouched.

More usefully, I then checked this branch's tree out inside the gfx1250 container and ran the test against it, rather than relying on the signature-only API comparison from the original description. That found three real harness drifts the signature check had missed (fixed in b45036b4):

  • DeepseekV4AttentionMetadataBuilder.block_size moved 128 -> 256. CommonAttentionBuilder asserts model_runner.block_size % builder.block_size == 0, and the harness hardcoded 128, so it asserted immediately. --block-size now defaults to the builder's own class attribute — resolves to 128 on the pinned tree and 256 on main, both verified.
  • config.index_cache_dtype (new with the fp4 indexer, feat(deepseek_v4): Add fp4 indexer fo dsv4 #1709) is read unguarded by the indexer; the stub config now mirrors Config.__post_init__'s kv_cache_dtype default.
  • config.dspark (new with the DSpark refactor, refactor(spec_decode): drafter abstraction + align DSpark with the HF reference (acceptance 34.9% -> 48%) #1700) is read unguarded by prepare_decode as config.dspark.ragged; the stub uses the real DSparkConfig when the tree has one.

Result on this branch's tree: Dense and HCA pass in both prefill and decode, 4 of 6 configs.

CSA can't run there — main's ATOM calls aiter::rope_rotate_activation() with 9 args while the container's pinned aiter fork declares 8. That's an ATOM<->aiter version skew in the bring-up image (its aiter is a gfx1250-tuning fork, not upstream), not something this change can address. All 6 configs including CSA still pass on the pinned d13bfb0c tree, with numbers unchanged before and after the merge.

Description updated accordingly. CSA on current main still wants a CI run or a box with a matching aiter.

The Pre Checkin ruff job installs ruff from pip, so it picks up 0.16.0,
whose default rule set is broader than 0.15.x: C408 (dict() call ->
literal), I001 (import sorting) and RUF010 (explicit f-string conversion
flag) are now on by default. Locally 0.15.7 reported the file clean, so
this only showed up in CI.

All four fixes are semantically identical rewrites:
  - the two Capture.wrap_* dict(...) calls become dict literals
  - the local imports in run_decode are sorted
  - f"{str(ok):>6}" becomes f"{ok!s:>6}"

Verified with ruff 0.16.0 (clean) and black (repo-wide, 425 files
unchanged). Test re-run on gfx1250 afterwards: all 6 configs still pass
with identical numbers.
@zufayu
zufayu requested a review from yhl-amd July 30, 2026 02:37

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant