Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
90 changes: 90 additions & 0 deletions flash_attn/cute/block_sparse_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
233 changes: 183 additions & 50 deletions flash_attn/cute/flash_fwd.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -735,6 +737,7 @@ def __call__(
TileScheduler,
aux_data,
fastdiv_mods,
blocksparse_tensors,
).launch(
grid=grid_dim,
block=[self.num_threads, 1, 1],
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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,
Expand Down