Skip to content

Instantly share code, notes, and snippets.

@S1ro1
Last active December 22, 2025 22:19
Show Gist options
  • Select an option

  • Save S1ro1/2b1e1e139a2e0b8cf4abfc65a94ed874 to your computer and use it in GitHub Desktop.

Select an option

Save S1ro1/2b1e1e139a2e0b8cf4abfc65a94ed874 to your computer and use it in GitHub Desktop.
# Install fa4 with `uv add git+https://github.com/Dao-AILab/flash-attention.git@main#subdirectory=flash_attn/cute`
import torch
from torch.library import Library
from flash_attn.cute.interface import _flash_attn_fwd, _flash_attn_bwd
# I guess the garbage collector removes the library object so it unregisters the implementation if not kept alive?
_lib = None
def register_fa4():
global _lib
if _lib is not None:
return
_lib = Library("aten", "IMPL", "CUDA")
_lib.impl("_flash_attention_forward", _fa4_forward_impl, "CUDA")
_lib.impl("_flash_attention_backward", _fa4_backward_impl, "CUDA")
_lib.impl(
"_scaled_dot_product_flash_attention",
_fa4_sdpa_forward_impl,
"CUDA",
)
_lib.impl(
"_scaled_dot_product_flash_attention_backward",
_fa4_sdpa_backward_impl,
"CUDA",
)
def _run_forward(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
cu_seq_q: torch.Tensor | None,
cu_seq_k: torch.Tensor | None,
scale: float | None,
is_causal: bool,
window_size_left: int | None,
window_size_right: int | None,
seqused_k: torch.Tensor | None,
out: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
kwargs = {
"softmax_scale": scale,
"causal": is_causal,
"window_size_left": window_size_left,
"window_size_right": window_size_right,
"return_lse": True,
"cu_seqlens_q": cu_seq_q,
"cu_seqlens_k": cu_seq_k,
"seqused_k": seqused_k.contiguous() if seqused_k is not None else None,
}
if out is not None:
kwargs["out"] = out
out, lse = _flash_attn_fwd(query, key, value, **kwargs)
return out, lse.contiguous()
def _run_backward(
grad_out: torch.Tensor,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
out: torch.Tensor,
logsumexp: torch.Tensor,
cu_seq_q: torch.Tensor | None,
cu_seq_k: torch.Tensor | None,
scale: float | None,
is_causal: bool,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
dq, dk, dv = _flash_attn_bwd(
query,
key,
value,
out,
grad_out,
logsumexp.contiguous(),
softmax_scale=scale,
causal=is_causal,
cu_seqlens_q=cu_seq_q,
cu_seqlens_k=cu_seq_k,
)
return dq, dk, dv
def _fa4_forward_impl(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
cum_seq_q: torch.Tensor | None,
cum_seq_k: torch.Tensor | None,
max_q: int,
max_k: int,
dropout_p: float,
is_causal: bool,
return_debug_mask: bool,
*,
scale: float | None = None,
window_size_left: int | None = None,
window_size_right: int | None = None,
seqused_k: torch.Tensor | None = None,
alibi_slopes: torch.Tensor | None = None,
out: torch.Tensor | None = None,
):
out, lse = _run_forward(
query,
key,
value,
cum_seq_q,
cum_seq_k,
scale,
is_causal,
window_size_left,
window_size_right,
seqused_k,
out,
)
rng_state = torch.zeros((2,), dtype=torch.uint64, device=query.device)
philox_offset = torch.zeros((), dtype=torch.uint64, device=query.device)
debug_mask = torch.empty(0, dtype=query.dtype, device=query.device)
return out, lse, rng_state, philox_offset, debug_mask
def _fa4_backward_impl(
grad_out: torch.Tensor,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
out: torch.Tensor,
logsumexp: torch.Tensor,
cum_seq_q: torch.Tensor | None,
cum_seq_k: torch.Tensor | None,
max_q: int,
max_k: int,
dropout_p: float,
is_causal: bool,
rng_state: torch.Tensor,
unused: torch.Tensor,
*,
scale: float | None = None,
window_size_left: int | None = None,
window_size_right: int | None = None,
):
dq, dk, dv = _run_backward(
grad_out,
query,
key,
value,
out,
logsumexp,
cum_seq_q,
cum_seq_k,
scale,
is_causal,
)
return dq, dk, dv
def _fa4_sdpa_forward_impl(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
dropout_p: float = 0.0,
is_causal: bool = False,
return_debug_mask: bool = False,
*,
scale: float | None = None,
):
q, k, v= (
query.transpose(1, 2),
key.transpose(1, 2),
value.transpose(1, 2),
)
out_bhsd = torch.empty_like(query)
out_bshd = out_bhsd.transpose(1, 2)
_, lse, rng_state, philox_offset, debug_mask = _fa4_forward_impl(
q,
k,
v,
None,
None,
query.size(1),
key.size(1),
dropout_p,
is_causal,
return_debug_mask,
scale=scale,
out=out_bshd,
)
return (
out_bhsd,
lse,
None,
None,
query.size(2),
key.size(2),
rng_state,
philox_offset,
debug_mask,
)
def _fa4_sdpa_backward_impl(
grad_out: torch.Tensor,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
out: torch.Tensor,
logsumexp: torch.Tensor,
cum_seq_q: torch.Tensor | None,
cum_seq_k: torch.Tensor | None,
max_q: int,
max_k: int,
dropout_p: float,
is_causal: bool,
philox_seed: torch.Tensor,
philox_offset: torch.Tensor,
*,
scale: float | None = None,
):
query, key, value = (
query.transpose(1, 2),
key.transpose(1, 2),
value.transpose(1, 2),
)
out, grad_out = out.transpose(1, 2), grad_out.transpose(1, 2)
dq, dk, dv = _fa4_backward_impl(
grad_out,
query,
key,
value,
out,
logsumexp,
None,
None,
max_q,
max_k,
dropout_p,
is_causal,
philox_seed,
philox_offset,
scale=scale,
)
dq, dk, dv = dq.transpose(1, 2), dk.transpose(1, 2), dv.transpose(1, 2)
return dq, dk, dv
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment