diff --git a/flash_attn/cute/cute_dsl_utils.py b/flash_attn/cute/cute_dsl_utils.py index 6dfad6606ef..f97858ab1a1 100644 --- a/flash_attn/cute/cute_dsl_utils.py +++ b/flash_attn/cute/cute_dsl_utils.py @@ -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 diff --git a/flash_attn/cute/interface.py b/flash_attn/cute/interface.py index 22ae840d7a2..e73c3709496 100644 --- a/flash_attn/cute/interface.py +++ b/flash_attn/cute/interface.py @@ -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: @@ -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. @@ -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)] @@ -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] @@ -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, @@ -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 @@ -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,