Skip to content

[WIP]: feat(cutedsl-sm100): BWD LoopK with warp-specialized pipeline + IndexSparse/BlockSparse support#334

Open
cennn wants to merge 7 commits into
mainfrom
feat/cutedsl-sm100-sparse
Open

[WIP]: feat(cutedsl-sm100): BWD LoopK with warp-specialized pipeline + IndexSparse/BlockSparse support#334
cennn wants to merge 7 commits into
mainfrom
feat/cutedsl-sm100-sparse

Conversation

@cennn

@cennn cennn commented Jul 13, 2026

Copy link
Copy Markdown
Collaborator

Summary

在 CuTe-DSL SM100 BWD 内核中新增 InnerLoopK 路径(swap_bwd_qk_loop=True),支持外层遍历 Q blocks、内层流式迭代 K/V blocks。在此基础上支持 IndexSparseBlockSparse 的 BWD LoopK,使得 IndexAttn 等稀疏注意力场景可以选择 LoopK 方向以获得更好的 dQ 存储效率(dQ TMEM 累积 + 单次 store,而非 LoopQ 的 atomicAdd dQ)。

核心改动

  1. 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 冲突)
    • TMEM 布局重新编排:dQ persistent @0,S/P→dK @hdim,dP/dS @hdim+tile_m,dV @hdim+2*tile_m
    • dKV per-K-iter 通过 reduce warps 做 TMEM→SMEM→cpasync_reduce_bulk_add_f32 原子归约到 GMEM
    • dQ 在外层 tile 结束后一次性 reduce 到 GMEM
  2. IndexSparse + BlockSparse BWD LoopK

    • 新增 index_attn_indices_to_block_sparse():token-level indices (total_q, NHK, topk) → block-level BlockSparseTensors(GPU 矢量化,scatter_ 构建 presence matrix)
    • 新增 prepare_block_sparse_bwd_loopk():forward-direction normalize(per M-block → N-block list)
    • LoopK 所有 warp(load/mma/compute/reduce)的 sparse 路径共享 get_curr_blocksparse_tensors + get_m_block_from_iter_bwd
  3. 其他

    • LoopK 强制 dKV_postprocess=True(fp32 累积 → dtype 转换 + softmax_scale 缩放)
    • LoopK 禁用 2-CTA mode(disable_2cta 条件扩展)
    • K_stage=1 for D≥128(避免 SMEM 超限);K_stage=2 for D<128

正确性验证

B300 SM103 环境,D=128, bf16:

Dense LoopK(vs LoopQ, bit-identical, max_diff=0.000000):

配置 dQ dK dV
MHA full B=2 Sq=256 PASS PASS PASS
MHA full B=1 Sq=1024 PASS PASS PASS
MHA causal B=2 Sq=256 PASS PASS PASS
GQA full B=2 NHQ=8 NHK=2 PASS PASS PASS
GQA causal B=2 NHQ=8 NHK=2 PASS PASS PASS

BlockSparse LoopK(vs LoopQ reference):

配置 dQ dK dV
MHA B=1 Sq=512 sparsity=0.5 PASS (0.0005) PASS (0.001) PASS (0.0005)
MHA B=2 Sq=512 sparsity=0.75 PASS (0.001) PASS (0.001) PASS (0.001)

IndexSparse LoopK(vs LoopQ reference):

配置 dQ dK dV
MHA B=1 Sq=512 topk=256 PASS (0.002) PASS (0.0005) PASS (0.0005)
MHA B=2 Sq=512 topk=256 PASS (0.001) PASS (0.001) PASS (0.001)

限制

  • LoopK 当前仅支持 1-CTA、non-varlen、hdim≤128
  • LoopK 始终走 dKV_postprocess 路径(fp32 原子累加),尚未做 direct TMA store 优化

cennn added 4 commits July 11, 2026 02:37
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.
cennn added 3 commits July 17, 2026 15:10
…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).
@cennn cennn changed the title feat(cutedsl-sm100): BWD LoopK with warp-specialized pipeline + IndexSparse/BlockSparse support [WIP]: feat(cutedsl-sm100): BWD LoopK with warp-specialized pipeline + IndexSparse/BlockSparse support Jul 17, 2026
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