Conversation
Add structural foundation for SM100 BWD InnerLoopK: - TMEM layout: dQ persistent, S/P→dK overlapped, dP/dS, dV - SMEM: multi-stage sK/sV/sKt for K streaming - Scheduler: outer Q blocks, inner K range via get_n_block_min_max - load_loop_k: Q/dO fixed, K/V pipeline streaming - mma_loop_k: dQ accumulated, dK/dV fresh per iter - kernel dispatch: conditional LoopK vs LoopQ for load/mma warps 1-CTA only, hdim<=128. compute_loop and reduce warps pending.
- compute_loop: LoopK conditionals for block extraction, LSE/dPsum pre-wait/post-release, per-iter K-block mask, skip dKV epilogue - dKV_reduce_loop_k: per-K-iter dK/dV TMA atomic reduce + end-of-tile dQ reduce (mirrors dQacc_reduce pattern with 3 TMEM T2R copies) - kernel dispatch: route reduce warps to dKV_reduce_loop_k for LoopK - force dKV_postprocess=True for LoopK (fp32 dKacc/dVacc accumulation) - frontend: add swap_bwd_qk_loop param to _flex_flash_attn_bwd + compile key
…se support - Fix 5 pipeline deadlocks in BWD LoopK warp-specialized kernel: 1. Separate pipeline_K_load from pipeline_Q to avoid barrier conflict 2. Reorder K release before acquire in single-stage buffer 3. Remove unnecessary pipeline_dQ empty wait in LoopK 4. Fix pipeline_dKV consumer group to use reduce_warps for LoopK 5. Add producer_state_Q.advance() after commit in load_loop_k - Add IndexSparse BWD LoopK: indices-driven K/V load with sparse dKV store - Add BlockSparse BWD LoopK: block-level sparse mask support - Add index_attn_indices_to_block_sparse conversion in sparse_utils.py - Fix int64->int32 dtype for mask_block_idx in sparse_utils.py - Disable 2CTA mode when swap_bwd_qk_loop is enabled - Remove unused variables (tidx, producer_state_dO) in load_loop_k Verified: Dense LoopK bit-identical to LoopQ, BlockSparse MHA pass, IndexSparse all pass on B300 SM103.
…env vars - Enable PackGQA for BWD SM100 (was globally disabled). Host-side reorder lse/dpsum/dq_accum to match CuTe column-major packed layout. Fix causal mask and block boundary computations for packed coordinates. - Remove LoopK-only restriction for IndexSparse BWD, allowing LoopQ path. - Add MAGI_ATTENTION_FFA_CUTEDSL_INNER_DIR_MAX_TO_MIN env var to control BWD inner loop direction (both LoopQ and LoopK). - Add MAGI_ATTENTION_FFA_CUTEDSL_MASK_MODE=dispatch env var for causal zone-based mask dispatch, skipping R2P on fully-valid tiles.
…guard - Add transpose_block_sparse_tensors() to convert forward-direction (M→N) block sparse tensors to backward-direction (N→M) with correct coarse Q-block granularity (subtile_factor * tile_m). - When index_attn_indices is used with LoopQ (swap_bwd_qk_loop=False), auto-transpose the block sparse tensors for the backward path. - Add index_attn_indices to disable_2cta conditions to prevent 2-CTA mode with block sparsity.
When block_sparse_tensors have head_dim=NHK but the kernel indexes by head_idx in NHQ space (pack_gqa=False), expand NHK→NHQ via repeat_interleave. When pack_gqa=True (head_idx maps to KV space), pass num_head_kv to normalization instead. Verified: BlockSparse + IndexSparse × LoopQ/LoopK × PackGQA/NoPG all 8 combinations PASS (B=1, Sq=512, NHQ=8, NHK=2, HD=128).
…Sparse BWD The MMA warp could overwrite dV(j) in TMEM before the reduce warp finished scatter-reducing it to global memory. For IndexSparse paths, add a pipeline_dKV.sync_object_empty.wait after phase flip to stall the MMA until the reduce warp has consumed dV from TMEM. Fixes multi_kiter and multi_3iter dV correctness (all 6/6 tests PASS).
IndexSparse uses its own per-KV-head M-block scheduling with 128-tile geometry, which conflicts with pack_gqa q_stage=2 requiring 256-aligned block sizes. Disable pack_gqa when IS is active in both FWD and BWD, and condition seqlen_q_packgqa on the actual pack_gqa state. All 8/8 IS tests PASS (including gqa8 and gqa8_partial). Non-IS regression: 3/3 PASS (dense_mha, dense_gqa, dense_short).
IS builds BST with m_block_size=128. When qseqlen > tile_m, q_stage=2 makes base_m_block=256 which fails the block-sparse normalization check. IS manages its own Q-tile scheduling so q_stage double-buffering is unnecessary. Force q_stage=1 for IS paths. 8/8 IS tests PASS, 3/3 non-IS regression PASS.
- BWD: early T2R dV before dK scatter to unblock MMA sooner - BWD: all 128 reduce threads participate in IS scatter (was 32) - BWD: SMEM transpose + 128B bulk scatter replacing per-vector 16B ops - BWD: independent sScatter SMEM buffer eliminates in-place transpose barrier, saving 1 barrier per scatter call (3 vs 4) - FWD/BWD: restore q_stage=2 for IS, enable PackGQA with IS - FWD/BWD: GQA BST head-dim expansion for NHK>1 scenarios - BWD: row-major IS dk/dv_accum postprocess (simple scale+convert) - BWD: PackGQA dQ postprocess on packed accumulators before unpacking BWD IS: 103-117 TFLOPS (1.47-1.57x vs prior), FWD IS: 581-685 TFLOPS
…en indices, E2 atomic scatter
R1: Deferred dV T2R - move dV TMEM-to-register after dK scatter to reduce
simultaneous register pressure (dK and dV no longer live at same time)
R2: SMEM token indices preload - preload token_indices from GMEM to SMEM
before scatter loops, eliminating repeated GMEM reads during scatter
E2: Add per-element atomicAdd scatter mode (env var gated) as alternative
to bulk cp.reduce.async scatter (MAGI_ATTENTION_FFA_CUTEDSL_IS_SCATTER_ATOMIC=1)
…LOPS) - Fuse scatter transpose into direct sdQacc->sScatter 16B vector copies, eliminating the 32-float register staging buffer and its spill traffic - Defer cp.async.bulk wait to just before next sScatter overwrite so the bulk DRAM round-trip overlaps with R2S work of the current stage - Pad sScatter row stride 32->36 fp32 to break 32-way SMEM bank conflict - Split dK R2S into 2-stage batches and move dV T2R before dK batch-2 scatter to release the dKV pipeline earlier BWD IS: 113-129 -> 165-195 TFLOPS. Correctness 8/8 (bulk + atomic modes), dense BWD regression unaffected.
Bypass the accumulator-layout sdQacc round-trip in sparse dK/dV reduce. Each reduce thread already owns one token row in its T2R fragment, so stage that row directly into padded row-major sScatter for bulk reduce, or issue per-element atomic adds directly from registers in atomic mode. This removes the conflicting sdQacc gather, sparse R2S barriers, and 175 lines of batching/transpose code. BWD IndexSparse improves from 165-195 to 273-336 TFLOPS and exceeds the SM90 baseline across 32K-256K. Validation: bulk 8/8, atomic 8/8, PackGQA 4/4; dense BWD unchanged.
Bug: sparse_tensor_m_block() divided m_block by qhead_per_kvhead for IndexSparse, but tile_token_indices is already built in packed M_blocks space by prepare_index_sparse_tiles(pack_gqa=True). This mapped all packed tiles to sparse index 0, making every tile read tile-0 token indices — hidden when all tiles share uniform indices (old test), exposed when per-tile random indices are used (real workloads). Fix: Remove qhpk division in all IS-specific m_block→tile_token_indices mappings (FWD ffa_fwd_sm100.py, BWD ffa_bwd_sm100.py, FWD paged_kv.py IndexSparseKVLoader). Test: Add tests/test_kernel/cutedsl/test_ffa_index_sparse.py — a proper sparse sweep covering seqlen x topk x head_config x PackGQA x batch x headdim x unaligned_seqlen, with random per-tile token indices. Both bulk and atomic scatter modes: 18 passed + 1 xfail (pre-existing topk non-multiple-of-128 multi-iter issue).
cennn
force-pushed
the
feat/cutedsl-sm100-sparse
branch
from
July 22, 2026 13:38
5525bf0 to
d1716b9
Compare
…M90 convention) - prepare_index_sparse_tiles: init tile_token_indices with -1 instead of 0 - IndexSparseKVLoader.preload_token_indices: clamp negative indices to 0 (safe dummy load, avoids OOB from -1 sentinel) - Aligns with SM90 CUTLASS IndexAttnBlockMeta convention where: * Host fills padding with -1 * Kernel maps -1 to safe position 0 * Softmax masks via is_valid_total/seqlen_k
…rrival softmax_block_sparse_sm100 consumed S-tiles in ascending block order (reverse of load order) but labeled mask steps descending and applied mask_seqlen only on the first step. For IndexSparse with a partial last block (topk %% 128 != 0, mask_block_cnt > 1), the seqlen_k/padding mask landed on the wrong physical block, leaving the partial tail unmasked. Fix: ascending per-step mask labels + mask_seqlen=True on every mask step (full blocks no-op), mirroring SM90 block-identity-driven padding mask. Rewrite test_ffa_index_sparse.py to mirror SM90 test_index_attn.py tier/config structure and parameter ranges; add SM100-only partial-topk tier. Switch grad relative-error metric to normalized L1 (robust to near-zero dK/dV at high GQA pack ratios). 25 passed.
Mirror SM90 tests/test_attn/test_index_attn.py: gate on max absolute difference (atol=0.01 forward O, 0.05 gradients) instead of cosine + per-element relative error. max-abs is more interpretable (cosine stays ~1 under systematic scale error; per-element rel blows up on near-zero dK/dV at high GQA pack ratios). All 25 cases pass with max_abs <= 0.031.
Reuse magi_attention.testing.precision.assert_close and mirror the exact atol/rtol/mismatch constants from remote main tests/test_attn/sparse_test_utils.py (FWD 0.01/0.05; dQ 0.02/0.3; dK 0.02/0.15; dV 0.02/0.05; mismatch 1%) instead of the ad-hoc max-abs thresholds. Validates O/dQ/dK/dV exactly like main compare_sdpa_fwd/compare_sdpa_bwd_all. 25 passed.
Post-merge cleanup: rename internal index_attn_indices -> index_sparse_indices (incl. index_attn_indices_to_block_sparse helper) in the SM100 cutedsl kernel to match main PR #322 taxonomy. No functional change; cutedsl index sparse sweep 25 passed.
… SM80/SM90 compat The DSL kernel signature for SM80/SM90 does not include descale_tensors or mTileTokenIndices. Gate these compile_args/call_args entries behind major_arch in [10, 11] so that SM80 force-fallback tests (test_ffa_simple) no longer get too-many-positional-arguments errors. Verified: test_ffa_simple 2 passed, test_ffa_index_sparse 25 passed.
…nd_bst_heads helper - Rename remaining IndexAttn references to IndexSparse in comments/docstrings - Update test file references from test_index_attn.py to test_index_sparse.py - Extract duplicated BST head-expansion logic (FWD/BWD) into _expand_bst_heads() - No logic changes; all 25 index_sparse + 2 ffa_simple tests pass
Mirrors exps/attn/sparse/bench_sparse_analysis/phase6_video_production.py but drives _flex_flash_attn_fwd/_bwd + prepare_index_sparse_tiles directly on B300. Compares Dense (K=topk) vs IndexSparse (scatter from full KV), PackGQA, bf16. Same scenario config: kvseqlen/64=qseqlen, kvseqlen/8=topk. Results (B300): FWD Dense: 1509-1941 TFLOPS | FWD IS: 464-541 TFLOPS BWD Dense: 325-1155 TFLOPS | BWD IS: 273-342 TFLOPS
- Add InnerLoadMode/InnerStoreMode/OuterStoreMode enums to sparse_utils.py - FWD IS-TMA: dispatch TMA load when sparse_k_block_size >= 128 - BWD IS-TMA: dispatch TMA load for K/V + TMA reduce-add store for dK/dV - Wire inner_load_tma through flex_flash_attn.py to FWD/BWD kernel constructors - Expose sparse_k_block_size in prepare_index_sparse_tiles (sets inner_load_mode) - Refactor test_ffa_index_sparse.py to SM90-style sweep: 3 test functions covering 44 configs (18 classic + 21 comprehensive + 5 partial_topk) - Reuse build_index_sparse_indices from shared magi_attention.utils.sparse_utils Correctness: IS-TMA FWD/BWD cosine=1.0 vs scatter reference (25/25 + 3/3 sweep) Perf: IS-TMA kbs=128 FWD achieves 71-94% of Dense; BWD 18-45% (sparse scheduler overhead)
…r to InnerLoadMode enum Guard IS-scatter-specific code (IndexSparseKVLoader, sScatter/sTokenIndices smem, mK_raw/mV_raw tensors, extra pipeline barrier) with InnerLoadMode != Tma so that IS-TMA (kbs>=128) kernels no longer pay for unused scatter infrastructure. Refactor inner_load_tma bool to InnerLoadMode enum for clarity. BWD IS-TMA vs Dense-LoopK: 78-83% -> 86-97% (+5-17pp).
…TMA (kbs=128) Merge v2 temp script into 0-run_video_production.py. Drop IS-scatter \(legacy per-token path\), keep only the three production-relevant methods: Dense \(gathered KV\), BlockSparse \(kbs=128\), IndexSparse-TMA \(kbs=128\). Add JIT cache clearing between methods to avoid type mismatches.
FWD compile_key was missing is_index_sparse and inner_load_mode, causing BS and IS-TMA kernels to share the same cache entry and type-mismatch at call time. BWD already had is_index_sparse in its key. Also refactor video production benchmark: Dense BWD uses LoopK for fair comparison, split IS tile creation to reduce peak memory (fixes 256K OOM), remove compile_cache.clear() workaround.
IS indices (B, NHQ, SQ, topk) were materialized with .contiguous(), consuming NHQ * SQ * topk * 4B (256 GB at 512K). Since all Q heads share the same pattern (NHK=1), use expand() view (zero memory) instead. Also rename IS-TMA -> IS kbs=128 (TMA is the default for kbs>=128).
…est params with SM90
…ssic+comprehensive only
…ode enum dispatch
…55 to avoid flakiness
…1-2% Root cause: IS-TMA classified all K-blocks as mask_block (full_block_cnt=0), forcing per-tile softmax masking on every iteration. BS uses full_block for most blocks, skipping the masking loop entirely. Fix: sparse_utils.py classifies IS-TMA blocks as full_block (tail block is mask only when TOPK pct 128 != 0). ffa_fwd_sm100.py IS-TMA producer uses produce_block_sparse_inner_iters_sm100 with wrapper load functions for logical to physical block translation. FWD IS-TMA now reaches 97.9-99.5% of BS (was 84-92%). BWD IS-TMA unchanged at 99.3-99.9% of BS. All 4 sparse tests pass.
- Relax LoopK-only assert: IS-scatter still requires LoopK, IS-TMA supports both - Add _build_is_tma_bwd_loopq_bst: builds transposed BST for InnerLoopQ - IS-TMA InnerLoopQ runs as pure BlockSparse (no mTileTokenIndices needed) - Fix PackGQA seqlen_q for InnerLoopQ BST validation (all methods) - Benchmark: add BWD InnerLoopQ subplot, fix BS InnerLoopQ BST for PackGQA
The BWD InnerLoopQ benchmark used per-Q-position independent random K-block selection, causing ~47%% of coarse Q-block tiles to be partially valid (wasted computation within subtile_factor=2 blocks). Fix: align selection to coarse-block boundaries (pairs of 2 Q positions share the same K-block selection). This eliminates amplification and represents realistic production workloads. Results at 32K: BS BWD-Q improves from 534T to 890T (1.67x), achieving 69%% of Dense (up from 41%%). At 64K+: sparse reaches 85-97%% of Dense.
…lder for RangeMerge
…bst_heads, remove dead _is_tma_loopq, parameterize >>7 shift
…e enum dispatch - FWD: outer_store_mode param replaces use_tma_O boolean - BWD: outer_store_mode param replaces use_tma_store boolean - flex_flash_attn: derive_outer_store_mode() computes mode from config - sparse_utils: add derive_outer_store_mode helper and OuterStoreMode.Tma1d docs - compile_key: include outer_store_mode for proper JIT cache isolation - Tma1d path shares STG codepath (future: per-row cp.async.bulk S2G)
- Add cpasync_bulk_copy_s2g @dsl_user_op for raw cp.async.bulk S2G PTX - Construct linear row-major sO layout (identity swizzle) for Tma1d mode - Implement per-row bulk S2G in both epilogue warp and correction_warp_epi paths - Adapt correction_epilogue R2S: use plain partition_D for linear layout - Update derive_outer_store_mode: Tma1d only for varlen without PackGQA (PackGQA makes row->gmem mapping non-trivial, falls back to Stg) - Verified: 6/6 varlen configs pass (MHA/GQA, D=64/128, multi-batch) - Verified: 4/4 sparse tests pass (IndexSparse + BlockSparse)
IS-TMA (kbs>=128) now uses InnerLoopQ (swap_bwd_qk_loop=False) where each CTA exclusively owns one K-block, eliminating atomic reduce-add contention for dKV. This makes dKV accumulation inherently deterministic without needing the semaphore-based ordering mechanism. Changes: - tile_scheduler: add missing producer_tail() to SingleTileLPTBwdScheduler (TileSchedulerProtocol compliance) - test_ffa_index_sparse: switch BWD to InnerLoopQ for kbs>=128, remove the relaxed mismatch_threshold=0.50 hack, use uniform 1% threshold for all dV
…ation - Add test_block_sparse_comprehensive_sweep_loopq: 12 configs (6 head_configs x 2 dims) verifying BS BWD LoopQ produces deterministic dKV (strict 1% mismatch threshold) - Refactor _run_bs_config to accept swap_bwd_qk_loop parameter - Add _build_loopq_bst helper for transposed BST construction - Verified Dense BWD deterministic support: - LoopQ + deterministic=True: dK, dV, dQ all bit-exact - LoopK + deterministic=True: dQ bit-exact (dKV non-deterministic as expected)
…ize tests - When deterministic=True + swap_bwd_qk_loop=True + qhpk>1: internally force LoopQ to guarantee bit-exact dK/dV/dQ (LoopK reduce warps lack semaphore) - Fix transpose_block_sparse_tensors to include full_block in presence matrix - Refactor BS/IS sweep tests to @pytest.mark.parametrize style (74 test cases)
Match CUTLASS SM90 static_assert: Deterministic mode is not supported yet when BwdInnerLoopK is true. Remove silent/auto LoopQ switch.
This branch was successfully deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
在 CuTe-DSL SM100 BWD 内核中新增 InnerLoopK 路径(
swap_bwd_qk_loop=True),支持外层遍历 Q blocks、内层流式迭代 K/V blocks。在此基础上支持 IndexSparse 和 BlockSparse 的 BWD LoopK,使得 IndexAttn 等稀疏注意力场景可以选择 LoopK 方向以获得更好的 dQ 存储效率(dQ TMEM 累积 + 单次 store,而非 LoopQ 的 atomicAdd dQ)。核心改动
BWD LoopK warp-specialized pipeline
load_loop_k(producer)、mma_loop_k(consumer)、dKV_reduce_loop_k(reduce warps)三个 LoopK 专用方法pipeline_K_load(与pipeline_Q分离,避免 mbarrier 冲突)cpasync_reduce_bulk_add_f32原子归约到 GMEMIndexSparse + BlockSparse BWD LoopK
index_attn_indices_to_block_sparse():token-level indices(total_q, NHK, topk)→ block-levelBlockSparseTensors(GPU 矢量化,scatter_构建 presence matrix)prepare_block_sparse_bwd_loopk():forward-direction normalize(per M-block → N-block list)get_curr_blocksparse_tensors+get_m_block_from_iter_bwd其他
dKV_postprocess=True(fp32 累积 → dtype 转换 + softmax_scale 缩放)disable_2cta条件扩展)正确性验证
B300 SM103 环境,D=128, bf16:
Dense LoopK(vs LoopQ, bit-identical, max_diff=0.000000):
BlockSparse LoopK(vs LoopQ reference):
IndexSparse LoopK(vs LoopQ reference):
限制
dKV_postprocess路径(fp32 原子累加),尚未做 direct TMA store 优化