test: add block-level DeepSeek-V4 attention test (real DeepseekV4Attention) - #1723
carlushuang wants to merge 4 commits into
Conversation
…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.
🏷️ CI GuideRuns automatically on every eligible PR before approval:
Heavy model tests:
|
…/dsv4-attn-block-test
…_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.
|
Merged current 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
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 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.
DSV4 sibling of #1402's
test_attention_block_gptoss.py. Drives the realDeepseekV4Attentionend 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_ratio0=Dense, 4=CSA, 128=HCA) in both prefill and decode.Why it drives the real metadata builder
The GPT-OSS test hand-fills a 5-field
AttentionMetaData.AttentionMetaData_DSV4carries compress plans, per-seq state slots,batch_id_per_tokenand three ragged paged index sets, so hand-filling it would mean reimplementing most ofDeepseekV4AttentionMetadataBuilder(~2700 lines).Instead the harness duck-types the slice of
ModelRunnerthe 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 realconfig.jsonviaDeepseekV4Args.from_hf_config.Three things the harness must get right, each of which fails confusingly:
cudagraph_modemust be present andNone, elseforward()takes the PIECEWISE custom-op split instead offorward_impl.args.quant_configmust come frommake_v4_quant_config. Without itq_norm'sfused_quantpath is inactive andqr, qr_scale = self.q_norm(...)fails to unpack.model_runner.max_per_req_cache_slotsmust be non-zero, or_build_paged_prefill_metahits its warmup guard and returns early, leavingkv_indices_prefix_*asNonefor 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.
check_dense_windowalso verifies the builder's gather list against the analytic sliding window.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.05idiom.checkAllclosereturns the fraction of mismatching elements, so that idiom silently accepts a 5%-wrong tensor — verified: ax1.05perturbation of the reference passes under it and fails here.Validation
gfx1250 (MI400-class, single GPU), bf16 KV, whole matrix in ~14 s:
Plus batch 1 / 2 / 4, and a long-context CSA run where the top-k actually truncates (ctx 2560 -> 640 committed vs
index_topk512), which also passes. Worst error throughout is bf16 rounding.--kv-cache-dtype fp8skips on gfx1250 with an explanation rather than a wall of template errors:kv_fp8setsfp8_2buff=True, routingqk_norm_ropeinto aiter'sfused_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(merge6cf4ac12), 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
mainsurfaced three harness assumptions that my earlier signature-only API check had missed — all fixed inb45036b4, all stub-config/geometry issues in the test rather than problems with the module under test:DeepseekV4AttentionMetadataBuilder.block_sizemoved 128 -> 256CommonAttentionBuilderassertsmodel_runner.block_size % builder.block_size == 0; the harness hardcoded 128--block-sizenow 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_dtypeis new (fp4 indexer, #1709)get_current_atom_config()Config.__post_init__default ofkv_cache_dtypeconfig.dsparkis new (DSpark refactor, #1700)prepare_decodereadsconfig.dspark.raggedunguardedDSparkConfigwhen present, else an all-defaults stand-inOn 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-traceskill, 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.
tis µs relative to the first kernel of the layer.aiter::mhc_fused_post_pre_gemm_sqrsumaiter::mhc_pre_big_fuse_rmsnorm_gemm_a16w16…BLOCK_M_16_gemm_a16w16…BLOCK_M_16aiter::dynamic_per_group_scaled_quantfused_compress_attn_w32_D128…Q_ue8m0fused_compress_attn_w32_D512…_update_compressor_states_kernel_gemm_a8w8_blockscale_preshufflewqkv_a_update_compressor_states_kernel_gemm_a8w8_blockscale_reducewqkv_aaiter::add_rmsnorm_quantq_norm+ fp8 quant_gemm_a8w8_blockscale_preshufflewq_bqk_norm_rope…kvw_paged…flydsl_gemm_a8w8_blockscale_preshuffleaiter::rope_hadamard_rotate_activation_quant_gemm_a16w16weights_proj_scale_indexer_weights_kernel_gluon_deepgemm_fp8_paged_mqa_logits_preshuffleaiter::ob::radix_topk_one_block_kernel_csa_translate_pack_kernel_pa_decode_sparse…KV_SPLITS_16_pa_decode_sparse_reduce…_inverse_rope_gptj_kernel_batched_gemm_bf16…BLOCK_M_16wo_agrouped LoRA_batched_gemm_bf16_reduceaiter::dynamic_per_group_scaled_quantwo_b_gemm_a8w8_blockscale_preshufflewo_b_gemm_a8w8_blockscale_reduceImplementation mix — 122.0 µs of kernel time in a ~162 µs layer:
ops/triton/_triton_kernels/+ ATOMmodel_ops/v4_kernels/csrc/kernels/*.cuops/flydsl/kernels/Notes for reviewers:
ENABLE_CK=0plus theATOM_USE_TRITON_*recipe routes everything to Triton/FlyDSL. On gfx950 the same path would pull in hand-written.coassembly, so the test covers a materially different kernel set per arch.fused_compress_attn_D128and one state update run on stream 22; the main-KVD512compressor and the second state update stay on stream 0. Somaybe_compressors_asynchides part of the compressor cost, not all of it.swa_writeis fused into qk-norm+RoPE in decode (notekvw_pagedin the kernel name). The standaloneswa_writekernel only runs in prefill — which is why the test's stage accounting shows it under prefill only.csa_translate_pack); HCA keeps the compressor but not the indexer. That is the per-layer-type difference the three test configs cover.