Last active
December 22, 2025 22:19
-
-
Save S1ro1/2b1e1e139a2e0b8cf4abfc65a94ed874 to your computer and use it in GitHub Desktop.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| # 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