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
10 changes: 10 additions & 0 deletions flash_attn/cute/cute_dsl_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,16 @@ def assume_tensor_aligned(t):

def to_cute_tensor(t, assumed_align=16, leading_dim=-1, fully_dynamic=False, enable_tvm_ffi=True):
"""Convert torch tensor to cute tensor for TVM FFI. leading_dim=-1 defaults to t.ndim-1."""
if hasattr(t, "_cute_tensor"):
tensor = t._cute_tensor
if fully_dynamic:
marked = tensor.mark_layout_dynamic()
return tensor if marked is None else marked
if leading_dim == -1:
leading_dim = t.ndim - 1
marked = tensor.mark_layout_dynamic(leading_dim=leading_dim)
return tensor if marked is None else marked

# NOTE: torch 2.9.1 doesn't support fp8 via DLPack but 2.11.0 nightly does
# currently export raw bytes as uint8 and tell cutlass correct type
# can directly export as fp8 when torch supports it
Expand Down
151 changes: 148 additions & 3 deletions flash_attn/cute/interface.py
Original file line number Diff line number Diff line change
Expand Up @@ -255,6 +255,71 @@ def _validate_tensor(t, name, expected_shape, expected_dtype, expected_device):
}


def _contiguous_stride(shape):
stride = []
running = 1
for dim in reversed(shape):
stride.append(running)
running *= dim
return tuple(reversed(stride))


class _CompileOnlyTensorSpec:
def __init__(self, shape, dtype, assumed_align=16, stride=None):
self.shape = tuple(shape)
self.dtype = dtype
self.device = torch.device("cuda")
self.requires_grad = False
self.is_cuda = True
self._stride = (
tuple(
cute.sym_int64(divisibility=1) if item is None else item
for item in stride
)
if stride is not None
else _contiguous_stride(self.shape)
)
self._cute_tensor = cute.runtime.make_fake_tensor(
_compile_only_cute_dtype(dtype),
tuple(cute.sym_int() for _ in self.shape),
stride=self._stride,
assumed_align=assumed_align,
)

@property
def ndim(self):
return len(self.shape)

def stride(self, dim=None):
if dim is None:
return self._stride
return self._stride[dim]

def element_size(self):
return torch.empty((), dtype=self.dtype).element_size()

def _compile_only_cute_dtype(dtype):
if dtype == torch.int32:
return cutlass.Int32
return torch2cute_dtype_map[dtype]


def _make_compile_only_tensor_spec(
shape,
dtype,
assumed_align=16,
stride=None,
):
if shape is None:
return None
return _CompileOnlyTensorSpec(
shape,
dtype,
assumed_align=assumed_align,
stride=stride,
)


def num_splits_heuristic(total_mblocks, num_SMs, num_n_blocks, max_splits):
# If num_n_blocks is too small, use 1 split. For example, we never split for hdim = 128 and seqlen_k = 512.
if num_n_blocks <= 4:
Expand Down Expand Up @@ -331,6 +396,7 @@ def _flash_attn_fwd(
v_descale: Optional[torch.Tensor] = None,
gather_kv_indices: Optional[torch.Tensor] = None,
output_scale: Optional[torch.Tensor] = None,
compile_only: bool = False,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Forward pass for FlashAttention.

Expand All @@ -348,6 +414,8 @@ def _flash_attn_fwd(
output_scale: 0-d FP32 GPU tensor. Presence opts into the static per-tensor
FP8 (e4m3fn) fused-quant output: the kernel writes FP8 with
dequant = out_fp8 * output_scale. SM100/SM110 only.
compile_only: If True, compile the selected kernel and return without
launching it.
"""
q, k, v = [maybe_contiguous(t) for t in (q, k, v)]
q_descale, k_descale, v_descale = [maybe_contiguous(t) for t in (q_descale, k_descale, v_descale)]
Expand Down Expand Up @@ -587,8 +655,19 @@ def _flash_attn_fwd(

is_split_kv = num_splits > 1
if is_split_kv:
out_partial = torch.empty(num_splits, *q_batch_seqlen_shape, num_head, head_dim_v, dtype=torch.float32, device=device)
lse_partial = torch.empty(num_splits, *lse_shape, dtype=torch.float32, device=device)
if isinstance(q, _CompileOnlyTensorSpec):
out_partial = _make_compile_only_tensor_spec(
(num_splits, *q_batch_seqlen_shape, num_head, head_dim_v),
torch.float32,
)
lse_partial = _make_compile_only_tensor_spec(
(num_splits, *lse_shape),
torch.float32,
assumed_align=4,
)
else:
out_partial = torch.empty(num_splits, *q_batch_seqlen_shape, num_head, head_dim_v, dtype=torch.float32, device=device)
lse_partial = torch.empty(num_splits, *lse_shape, dtype=torch.float32, device=device)

use_2cta_instrs = (
arch // 10 in [10, 11]
Expand Down Expand Up @@ -751,7 +830,6 @@ def _flash_attn_fwd(
fa_logging.get_fa_log_level(),
output_quant_key,
)

if compile_key not in _flash_attn_fwd.compile_cache:
(
cu_seqlens_q_tensor,
Expand Down Expand Up @@ -1058,6 +1136,9 @@ def _flash_attn_fwd(
*compile_args, options="--enable-tvm-ffi"
)

if compile_only:
return out, lse

if not is_fake_mode():
q_call, k_call, v_call = q.detach(), k.detach(), v.detach()
qv_call = qv.detach() if qv is not None else None
Expand Down Expand Up @@ -2347,6 +2428,70 @@ def flash_attn_varlen_func(
)


def compile_flash_attn_varlen_func_from_specs(
*,
q_shape: Tuple[int, ...],
k_shape: Tuple[int, ...],
v_shape: Tuple[int, ...],
q_dtype: torch.dtype,
v_stride: Optional[Tuple[int, ...]] = None,
cu_seqlens_q_shape: Optional[Tuple[int, ...]] = None,
cu_seqlens_k_shape: Optional[Tuple[int, ...]] = None,
max_seqlen_q: Optional[int] = None,
max_seqlen_k: Optional[int] = None,
softmax_scale: Optional[float] = None,
causal: bool = False,
window_size: Tuple[Optional[int], Optional[int]] = (None, None),
num_splits: int = 1,
return_lse: bool = False,
):
q = _make_compile_only_tensor_spec(q_shape, q_dtype)
k = _make_compile_only_tensor_spec(k_shape, q_dtype)
v = _make_compile_only_tensor_spec(v_shape, q_dtype, stride=v_stride)
out = _make_compile_only_tensor_spec(
(*q_shape[:-1], v_shape[-1]),
q_dtype,
)
lse = None
if return_lse:
assert q is not None
if cu_seqlens_q_shape is None:
lse_shape = (*q.shape[:-3], q.shape[-2], q.shape[-3])
lse_stride = None
else:
lse_shape = (q.shape[-2], q.shape[0])
lse_stride = (None, 1)
lse = _make_compile_only_tensor_spec(
lse_shape,
torch.float32,
4,
stride=lse_stride,
)

return _flash_attn_fwd(
q=q,
k=k,
v=v,
cu_seqlens_q=_make_compile_only_tensor_spec(
cu_seqlens_q_shape, torch.int32, 4
),
cu_seqlens_k=_make_compile_only_tensor_spec(
cu_seqlens_k_shape, torch.int32, 4
),
max_seqlen_q=max_seqlen_q,
max_seqlen_k=max_seqlen_k,
softmax_scale=softmax_scale,
causal=causal,
window_size_left=window_size[0],
window_size_right=window_size[1],
num_splits=num_splits,
return_lse=return_lse,
out=out,
lse=lse,
compile_only=True,
)


def _compile_fwd_combine(
dtype, dtype_partial, head_dim, tile_m, k_block_size, log_max_splits,
has_cu_seqlens, has_seqused, has_lse, has_varlen_batch_idx, output_quant_key,
Expand Down