Repository navigation
dsv4.1: compression, KV I/O, and metadata kernels - #39652
Conversation
755f76b to
0ab8faf
Compare
8b59cbc to
beaa5af
Compare
0ab8faf to
3186760
Compare
3186760 to
cd871a9
Compare
…ame low-ratio metadata fields
flash_c2_decode_kernel loads the RoPE frequencies twice for the same thread and the same position: once right after the norm weight, and again inside the `tx >= kNopeThreads` branch that consumes them. `freq` is never read between the two, so the first load is dead and the second re-issues the same 8 bytes per lane. Keep the early load -- it is issued before the softmax and the RMSNorm reduction, so its latency overlaps that work -- and remove the reload inside the branch. Verified bit-exact on GB300 (sm_103) against the pre-change kernel: identical sha256 of `out`, the paged cache and the pair-state ring for V4, V41 and V41_FP4 decode plus a draft_len=4 verify batch, with the JIT cache cleared between runs. A deliberate eps perturbation was used as a negative control to confirm the harness recompiles and detects kernel changes. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
|
Read the whole diff. The kernel work is high quality and the header comments are unusually good — a few things I checked and can confirm rather than ask about:
Two things on the diff. 1. This PR deletes the draft-pad regression tests and replaces them with tests for something else
They do not move anywhere else — I grepped the whole branch for The logic they pinned is still live and unchanged on this branch: // c_plan.cuh:524
const auto mtp_pad = ring_size > window_size ? ring_size - window_size + 2 : 0;used at
Could the new boundary tests be added to the file rather than replacing it? If the old ones had to go for runtime or flakiness reasons, could that be said explicitly so the trade is on the record? 2. None of the new entry points have a testEvery new public surface in this PR is untested on the branch:
This is the cheapest possible test to write, because the PR ships both sides: Pushed one cleanup
// c2.cuh:141
if (tx >= kNopeThreads) freq.load(params.freqs_cis + (pos - 1) * kRopeDim, tx - kNopeThreads);
...
// c2.cuh:202, inside the branch that consumes it
freq.load(params.freqs_cis + (pos - 1) * kRopeDim, tx - kNopeThreads);
Note this branch is the base of the stack, so the commit is not in (Review by Claude Opus 5, run by @BBuf.) |
…_quant next to the kernels; align compress kernel names
|
/tag-and-rerun-ci |
|
/rerun-test test/registered/kernels/ops/attention/test_c4_v2.py test/registered/kernels/ops/attention/test_c128_v2.py test/registered/kernel/attention/test_deepseek_v4_compress_plan_bounds.py test/registered/kernels/ops/attention/test_deepseek_v4_compress_state_runtime_shapes.py test/registered/kernels/ops/attention/test_fp4_indexer.py test/registered/kernels/ops/attention/test_paged_mqa_metadata.py |
|
Results for 🚀 |
|
/rerun-test test/registered/attention/unittests/dsv4/test_deepseek_v4.py test/registered/e2e/models/test_deepseek_v4_flash_fp4_b200.py test/registered/attention/unittests/dsa/test_dsa.py |
|
Results for 🚀 🚀 |
Summary
KVLayout(kernels/ops/attention/dsv4/kv_layout.py,include/sgl_kernel/deepseek_v4/kv_layout.cuh): the paged FlashMLA main-KV formats as one enum with their per-token data / scale bytes, tile size and page alignment.V4is the existing 584-byte layout (448 fp8 nope + 64 bf16 rope + ue8m0 scales);V41(528 B, every dim fp8 with one ue8m0 scale per 32) andV41_FP4(288 B, e2m1 codes with one e4m3 scale per 16) are the DeepSeek-V4.1 formats.PagedKV<layout, page_bits>::rowgives every writer the same address arithmetic, andv41::store_rowis the single quantizer for the V4.1 rows.layoutthrough the existing cache writers and reader:fused_store_cache(with an optional in-kernel RoPE for the V4.1 layouts),compress_norm_rope_store,fused_k_norm_rope_flashmla,dequantize_k_cache_paged. The default staysV4, so current callers are unchanged.csrc/deepseek_v4/c1.cuh,c2.cuh,ops/attention/dsv4/low_ratio_compress.py): pair-pooling (ratio 2), RMSNorm, RoPE and the main-KV write in one launch, returning the pre-RoPE latent;out_loc == 0marks a padded graph row.dequant_k_cache.py),fp4_utils.cuh(e2m1 fake-quant helpers next to the existingfp8_utils.cuh), the pure-torch reference quantizer (torch_quant.py), and two small-batch metadata builders (build_page_table_positions_small,build_low_ratio_metadata).Changes to existing kernels
csrc/deepseek_v4/store.cuh,fused_norm_rope_v2.cuh,main_norm_rope.cuh: the V4 path now addresses pages throughPagedKV; same bytes as before. The kernel templates gain aKVLayoutparameter.csrc/deepseek_v4/c_plan.cuh:ragged_idis a zero-based uint16, so a 65536-token batch is valid;batch_sizemust stay below 65535 becausepack_w(65535, 65535)is the invalid-write sentinel.test_deepseek_v4_compress_plan_bounds.pycovers both bounds and the sentinel collision.get_paged_mqa_logits_metadataaccepts page size 128 as well as 64 (the schedule depends only on the sequence lengths).Verification
test_c4_v2.py,test_c128_v2.py,test_deepseek_v4_compress_state_runtime_shapes.py,test_fp4_indexer.py,test_paged_mqa_metadata.py, plustest_deepseek_v4_compress_plan_bounds.py.CI States
Latest PR Test (Base): 🚫 Run #35147198293
Latest PR Test (Extra): ❌ Run #35147197934
Latest PR Test (AMD ROCm 10): ⏳ Run #35147198240