Skip to content
Open
62 changes: 61 additions & 1 deletion flash_attn/cute/flash_bwd.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,11 +11,12 @@
import cutlass
import cutlass.cute as cute
from cutlass.cute.nvgpu import cpasync, warp
from cutlass import Int32
from cutlass import Int32, Int64
import cutlass.utils as utils_basic

from quack import layout_utils
from flash_attn.cute import ampere_helpers as sm80_utils
from flash_attn.cute import philox
from flash_attn.cute.cute_dsl_utils import assume_tensor_aligned
from flash_attn.cute import utils
from flash_attn.cute.mask import AttentionMask
Expand Down Expand Up @@ -52,6 +53,7 @@ def __init__(
V_in_regs: bool = False,
score_mod: cutlass.Constexpr | None = None,
score_mod_bwd: cutlass.Constexpr | None = None,
dropout_p: float = 0.0,
):
"""Initializes the configuration for a flash attention v2 kernel.

Expand Down Expand Up @@ -100,6 +102,7 @@ def __init__(
self.share_QV_smem = V_in_regs
self.score_mod = score_mod
self.score_mod_bwd = score_mod_bwd
self.dropout_p = dropout_p

@staticmethod
def can_implement(
Expand Down Expand Up @@ -394,6 +397,8 @@ def __call__(
mdV_semaphore: Optional[cute.Tensor] = None,
aux_data: AuxData = AuxData(),
blocksparse_tensors: Optional[BlockSparseTensors] = None,
dropout_seed: Optional[Int64] = None,
dropout_offset: Optional[Int64] = None,
# Always keep stream as the last parameter (EnvStream: obtained implicitly via TVM FFI).
stream: cuda.CUstream = None,
):
Expand Down Expand Up @@ -480,6 +485,8 @@ def __call__(
tile_sched_params,
TileScheduler,
aux_data,
dropout_seed,
dropout_offset,
).launch(
grid=grid_dim,
block=[self.num_threads, 1, 1],
Expand Down Expand Up @@ -527,10 +534,16 @@ def kernel(
tile_sched_params: ParamsBase,
TileScheduler: cutlass.Constexpr[Callable],
aux_data: AuxData = AuxData(),
dropout_seed: Optional[Int64] = None,
dropout_offset: Optional[Int64] = None,
):
# Thread index, block index
tidx, _, _ = cute.arch.thread_idx()

# Stash dropout RNG state for compute_one_m_block (runtime SSA values).
self._dropout_seed = dropout_seed
self._dropout_offset = dropout_offset

tile_scheduler = TileScheduler.create(tile_sched_params)
work_tile = tile_scheduler.initial_work_tile_info()

Expand Down Expand Up @@ -847,6 +860,10 @@ def kernel(
batch_idx=batch_idx, head_idx=head_idx,
mask_seqlen=True, mask_causal=self.is_causal, mask_local=self.is_local
)
dropout_fn = partial(
self.apply_dropout, n_block=n_block, batch_idx=batch_idx,
head_idx=head_idx, seqlen_info=seqlen, thr_mma_SdP=thr_mma_sdp,
)
smem_pipe_read_q = cutlass.Int32(0)
smem_pipe_read_do = cutlass.Int32(0)
smem_pipe_write_q = cutlass.Int32(self.num_stages_Q - 1)
Expand All @@ -855,6 +872,7 @@ def kernel(
compute_one_m_block(
m_tile, smem_pipe_read_q, smem_pipe_read_do, smem_pipe_write_q, smem_pipe_write_do,
mask_fn=mask_fn,
dropout_fn=dropout_fn,
)
smem_pipe_read_q = self.advance_pipeline(smem_pipe_read_q, self.num_stages_Q)
smem_pipe_read_do = self.advance_pipeline(smem_pipe_read_do, self.num_stages_dO)
Expand Down Expand Up @@ -894,6 +912,7 @@ def compute_one_m_block(
softmax_scale_log2: cutlass.Float32,
aux_data: AuxData = AuxData(),
mask_fn: Optional[Callable] = None,
dropout_fn: Optional[Callable] = None,
):
def load_Q_next():
m_block_next = m_block + (self.num_stages_Q - 1 if cutlass.const_expr(self.num_stages_Q > 1) else 1)
Expand Down Expand Up @@ -951,6 +970,12 @@ def load_dO_next():
assert cute.size(acc_S_mn, mode=[0]) == cute.size(tLSErLSE)
for r in cutlass.range(cute.size(acc_S_mn, mode=[0]), unroll_full=True):
acc_S_mn[r, None].store(cute.math.exp2(acc_S_mn[r, None].load() * softmax_scale_log2 - tLSErLSE[r], fastmath=True))
# Dropout: regenerate the IDENTICAL keep-mask used in the forward pass and
# apply it to the recomputed P (zero dropped positions, scale kept by 1/(1-p)).
# Zeroing P here automatically zeroes the dV and dS contributions of dropped
# entries, matching the forward where those entries did not contribute to O.
if cutlass.const_expr(self.dropout_p > 0.0):
dropout_fn(acc_S, m_block=m_block)
# if cute.arch.thread_idx()[0] == 0 and cute.arch.block_idx()[0] == bidx: cute.print_tensor(acc_S_mn)

# MMA dP
Expand Down Expand Up @@ -1065,6 +1090,41 @@ def dQ_mma(hook_fn):
cute.arch.barrier()
dQ_mma(load_Q_next)

@cute.jit
def apply_dropout(
self,
acc_S: cute.Tensor,
thr_mma_SdP: cute.ThrMma,
batch_idx,
head_idx,
m_block,
n_block,
seqlen_info: SeqlenInfoQK,
):
"""Regenerate the forward dropout keep-mask on the recomputed P and apply it.

Must produce the IDENTICAL mask as the forward pass: keyed by global
(batch, head, q_idx, kv_idx). The bwd S tile is (n, m) transposed when
SdP_swapAB, so the identity tensor coords are swapped accordingly and we
pass q_idx/kv_idx in the correct order to the shared applicator.
"""
acc_shape = (self.m_block_size, self.n_block_size)
cS = cute.make_identity_tensor(acc_shape if not self.SdP_swapAB else acc_shape[::-1])
offset = (m_block * self.m_block_size, n_block * self.n_block_size)
cS = cute.domain_offset(offset if not self.SdP_swapAB else offset[::-1], cS)
tScS = thr_mma_SdP.partition_C(cS)
p_keep = cutlass.Float32(1.0 - self.dropout_p)
scale = cutlass.Float32(1.0 / (1.0 - self.dropout_p))
philox.apply_dropout(
acc_S, tScS,
self._dropout_seed, self._dropout_offset,
p_keep, scale,
batch_idx, head_idx,
num_heads=1,
seqlen_k=seqlen_info.seqlen_k,
transpose_indices=self.SdP_swapAB,
)

@cute.jit
def epilogue(
self,
Expand Down
61 changes: 60 additions & 1 deletion flash_attn/cute/flash_fwd.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@

import cutlass
import cutlass.cute as cute
from cutlass import Float32, Int32, const_expr
from cutlass import Float32, Int32, Int64, const_expr
from cutlass.cute.nvgpu import cpasync, warp
import cutlass.utils as utils_basic
from cutlass.base_dsl.arch import Arch
Expand All @@ -31,6 +31,7 @@
from flash_attn.cute.seqlen_info import SeqlenInfoQK
from flash_attn.cute.block_info import BlockInfo
from flash_attn.cute.pack_gqa import PackGQA, pack_gqa_layout
from flash_attn.cute import philox
from flash_attn.cute.named_barrier import NamedBarrierFwd
from flash_attn.cute.block_sparsity import BlockSparseTensors
from flash_attn.cute.tile_scheduler import SingleTileScheduler, SingleTileVarlenScheduler, TileSchedulerArguments
Expand All @@ -57,6 +58,7 @@ def __init__(
mask_mod: Optional[cutlass.Constexpr] = None,
has_aux_tensors: bool = False,
q_subtile_factor: int = 1,
dropout_p: float = 0.0,
):
"""Initializes the configuration for a flash attention kernel.

Expand Down Expand Up @@ -99,6 +101,7 @@ def __init__(
self.Q_in_regs = Q_in_regs
self.score_mod = score_mod
self.mask_mod = mask_mod
self.dropout_p = dropout_p
self.qk_acc_dtype = Float32
self.score_vec_size: cutlass.Constexpr = getattr(
score_mod, "__vec_size__", 1 if cutlass.const_expr(has_aux_tensors) else 2
Expand Down Expand Up @@ -638,6 +641,8 @@ def __call__(
learnable_sink: Optional[cute.Tensor] = None,
blocksparse_tensors: Optional[BlockSparseTensors] = None,
aux_data: AuxData = AuxData(),
dropout_seed: Optional[Int64] = None,
dropout_offset: Optional[Int64] = None,
# Always keep stream as the last parameter (EnvStream: obtained implicitly via TVM FFI).
stream: cuda.CUstream = None,
):
Expand All @@ -655,6 +660,10 @@ def __call__(
self.num_producer_threads = self.num_threads
self.num_Q_load_threads = self.num_threads
self.num_epilogue_threads = self.num_threads
# SM120 (Blackwell GeForce / RTX PRO / DGX Spark) lacks the WGMMA-era TMA-store
# epilogue path used here; only sm_90..sm_119 use TMA for the O store. On sm_120
# tma_atom_O is None, so using it crashes in tma_partition with
# 'NoneType' object has no attribute '_trait'. (issue #2649)
self.use_tma_O = Arch.sm_90 <= self.arch < Arch.sm_120
self._setup_attributes()
SharedStorage = self._get_shared_storage_cls()
Expand Down Expand Up @@ -740,6 +749,8 @@ def __call__(
TileScheduler,
aux_data,
fastdiv_mods,
dropout_seed,
dropout_offset,
).launch(
grid=grid_dim,
block=[self.num_threads, 1, 1],
Expand Down Expand Up @@ -779,10 +790,16 @@ def kernel(
TileScheduler: cutlass.Constexpr[Callable],
aux_data: AuxData = AuxData(),
fastdiv_mods=None,
dropout_seed: Optional[Int64] = None,
dropout_offset: Optional[Int64] = None,
):
# Thread index, block index
tidx, _, _ = cute.arch.thread_idx()

# Stash dropout RNG state for compute_one_n_block (runtime SSA values).
self._dropout_seed = dropout_seed
self._dropout_offset = dropout_offset

tile_scheduler = TileScheduler.create(tile_sched_params)
work_tile = tile_scheduler.initial_work_tile_info()
m_block, num_head, batch_size, _ = work_tile.tile_idx
Expand Down Expand Up @@ -1180,6 +1197,16 @@ def load_K_next():
mask_fn(acc_S, n_block=n_block)
row_scale = softmax.online_softmax(acc_S, is_first=is_first_n_block, check_inf=check_inf)
softmax.rescale_O(mma_params.acc_O, row_scale)
if const_expr(self.dropout_p > 0.0):
self.apply_dropout(
mma_params.thr_mma_qk,
batch_idx,
head_idx,
m_block,
acc_S,
n_block,
seqlen=seqlen,
)
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)
Expand Down Expand Up @@ -1234,6 +1261,38 @@ def apply_score_mod(
qhead_per_kvhead=self.qhead_per_kvhead if const_expr(self.pack_gqa) else 1,
)

@cute.jit
def apply_dropout(
self,
thr_mma_qk,
batch_idx,
head_idx,
m_block,
acc_S,
n_block,
seqlen,
):
# Build the same global (q_idx, kv_idx) identity coords used by mask/score_mod,
# so the forward and the backward recompute draw the IDENTICAL keep-mask.
cS = cute.make_identity_tensor((self.tile_m, self.tile_n))
cS = cute.domain_offset((m_block * self.tile_m, n_block * self.tile_n), cS)
tScS = thr_mma_qk.partition_C(cS)
p_keep = Float32(1.0 - self.dropout_p)
scale = Float32(1.0 / (1.0 - self.dropout_p))
philox.apply_dropout(
acc_S,
tScS,
self._dropout_seed,
self._dropout_offset,
p_keep,
scale,
batch_idx,
head_idx,
num_heads=1,
seqlen_k=seqlen.seqlen_k,
transpose_indices=False,
)


# SM90 forward pass moved to flash_fwd_sm90.py; re-export for backward compatibility
def __getattr__(name):
Expand Down
Loading