Skip to content
Merged
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
26 changes: 19 additions & 7 deletions flash_attn/cute/flash_bwd_preprocess.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@ def __init__(
head_dim_v: int,
tile_m: int = 128,
num_threads: int = 256,
use_padded_offsets: bool = True,
):
"""
All contiguous dimensions must be at least 16 bytes aligned which indicates the head dimension
Expand All @@ -64,6 +65,7 @@ def __init__(
self.head_dim_v_padded = int(math.ceil(head_dim_v / hdim_multiple_of) * hdim_multiple_of)
self.check_hdim_v_oob = head_dim_v != self.head_dim_v_padded
self.num_threads = num_threads
self.use_padded_offsets = use_padded_offsets

@staticmethod
def can_implement(dtype, head_dim, tile_m, num_threads) -> bool:
Expand Down Expand Up @@ -250,9 +252,15 @@ def kernel(
)
mO_cur = seqlen.offset_batch(mO, batch_idx, dim=0)[None, head_idx, None]
mdO_cur = seqlen.offset_batch(mdO, batch_idx, dim=0)[None, head_idx, None]
mPdPsum_cur = seqlen.offset_batch(mPdPsum, batch_idx, dim=2, padded=True)[
None, head_idx
]
# Stats buffers (dpsum/lse_log2) are always consumed with padded q-offsets
# on the generic backward path (mdQaccum is present). Keep dedicated hd256
# behavior controlled by self.use_padded_offsets.
stats_use_padded_offsets = self.use_padded_offsets
if const_expr(mdQaccum is not None):
stats_use_padded_offsets = True
mPdPsum_cur = seqlen.offset_batch(
mPdPsum, batch_idx, dim=2, padded=stats_use_padded_offsets
)[None, head_idx]
headdim_v = mO_cur.shape[cute.rank(mO_cur) - 1]
seqlen_q = seqlen.seqlen
seqlen_q_rounded = cute.round_up(seqlen_q, self.tile_m)
Expand Down Expand Up @@ -330,7 +338,11 @@ def kernel(
# Clear dQaccum
if const_expr(mdQaccum is not None):
mdQaccum_cur = seqlen.offset_batch(
mdQaccum, batch_idx, dim=2, padded=True, multiple=self.head_dim_padded
mdQaccum,
batch_idx,
dim=2,
padded=True,
multiple=self.head_dim_padded,
)[None, head_idx]
blkdQaccum_shape = (self.tile_m * self.head_dim_padded,)
gdQaccum = cute.local_tile(mdQaccum_cur, blkdQaccum_shape, (m_block,))
Expand All @@ -341,9 +353,9 @@ def kernel(
cute.copy(gmem_tiled_copy_dQaccum, zero, tdQgdQaccum)

if const_expr(mLSE is not None):
mLSElog2_cur = seqlen.offset_batch(mLSElog2, batch_idx, dim=2, padded=True)[
None, head_idx
]
mLSElog2_cur = seqlen.offset_batch(
mLSElog2, batch_idx, dim=2, padded=stats_use_padded_offsets
)[None, head_idx]
gLSElog2 = cute.local_tile(mLSElog2_cur, (self.tile_m,), (m_block,))
LOG2_E = math.log2(math.e)
if tidx < seqlen_q_rounded - m_block * self.tile_m:
Expand Down
224 changes: 155 additions & 69 deletions flash_attn/cute/interface.py

Large diffs are not rendered by default.

Loading
Loading