Skip to content
Closed
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
20 changes: 11 additions & 9 deletions flash_attn/cute/flash_bwd_mla_dq_dqv_sm100.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@
Performs both dQ = dS @ K and dQv = dS @ V, where K and V are
gathered according to index tensor mIdxTopK.

This uses MQA with 128 heads.
This uses MQA with 1..128 heads and a fixed 128-row GEMM tile.

Inputs:
- dS: [batch, seqlen_q, nheads, top_k] or [total_q, nheads, top_k]
Expand Down Expand Up @@ -59,17 +59,17 @@ def __init__(
):
self.acc_dtype: Type[cutlass.Numeric] = acc_dtype
self.nheads = nheads
assert self.nheads == 128, (
"only 128 heads supported; will expand to include 64 heads in a future PR."
)
# Head padding: see pack_gqa.padded_qheads_tma_source.
self.tile_m = 128
assert 0 < self.nheads <= self.tile_m, f"at most {self.tile_m} heads, got {nheads}"
self.head_dim_k = head_dim_k or 0 # when head_dim_k not provided, dQ is not computed
self.head_dim_v = head_dim_v
self.top_k = top_k
self.tile_k = 128

self.cluster_shape_mn = (1, 2)
self.mma_tiler_dQ = (self.nheads, self.head_dim_k, self.tile_k)
self.mma_tiler_dQv = (self.nheads, self.head_dim_v // 2, self.tile_k)
self.mma_tiler_dQ = (self.tile_m, self.head_dim_k, self.tile_k)
self.mma_tiler_dQv = (self.tile_m, self.head_dim_v // 2, self.tile_k)
self.num_mainloop_iters = self.top_k // self.tile_k
self.arch = "sm_100"

Expand Down Expand Up @@ -187,12 +187,14 @@ def static_reshape(t: cute.Tensor, *static_shapes) -> cute.Tensor:
),
)

mdS = static_reshape(mdS, self.nheads, self.top_k)
mdQv = static_reshape(mdQv, self.nheads, self.head_dim_v)
# Dynamic head extent for the fixed tile; see pack_gqa.padded_qheads_tma_source.
nheads = self.nheads if self.nheads == self.tile_m else Int32(self.nheads)
mdS = static_reshape(mdS, nheads, self.top_k)
mdQv = static_reshape(mdQv, nheads, self.head_dim_v)
mV = static_reshape(mV, self.head_dim_v)
mIdxTopK = static_reshape(mIdxTopK, self.top_k)
if const_expr(self.compute_dQ):
mdQ = static_reshape(mdQ, self.nheads, self.head_dim_k)
mdQ = static_reshape(mdQ, nheads, self.head_dim_k)
mK = static_reshape(mK, self.head_dim_k)

# ---- layout info ----
Expand Down
124 changes: 91 additions & 33 deletions flash_attn/cute/flash_bwd_mla_sm100.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,11 @@

from quack import copy_utils, layout_utils

from flash_attn.cute.pack_gqa import pack_gqa_layout
from flash_attn.cute.pack_gqa import (
pack_gqa_layout,
padded_qheads_tma_source,
regroup_padded_qheads,
)
from flash_attn.cute.seqlen_info import SeqlenInfoQK
from flash_attn.cute.block_info import BlockInfo
import flash_attn.cute.blackwell_helpers as fa_sm100_utils
Expand Down Expand Up @@ -52,20 +56,27 @@ def __init__(
has_seqused_q: bool = False,
disable_bitmask: bool = False,
use_clc_scheduler: bool = True,
qhead_per_kvhead_valid: Optional[int] = None,
):
use_cpasync_load_KV = True
self.is_causal = is_causal
self.is_local = False
self.pack_gqa = True
self.qhead_per_kvhead = qhead_per_kvhead
# Head padding and scaleP/dPsum contract: see pack_gqa.padded_qheads_tma_source.
if qhead_per_kvhead_valid is None:
qhead_per_kvhead_valid = qhead_per_kvhead
assert 0 < qhead_per_kvhead_valid <= qhead_per_kvhead
self.qhead_per_kvhead_valid = qhead_per_kvhead_valid
self.pad_qheads = qhead_per_kvhead_valid != qhead_per_kvhead
self.nheads_kv = nheads_kv
self.has_seqused_q = has_seqused_q
self.use_tma_O = True
self.use_cpasync_load_KV = True
self.use_tma_KV = False
self.topk_length = topk_length
self.is_topk_gather = True
assert qhead_per_kvhead == 128 or qhead_per_kvhead == 64
assert qhead_per_kvhead in (64, 128), f"sparse MLA bwd supports 64 or 128 heads, got {qhead_per_kvhead}"

# user-provided option if topk indices guaranteed in bounds
self.disable_bitmask = disable_bitmask
Expand Down Expand Up @@ -211,7 +222,8 @@ def __init__(
self.num_stages_dP = 1
self.num_stages_dPt = 1
self.num_stages_dV = 2 # == hdimv splits, for Umma <-> Async
self.num_epi_stages_dV = 8 # == 2 splits x 4 slots/split
# Per hdimv split: 2 warpgroup halves x 2 subtile parities.
self.num_epi_stages_dV = 4

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Image

asked fables to make a viz for. why this change


self.num_stages_scaleP = 1
self.num_stages_dPsum = 1
Expand Down Expand Up @@ -251,7 +263,6 @@ def mbar_struct(num_stages):
sdO_struct,
sP_struct,
sdS_struct,
sQvt_struct,
sScaleP_struct,
sdPsum_struct,
) = (
Expand All @@ -261,11 +272,28 @@ def mbar_struct(num_stages):
(self.dtype, self.sdO_layout_staged),
(self.dtype, self.sPt_layout_staged),
(self.dtype, self.sdSt_layout_staged),
(self.dtype, self.sQvt_layout_staged),
(self.dtype_scale, self.sScaleP_layout_staged),
(self.dtype_scale, self.sdPsum_layout_staged),
]
)
# dV staging must not overwrite operands still in use by another split's MMA.
# If staging equals one dOt/Qvt stage, split s stages in place over stage s.
# Otherwise, both splits take turns using one staging area after the operands.
staging_elems = (
math.prod(self.tile_dV) * self.num_epi_stages_dV * self.dtype_acc.width // self.dtype.width
)
operand_elems = cute.cosize(self.sQvt_layout_staged)
assert self.num_stages_Qvt == self.num_hdimv_splits
stage_elems = operand_elems // self.num_hdimv_splits
if staging_elems == stage_elems:
self.sdV_split_offsets = [split * stage_elems for split in range(self.num_hdimv_splits)]
sQv_elems = operand_elems
else:
self.sdV_split_offsets = [operand_elems] * self.num_hdimv_splits
sQv_elems = operand_elems + staging_elems
sQvt_struct = cute.struct.Align[
cute.struct.MemRange[self.dtype, sQv_elems], self.buffer_align_bytes
]

(
mbar_ptr_V_struct, # load V
Expand Down Expand Up @@ -415,6 +443,12 @@ def __call__(
)
topk_length_dynamic = mIndexTopk.shape[0]

# TMA source contract: see pack_gqa.padded_qheads_tma_source.
if const_expr(self.pad_qheads):
mQv_valid, mdO_valid, mP_valid, mdS_valid = [
padded_qheads_tma_source(mX, self.qhead_per_kvhead_valid, head_idx=2)
for mX in (mQv, mdO, mP, mdS)
]
if const_expr(self.pack_gqa):
mQv, mdO, mP, mdS, mScaleP = [
pack_gqa_layout(mX, self.qhead_per_kvhead, self.nheads_kv, head_idx=2)
Expand All @@ -424,15 +458,18 @@ def __call__(
]
if const_expr(mdPsum is not None):
mdPsum = pack_gqa_layout(mdPsum, self.qhead_per_kvhead, self.nheads_kv, head_idx=1)
if const_expr(not self.pad_qheads):
mQv_valid, mdO_valid, mP_valid, mdS_valid = mQv, mdO, mP, mdS

# ((h/h_k, s_q), dv, h_k, b) -> (dv, (h/h_k, s_q), h_k, b)
# or ((h/h_k, total_q), dv, h_k) -> (dv, (h/h_k, total_q), h_k)
# *_valid: (h, dv, s_q, b) -> (dv, h, s_q, b).
mma_operand_layout_transpose = (
[1, 0, 2, 3] if const_expr(mCuSeqlensQ is None) else [1, 0, 2]
)
mQvt, mdOt = [
mQvt, mdOt, mQvt_valid, mdOt_valid = [
cute.make_tensor(mX.iterator, cute.select(mX.layout, mode=mma_operand_layout_transpose))
for mX in (mQv, mdO)
for mX in (mQv, mdO, mQv_valid, mdO_valid)
]

# fmt: off
Expand Down Expand Up @@ -502,21 +539,37 @@ def __call__(
)
cta_shape = cta_layout_vmnk.shape

def make_tma(make_fn, mX, smem_layout, mma_tiler, tiled_mma):
return make_fn(tma_load_op, mX, smem_layout, mma_tiler, tiled_mma, cta_shape)
def regroup(tensor, transposed=False):
if const_expr(self.pad_qheads):
# Transposed operands: undo the transpose, fold, and transpose back, mirroring
# how mdOt/mQvt derive from mdO/mQv.
if const_expr(transposed):
tensor = cute.make_tensor(
tensor.iterator, cute.select(tensor.layout, mode=mma_operand_layout_transpose)
)
tensor = regroup_padded_qheads(tensor, self.qhead_per_kvhead, head_idx=2)
if const_expr(transposed):
tensor = cute.make_tensor(
tensor.iterator, cute.select(tensor.layout, mode=mma_operand_layout_transpose)
)
return tensor

def make_tma(make_fn, mX, smem_layout, mma_tiler, tiled_mma, transposed):
atom, tensor = make_fn(tma_load_op, mX, smem_layout, mma_tiler, tiled_mma, cta_shape)
return atom, regroup(tensor, transposed)

A, B = cute.nvgpu.make_tiled_tma_atom_A, cute.nvgpu.make_tiled_tma_atom_B

# (atom_name, tensor_name, make_fn, m, smem_layout, mma_tiler, tiled_mma)
# (atom_name, tensor_name, make_fn, m, smem_layout, mma_tiler, tiled_mma, transposed)
_tma_specs = [
("tma_atom_dO", "tma_tensor_dO", B, mdO, self.sdO_layout, self.mma_tiler_VdO, tiled_mma_VdO),
("tma_atom_dOt", "tma_tensor_dOt", B, mdOt, self.sdOt_layout, self.mma_tiler_PtdOt, tiled_mma_PtdOt),
("tma_atom_Qvt", "tma_tensor_Qvt", B, mQvt, self.sQvt_layout, self.mma_tiler_dStQvt, tiled_mma_dStQvt),
("tma_atom_dO", "tma_tensor_dO", B, mdO_valid, self.sdO_layout, self.mma_tiler_VdO, tiled_mma_VdO, False),
("tma_atom_dOt", "tma_tensor_dOt", B, mdOt_valid, self.sdOt_layout, self.mma_tiler_PtdOt, tiled_mma_PtdOt, True),
("tma_atom_Qvt", "tma_tensor_Qvt", B, mQvt_valid, self.sQvt_layout, self.mma_tiler_dStQvt, tiled_mma_dStQvt, True),
]
_tmas = {}
for atom_name, tensor_name, make_fn, m, smem_layout, mma_tiler, tiled_mma in _tma_specs:
for atom_name, tensor_name, make_fn, m, smem_layout, mma_tiler, tiled_mma, transposed in _tma_specs:
_tmas[atom_name], _tmas[tensor_name] = (
make_tma(make_fn, m, smem_layout, mma_tiler, tiled_mma)
make_tma(make_fn, m, smem_layout, mma_tiler, tiled_mma, transposed)
)

(tma_atom_dO, tma_tensor_dO,
Expand All @@ -526,10 +579,11 @@ def make_tma(make_fn, mX, smem_layout, mma_tiler, tiled_mma):
# Make TMA load for P separately
tma_atom_P, tma_tensor_P = cute.nvgpu.cpasync.make_tiled_tma_atom(
cpasync.CopyBulkTensorTileG2SOp(),
mP,
mP_valid,
self.sP_layout,
self.tile_P,
)
tma_tensor_P = regroup(tma_tensor_P)

# ==== TMA store ====
tma_store_op = cpasync.CopyBulkTensorTileS2GOp()
Expand All @@ -545,8 +599,9 @@ def make_tma(make_fn, mX, smem_layout, mma_tiler, tiled_mma):
self.dtype_dV, self.dV_layout_major, self.tile_dV, self.num_epi_stages_dV
)
tma_atom_dS, tma_tensor_dS = cpasync.make_tiled_tma_atom(
tma_store_op, mdS, cute.select(sdS_layout_staged, mode=[0, 1]), self.tile_dS
tma_store_op, mdS_valid, cute.select(sdS_layout_staged, mode=[0, 1]), self.tile_dS
)
tma_tensor_dS = regroup(tma_tensor_dS)
# fmt: on

# ==== Allocate shared memory ====
Expand Down Expand Up @@ -805,10 +860,14 @@ def make_pipeline(cls, mbar_ptr, num_stages, producer, consumer, tx_count=None):
(storage.sQv, sQvt_layout_staged), # {dOt, Qvt, dV} overlap
]
)
sdV = cute.make_tensor(
cute.recast_ptr(sdOt.iterator, sdV_layout_staged.inner, self.dtype_acc), sdV_layout_staged.outer
)
assert cute.cosize(sdV) * self.dtype_acc.width // self.dtype.width == cute.cosize(sdOt)
# Staging placement: see _get_shared_storage_cls.
sdVs = [
cute.make_tensor(
cute.recast_ptr(sdOt.iterator + offset, sdV_layout_staged.inner, self.dtype_acc),
sdV_layout_staged.outer,
)
for offset in self.sdV_split_offsets
]

sScaleP = storage.sScaleP.get_tensor(sScaleP_layout_staged)
sdPsum = storage.sdPsum.get_tensor(sdPsum_layout_staged)
Expand Down Expand Up @@ -1053,7 +1112,7 @@ def make_pipeline(cls, mbar_ptr, num_stages, producer, consumer, tx_count=None):
self.dVacc_store(
mIndexTopk,
mdV,
sdV,
sdVs,
tdVtdV0,
tdVtdV1,
thr_mma_PtdOt,
Expand Down Expand Up @@ -1972,7 +2031,7 @@ def dVacc_store(
self,
mIndexTopk: cute.Tensor,
mdV: cute.Tensor,
sdV: cute.Tensor,
sdVs: list[cute.Tensor],
tdVtdV0: cute.Tensor,
tdVtdV1: cute.Tensor,
thr_mma_PtdOt: cute.ThrMma,
Expand Down Expand Up @@ -2020,15 +2079,14 @@ def dVacc_store(
tiled_copy_r2s = tiled_copy_2d(self.dtype_acc, 4, 64)
thr_copy_r2s = tiled_copy_r2s.get_slice(tidx % 64)

# ((4,1),1,8,(1,8)):((1,0),0,4,(0,2048))
tRS_sdV = thr_copy_r2s.partition_D(sdV)

tiled_copy_s2r = copy_utils.tiled_copy_2d(self.dtype_acc, 8, self.num_epilogue_threads, 4)
thr_copy_s2r = tiled_copy_s2r.get_slice(tidx)
# (V, M, N, STAGE)
tSR_sdV = thr_copy_s2r.partition_S(sdV)
# ((4,1),1,8,(1,8)):((1,0),0,4,(0,2048)) and (V, M, N, STAGE), per split
tRS_sdVs = [thr_copy_r2s.partition_D(sdV) for sdV in sdVs]
tSR_sdVs = [thr_copy_s2r.partition_S(sdV) for sdV in sdVs]
tRS_sdV, tSR_sdV = tRS_sdVs[0], tSR_sdVs[0]

cdV = cute.make_identity_tensor(cute.product_each(sdV.shape[:2]))
cdV = cute.make_identity_tensor(cute.product_each(sdVs[0].shape[:2]))
# (V, M, N)
tdVcdV = thr_copy_s2r.partition_S(cdV)

Expand Down Expand Up @@ -2086,18 +2144,18 @@ def dVacc_store(

tRS_rdV_cur = cute.make_tensor(tdVrdV_cur.iterator, tRS_rdV_cur_shape)

stage = 4 * split + 2 * wg_half + (i % 2)
cute.copy(tiled_copy_r2s, tRS_rdV_cur, tRS_sdV[None, None, None, stage])
stage = 2 * wg_half + (i % 2)
cute.copy(tiled_copy_r2s, tRS_rdV_cur, tRS_sdVs[split][None, None, None, stage])
cute.arch.fence_view_async_shared()
self.epi_barrier.arrive_and_wait()

tSR_rdV = cute.make_rmem_tensor(tdVrdV_out_shape, dtype=self.dtype_acc)

for w in cutlass.range_constexpr(2):
stage_out = 4 * split + 2 * w + (i % 2)
stage_out = 2 * w + (i % 2)
cute.copy(
tiled_copy_s2r,
tSR_sdV[None, None, None, stage_out],
tSR_sdVs[split][None, None, None, stage_out],
tSR_rdV[None, None, None, w],
)

Expand Down
Loading
Loading