Skip to content
Merged
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
149 changes: 121 additions & 28 deletions flashmask/flash_mask/cute/interface.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,11 @@
normalize_block_sparse_tensors,
)

try:
from ..utils import accum_zero_axis1_kv
except ImportError:
accum_zero_axis1_kv = None


def maybe_contiguous(x):
return x.contiguous() if x is not None and x.strides[-1] != 1 else x
Expand Down Expand Up @@ -693,6 +698,8 @@ def _flash_attn_bwd(
seqused_k: Optional[paddle.Tensor] = None,
learnable_sink: Optional[paddle.Tensor] = None,
deterministic: bool = False,
kv_postprocess_start: Optional[int] = None,
kv_postprocess_end: Optional[int] = None,
) -> Tuple[paddle.Tensor, paddle.Tensor, paddle.Tensor, Optional[paddle.Tensor]]:
compute_capability = paddle.device.cuda.get_device_capability()[0]
assert compute_capability in [10], "Unsupported compute capability. Supported: 10.x"
Expand Down Expand Up @@ -919,22 +926,40 @@ def _flash_attn_bwd(
dk_accum_shape = [num_head_kv, total_k_rounded_padded * head_dim_rounded]
dv_accum_shape = [num_head_kv, total_k_rounded_padded * head_dim_v_rounded]

# ---- Consolidate fp32 zero-fill workspaces into a single zeros() launch ----
# In Split-D BWD dq_accum must be zero-initialized (preprocess does not fully zero
# the [low | high] split layout, see PR #2447). dk_accum/dv_accum need zeros for
# bulk reduce-add accumulation. Folding multiple zeros() into one large zeros()
# cuts host-side bf16/fp32 fill overhead from 3 launches to 1 (~2/3 of fp32 fill
# time, ~567us/bwd in 4k benchmark).
kv_postprocess_enabled = kv_postprocess_start is not None or kv_postprocess_end is not None

def _kv_postprocess_range():
if not kv_postprocess_enabled:
return 0, seqlen_k
assert fixed_seqlen, "kv_postprocess range only supports fixed seqlen"
start = 0 if kv_postprocess_start is None else int(kv_postprocess_start)
end = seqlen_k if kv_postprocess_end is None else int(kv_postprocess_end)
assert 0 <= start <= end <= seqlen_k, (
f"invalid kv_postprocess range [{start}, {end}) for seqlen_k={seqlen_k}"
)
assert start % n_block_size == 0, (
f"kv_postprocess_start must be aligned to n_block_size={n_block_size}, got {start}"
)
assert end % n_block_size == 0, (
f"kv_postprocess_end must be aligned to n_block_size={n_block_size}, got {end}"
)
return start, end

kv_post_start, kv_post_end = _kv_postprocess_range()


# ---- Compute shapes for fp32 accum workspaces ----
def _numel(shape):
n = 1
for d in shape:
n *= d
return n

zero_kv_accum_range = kv_postprocess_enabled and need_kv_accum and accum_zero_axis1_kv is not None
zero_specs = [] # list of (key, shape, numel)
if is_split_d_bwd:
zero_specs.append(("dq_accum", dq_accum_shape, _numel(dq_accum_shape)))
if need_kv_accum:
if need_kv_accum and not zero_kv_accum_range:
zero_specs.append(("dk_accum", dk_accum_shape, _numel(dk_accum_shape)))
zero_specs.append(("dv_accum", dv_accum_shape, _numel(dv_accum_shape)))

Expand All @@ -956,8 +981,32 @@ def _numel(shape):
else:
dq_accum = paddle.empty(shape=dq_accum_shape, dtype=paddle.float32)
if need_kv_accum:
dk_accum = _accum_buffers["dk_accum"]
dv_accum = _accum_buffers["dv_accum"]
if zero_kv_accum_range:
# if start/end range is given, we can initialize only part of the accum buffer
dk_accum = paddle.empty(shape=dk_accum_shape, dtype=paddle.float32)
dv_accum = paddle.empty(shape=dv_accum_shape, dtype=paddle.float32)
if is_split_d_bwd:
dk_zero_hdim, dv_zero_hdim = head_dim // 2, head_dim_v // 2
dk_split, dv_split = True, True
elif is_split_dv_bwd:
dk_zero_hdim, dv_zero_hdim = head_dim_rounded, head_dim_v // 2
dk_split, dv_split = False, True
else:
dk_zero_hdim, dv_zero_hdim = head_dim_rounded, head_dim_v_rounded
dk_split, dv_split = False, False
accum_zero_axis1_kv(
dk_accum,
dv_accum,
kv_post_start,
kv_post_end - kv_post_start,
dk_zero_hdim,
dv_zero_hdim,
dk_split,
dv_split,
)
else:
dk_accum = _accum_buffers["dk_accum"]
dv_accum = _accum_buffers["dv_accum"]

dtype = paddle2cute_dtype_map[q.dtype]
q_tensor, k_tensor, v_tensor, o_tensor, do_tensor, dq_tensor, dk_tensor, dv_tensor = [
Expand Down Expand Up @@ -1202,6 +1251,45 @@ def _numel(shape):
num_threads = 256 if compute_capability == 9 else 128
arch = compute_capability * 10

kv_postprocess_enabled = kv_postprocess_start is not None or kv_postprocess_end is not None

def _kv_postprocess_range():
if not kv_postprocess_enabled:
return 0, seqlen_k
assert fixed_seqlen, "kv_postprocess range only supports fixed seqlen"
start = 0 if kv_postprocess_start is None else int(kv_postprocess_start)
end = seqlen_k if kv_postprocess_end is None else int(kv_postprocess_end)
assert 0 <= start <= end <= seqlen_k, (
f"invalid kv_postprocess range [{start}, {end}) for seqlen_k={seqlen_k}"
)
assert start % n_block_size == 0, (
f"kv_postprocess_start must be aligned to n_block_size={n_block_size}, got {start}"
)
return start, end

kv_post_start, kv_post_end = _kv_postprocess_range()

def _slice_kv_out(t):
if not kv_postprocess_enabled:
return t
return t[:, kv_post_start:kv_post_end, :, :]

def _slice_kv_accum(t, accum_hdim):
if not kv_postprocess_enabled:
return t
return t[..., kv_post_start * accum_hdim:kv_post_end * accum_hdim]

def _to_cute(t):
return from_dlpack(t.detach(), assumed_align=16).mark_layout_dynamic(
leading_dim=t.ndim - 1
)

def _kv_out_cute(t, tensor):
return _to_cute(_slice_kv_out(t)) if kv_postprocess_enabled else tensor

def _kv_accum_cute(t, tensor, accum_hdim):
return _to_cute(_slice_kv_accum(t, accum_hdim)) if kv_postprocess_enabled else tensor

def _postprocess_run(d_accum_t, d_out_t, scale, hd, block_size, atom_layout, swapAB,
use_2cta, cluster, cu_seqlens_t, seqused_t, cache_tag):
compile_key_post = (dtype, hd, arch, block_size, num_threads, atom_layout, swapAB,
Expand Down Expand Up @@ -1229,11 +1317,6 @@ def _slice_accum(t):
n = t.shape[-1] // 2
return t[..., :n], t[..., n:]

def _to_cute(t):
return from_dlpack(t.detach(), assumed_align=16).mark_layout_dynamic(
leading_dim=t.ndim - 1
)

# dQ split [low | high] postprocess.
# Pass non-contiguous views directly: kernel uses universal gmem copy (not TMA)
# so strided last-dim-slice writes are fine. This avoids ~2-3 extra full-tensor
Expand All @@ -1251,9 +1334,12 @@ def _to_cute(t):

# dK split [low | high] postprocess
dk_accum_low, dk_accum_high = _slice_accum(dk_accum)
dk_accum_low = _slice_kv_accum(dk_accum_low, half_hdim)
dk_accum_high = _slice_kv_accum(dk_accum_high, half_hdim)
dk_post = _slice_kv_out(dk)
for accum_part, out_part in (
(dk_accum_low, dk[..., :half_hdim]),
(dk_accum_high, dk[..., half_hdim:]),
(dk_accum_low, dk_post[..., :half_hdim]),
(dk_accum_high, dk_post[..., half_hdim:]),
):
_postprocess_run(
_to_cute(accum_part), _to_cute(out_part), softmax_scale,
Expand All @@ -1263,9 +1349,12 @@ def _to_cute(t):

# dV split [low | high] postprocess
dv_accum_low, dv_accum_high = _slice_accum(dv_accum)
dv_accum_low = _slice_kv_accum(dv_accum_low, half_hdim_v)
dv_accum_high = _slice_kv_accum(dv_accum_high, half_hdim_v)
dv_post = _slice_kv_out(dv)
for accum_part, out_part in (
(dv_accum_low, dv[..., :half_hdim_v]),
(dv_accum_high, dv[..., half_hdim_v:]),
(dv_accum_low, dv_post[..., :half_hdim_v]),
(dv_accum_high, dv_post[..., half_hdim_v:]),
):
_postprocess_run(
_to_cute(accum_part), _to_cute(out_part), cutlass.Float32(1.0),
Expand All @@ -1279,28 +1368,28 @@ def _slice_accum(t):
n = t.shape[-1] // 2
return t[..., :n], t[..., n:]

def _to_cute(t):
return from_dlpack(t.detach(), assumed_align=16).mark_layout_dynamic(
leading_dim=t.ndim - 1
)

_postprocess_run(
dq_accum_tensor, dq_tensor, softmax_scale,
head_dim, m_block_size, AtomLayoutMdQ, dQ_swapAB,
use_2cta_instrs, 1, cu_seqlens_q_tensor, seqused_q_tensor, "dq",
)

_postprocess_run(
dk_accum_tensor, dk_tensor, softmax_scale,
_kv_accum_cute(dk_accum, dk_accum_tensor, head_dim_rounded),
_kv_out_cute(dk, dk_tensor),
softmax_scale,
head_dim, n_block_size, AtomLayoutNdKV, dKV_swapAB,
False, cluster_size, cu_seqlens_k_tensor, seqused_k_tensor, "dk",
)

# dV split [low | high] postprocess
dv_accum_low, dv_accum_high = _slice_accum(dv_accum)
dv_accum_low = _slice_kv_accum(dv_accum_low, half_hdim_v)
dv_accum_high = _slice_kv_accum(dv_accum_high, half_hdim_v)
dv_post = _slice_kv_out(dv)
for accum_part, out_part in (
(dv_accum_low, dv[..., :half_hdim_v]),
(dv_accum_high, dv[..., half_hdim_v:]),
(dv_accum_low, dv_post[..., :half_hdim_v]),
(dv_accum_high, dv_post[..., half_hdim_v:]),
):
_postprocess_run(
_to_cute(accum_part), _to_cute(out_part), cutlass.Float32(1.0),
Expand All @@ -1317,12 +1406,16 @@ def _to_cute(t):

if qhead_per_kvhead > 1:
_postprocess_run(
dk_accum_tensor, dk_tensor, softmax_scale,
_kv_accum_cute(dk_accum, dk_accum_tensor, head_dim_rounded),
_kv_out_cute(dk, dk_tensor),
softmax_scale,
head_dim, n_block_size, AtomLayoutNdKV, dKV_swapAB,
False, cluster_size, cu_seqlens_k_tensor, seqused_k_tensor, "dk",
)
_postprocess_run(
dv_accum_tensor, dv_tensor, cutlass.Float32(1.0),
_kv_accum_cute(dv_accum, dv_accum_tensor, head_dim_v_rounded),
_kv_out_cute(dv, dv_tensor),
cutlass.Float32(1.0),
head_dim_v, n_block_size, AtomLayoutNdKV, dKV_swapAB,
False, cluster_size, cu_seqlens_k_tensor, seqused_k_tensor, "dv",
)
Expand Down