Skip to content

Instantly share code, notes, and snippets.

@speedcell4
Created July 11, 2026 17:11
Show Gist options
  • Select an option

  • Save speedcell4/8755169dffcb5dfc57ef2918d712a939 to your computer and use it in GitHub Desktop.

Select an option

Save speedcell4/8755169dffcb5dfc57ef2918d712a939 to your computer and use it in GitHub Desktop.
@torch.library.custom_op("qwen3_demo::flash_attn", mutates_args=())
def _flash_attn(q: Tensor, k: Tensor, v: Tensor, causal: bool, softmax_scale: float) -> Tensor:
try:
from flash_attn.flash_attn_interface import flash_attn_func
except ImportError as exc:
raise RuntimeError("flash-attn is required for attention_backend=flash") from exc
scale = None if softmax_scale <= 0 else softmax_scale
return flash_attn_func(q, k, v, dropout_p=0.0, softmax_scale=scale, causal=causal)
@_flash_attn.register_fake
def _flash_attn_fake(q: Tensor, k: Tensor, v: Tensor, causal: bool, softmax_scale: float) -> Tensor:
return torch.empty_like(q)
def _flash_attn_setup(ctx, inputs, output) -> None:
q, k, v, causal, softmax_scale = inputs
ctx.save_for_backward(q, k, v)
ctx.causal = causal
ctx.softmax_scale = softmax_scale
def _flash_attn_backward(ctx, grad_output: Tensor):
q, k, v = ctx.saved_tensors
with torch.enable_grad():
q_input = q.detach().requires_grad_(True)
k_input = k.detach().requires_grad_(True)
v_input = v.detach().requires_grad_(True)
scale = None if ctx.softmax_scale <= 0 else ctx.softmax_scale
output = sdpa_attention(q_input, k_input, v_input, causal=ctx.causal, scale=scale)
gradients = torch.autograd.grad(
output,
(q_input, k_input, v_input),
grad_output.contiguous(),
)
return gradients[0], gradients[1], gradients[2], None, None
_flash_attn.register_autograd(_flash_attn_backward, setup_context=_flash_attn_setup)
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment