diff --git a/flash_attn/cute/block_sparse_utils.py b/flash_attn/cute/block_sparse_utils.py index dd95395ed04..d55e967b9cb 100644 --- a/flash_attn/cute/block_sparse_utils.py +++ b/flash_attn/cute/block_sparse_utils.py @@ -81,6 +81,96 @@ def get_curr_blocksparse_tensors( return _get_curr_blocksparse_tensors(batch_idx, head_idx, m_block, blocksparse_tensors) +@cute.jit +def run_block_sparse_mainloop_sm80( + blocksparse_tensors: BlockSparseTensors, + batch_idx: cutlass.Int32, + head_idx: cutlass.Int32, + m_block: cutlass.Int32, + seqlen_info: SeqlenInfoQK, + mma_one_n_block: Callable, + mask_fn: Callable, + mask_mod: cutlass.Constexpr, + fastdiv_mods=None, +): + """Block-sparse mainloop for the SM80/SM120 (non-warp-specialized) forward kernel. + + Mask blocks are processed first (applying ``mask_mod`` + seqlen masking), then full + blocks (seqlen masking only, no ``mask_mod``). Blocks are visited in descending + ``n_block`` order to match the dense mainloop. The first full block always receives + seqlen masking even when mask blocks preceded it, since a full block may sit at the + highest ``n`` position (the seqlen_kv boundary) regardless of mask-block positions. + + This mirrors the masking contract of ``consume_block_sparse_loads`` (SM90/SM100) but + drives the non-warp-specialized ``mma_one_n_block`` (one load+compute per block); the + cp.async pipeline has no producer warp to run the warp-specialized produce/consume + helpers, and sparse blocks are not contiguous so cross-block prefetch does not apply. + + Args: + mma_one_n_block: callable ``(n_block, mask_fn, is_first_n_block) -> None`` that + loads K/V for ``n_block`` and runs QK GEMM, mask, softmax, and PV GEMM. + mask_fn: ``mask.apply_mask`` partial with batch/head/m_block/thr_mma bound. + mask_mod: the user ``mask_mod`` constexpr (``None`` -> no mask_mod application). + fastdiv_mods: fast-division helpers, used only when ``mask_mod is not None``. + """ + ( + curr_mask_block_cnt, + curr_mask_block_idx, + curr_full_block_cnt, + curr_full_block_idx, + ) = get_curr_blocksparse_tensors( + batch_idx, head_idx, m_block, blocksparse_tensors, seqlen_info + ) + + # Mask blocks: first gets is_first=True + seqlen masking; the rest do not. + if curr_mask_block_cnt > 0: + n_block = curr_mask_block_idx[curr_mask_block_cnt - 1] + mma_one_n_block( + n_block=n_block, + mask_fn=partial( + mask_fn, + mask_mod=mask_mod, + mask_seqlen=True, + fastdiv_mods=fastdiv_mods if const_expr(mask_mod is not None) else None, + ), + is_first_n_block=True, + ) + for i in cutlass.range(1, curr_mask_block_cnt): + n_block = curr_mask_block_idx[curr_mask_block_cnt - 1 - i] + mma_one_n_block( + n_block=n_block, + mask_fn=partial(mask_fn, mask_mod=mask_mod, mask_seqlen=False), + is_first_n_block=False, + ) + + # Full blocks: no mask_mod. The first full block always gets seqlen masking; whether + # it is the very first block overall (is_first) depends on there being no mask blocks. + if const_expr(curr_full_block_idx is not None): + if curr_full_block_cnt > 0: + n_block = curr_full_block_idx[curr_full_block_cnt - 1] + # is_first_n_block is a compile-time flag, so branch on the runtime + # "any mask blocks?" predicate and pass a literal in each arm. + if curr_mask_block_cnt == 0: + mma_one_n_block( + n_block=n_block, + mask_fn=partial(mask_fn, mask_mod=None, mask_seqlen=True), + is_first_n_block=True, + ) + else: + mma_one_n_block( + n_block=n_block, + mask_fn=partial(mask_fn, mask_mod=None, mask_seqlen=True), + is_first_n_block=False, + ) + for j in cutlass.range(1, curr_full_block_cnt): + n_block = curr_full_block_idx[curr_full_block_cnt - 1 - j] + mma_one_n_block( + n_block=n_block, + mask_fn=partial(mask_fn, mask_mod=None, mask_seqlen=False), + is_first_n_block=False, + ) + + # NOTE [SM100 block-sparse empty tiles: mbarrier contract] # # For block-sparse SM100 forward, a given (m_block, stage) Q tile can have zero active diff --git a/flash_attn/cute/flash_fwd.py b/flash_attn/cute/flash_fwd.py index 73c50ec9e8f..b4d84f67fba 100644 --- a/flash_attn/cute/flash_fwd.py +++ b/flash_attn/cute/flash_fwd.py @@ -33,6 +33,7 @@ from flash_attn.cute.pack_gqa import PackGQA from flash_attn.cute.named_barrier import NamedBarrierFwd from flash_attn.cute.block_sparsity import BlockSparseTensors +from flash_attn.cute.block_sparse_utils import run_block_sparse_mainloop_sm80 from flash_attn.cute.tile_scheduler import SingleTileScheduler, SingleTileVarlenScheduler, TileSchedulerArguments from flash_attn.cute.utils import AuxData @@ -704,6 +705,7 @@ def __call__( grid_dim = TileScheduler.get_grid_shape(tile_sched_params) softmax_scale_log2, softmax_scale = utils.compute_softmax_scale_log2(softmax_scale, self.score_mod) fastdiv_mods = utils.compute_fastdiv_mods(mQ, mK, self.qhead_per_kvhead, self.pack_gqa, aux_data.tensors) + self.use_block_sparsity = const_expr(blocksparse_tensors is not None) self.kernel( mQ, @@ -735,6 +737,7 @@ def __call__( TileScheduler, aux_data, fastdiv_mods, + blocksparse_tensors, ).launch( grid=grid_dim, block=[self.num_threads, 1, 1], @@ -774,6 +777,7 @@ def kernel( TileScheduler: cutlass.Constexpr[Callable], aux_data: AuxData = AuxData(), fastdiv_mods=None, + blocksparse_tensors: Optional[BlockSparseTensors] = None, ): # Thread index, block index tidx, _, _ = cute.arch.thread_idx() @@ -951,6 +955,23 @@ def kernel( aux_data=aux_data, fastdiv_mods=fastdiv_mods, ) + # Non-pipelined per-block compute for the block-sparse mainloop (binds seqlen, + # since block-sparse blocks are dispatched by index rather than by the dense loop). + compute_one_n_block_bs = partial( + self.mma_one_n_block_bs, + mma_params=mma_params, + smem_copy_params=smem_copy_params, + softmax=softmax, + load_K=load_K, + load_V=load_V, + score_mod=self.score_mod, + batch_idx=batch_size, + head_idx=num_head, + m_block=m_block, + seqlen=seqlen, + aux_data=aux_data, + fastdiv_mods=fastdiv_mods, + ) # /////////////////////////////////////////////////////////////////////////////// # Prologue @@ -970,23 +991,35 @@ def preprocess_Q(): # If Q_in_regs, we load Q, then load 1 stage of K, then (optionally) rotate Q and # read from smem_q to registers, then load V. # If !Q_in_regs, we load Q, load all stages of K & V, then (optionally) rotate Q. - if const_expr(self.Q_in_regs): - load_K(n_block, smem_pipe_write=0, need_predicates=True) - cute.arch.cp_async_commit_group() - preprocess_Q() - cute.arch.barrier() # Make sure all threads have read smem_q before loading V - - for stage in cutlass.range_constexpr(self.num_stages): - if const_expr(not self.Q_in_regs or stage > 0): - if stage == 0 or n_block - stage >= 0: - load_K(n_block - stage, smem_pipe_write=stage, need_predicates=stage == 0) - cute.arch.cp_async_commit_group() - if const_expr(stage < self.num_stages - 1): - if stage == 0 or n_block - stage >= 0: - load_V(n_block - stage, smem_pipe_write=stage, need_predicates=stage == 0) + if const_expr(self.use_block_sparsity): + # Block-sparse: each n_block loads its own K/V (blocks are not contiguous), + # so we skip the dense contiguous prefetch and just make Q available. Drain + # all pending async copies (only Q is in flight) so the per-block loads in + # mma_one_n_block_bs start from a clean cp.async group count. + cute.arch.cp_async_wait_group(0) + cute.arch.barrier() + if const_expr(self.Q_in_regs): + tSrQ_copy_view = smem_thr_copy_Q.retile(tSrQ) + cute.copy(smem_thr_copy_Q, tSsQ, tSrQ_copy_view) + cute.arch.barrier() + else: + if const_expr(self.Q_in_regs): + load_K(n_block, smem_pipe_write=0, need_predicates=True) cute.arch.cp_async_commit_group() - if const_expr(not self.Q_in_regs): - preprocess_Q() + preprocess_Q() + cute.arch.barrier() # Make sure all threads have read smem_q before loading V + + for stage in cutlass.range_constexpr(self.num_stages): + if const_expr(not self.Q_in_regs or stage > 0): + if stage == 0 or n_block - stage >= 0: + load_K(n_block - stage, smem_pipe_write=stage, need_predicates=stage == 0) + cute.arch.cp_async_commit_group() + if const_expr(stage < self.num_stages - 1): + if stage == 0 or n_block - stage >= 0: + load_V(n_block - stage, smem_pipe_write=stage, need_predicates=stage == 0) + cute.arch.cp_async_commit_group() + if const_expr(not self.Q_in_regs): + preprocess_Q() # /////////////////////////////////////////////////////////////////////////////// # Mainloop @@ -1016,45 +1049,62 @@ def preprocess_Q(): fastdiv_mods=fastdiv_mods if const_expr(self.mask_mod is not None) else None, ) - # First iteration with seqlen masking - smem_pipe_read = Int32(0) - smem_pipe_write = Int32(self.num_stages - 1) - compute_one_n_block( - n_block, - smem_pipe_read, - smem_pipe_write, - is_first_n_block=True, - seqlen=seqlen, - mask_fn=partial(mask_fn, mask_mod=self.mask_mod, mask_seqlen=True), - ) - smem_pipe_read = self.advance_pipeline(smem_pipe_read) - smem_pipe_write = self.advance_pipeline(smem_pipe_write) - # Next couple of iterations with causal masking - if const_expr(self.is_causal or self.is_local): - n_block_min_causal_local_mask = block_info.get_n_block_min_causal_local_mask( - seqlen, m_block, n_block_min + if const_expr(self.use_block_sparsity): + # Block-sparse mainloop: visit only the active mask/full blocks for this + # (batch, head, m_block). The Sm80 cp.async pipeline is not warp-specialized, + # so we use a per-block load+compute (mma_one_n_block_bs) rather than the + # warp-specialized produce/consume_block_sparse_loads path used by SM90/SM100. + run_block_sparse_mainloop_sm80( + blocksparse_tensors, + batch_size, + num_head, + m_block, + seqlen, + compute_one_n_block_bs, + mask_fn, + self.mask_mod, + fastdiv_mods if const_expr(self.mask_mod is not None) else None, ) - for n_tile in cutlass.range(n_block_max - 1 - n_block_min_causal_local_mask, unroll=1): - n_block = n_block_max - 2 - n_tile - compute_one_n_block( - n_block, - smem_pipe_read, - smem_pipe_write, - seqlen=seqlen, - mask_fn=partial(mask_fn, mask_mod=self.mask_mod, mask_seqlen=True), - ) - smem_pipe_read = self.advance_pipeline(smem_pipe_read) - smem_pipe_write = self.advance_pipeline(smem_pipe_write) - # The remaining iterations have no masking - for n_tile in cutlass.range(n_block, unroll=1): + else: + # First iteration with seqlen masking + smem_pipe_read = Int32(0) + smem_pipe_write = Int32(self.num_stages - 1) compute_one_n_block( - n_block - n_tile - 1, smem_pipe_read, smem_pipe_write, - seqlen=seqlen, is_first_n_block=False, - mask_fn=partial(mask_fn, mask_mod=self.mask_mod, mask_seqlen=False) + n_block, + smem_pipe_read, + smem_pipe_write, + is_first_n_block=True, + seqlen=seqlen, + mask_fn=partial(mask_fn, mask_mod=self.mask_mod, mask_seqlen=True), ) smem_pipe_read = self.advance_pipeline(smem_pipe_read) smem_pipe_write = self.advance_pipeline(smem_pipe_write) - # TODO: local + # Next couple of iterations with causal masking + if const_expr(self.is_causal or self.is_local): + n_block_min_causal_local_mask = block_info.get_n_block_min_causal_local_mask( + seqlen, m_block, n_block_min + ) + for n_tile in cutlass.range(n_block_max - 1 - n_block_min_causal_local_mask, unroll=1): + n_block = n_block_max - 2 - n_tile + compute_one_n_block( + n_block, + smem_pipe_read, + smem_pipe_write, + seqlen=seqlen, + mask_fn=partial(mask_fn, mask_mod=self.mask_mod, mask_seqlen=True), + ) + smem_pipe_read = self.advance_pipeline(smem_pipe_read) + smem_pipe_write = self.advance_pipeline(smem_pipe_write) + # The remaining iterations have no masking + for n_tile in cutlass.range(n_block, unroll=1): + compute_one_n_block( + n_block - n_tile - 1, smem_pipe_read, smem_pipe_write, + seqlen=seqlen, is_first_n_block=False, + mask_fn=partial(mask_fn, mask_mod=self.mask_mod, mask_seqlen=False) + ) + smem_pipe_read = self.advance_pipeline(smem_pipe_read) + smem_pipe_write = self.advance_pipeline(smem_pipe_write) + # TODO: local # normalize acc_O by row_sum and calculate the lse row_scale = softmax.finalize() @@ -1192,6 +1242,89 @@ def load_K_next(): ) # if const_expr(self.num_stages > 1): # load_K_next() + + @cute.jit + def mma_one_n_block_bs( + self, + n_block: Int32, + mma_params: SimpleNamespace, + smem_copy_params: SimpleNamespace, + softmax: Softmax, + load_K: Callable, + load_V: Callable, + score_mod: Callable | None, + batch_idx: cutlass.Int32, + head_idx: cutlass.Int32, + m_block: cutlass.Int32, + seqlen: SeqlenInfoQK, + aux_data: AuxData = AuxData(), + fastdiv_mods=None, + mask_fn: Optional[Callable] = None, + is_first_n_block: cutlass.Constexpr = False, + ): + """Process one KV block for block-sparse attention (load, QK GEMM, mask, softmax, PV GEMM). + + Unlike ``compute_one_n_block``, this does not overlap the loads with the next block: + in the block-sparse case the next block index is not known ahead of time (blocks are + not contiguous), so each block loads its own K/V, waits, then computes. + """ + acc_S = cute.make_rmem_tensor( + mma_params.thr_mma_qk.partition_shape_C((self.tile_m, self.tile_n)), Float32 + ) + acc_S.fill(0.0) + + load_K(n_block, smem_pipe_write=0, need_predicates=True) + cute.arch.cp_async_commit_group() + load_V(n_block, smem_pipe_write=0, need_predicates=True) + cute.arch.cp_async_commit_group() + cute.arch.cp_async_wait_group(1) + cute.arch.barrier() + + sm80_utils.gemm( + mma_params.thr_mma_qk, + acc_S, + mma_params.tSrQ, + mma_params.tSrK, + smem_copy_params.tSsQ, + smem_copy_params.tSsK[None, None, None, 0], + smem_copy_params.smem_thr_copy_Q, + smem_copy_params.smem_thr_copy_K, + A_in_regs=self.Q_in_regs, + ) + if const_expr(score_mod is not None): + self.apply_score_mod( + mma_params.thr_mma_qk, + batch_idx, + head_idx, + m_block, + acc_S, + n_block, + softmax_scale=softmax.softmax_scale, + seqlen=seqlen, + aux_data=aux_data, + fastdiv_mods=fastdiv_mods, + ) + + cute.arch.cp_async_wait_group(0) + cute.arch.barrier() + + if const_expr(mask_fn is not None): + mask_fn(acc_S, n_block=n_block) + + row_scale = softmax.online_softmax(acc_S, is_first=is_first_n_block, check_inf=True) + softmax.rescale_O(mma_params.acc_O, row_scale) + rP = cute.make_fragment_like(acc_S, self.dtype) + rP.store(acc_S.load().to(self.dtype)) + tOrP = layout_utils.reshape_acc_to_frgA(rP) + sm80_utils.gemm_rs( + mma_params.thr_mma_pv, + mma_params.acc_O, + tOrP, + mma_params.tOrVt, + smem_copy_params.tOsVt[None, None, None, 0], + smem_copy_params.smem_thr_copy_V, + ) + @cute.jit def apply_score_mod( self,