feat: import CUDA kernels from xllm/CCCL/FLA upstream repos

Sources cloned and tree'd (no --depth):
  - jd-opensource/xllm: ILU kernels, CUDA kernels, MoE kernels
  - NVIDIA/cccl: CUB tuning/dispatch headers (block-level primitives)
  - fla-org/flash-linear-attention: Triton GDN kernels
  - NVIDIA/cutlass: grouped GEMM reference (read, not copied)
  - Dao-AILab/flash-attention: attention kernel reference (SM80+, read only)

New CUDA kernels (from xllm, SM-agnostic, portable to BI-V100):
  ex_engine/xllm_kernels/cuda/activation.cu    (188 lines) — silu_and_mul, gelu
  ex_engine/xllm_kernels/cuda/norm.cu          (600 lines) — rms_norm, fused_add_rms_norm
  ex_engine/xllm_kernels/cuda/rope.cu          (258 lines) — rotary_embedding
  ex_engine/xllm_kernels/cuda/block_copy.cu    (209 lines) — copy_blocks, swap_blocks
  ex_engine/xllm_kernels/cuda/reshape_paged_cache.cu (101 lines) — KV cache ops
  ex_engine/xllm_kernels/cuda/headers/         (5 headers for compilation)

ILU bridge kernel sources (from xllm, verified SAME as upstream):
  ex_engine/xllm_kernels/ilu/    (10 files, 925 lines total)
  — activation.cpp, attention.cpp, fused_moe.cpp, group_gemm.cpp,
    matmul.cpp, norm.cpp, rope.cpp, ilu_ops_api.h, ixformer.h, utils.h

FLA Triton GDN kernels (for GatedDeltaNet without SM90+ FlashQLA):
  ex_engine/fla_kernels/gated_delta_rule/  (7 files, 2370 lines)
  — chunk_fwd.py (428), chunk.py (487), wy_fast.py (409),
    fused_recurrent.py (392), naive.py (161), gate.py (380)

CCCL sync (12 tuning + 14 dispatch headers updated from NVIDIA/cccl):
  cccl_upstream/cub/cub/device/dispatch/tuning/ — 12 changed files synced
  cccl_upstream/cub/cub/device/dispatch/ — 14 changed dispatch files synced

Compilation targets for real machine (ivcore10):
  1. CUDA kernels: --cuda-gpu-arch=ivcore10 via corex clang/16
  2. ILU bridges: torch.utils.cpp_extension linking ixformer .so
  3. FLA kernels: Triton JIT (if Triton works on BI-V100)
This commit is contained in:
claude
2026-08-14 07:48:52 +00:00
parent 051b02d3cd
commit 8d75652949
66 changed files with 12834 additions and 5458 deletions

View File

@@ -0,0 +1,17 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
from .chunk import chunk_gated_delta_rule, chunk_gdn
from .fused_recurrent import fused_recurrent_gated_delta_rule, fused_recurrent_gdn
from .naive import naive_chunk_gated_delta_rule, naive_recurrent_gated_delta_rule
__all__ = [
"chunk_gated_delta_rule", "chunk_gdn",
"fused_recurrent_gated_delta_rule", "fused_recurrent_gdn",
"naive_chunk_gated_delta_rule",
"naive_recurrent_gated_delta_rule",
]

View File

@@ -0,0 +1,591 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import warnings
import torch
from fla.modules.l2norm import l2norm_bwd, l2norm_fwd
from fla.ops.backends import dispatch
from fla.ops.common.chunk_delta_h import chunk_gated_delta_rule_bwd_dhu, chunk_gated_delta_rule_fwd_h
from fla.ops.common.chunk_o import chunk_bwd_dqkwg, chunk_bwd_dv_local, chunk_fwd_o
from fla.ops.common.gate import fused_beta_sigmoid, fused_beta_sigmoid_bwd
from fla.ops.cp import FLACPContext
from fla.ops.cp.chunk_delta_h import (
chunk_gated_delta_rule_bwd_dhu_pre_process,
chunk_gated_delta_rule_fwd_h_pre_process,
compress_h0,
expand_h0,
)
from fla.ops.gated_delta_rule.chunk_fwd import chunk_gated_delta_rule_fwd_intra
from fla.ops.gated_delta_rule.gate import gdn_gate_bwd, gdn_gate_chunk_cumsum
from fla.ops.gated_delta_rule.wy_fast import prepare_wy_repr_bwd, recompute_w_u_fwd
from fla.ops.utils import chunk_local_cumsum
from fla.ops.utils.constant import RCP_LN2
from fla.ops.utils.index import prepare_chunk_indices
from fla.utils import autocast_custom_bwd, autocast_custom_fwd, input_guard
def chunk_gated_delta_rule_fwd(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
scale: float,
initial_state: torch.Tensor,
output_final_state: bool,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
cp_context: FLACPContext | None = None,
chunk_indices: torch.LongTensor | None = None,
use_gate_in_kernel: bool = False,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
chunk_size: int = 64,
):
g_input = g if use_gate_in_kernel else None
if use_gate_in_kernel:
g = gdn_gate_chunk_cumsum(
g=g,
A_log=A_log,
chunk_size=chunk_size,
scale=RCP_LN2,
dt_bias=dt_bias,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
)
else:
g = chunk_local_cumsum(
g,
chunk_size=chunk_size,
scale=RCP_LN2,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
)
# obtain WY representation. u is actually the new v.
# fused kkt + solve_tril + recompute_w_u
w, u, A = chunk_gated_delta_rule_fwd_intra(
k=k,
v=v,
g=g,
beta=beta,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
chunk_size=chunk_size,
)
if cp_context is not None:
initial_state = chunk_gated_delta_rule_fwd_h_pre_process(
k=k,
w=w,
u=u,
g=g,
cu_seqlens=cu_seqlens,
initial_state=initial_state,
context=cp_context,
state_v_first=state_v_first,
chunk_size=chunk_size,
)
h, v_new, final_state = chunk_gated_delta_rule_fwd_h(
k=k,
w=w,
u=u,
g=g,
initial_state=initial_state,
output_final_state=output_final_state,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
state_v_first=state_v_first,
chunk_size=chunk_size,
)
if cp_context is not None:
initial_state = compress_h0(initial_state, context=cp_context)
o = chunk_fwd_o(
q=q,
k=k,
v=v_new,
h=h,
g=g,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
state_v_first=state_v_first,
chunk_size=chunk_size,
)
return g, o, A, final_state, initial_state, g_input
def chunk_gated_delta_rule_bwd(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
A: torch.Tensor,
scale: float,
initial_state: torch.Tensor,
do: torch.Tensor,
dht: torch.Tensor,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
cp_context: FLACPContext | None = None,
chunk_indices: torch.LongTensor | None = None,
use_gate_in_kernel: bool = False,
g_input: torch.Tensor | None = None,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
chunk_size: int = 64,
):
w, u = recompute_w_u_fwd(
k=k,
v=v,
beta=beta,
A=A,
g=g,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
)
if cp_context is not None:
initial_state = expand_h0(initial_state, context=cp_context)
h, v_new, _ = chunk_gated_delta_rule_fwd_h(
k=k,
w=w,
u=u,
g=g,
initial_state=initial_state,
output_final_state=False,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
state_v_first=state_v_first,
chunk_size=chunk_size,
)
dv = chunk_bwd_dv_local(
q=q,
k=k,
g=g,
do=do,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
chunk_size=chunk_size,
)
if cp_context is not None:
# initial_state is None in the CP mode
# We only need to compute dht of current rank and pass it to the backward kernel
dht, initial_state = chunk_gated_delta_rule_bwd_dhu_pre_process(
q=q,
k=k,
w=w,
do=do,
dv=dv,
g=g,
scale=scale,
cu_seqlens=cu_seqlens,
dht=dht,
initial_state=initial_state,
context=cp_context,
state_v_first=state_v_first,
chunk_size=chunk_size,
)
dh, dh0, dv = chunk_gated_delta_rule_bwd_dhu(
q=q,
k=k,
w=w,
g=g,
h0=initial_state,
dht=dht,
do=do,
dv=dv,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
state_v_first=state_v_first,
chunk_size=chunk_size,
)
dq, dk, dw, dg = chunk_bwd_dqkwg(
q=q,
k=k,
v=v_new,
w=w,
g=g,
h=h,
dv=dv,
do=do,
dh=dh,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
state_v_first=state_v_first,
chunk_size=chunk_size,
)
dk2, dv, db, dg2 = prepare_wy_repr_bwd(
k=k,
v=v,
beta=beta,
g=g,
A=A,
dw=dw,
du=dv,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
)
dk.add_(dk2)
dg.add_(dg2)
dg = chunk_local_cumsum(dg, chunk_size=chunk_size, reverse=True, cu_seqlens=cu_seqlens, chunk_indices=chunk_indices)
dA_log, ddt_bias = None, None
if use_gate_in_kernel:
dg, dA_log, ddt_bias = gdn_gate_bwd(g=g_input, A_log=A_log, dt_bias=dt_bias, dyg=dg)
return dq, dk, dv, db, dg, dh0, dA_log, ddt_bias
class ChunkGatedDeltaRuleFunction(torch.autograd.Function):
@staticmethod
@input_guard
@autocast_custom_fwd
def forward(
ctx,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
scale: float,
initial_state: torch.Tensor,
output_final_state: bool,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
cu_seqlens_cpu: torch.LongTensor | None = None,
use_qk_l2norm_in_kernel: bool = False,
use_gate_in_kernel: bool = False,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
use_beta_sigmoid_in_kernel: bool = False,
allow_neg_eigval: bool = False,
cp_context: FLACPContext | None = None,
chunk_size: int = 64,
):
q_rstd, k_rstd = None, None
if use_qk_l2norm_in_kernel:
q, q_rstd = l2norm_fwd(q)
k, k_rstd = l2norm_fwd(k)
beta_raw = beta
if use_beta_sigmoid_in_kernel:
beta = fused_beta_sigmoid(beta_raw, scale=2.0 if allow_neg_eigval else 1.0)
chunk_indices = None
if cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, chunk_size, cu_seqlens_cpu=cu_seqlens_cpu)
g, o, A, final_state, initial_state, g_input = chunk_gated_delta_rule_fwd(
q=q,
k=k,
v=v,
g=g,
beta=beta,
scale=scale,
initial_state=initial_state,
output_final_state=output_final_state,
cu_seqlens=cu_seqlens,
cp_context=cp_context,
chunk_indices=chunk_indices,
state_v_first=state_v_first,
use_gate_in_kernel=use_gate_in_kernel,
A_log=A_log,
dt_bias=dt_bias,
chunk_size=chunk_size,
)
ctx.save_for_backward(
q,
q_rstd,
k,
k_rstd,
v,
g,
beta_raw,
beta,
A,
initial_state,
cu_seqlens,
chunk_indices,
g_input,
A_log,
dt_bias,
)
ctx.scale = scale
ctx.chunk_size = chunk_size
ctx.use_qk_l2norm_in_kernel = use_qk_l2norm_in_kernel
ctx.use_beta_sigmoid_in_kernel = use_beta_sigmoid_in_kernel
ctx.allow_neg_eigval = allow_neg_eigval
ctx.cp_context = cp_context
ctx.state_v_first = state_v_first
ctx.use_gate_in_kernel = use_gate_in_kernel
return o.to(q.dtype), final_state
@staticmethod
@input_guard
@autocast_custom_bwd
def backward(
ctx,
do: torch.Tensor,
dht: torch.Tensor,
):
(
q,
q_rstd,
k,
k_rstd,
v,
g,
beta_raw,
beta,
A,
initial_state,
cu_seqlens,
chunk_indices,
g_input,
A_log,
dt_bias,
) = ctx.saved_tensors
dq, dk, dv, db, dg, dh0, dA_log, ddt_bias = chunk_gated_delta_rule_bwd(
q=q,
k=k,
v=v,
g=g,
beta=beta,
A=A,
scale=ctx.scale,
initial_state=initial_state,
do=do,
dht=dht,
cu_seqlens=cu_seqlens,
cp_context=ctx.cp_context,
chunk_indices=chunk_indices,
state_v_first=ctx.state_v_first,
use_gate_in_kernel=ctx.use_gate_in_kernel,
g_input=g_input,
A_log=A_log,
dt_bias=dt_bias,
chunk_size=ctx.chunk_size,
)
if ctx.use_qk_l2norm_in_kernel:
dq = l2norm_bwd(q, q_rstd, dq)
dk = l2norm_bwd(k, k_rstd, dk)
if ctx.use_beta_sigmoid_in_kernel:
db = fused_beta_sigmoid_bwd(beta_raw, db, scale=2.0 if ctx.allow_neg_eigval else 1.0)
return (
dq.to(q), dk.to(k), dv.to(v), dg.to(g), db.to(beta_raw),
None, dh0, None, None, None, None, None, None, dA_log, ddt_bias,
None, None, None, None,
)
@dispatch('gated_delta_rule')
@torch.compiler.disable
def chunk_gated_delta_rule(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
scale: float | None = None,
initial_state: torch.Tensor | None = None,
output_final_state: bool = False,
use_qk_l2norm_in_kernel: bool = False,
use_beta_sigmoid_in_kernel: bool = False,
allow_neg_eigval: bool = False,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
cu_seqlens_cpu: torch.LongTensor | None = None,
cp_context: FLACPContext | None = None,
**kwargs,
):
r"""
Args:
q (torch.Tensor):
queries of shape `[B, T, H, K]`.
k (torch.Tensor):
keys of shape `[B, T, H, K]`.
v (torch.Tensor):
values of shape `[B, T, HV, V]`.
GVA (Grouped Value Attention) is applied if `HV > H`, where `HV` must be divisible by `H`.
g (torch.Tensor):
(forget) gating tensor of shape `[B, T, HV]`.
When `use_gate_in_kernel=False` (default), `g` should be in log space (pre-computed decay).
When `use_gate_in_kernel=True`, `g` is the raw input before gate activation;
the kernel fuses `-exp(A_log) * softplus(g + dt_bias)` + chunk cumsum internally.
beta (torch.Tensor):
betas of shape `[B, T, HV]`.
scale (Optional[float]):
Scale factor for the RetNet attention scores.
If not provided, it will default to `1 / sqrt(K)`. Default: `None`.
initial_state (Optional[torch.Tensor]):
Initial state of shape `[N, HV, K, V]` for `N` input sequences.
For equal-length input sequences, `N` equals the batch size `B`.
Default: `None`.
output_final_state (Optional[bool]):
Whether to output the final state of shape `[N, HV, K, V]`. Default: `False`.
use_qk_l2norm_in_kernel (bool):
Whether to apply L2norm to the q/k tensor internally. Default: `False`.
use_gate_in_kernel (bool):
Whether to compute the log-space GDN decay internally.
When `True`, the passed `g` is the raw input, and `A_log` must be provided.
The kernel fuses gate activation + chunk cumsum in a single pass.
Default: `False`.
A_log (Optional[torch.Tensor]):
Decay parameter of shape `[HV]`. Required when `use_gate_in_kernel=True`.
dt_bias (Optional[torch.Tensor]):
Bias added to `g` before activation, of shape `[HV]`.
Only used when `use_gate_in_kernel=True`.
use_beta_sigmoid_in_kernel (bool):
Whether to apply `torch.sigmoid(beta)` before launching the chunk kernel.
- If `True`, the passed `beta` acts as the raw beta logits.
- If `False`, `beta` is expected to already be in post-sigmoid space.
Default: `False`.
allow_neg_eigval (bool):
Whether to allow negative eigenvalues by scaling `beta` to `[0, 2)`.
Only takes effect together with `use_beta_sigmoid_in_kernel=True`, in which case
the kernel computes `2 * sigmoid(beta)` instead of `sigmoid(beta)`. Default: `False`.
state_v_first (Optional[bool]):
Store the recurrent state in V-first ``[V, K]`` layout instead of the default ``[K, V]``. Default: ``False``.
cu_seqlens (torch.LongTensor):
Cumulative sequence lengths of shape `[N+1]` used for variable-length training,
consistent with the FlashAttention API.
cp_context (Optional[FLACPContext]):
Context parallel context for distributed training across multiple devices.
When provided, `initial_state` and `output_final_state` are not supported,
and `cu_seqlens` will be overridden by the context. Default: `None`.
Returns:
o (torch.Tensor):
Outputs of shape `[B, T, HV, V]`.
final_state (torch.Tensor):
Final state of shape `[N, HV, K, V]` if `output_final_state=True` else `None`.
Examples::
>>> import torch
>>> import torch.nn.functional as F
>>> from einops import rearrange
>>> from fla.ops.gated_delta_rule import chunk_gated_delta_rule
# inputs with equal lengths
>>> B, T, H, HV, K, V = 4, 2048, 4, 8, 512, 512
>>> q = torch.randn(B, T, H, K, dtype=torch.bfloat16, device='cuda')
>>> k = F.normalize(torch.randn(B, T, H, K, dtype=torch.bfloat16, device='cuda'), p=2, dim=-1)
>>> v = torch.randn(B, T, HV, V, dtype=torch.bfloat16, device='cuda')
>>> beta = torch.rand(B, T, HV, dtype=torch.bfloat16, device='cuda').sigmoid()
>>> g = F.logsigmoid(torch.rand(B, T, HV, dtype=torch.bfloat16, device='cuda'))
>>> h0 = torch.randn(B, HV, K, V, dtype=torch.bfloat16, device='cuda')
>>> o, ht = chunk_gated_delta_rule(
q, k, v, g, beta,
initial_state=h0,
output_final_state=True
)
# for variable-length inputs, the batch size `B` is expected to be 1 and `cu_seqlens` is required
>>> q, k, v, beta, g = map(lambda x: rearrange(x, 'b t ... -> 1 (b t) ...'), (q, k, v, beta, g))
# for a batch with 4 sequences, `cu_seqlens` with 5 start/end positions are expected
>>> cu_seqlens = q.new_tensor([0, 2048, 4096, 6144, 8192], dtype=torch.long)
>>> o, ht = chunk_gated_delta_rule(
q, k, v, g, beta,
initial_state=h0,
output_final_state=True,
cu_seqlens=cu_seqlens
)
"""
if 'transpose_state_layout' in kwargs:
if state_v_first:
raise ValueError("Cannot pass both `state_v_first` and the deprecated `transpose_state_layout`.")
warnings.warn(
"`transpose_state_layout` is deprecated and renamed to `state_v_first`.",
DeprecationWarning,
stacklevel=2,
)
state_v_first = kwargs.pop('transpose_state_layout')
# Validate head dimensions
if q.shape[2] != k.shape[2]:
raise ValueError(
f"q and k must have the same number of heads, "
f"but got q.shape[2]={q.shape[2]} and k.shape[2]={k.shape[2]}"
)
H, HV = q.shape[2], v.shape[2]
if HV % H != 0:
raise ValueError(
f"For GVA, num_v_heads (HV={HV}) must be evenly divisible by "
f"num_heads (H={H}), but got HV % H = {HV % H}"
)
if 'head_first' in kwargs:
raise DeprecationWarning(
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
)
chunk_size = kwargs.pop('chunk_size', 64)
if chunk_size not in (16, 32, 64):
raise ValueError(f"`chunk_size` must be 16, 32, or 64 for Gated Delta Rule, got {chunk_size}.")
if cp_context is not None:
assert initial_state is None, "Initial state is not supported for CP"
assert output_final_state is False, "Output final state is not supported for CP"
assert cp_context.cu_seqlens is not None, "cu_seqlens is required for CP"
cu_seqlens = cp_context.cu_seqlens
if cp_context.cu_seqlens_cpu is not None:
cu_seqlens_cpu = cp_context.cu_seqlens_cpu
if cu_seqlens is not None:
if q.shape[0] != 1:
raise ValueError(
f"The batch size is expected to be 1 rather than {q.shape[0]} when using `cu_seqlens`."
f"Please flatten variable-length inputs before processing.",
)
if initial_state is not None and initial_state.shape[0] != len(cu_seqlens) - 1:
raise ValueError(
f"The number of initial states is expected to be equal to the number of input sequences, "
f"i.e., {len(cu_seqlens) - 1} rather than {initial_state.shape[0]}.",
)
use_gate_in_kernel = kwargs.get('use_gate_in_kernel', False)
A_log = kwargs.get('A_log')
dt_bias = kwargs.get('dt_bias')
if use_gate_in_kernel:
assert A_log is not None, "A_log must be provided when use_gate_in_kernel=True."
if allow_neg_eigval and not use_beta_sigmoid_in_kernel:
raise ValueError("`allow_neg_eigval=True` requires `use_beta_sigmoid_in_kernel=True`.")
if scale is None:
scale = k.shape[-1] ** -0.5
o, final_state = ChunkGatedDeltaRuleFunction.apply(
q,
k,
v,
g,
beta,
scale,
initial_state,
output_final_state,
state_v_first,
cu_seqlens,
cu_seqlens_cpu,
use_qk_l2norm_in_kernel,
use_gate_in_kernel,
A_log,
dt_bias,
use_beta_sigmoid_in_kernel,
allow_neg_eigval,
cp_context,
chunk_size,
)
return o, final_state
chunk_gdn = chunk_gated_delta_rule

View File

@@ -0,0 +1,428 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import torch
import triton
import triton.language as tl
from fla.ops.backends import dispatch
from fla.ops.common.chunk_scaled_dot_kkt import chunk_scaled_dot_kkt_fwd
from fla.ops.gated_delta_rule.wy_fast import recompute_w_u_fwd
from fla.ops.utils import prepare_chunk_indices, solve_tril
from fla.ops.utils.cache import fla_cache_autotune
from fla.ops.utils.op import exp2
from fla.utils import IS_INTEL, IS_TF32_SUPPORTED, autotune_cache_kwargs
if IS_TF32_SUPPORTED:
SOLVE_TRIL_DOT_PRECISION = tl.constexpr('tf32')
else:
SOLVE_TRIL_DOT_PRECISION = tl.constexpr('ieee')
@triton.heuristics({
'USE_G': lambda args: args['g'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({'BK': BK}, num_warps=num_warps)
for BK in [32, 64]
for num_warps in [1, 2, 4]
],
key=['H', 'HV', 'K', 'BC'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_gated_delta_rule_fwd_kkt_solve_kernel(
k,
g,
beta,
A,
cu_seqlens,
chunk_indices,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
BT: tl.constexpr,
BC: tl.constexpr,
BK: tl.constexpr,
USE_G: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
"""
Fused kernel: compute beta * K @ K^T (lower triangular) + solve_tril (I+A)^{-1} in one pass.
This kernel fuses chunk_scaled_dot_kkt_fwd and solve_tril into a single kernel,
avoiding the HBM round-trip for the intermediate A matrix.
Steps:
1. Compute all 10 lower-triangular [BC, BC] blocks of beta * K @ K^T in registers
2. Apply gate and beta scaling
3. Forward substitution on diagonal blocks
4. Block merge to get full (I+A)^{-1}
5. Write result to A (output)
"""
i_t, i_bh = tl.program_id(0).to(tl.int64), tl.program_id(1).to(tl.int64)
i_b, i_h = i_bh // HV, i_bh % HV
if IS_VARLEN:
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
else:
bos, eos = i_b * T, i_b * T + T
if i_t * BT >= T:
return
i_tc0 = i_t * BT
i_tc1 = i_t * BT + BC
i_tc2 = i_t * BT + 2 * BC
i_tc3 = i_t * BT + 3 * BC
k += (bos * H + i_h // (HV // H)) * K
A += (bos * HV + i_h) * BT
o_i = tl.arange(0, BC)
m_tc0 = (i_tc0 + o_i) < T
m_tc1 = (i_tc1 + o_i) < T
m_tc2 = (i_tc2 + o_i) < T
m_tc3 = (i_tc3 + o_i) < T
# load beta for each sub-chunk
p_b0 = beta + bos * HV + i_h + (i_tc0 + o_i) * HV
p_b1 = beta + bos * HV + i_h + (i_tc1 + o_i) * HV
p_b2 = beta + bos * HV + i_h + (i_tc2 + o_i) * HV
p_b3 = beta + bos * HV + i_h + (i_tc3 + o_i) * HV
b_b0 = tl.load(p_b0, mask=m_tc0, other=0.0).to(tl.float32)
b_b1 = tl.load(p_b1, mask=m_tc1, other=0.0).to(tl.float32)
b_b2 = tl.load(p_b2, mask=m_tc2, other=0.0).to(tl.float32)
b_b3 = tl.load(p_b3, mask=m_tc3, other=0.0).to(tl.float32)
# load gate if used
if USE_G:
p_g0 = g + bos * HV + i_h + (i_tc0 + o_i) * HV
p_g1 = g + bos * HV + i_h + (i_tc1 + o_i) * HV
p_g2 = g + bos * HV + i_h + (i_tc2 + o_i) * HV
p_g3 = g + bos * HV + i_h + (i_tc3 + o_i) * HV
b_g0 = tl.load(p_g0, mask=m_tc0, other=0.0).to(tl.float32)
b_g1 = tl.load(p_g1, mask=m_tc1, other=0.0).to(tl.float32)
b_g2 = tl.load(p_g2, mask=m_tc2, other=0.0).to(tl.float32)
b_g3 = tl.load(p_g3, mask=m_tc3, other=0.0).to(tl.float32)
############################################################################
# Step 1: compute all 10 lower-triangular [BC, BC] blocks of K @ K^T
############################################################################
# 4 diagonal blocks
b_A00 = tl.zeros([BC, BC], dtype=tl.float32)
b_A11 = tl.zeros([BC, BC], dtype=tl.float32)
b_A22 = tl.zeros([BC, BC], dtype=tl.float32)
b_A33 = tl.zeros([BC, BC], dtype=tl.float32)
# 6 off-diagonal blocks
b_A10 = tl.zeros([BC, BC], dtype=tl.float32)
b_A20 = tl.zeros([BC, BC], dtype=tl.float32)
b_A21 = tl.zeros([BC, BC], dtype=tl.float32)
b_A30 = tl.zeros([BC, BC], dtype=tl.float32)
b_A31 = tl.zeros([BC, BC], dtype=tl.float32)
b_A32 = tl.zeros([BC, BC], dtype=tl.float32)
for i_k in range(tl.cdiv(K, BK)):
o_k = i_k * BK + tl.arange(0, BK)
p_k0 = k + (i_tc0 + o_i)[:, None] * (H*K) + o_k[None, :]
b_k0 = tl.load(p_k0, mask=m_tc0[:, None] & (o_k[None, :] < K), other=0.0)
# diagonal block 0
b_A00 += tl.dot(b_k0, tl.trans(b_k0))
if i_tc1 < T:
p_k1 = k + (i_tc1 + o_i)[:, None] * (H*K) + o_k[None, :]
b_k1 = tl.load(p_k1, mask=m_tc1[:, None] & (o_k[None, :] < K), other=0.0)
# diagonal block 1
b_A11 += tl.dot(b_k1, tl.trans(b_k1))
# off-diagonal (1,0)
b_A10 += tl.dot(b_k1, tl.trans(b_k0))
if i_tc2 < T:
p_k2 = k + (i_tc2 + o_i)[:, None] * (H*K) + o_k[None, :]
b_k2 = tl.load(p_k2, mask=m_tc2[:, None] & (o_k[None, :] < K), other=0.0)
# diagonal block 2
b_A22 += tl.dot(b_k2, tl.trans(b_k2))
# off-diagonal (2,0), (2,1)
b_A20 += tl.dot(b_k2, tl.trans(b_k0))
b_A21 += tl.dot(b_k2, tl.trans(b_k1))
if i_tc3 < T:
p_k3 = k + (i_tc3 + o_i)[:, None] * (H*K) + o_k[None, :]
b_k3 = tl.load(p_k3, mask=m_tc3[:, None] & (o_k[None, :] < K), other=0.0)
# diagonal block 3
b_A33 += tl.dot(b_k3, tl.trans(b_k3))
# off-diagonal (3,0), (3,1), (3,2)
b_A30 += tl.dot(b_k3, tl.trans(b_k0))
b_A31 += tl.dot(b_k3, tl.trans(b_k1))
b_A32 += tl.dot(b_k3, tl.trans(b_k2))
############################################################################
# Step 2: apply gate and beta scaling
############################################################################
# apply gate, beta scaling, and masking
# m_d: strictly lower triangular mask for diagonal blocks
# m_tc: boundary mask to prevent NaN from 0 * inf (IEEE 754) when
# out-of-bounds g loads as 0 via boundary_check and exp2(0 - g_inbounds) overflows
m_d = o_i[:, None] > o_i[None, :]
m_I = o_i[:, None] == o_i[None, :]
if USE_G:
b_A00 *= tl.where(m_d & m_tc0[:, None] & m_tc0[None, :], exp2(b_g0[:, None] - b_g0[None, :]), 0.)
b_A11 *= tl.where(m_d & m_tc1[:, None] & m_tc1[None, :], exp2(b_g1[:, None] - b_g1[None, :]), 0.)
b_A22 *= tl.where(m_d & m_tc2[:, None] & m_tc2[None, :], exp2(b_g2[:, None] - b_g2[None, :]), 0.)
b_A33 *= tl.where(m_d & m_tc3[:, None] & m_tc3[None, :], exp2(b_g3[:, None] - b_g3[None, :]), 0.)
b_A10 *= tl.where(m_tc1[:, None] & m_tc0[None, :], exp2(b_g1[:, None] - b_g0[None, :]), 0.)
b_A20 *= tl.where(m_tc2[:, None] & m_tc0[None, :], exp2(b_g2[:, None] - b_g0[None, :]), 0.)
b_A21 *= tl.where(m_tc2[:, None] & m_tc1[None, :], exp2(b_g2[:, None] - b_g1[None, :]), 0.)
b_A30 *= tl.where(m_tc3[:, None] & m_tc0[None, :], exp2(b_g3[:, None] - b_g0[None, :]), 0.)
b_A31 *= tl.where(m_tc3[:, None] & m_tc1[None, :], exp2(b_g3[:, None] - b_g1[None, :]), 0.)
b_A32 *= tl.where(m_tc3[:, None] & m_tc2[None, :], exp2(b_g3[:, None] - b_g2[None, :]), 0.)
else:
b_A00 = tl.where(m_d, b_A00, 0.)
b_A11 = tl.where(m_d, b_A11, 0.)
b_A22 = tl.where(m_d, b_A22, 0.)
b_A33 = tl.where(m_d, b_A33, 0.)
# diagonal blocks: scaled by beta
b_A00 = b_A00 * b_b0[:, None]
b_A11 = b_A11 * b_b1[:, None]
b_A22 = b_A22 * b_b2[:, None]
b_A33 = b_A33 * b_b3[:, None]
# off-diagonal blocks: full block, scaled by beta
b_A10 = b_A10 * b_b1[:, None]
b_A20 = b_A20 * b_b2[:, None]
b_A21 = b_A21 * b_b2[:, None]
b_A30 = b_A30 * b_b3[:, None]
b_A31 = b_A31 * b_b3[:, None]
b_A32 = b_A32 * b_b3[:, None]
############################################################################
# Step 3: forward substitution on diagonal blocks -> (I + A_diag)^{-1}
#
# Same algorithm as solve_tril, but rows are extracted from in-register
# [BC, BC] tensor via tl.sum(tl.where(mask, tensor, 0), 0) instead of
# tl.load from HBM.
############################################################################
b_Ai00 = -b_A00
b_Ai11 = -b_A11
b_Ai22 = -b_A22
b_Ai33 = -b_A33
for i in range(2, min(BC, T - i_tc0)):
b_a00 = tl.sum(tl.where((o_i == i)[:, None], -b_A00, 0.), 0)
b_a00 = tl.where(o_i < i, b_a00, 0.)
b_a00 = b_a00 + tl.sum(b_a00[:, None] * b_Ai00, 0)
b_Ai00 = tl.where((o_i == i)[:, None], b_a00, b_Ai00)
for i in range(2, min(BC, T - i_tc1)):
b_a11 = tl.sum(tl.where((o_i == i)[:, None], -b_A11, 0.), 0)
b_a11 = tl.where(o_i < i, b_a11, 0.)
b_a11 = b_a11 + tl.sum(b_a11[:, None] * b_Ai11, 0)
b_Ai11 = tl.where((o_i == i)[:, None], b_a11, b_Ai11)
for i in range(2, min(BC, T - i_tc2)):
b_a22 = tl.sum(tl.where((o_i == i)[:, None], -b_A22, 0.), 0)
b_a22 = tl.where(o_i < i, b_a22, 0.)
b_a22 = b_a22 + tl.sum(b_a22[:, None] * b_Ai22, 0)
b_Ai22 = tl.where((o_i == i)[:, None], b_a22, b_Ai22)
for i in range(2, min(BC, T - i_tc3)):
b_a33 = tl.sum(tl.where((o_i == i)[:, None], -b_A33, 0.), 0)
b_a33 = tl.where(o_i < i, b_a33, 0.)
b_a33 = b_a33 + tl.sum(b_a33[:, None] * b_Ai33, 0)
b_Ai33 = tl.where((o_i == i)[:, None], b_a33, b_Ai33)
b_Ai00 += m_I
b_Ai11 += m_I
b_Ai22 += m_I
b_Ai33 += m_I
############################################################################
# Step 4: block merge -> full (I + A)^{-1}
############################################################################
b_Ai10 = -tl.dot(
tl.dot(b_Ai11, b_A10, input_precision=SOLVE_TRIL_DOT_PRECISION),
b_Ai00,
input_precision=SOLVE_TRIL_DOT_PRECISION
)
b_Ai21 = -tl.dot(
tl.dot(b_Ai22, b_A21, input_precision=SOLVE_TRIL_DOT_PRECISION),
b_Ai11,
input_precision=SOLVE_TRIL_DOT_PRECISION
)
b_Ai32 = -tl.dot(
tl.dot(b_Ai33, b_A32, input_precision=SOLVE_TRIL_DOT_PRECISION),
b_Ai22,
input_precision=SOLVE_TRIL_DOT_PRECISION
)
b_Ai20 = -tl.dot(
b_Ai22,
tl.dot(b_A20, b_Ai00, input_precision=SOLVE_TRIL_DOT_PRECISION) +
tl.dot(b_A21, b_Ai10, input_precision=SOLVE_TRIL_DOT_PRECISION),
input_precision=SOLVE_TRIL_DOT_PRECISION,
)
b_Ai31 = -tl.dot(
b_Ai33,
tl.dot(b_A31, b_Ai11, input_precision=SOLVE_TRIL_DOT_PRECISION) +
tl.dot(b_A32, b_Ai21, input_precision=SOLVE_TRIL_DOT_PRECISION),
input_precision=SOLVE_TRIL_DOT_PRECISION,
)
b_Ai30 = -tl.dot(
b_Ai33,
tl.dot(b_A30, b_Ai00, input_precision=SOLVE_TRIL_DOT_PRECISION) +
tl.dot(b_A31, b_Ai10, input_precision=SOLVE_TRIL_DOT_PRECISION) +
tl.dot(b_A32, b_Ai20, input_precision=SOLVE_TRIL_DOT_PRECISION),
input_precision=SOLVE_TRIL_DOT_PRECISION,
)
############################################################################
# Step 5: store full (I + A)^{-1} to output A
############################################################################
p_A00 = A + (i_tc0 + o_i)[:, None] * (HV*BT) + o_i[None, :]
p_A10 = A + (i_tc1 + o_i)[:, None] * (HV*BT) + o_i[None, :]
p_A11 = A + (i_tc1 + o_i)[:, None] * (HV*BT) + (BC + o_i)[None, :]
p_A20 = A + (i_tc2 + o_i)[:, None] * (HV*BT) + o_i[None, :]
p_A21 = A + (i_tc2 + o_i)[:, None] * (HV*BT) + (BC + o_i)[None, :]
p_A22 = A + (i_tc2 + o_i)[:, None] * (HV*BT) + (2*BC + o_i)[None, :]
p_A30 = A + (i_tc3 + o_i)[:, None] * (HV*BT) + o_i[None, :]
p_A31 = A + (i_tc3 + o_i)[:, None] * (HV*BT) + (BC + o_i)[None, :]
p_A32 = A + (i_tc3 + o_i)[:, None] * (HV*BT) + (2*BC + o_i)[None, :]
p_A33 = A + (i_tc3 + o_i)[:, None] * (HV*BT) + (3*BC + o_i)[None, :]
m_A0 = m_tc0[:, None] & (o_i[None, :] < BT)
m_A1 = m_tc1[:, None] & (o_i[None, :] < BT)
m_A2 = m_tc2[:, None] & (o_i[None, :] < BT)
m_A3 = m_tc3[:, None] & (o_i[None, :] < BT)
m_A11 = m_tc1[:, None] & ((BC + o_i)[None, :] < BT)
m_A21 = m_tc2[:, None] & ((BC + o_i)[None, :] < BT)
m_A22 = m_tc2[:, None] & ((2*BC + o_i)[None, :] < BT)
m_A31 = m_tc3[:, None] & ((BC + o_i)[None, :] < BT)
m_A32 = m_tc3[:, None] & ((2*BC + o_i)[None, :] < BT)
m_A33 = m_tc3[:, None] & ((3*BC + o_i)[None, :] < BT)
tl.store(p_A00, b_Ai00.to(A.dtype.element_ty), mask=m_A0)
tl.store(p_A10, b_Ai10.to(A.dtype.element_ty), mask=m_A1)
tl.store(p_A11, b_Ai11.to(A.dtype.element_ty), mask=m_A11)
tl.store(p_A20, b_Ai20.to(A.dtype.element_ty), mask=m_A2)
tl.store(p_A21, b_Ai21.to(A.dtype.element_ty), mask=m_A21)
tl.store(p_A22, b_Ai22.to(A.dtype.element_ty), mask=m_A22)
tl.store(p_A30, b_Ai30.to(A.dtype.element_ty), mask=m_A3)
tl.store(p_A31, b_Ai31.to(A.dtype.element_ty), mask=m_A31)
tl.store(p_A32, b_Ai32.to(A.dtype.element_ty), mask=m_A32)
tl.store(p_A33, b_Ai33.to(A.dtype.element_ty), mask=m_A33)
@dispatch('gated_delta_rule')
def chunk_gated_delta_rule_fwd_intra(
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor | None = None,
beta: torch.Tensor | None = None,
cu_seqlens: torch.LongTensor | None = None,
chunk_size: int = 64,
chunk_indices: torch.LongTensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
r"""
GDN intra-chunk forward: fused or unfused kkt + solve_tril + recompute_w_u.
For ``chunk_size == 64``, this uses the fused kkt + solve_tril path. For
other supported chunk sizes, it computes the mathematically equivalent
representation with ``chunk_scaled_dot_kkt_fwd`` followed by ``solve_tril``.
Args:
k (torch.Tensor):
The key tensor of shape `[B, T, H, K]`.
v (torch.Tensor):
The value tensor of shape `[B, T, HV, V]`.
g (torch.Tensor):
The cumulative sum of the gate tensor of shape `[B, T, HV]`. Default: `None`.
beta (torch.Tensor):
The beta tensor of shape `[B, T, HV]`.
cu_seqlens (torch.LongTensor):
The cumulative sequence lengths. Default: `None`.
chunk_size (int):
The chunk size. Default: 64.
chunk_indices (torch.LongTensor):
Precomputed chunk indices. Default: `None`.
Returns:
w (torch.Tensor): shape `[B, T, HV, K]`
u (torch.Tensor): shape `[B, T, HV, V]`
A (torch.Tensor): shape `[B, T, HV, BT]`, the solved (I+A)^{-1} matrix
"""
if chunk_size not in (16, 32, 64):
raise ValueError(f"`chunk_size` must be 16, 32, or 64, got {chunk_size}.")
B, T, H, K, HV = *k.shape, beta.shape[2]
BT = chunk_size
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
# The fused kernel keeps ten [BC, BC] fp32 accumulators live across the K loop.
# That fits NVIDIA's register file but spills on Intel GPUs, where the unfused
# two-kernel path measures 2.3-3.0x faster despite the extra HBM round-trip.
if BT == 64 and not IS_INTEL:
# Step 1: fused kkt + solve_tril
BC = 16
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
A = torch.zeros(B, T, HV, BT, device=k.device, dtype=k.dtype)
chunk_gated_delta_rule_fwd_kkt_solve_kernel[(NT, B * HV)](
k=k,
g=g,
beta=beta,
A=A,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
T=T,
H=H,
HV=HV,
K=K,
BT=BT,
BC=BC,
)
else:
# Step 1: mathematically equivalent unfused kkt + solve_tril
A = chunk_scaled_dot_kkt_fwd(
k=k,
g=g,
beta=beta,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
chunk_size=BT,
output_dtype=torch.float32,
)
A = solve_tril(
A=A,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
output_dtype=k.dtype,
)
# Step 2: recompute_w_u
w, u = recompute_w_u_fwd(
k=k,
v=v,
beta=beta,
A=A,
g=g,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
)
return w, u, A

View File

@@ -0,0 +1,478 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import warnings
import torch
import triton
import triton.language as tl
from fla.ops.utils.op import exp
from fla.ops.utils.softplus import softplus
from fla.utils import input_guard
@triton.heuristics({
'USE_G': lambda args: args['g'] is not None,
'USE_GK': lambda args: args['gk'] is not None,
'USE_GV': lambda args: args['gv'] is not None,
'USE_INITIAL_STATE': lambda args: args['h0'] is not None,
'STORE_FINAL_STATE': lambda args: args['ht'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
'USE_GATE_IN_KERNEL': lambda args: args['A_log'] is not None,
'HAS_DT_BIAS': lambda args: args['dt_bias'] is not None,
})
@triton.jit(do_not_specialize=['T'])
def fused_recurrent_gated_delta_rule_fwd_kernel(
q,
k,
v,
g,
gk,
gv,
beta,
A_log,
dt_bias,
o,
h0,
ht,
cu_seqlens,
scale,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BK: tl.constexpr,
BV: tl.constexpr,
USE_G: tl.constexpr,
USE_GK: tl.constexpr,
USE_GV: tl.constexpr,
USE_QK_L2NORM_IN_KERNEL: tl.constexpr,
IS_BETA_HEADWISE: tl.constexpr,
USE_INITIAL_STATE: tl.constexpr,
STORE_FINAL_STATE: tl.constexpr,
STATE_V_FIRST: tl.constexpr,
IS_VARLEN: tl.constexpr,
USE_GATE_IN_KERNEL: tl.constexpr,
HAS_DT_BIAS: tl.constexpr,
APPLY_BETA_SIGMOID: tl.constexpr,
ALLOW_NEG_EIGVAL: tl.constexpr,
):
pid = tl.program_id(0)
NV = tl.cdiv(V, BV)
i_v, i_nh = pid % NV, (pid // NV).to(tl.int64)
i_n, i_hv = i_nh // HV, i_nh % HV
i_h = i_hv // (HV // H)
if IS_VARLEN:
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
else:
bos, eos = i_n * T, i_n * T + T
o_k = tl.arange(0, BK)
o_v = i_v * BV + tl.arange(0, BV)
p_q = q + (bos * H + i_h) * K + o_k
p_k = k + (bos * H + i_h) * K + o_k
p_v = v + (bos * HV + i_hv) * V + o_v
if USE_G:
p_g = g + bos * HV + i_hv
if USE_GK:
p_gk = gk + (bos * HV + i_hv) * K + o_k
if USE_GV:
p_gv = gv + (bos * HV + i_hv) * V + o_v
if IS_BETA_HEADWISE:
p_beta = beta + bos * HV + i_hv
else:
p_beta = beta + (bos * HV + i_hv) * V + o_v
p_o = o + (bos * HV + i_hv) * V + o_v
mask_k = o_k < K
mask_v = o_v < V
if STATE_V_FIRST:
mask_h = mask_v[:, None] & mask_k[None, :]
else:
mask_h = mask_k[:, None] & mask_v[None, :]
if STATE_V_FIRST:
b_h = tl.zeros([BV, BK], dtype=tl.float32)
else:
b_h = tl.zeros([BK, BV], dtype=tl.float32)
if USE_INITIAL_STATE:
if STATE_V_FIRST:
p_h0 = h0 + i_nh * K*V + o_v[:, None] * K + o_k[None, :]
else:
p_h0 = h0 + i_nh * K*V + o_k[:, None] * V + o_v[None, :]
b_h += tl.load(p_h0, mask=mask_h, other=0).to(tl.float32)
for _ in tl.range(0, T):
b_q = tl.load(p_q, mask=mask_k, other=0).to(tl.float32)
b_k = tl.load(p_k, mask=mask_k, other=0).to(tl.float32)
b_v = tl.load(p_v, mask=mask_v, other=0).to(tl.float32)
if USE_QK_L2NORM_IN_KERNEL:
b_q = b_q / tl.sqrt(tl.sum(b_q * b_q) + 1e-6)
b_k = b_k / tl.sqrt(tl.sum(b_k * b_k) + 1e-6)
b_q = b_q * scale
if IS_BETA_HEADWISE:
b_beta = tl.load(p_beta).to(tl.float32)
else:
b_beta = tl.load(p_beta, mask=mask_v, other=0).to(tl.float32)
if APPLY_BETA_SIGMOID:
b_beta = tl.sigmoid(b_beta)
if ALLOW_NEG_EIGVAL:
b_beta = b_beta * 2
if USE_G:
b_g = tl.load(p_g).to(tl.float32)
if USE_GATE_IN_KERNEL:
b_A = tl.load(A_log + i_hv).to(tl.float32)
if HAS_DT_BIAS:
b_g = b_g + tl.load(dt_bias + i_hv).to(tl.float32)
b_g = -exp(b_A) * softplus(b_g)
b_h *= exp(b_g)
if USE_GK:
b_gk = tl.load(p_gk).to(tl.float32)
if STATE_V_FIRST:
b_h *= exp(b_gk[None, :])
else:
b_h *= exp(b_gk[:, None])
if USE_GV:
b_gv = tl.load(p_gv).to(tl.float32)
if STATE_V_FIRST:
b_h *= exp(b_gv[:, None])
else:
b_h *= exp(b_gv[None, :])
if STATE_V_FIRST:
b_v = b_beta * (b_v - tl.sum(b_h * b_k[None, :], 1))
b_h += b_v[:, None] * b_k[None, :]
b_o = tl.sum(b_h * b_q[None, :], 1)
else:
b_v = b_beta * (b_v - tl.sum(b_h * b_k[:, None], 0))
b_h += b_k[:, None] * b_v
b_o = tl.sum(b_h * b_q[:, None], 0)
tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=mask_v)
p_q += H*K
p_k += H*K
p_v += HV*V
if USE_G:
p_g += HV
if USE_GK:
p_gk += HV*K
if USE_GV:
p_gv += HV*V
p_beta += HV * (1 if IS_BETA_HEADWISE else V)
p_o += HV*V
if STORE_FINAL_STATE:
if STATE_V_FIRST:
p_ht = ht + i_nh * K*V + o_v[:, None] * K + o_k[None, :]
else:
p_ht = ht + i_nh * K*V + o_k[:, None] * V + o_v[None, :]
tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), mask=mask_h)
def fused_recurrent_gated_delta_rule_fwd(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor | None = None,
gk: torch.Tensor | None = None,
gv: torch.Tensor | None = None,
beta: torch.Tensor | None = None,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
scale: float = None,
initial_state: torch.Tensor = None,
output_final_state: bool = False,
use_qk_l2norm_in_kernel: bool = False,
use_beta_sigmoid_in_kernel: bool = False,
allow_neg_eigval: bool = False,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
B, T, H, K, V = *k.shape, v.shape[-1]
HV = v.shape[2]
N = B if cu_seqlens is None else len(cu_seqlens) - 1
BK = triton.next_power_of_2(K)
BV = min(8, triton.next_power_of_2(V)) if gv is None else triton.next_power_of_2(V)
NV = triton.cdiv(V, BV)
o = torch.empty_like(v)
if output_final_state:
if state_v_first:
final_state = q.new_empty(N, HV, V, K, dtype=torch.float32)
else:
final_state = q.new_empty(N, HV, K, V, dtype=torch.float32)
else:
final_state = None
grid = (NV * N * HV,)
fused_recurrent_gated_delta_rule_fwd_kernel[grid](
q=q,
k=k,
v=v,
g=g,
gk=gk,
gv=gv,
beta=beta,
A_log=A_log,
dt_bias=dt_bias,
o=o,
h0=initial_state,
ht=final_state,
cu_seqlens=cu_seqlens,
scale=scale,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BK=BK,
BV=BV,
IS_BETA_HEADWISE=beta.ndim != v.ndim,
USE_QK_L2NORM_IN_KERNEL=use_qk_l2norm_in_kernel,
APPLY_BETA_SIGMOID=use_beta_sigmoid_in_kernel,
ALLOW_NEG_EIGVAL=allow_neg_eigval,
STATE_V_FIRST=state_v_first,
num_warps=1,
num_stages=3,
)
return o, final_state
class FusedRecurrentFunction(torch.autograd.Function):
@staticmethod
@input_guard
def forward(
ctx,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor | None = None,
gk: torch.Tensor | None = None,
gv: torch.Tensor | None = None,
beta: torch.Tensor | None = None,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
scale: float = None,
initial_state: torch.Tensor = None,
output_final_state: bool = False,
use_qk_l2norm_in_kernel: bool = False,
use_beta_sigmoid_in_kernel: bool = False,
allow_neg_eigval: bool = False,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
):
o, final_state = fused_recurrent_gated_delta_rule_fwd(
q=q,
k=k,
v=v,
g=g,
gk=gk,
gv=gv,
beta=beta,
A_log=A_log,
dt_bias=dt_bias,
scale=scale,
initial_state=initial_state,
output_final_state=output_final_state,
use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel,
use_beta_sigmoid_in_kernel=use_beta_sigmoid_in_kernel,
allow_neg_eigval=allow_neg_eigval,
state_v_first=state_v_first,
cu_seqlens=cu_seqlens,
)
return o, final_state
@staticmethod
@input_guard
def backward(ctx, do, dht):
raise NotImplementedError(
"Backward pass is not implemented yet and we do not have plans to implement it "
"because we haven't figured out how to compute dg without materializing the full "
"hidden states for all time steps.",
)
def fused_recurrent_gated_delta_rule(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor | None = None,
gk: torch.Tensor | None = None,
gv: torch.Tensor | None = None,
beta: torch.Tensor | None = None,
scale: float = None,
initial_state: torch.Tensor = None,
output_final_state: bool = False,
use_qk_l2norm_in_kernel: bool = False,
use_gate_in_kernel: bool = False,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
use_beta_sigmoid_in_kernel: bool = False,
allow_neg_eigval: bool = False,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
**kwargs,
) -> tuple[torch.Tensor, torch.Tensor]:
r"""
Args:
q (torch.Tensor):
queries of shape `[B, T, H, K]`.
k (torch.Tensor):
keys of shape `[B, T, H, K]`.
v (torch.Tensor):
values of shape `[B, T, HV, V]`.
GVA (Grouped Value Attention) is applied if `HV > H`, where `HV` must be divisible by `H`.
g (torch.Tensor):
g (decays) of shape `[B, T, HV]`. Default: `None`.
When `use_gate_in_kernel=False` (default), `g` must be in log space (pre-computed decay).
When `use_gate_in_kernel=True`, `g` is the raw pre-activation input; the kernel fuses
`-exp(A_log) * softplus(g + dt_bias)` internally per step.
gk (torch.Tensor):
gk (decays) of shape `[B, T, HV, K]`. Default: `None`.
gv (torch.Tensor):
gv (decays) of shape `[B, T, HV, V]`. Default: `None`.
beta (torch.Tensor):
betas of shape `[B, T, HV]`.
scale (Optional[float]):
Scale factor for the RetNet attention scores.
If not provided, it will default to `1 / sqrt(K)`. Default: `None`.
initial_state (Optional[torch.Tensor]):
Initial state of shape `[N, HV, K, V]` for `N` input sequences.
For equal-length input sequences, `N` equals the batch size `B`.
Default: `None`.
output_final_state (Optional[bool]):
Whether to output the final state of shape `[N, HV, K, V]`. Default: `False`.
use_qk_l2norm_in_kernel (Optional[bool]):
Whether to use L2 normalization in the kernel. Default: `False`.
use_gate_in_kernel (bool):
Whether to compute the log-space GDN decay internally.
When `True`, `g` is the raw input and `A_log` must be provided; the kernel fuses
gate activation into the recurrence. Default: `False`.
A_log (Optional[torch.Tensor]):
Decay parameter of shape `[HV]`. Required when `use_gate_in_kernel=True`.
dt_bias (Optional[torch.Tensor]):
Bias added to `g` before activation, of shape `[HV]`.
Only used when `use_gate_in_kernel=True`.
use_beta_sigmoid_in_kernel (Optional[bool]):
Whether to apply `torch.sigmoid(beta)` inside the kernel.
- If `True`, the passed `beta` acts as the raw beta logits.
- If `False`, `beta` is expected to already be in post-sigmoid space.
Default: `False`.
allow_neg_eigval (Optional[bool]):
Whether to allow negative eigenvalues by scaling `beta` to `[0, 2)`.
Only takes effect together with `use_beta_sigmoid_in_kernel=True`, in which case
the kernel computes `2 * sigmoid(beta)` instead of `sigmoid(beta)`. Default: `False`.
state_v_first (Optional[bool]):
Store the recurrent state in V-first ``[V, K]`` layout instead of the default ``[K, V]``. Default: ``False``.
cu_seqlens (torch.LongTensor):
Cumulative sequence lengths of shape `[N+1]` used for variable-length training,
consistent with the FlashAttention API.
Returns:
o (torch.Tensor):
Outputs of shape `[B, T, HV, V]`.
final_state (torch.Tensor):
Final state of shape `[N, HV, K, V]` if `output_final_state=True` else `None`.
Examples::
>>> import torch
>>> import torch.nn.functional as F
>>> from einops import rearrange
>>> from fla.ops.gated_delta_rule import fused_recurrent_gated_delta_rule
# inputs with equal lengths
>>> B, T, H, HV, K, V = 4, 2048, 4, 8, 512, 512
>>> q = torch.randn(B, T, H, K, device='cuda')
>>> k = F.normalize(torch.randn(B, T, H, K, device='cuda'), p=2, dim=-1)
>>> v = torch.randn(B, T, HV, V, device='cuda')
>>> g = F.logsigmoid(torch.rand(B, T, HV, device='cuda'))
>>> beta = torch.rand(B, T, HV, device='cuda').sigmoid()
>>> h0 = torch.randn(B, HV, K, V, device='cuda')
>>> o, ht = fused_gated_recurrent_delta_rule(
q, k, v, g, beta,
initial_state=h0,
output_final_state=True
)
# for variable-length inputs, the batch size `B` is expected to be 1 and `cu_seqlens` is required
>>> q, k, v, g, beta = map(lambda x: rearrange(x, 'b t ... -> 1 (b t) ...'), (q, k, v, g, beta))
# for a batch with 4 sequences, `cu_seqlens` with 5 start/end positions are expected
>>> cu_seqlens = q.new_tensor([0, 2048, 4096, 6144, 8192], dtype=torch.long)
>>> o, ht = fused_gated_recurrent_delta_rule(
q, k, v, g, beta,
initial_state=h0,
output_final_state=True,
cu_seqlens=cu_seqlens
)
"""
if 'transpose_state_layout' in kwargs:
if state_v_first:
raise ValueError("Cannot pass both `state_v_first` and the deprecated `transpose_state_layout`.")
warnings.warn(
"`transpose_state_layout` is deprecated and renamed to `state_v_first`.",
DeprecationWarning,
stacklevel=2,
)
state_v_first = kwargs.pop('transpose_state_layout')
if cu_seqlens is not None:
if q.shape[0] != 1:
raise ValueError(
f"The batch size is expected to be 1 rather than {q.shape[0]} when using `cu_seqlens`."
f"Please flatten variable-length inputs before processing.",
)
if initial_state is not None and initial_state.shape[0] != len(cu_seqlens) - 1:
raise ValueError(
f"The number of initial states is expected to be equal to the number of input sequences, "
f"i.e., {len(cu_seqlens) - 1} rather than {initial_state.shape[0]}.",
)
if scale is None:
scale = k.shape[-1] ** -0.5
if beta is None:
beta = torch.ones_like(q[..., 0])
if use_gate_in_kernel:
if A_log is None:
raise ValueError("`A_log` must be provided when `use_gate_in_kernel=True`.")
if g is None:
raise ValueError("`g` (raw pre-activation) must be provided when `use_gate_in_kernel=True`.")
else:
A_log = None
dt_bias = None
if allow_neg_eigval and not use_beta_sigmoid_in_kernel:
raise ValueError("`allow_neg_eigval=True` requires `use_beta_sigmoid_in_kernel=True`.")
o, final_state = FusedRecurrentFunction.apply(
q,
k,
v,
g,
gk,
gv,
beta,
A_log,
dt_bias,
scale,
initial_state,
output_final_state,
use_qk_l2norm_in_kernel,
use_beta_sigmoid_in_kernel,
allow_neg_eigval,
state_v_first,
cu_seqlens,
)
return o, final_state
fused_recurrent_gdn = fused_recurrent_gated_delta_rule

View File

@@ -0,0 +1,344 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import torch
import torch.nn.functional as F
import triton
import triton.language as tl
from fla.ops.backends import dispatch
from fla.ops.utils.cache import fla_cache_autotune
from fla.ops.utils.index import prepare_chunk_indices
from fla.ops.utils.op import exp
from fla.ops.utils.softplus import softplus
from fla.utils import autocast_custom_bwd, autocast_custom_fwd, autotune_cache_kwargs, input_guard
def naive_gdn_gate(
g: torch.Tensor,
A_log: torch.Tensor,
dt_bias: torch.Tensor | None = None,
output_dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
"""
Torch reference implementation for GDN gate computation.
Computes: ``g = -A_log.exp() * softplus(g + dt_bias)``
Args:
g (torch.Tensor):
Input tensor of shape `[..., HV]`.
A_log (torch.Tensor):
Decay parameter tensor with `HV` elements.
dt_bias (torch.Tensor | None):
Optional bias tensor added to `g` before activation, shape `[HV]`.
Returns:
Output tensor of shape `[..., HV]`.
"""
g = g.float()
if dt_bias is not None:
g = g + dt_bias.float()
return (-A_log.float().exp() * F.softplus(g)).to(output_dtype)
@triton.heuristics({
'HAS_BIAS': lambda args: args['dt_bias'] is not None,
'HAS_SCALE': lambda args: args['scale'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({}, num_warps=num_warps)
for num_warps in [1, 2, 4, 8]
],
key=['H', 'BT', 'IS_VARLEN', 'REVERSE'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def gdn_gate_chunk_cumsum_scalar_kernel(
g,
A_log,
dt_bias,
o,
scale,
cu_seqlens,
chunk_indices,
T,
H: tl.constexpr,
BT: tl.constexpr,
REVERSE: tl.constexpr,
HAS_BIAS: tl.constexpr,
HAS_SCALE: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_t, i_bh = tl.program_id(0).to(tl.int64), tl.program_id(1).to(tl.int64)
i_b, i_h = i_bh // H, i_bh % H
if IS_VARLEN:
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
else:
bos, eos = i_b * T, i_b * T + T
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
p_g = g + bos * H + i_h + o_t * H
p_o = o + bos * H + i_h + o_t * H
b_g = tl.load(p_g, mask=m_t, other=0.0).to(tl.float32)
if HAS_BIAS:
b_g = b_g + tl.load(dt_bias + i_h).to(tl.float32)
b_A = tl.load(A_log + i_h).to(tl.float32)
b_gate = -exp(b_A) * softplus(b_g)
b_o = tl.cumsum(b_gate, axis=0)
if REVERSE:
b_z = tl.sum(b_gate, axis=0)
b_o = -b_o + b_z[None] + b_gate
if HAS_SCALE:
b_o *= scale
tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=m_t)
@triton.heuristics({
'HAS_BIAS': lambda args: args['dt_bias'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({}, num_warps=num_warps)
for num_warps in [1, 2, 4, 8]
],
key=['H', 'BT'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def gdn_gate_bwd_kernel(
g,
A_log,
dt_bias,
dyg,
dg,
dA,
T,
H: tl.constexpr,
BT: tl.constexpr,
HAS_BIAS: tl.constexpr,
):
i_t, i_h = tl.program_id(0).to(tl.int64), tl.program_id(1)
b_A = tl.load(A_log + i_h).to(tl.float32)
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
p_g = g + i_h + o_t * H
p_dg = dg + i_h + o_t * H
p_dyg = dyg + i_h + o_t * H
b_g = tl.load(p_g, mask=m_t, other=0.0).to(tl.float32)
b_dyg = tl.load(p_dyg, mask=m_t, other=0.0).to(tl.float32)
if HAS_BIAS:
b_g = b_g + tl.load(dt_bias + i_h).to(tl.float32)
# gate = -exp(A_log) * softplus(g + bias)
# d(gate)/d(g) = -exp(A_log) * sigmoid(g + bias) (softplus' = sigmoid)
# d(gate)/d(A_log) = -exp(A_log) * softplus(g + bias) = gate
b_neg_expA = -exp(b_A)
b_yg = b_neg_expA * softplus(b_g)
b_dg = b_neg_expA * (b_dyg * tl.sigmoid(b_g))
b_dA = tl.sum(b_dyg * b_yg, 0)
tl.store(p_dg, b_dg.to(p_dg.dtype.element_ty), mask=m_t)
tl.store(dA + i_t * H + i_h, b_dA)
@input_guard
@dispatch('gated_delta_rule')
def gdn_gate_chunk_cumsum(
g: torch.Tensor,
A_log: torch.Tensor,
chunk_size: int,
scale: float = None,
dt_bias: torch.Tensor | None = None,
cu_seqlens: torch.LongTensor | None = None,
chunk_indices: torch.LongTensor | None = None,
output_dtype: torch.dtype | None = torch.float,
) -> torch.Tensor:
B, T, H = g.shape
BT = chunk_size
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
o = torch.empty_like(g, dtype=output_dtype or g.dtype)
gdn_gate_chunk_cumsum_scalar_kernel[(NT, B * H)](
g=g,
A_log=A_log,
dt_bias=dt_bias,
o=o,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
T=T,
H=H,
BT=BT,
REVERSE=False,
)
return o
@dispatch('gated_delta_rule')
def gdn_gate_bwd(
g: torch.Tensor,
A_log: torch.Tensor,
dt_bias: torch.Tensor | None,
dyg: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None]:
H = g.shape[-1]
T = g.numel() // H
BT = 32
NT = triton.cdiv(T, BT)
dg = torch.empty_like(g, dtype=torch.float32)
dA = A_log.new_empty(NT, H, dtype=torch.float32)
gdn_gate_bwd_kernel[(NT, H)](
g=g,
A_log=A_log,
dt_bias=dt_bias,
dyg=dyg,
dg=dg,
dA=dA,
T=T,
H=H,
BT=BT,
)
dg = dg.view_as(g).type_as(g)
dA = dA.sum(0).view_as(A_log).type_as(A_log)
dbias = dg.view(-1, H).sum(0).to(dt_bias) if dt_bias is not None else None
return dg, dA, dbias
@triton.heuristics({
'HAS_BIAS': lambda args: args['dt_bias'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({'BT': BT}, num_warps=num_warps, num_stages=num_stages)
for BT in [32, 64, 128]
for num_warps in [1, 2, 4, 8]
for num_stages in [2, 3]
],
key=['H'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def gdn_gate_fwd_kernel(
g,
A_log,
dt_bias,
yg,
T,
H: tl.constexpr,
BT: tl.constexpr,
HAS_BIAS: tl.constexpr,
):
i_t, i_h = tl.program_id(0).to(tl.int64), tl.program_id(1)
b_A = tl.load(A_log + i_h).to(tl.float32)
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
p_g = g + i_h + o_t * H
p_yg = yg + i_h + o_t * H
b_g = tl.load(p_g, mask=m_t, other=0.0).to(tl.float32)
if HAS_BIAS:
b_g = b_g + tl.load(dt_bias + i_h).to(tl.float32)
b_yg = -exp(b_A) * softplus(b_g)
tl.store(p_yg, b_yg.to(p_yg.dtype.element_ty), mask=m_t)
@dispatch('gated_delta_rule')
def gdn_gate_fwd(
g: torch.Tensor,
A_log: torch.Tensor,
dt_bias: torch.Tensor | None = None,
output_dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
H = g.shape[-1]
T = g.numel() // H
yg = torch.empty_like(g, dtype=output_dtype)
def grid(meta):
return (triton.cdiv(T, meta['BT']), H)
gdn_gate_fwd_kernel[grid](
g=g,
A_log=A_log,
dt_bias=dt_bias,
yg=yg,
T=T,
H=H,
)
return yg
class GDNGateFunction(torch.autograd.Function):
@staticmethod
@input_guard
@autocast_custom_fwd
def forward(
ctx,
g: torch.Tensor,
A_log: torch.Tensor,
dt_bias: torch.Tensor | None = None,
output_dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
yg = gdn_gate_fwd(g=g, A_log=A_log, dt_bias=dt_bias, output_dtype=output_dtype)
ctx.save_for_backward(g, A_log, dt_bias)
return yg
@staticmethod
@input_guard
@autocast_custom_bwd
def backward(ctx, dyg: torch.Tensor):
g, A_log, dt_bias = ctx.saved_tensors
dg, dA, dbias = gdn_gate_bwd(g=g, A_log=A_log, dt_bias=dt_bias, dyg=dyg)
return dg, dA, dbias, None
@torch.compiler.disable
def fused_gdn_gate(
g: torch.Tensor,
A_log: torch.Tensor,
dt_bias: torch.Tensor | None = None,
output_dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
r"""
Fused GDN gate computation with autograd support.
Computes: ``g = -A_log.exp() * softplus(g + dt_bias)``
Args:
g (torch.Tensor):
Input tensor of shape `[..., HV]`.
A_log (torch.Tensor):
Decay parameter tensor with `HV` elements.
dt_bias (torch.Tensor | None):
Optional bias tensor added to `g` before activation, shape `[HV]`.
output_dtype (torch.dtype):
The dtype of the output tensor. Default: `torch.float32`.
Returns:
Output tensor of shape `[..., HV]`.
"""
return GDNGateFunction.apply(g, A_log, dt_bias, output_dtype)

View File

@@ -0,0 +1,161 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import torch
import torch.nn.functional as F
from einops import rearrange
def naive_recurrent_gated_delta_rule(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
beta: torch.Tensor,
g: torch.Tensor,
scale: float = None,
initial_state: torch.Tensor = None,
output_final_state: bool = False,
):
"""
Reference PyTorch implementation of recurrent gated delta rule.
Args:
q: [B, T, H, K]
k: [B, T, H, K]
v: [B, T, H, V]
beta: [B, T, H]
g: [B, T, H]
scale: float, optional
initial_state: [B, H, K, V], optional
output_final_state: bool
Returns:
o: [B, T, H, V]
final_state: [B, H, K, V] if output_final_state else None
"""
q, k, v, beta, g = map(lambda x: x.transpose(1, 2).contiguous().to(torch.float32), [q, k, v, beta, g])
B, H, T, K, V = *k.shape, v.shape[-1]
o = torch.zeros(B, H, T, V).to(v)
h = torch.zeros(B, H, K, V).to(v)
if initial_state is not None:
h = initial_state.to(torch.float32)
if scale is None:
scale = 1 / (q.shape[-1] ** 0.5)
q = q * scale
for i in range(T):
b_q = q[:, :, i]
b_k = k[:, :, i]
b_v = v[:, :, i].clone()
h = h.clone() * g[:, :, i].exp()[..., None, None]
b_beta = beta[:, :, i]
b_v = b_v - (h.clone() * b_k[..., None]).sum(-2)
b_v = b_v * b_beta[..., None]
h = h.clone() + b_k.unsqueeze(-1) * b_v.unsqueeze(-2)
o[:, :, i] = torch.einsum('bhd,bhdm->bhm', b_q, h)
if not output_final_state:
h = None
o = o.transpose(1, 2).contiguous()
return o, h
def naive_chunk_gated_delta_rule(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
chunk_size: int = 64,
scale: float = None,
initial_state: torch.Tensor = None,
output_final_state: bool = False,
):
"""
Reference PyTorch implementation of chunk gated delta rule.
Args:
q: [B, T, H, K]
k: [B, T, H, K]
v: [B, T, H, V]
g: [B, T, H]
beta: [B, T, H]
chunk_size: int
scale: float, optional
initial_state: [B, H, K, V], optional
output_final_state: bool
Returns:
o: [B, T, H, V]
final_state: [B, H, K, V] if output_final_state else None
"""
BT = chunk_size
if scale is None:
scale = 1 / (q.shape[-1] ** 0.5)
q, k, v, beta, g = map(lambda x: x.transpose(1, 2).contiguous().to(torch.float32), [q, k, v, beta, g])
T = q.shape[-2]
pad_len = (BT - (T % BT)) % BT
if pad_len > 0:
q = F.pad(q, (0, 0, 0, pad_len))
k = F.pad(k, (0, 0, 0, pad_len))
v = F.pad(v, (0, 0, 0, pad_len))
beta = F.pad(beta, (0, pad_len))
g = F.pad(g, (0, pad_len))
q, k, v, beta, g = map(lambda x: x.to(torch.float32), [q, k, v, beta, g])
decay = g
chunk_size = BT
b, h, l, d_k = q.shape
d_v = v.shape[-1]
q = q * scale
v = v * beta[..., None]
k_beta = k * beta[..., None]
assert l % chunk_size == 0
# note that diagonal is masked.
mask = torch.triu(torch.ones(chunk_size, chunk_size, dtype=torch.bool, device=q.device), diagonal=0)
q, k, v, k_beta, decay = map(
lambda x: rearrange(x, 'b h (n c) d -> b h n c d', c=chunk_size),
[q, k, v, k_beta, decay.unsqueeze(-1)],
)
decay = decay.squeeze(-1).cumsum(-1)
decay_exp = decay.exp()[..., None]
L_mask = ((decay.unsqueeze(-1) - decay.unsqueeze(-2)).tril().exp().float()).tril()
attn = -((k_beta @ k.transpose(-1, -2)) * L_mask).masked_fill(mask, 0)
for i in range(1, chunk_size):
attn[..., i, :i] = attn[..., i, :i].clone() + (attn[..., i, :i, None].clone() * attn[..., :i, :i].clone()).sum(-2)
attn = attn + torch.eye(chunk_size, dtype=torch.float, device=q.device)
attn = attn
k_cumsum = attn @ v
k_cumdecay = attn @ (k_beta * decay_exp)
v = k_cumsum
S = k.new_zeros(b, h, d_k, d_v)
if initial_state is not None:
S = initial_state.to(torch.float32)
o = torch.zeros_like(v)
mask = torch.triu(torch.ones(chunk_size, chunk_size, dtype=torch.bool, device=q.device), diagonal=1)
for i in range(0, l // chunk_size):
q_i, k_i, v_i = q[:, :, i], k[:, :, i], v[:, :, i]
attn = (q_i @ k_i.transpose(-1, -2) * L_mask[:, :, i]).masked_fill_(mask, 0)
v_prime = (k_cumdecay[:, :, i]) @ S
v_new = v_i - v_prime
o_inter = (q_i * decay[:, :, i, :, None].exp()) @ S
o[:, :, i] = o_inter + attn @ v_new
S = S * decay[:, :, i, -1, None, None].exp() + (k_i * (decay[:, :, i, -1, None] - decay[:, :, i]).exp()
[..., None]).transpose(-1, -2) @ v_new
if not output_final_state:
S = None
# unpad
o = rearrange(o, 'b h n c d -> b h (n c) d')
o = o[:, :, :T]
o = o.transpose(1, 2)
return o, S

View File

@@ -0,0 +1,351 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import torch
import triton
import triton.language as tl
from fla.ops.backends import dispatch
from fla.ops.utils import prepare_chunk_indices
from fla.ops.utils.cache import fla_cache_autotune
from fla.ops.utils.op import exp2
from fla.utils import IS_INTEL, IS_NVIDIA_BLACKWELL, autotune_cache_kwargs, check_shared_mem
# Blackwell can select unstable Triton configs for prepare_wy_repr_bwd_kernel
# during autotuning (see #913). Restrict it to the config that has been
# validated on B200 until the wider config space is re-validated.
PREPARE_WY_REPR_BWD_NUM_WARPS = [2] if IS_NVIDIA_BLACKWELL else [2, 4]
PREPARE_WY_REPR_BWD_NUM_STAGES = [4] if IS_NVIDIA_BLACKWELL else [2, 3, 4]
# Intel keeps scaling past the warp counts NVIDIA prefers: 16 warps is ~1.3x faster
# than 8 for recompute_w_u.
RECOMPUTE_W_U_NUM_WARPS = [2, 4, 8, 16] if IS_INTEL else [2, 4, 8]
@triton.heuristics({
'USE_G': lambda args: args['g'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({}, num_warps=num_warps, num_stages=num_stages)
for num_warps in RECOMPUTE_W_U_NUM_WARPS
for num_stages in [2, 3, 4]
],
key=['H', 'HV', 'K', 'V', 'BT', 'BK', 'BV', 'IS_VARLEN'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def recompute_w_u_fwd_kernel(
k,
v,
beta,
w,
u,
A,
g,
cu_seqlens,
chunk_indices,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BK: tl.constexpr,
BV: tl.constexpr,
USE_G: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_t, i_bh = tl.program_id(0).to(tl.int64), tl.program_id(1).to(tl.int64)
i_b, i_h = i_bh // HV, i_bh % HV
if IS_VARLEN:
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
else:
bos, eos = i_b * T, i_b * T + T
o_t = i_t * BT + tl.arange(0, BT)
o_A = tl.arange(0, BT)
m_t = o_t < T
m_A = m_t[:, None] & (o_A[None, :] < BT)
p_b = beta + bos*HV + i_h + o_t * HV
b_b = tl.load(p_b, mask=m_t, other=0.0)
p_A = A + (bos*HV + i_h) * BT + o_t[:, None] * (HV*BT) + o_A[None, :]
b_A = tl.load(p_A, mask=m_A, other=0.0)
for i_v in range(tl.cdiv(V, BV)):
o_v = i_v * BV + tl.arange(0, BV)
m_v = m_t[:, None] & (o_v[None, :] < V)
p_v = v + (bos*HV + i_h) * V + o_t[:, None] * (HV*V) + o_v[None, :]
p_u = u + (bos*HV + i_h) * V + o_t[:, None] * (HV*V) + o_v[None, :]
b_v = tl.load(p_v, mask=m_v, other=0.0)
b_vb = (b_v * b_b[:, None]).to(b_v.dtype)
b_u = tl.dot(b_A, b_vb, allow_tf32=False)
tl.store(p_u, b_u.to(p_u.dtype.element_ty), mask=m_v)
if USE_G:
p_g = g + (bos*HV + i_h) + o_t * HV
b_g = exp2(tl.load(p_g, mask=m_t, other=0.0))
for i_k in range(tl.cdiv(K, BK)):
o_k = i_k * BK + tl.arange(0, BK)
m_k = m_t[:, None] & (o_k[None, :] < K)
p_k = k + (bos*H + i_h // (HV // H)) * K + o_t[:, None] * (H*K) + o_k[None, :]
p_w = w + (bos*HV + i_h) * K + o_t[:, None] * (HV*K) + o_k[None, :]
b_k = tl.load(p_k, mask=m_k, other=0.0)
b_kb = b_k * b_b[:, None]
if USE_G:
b_kb *= b_g[:, None]
b_w = tl.dot(b_A, b_kb.to(b_k.dtype))
tl.store(p_w, b_w.to(p_w.dtype.element_ty), mask=m_k)
@triton.heuristics({
'USE_G': lambda args: args['g'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({}, num_warps=num_warps, num_stages=num_stages)
for num_warps in PREPARE_WY_REPR_BWD_NUM_WARPS
for num_stages in PREPARE_WY_REPR_BWD_NUM_STAGES
],
key=['H', 'HV', 'K', 'V', 'BT', 'BK', 'BV', 'IS_VARLEN'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def prepare_wy_repr_bwd_kernel(
k,
v,
beta,
g,
A,
dw,
du,
dk,
dv,
db,
dg,
cu_seqlens,
chunk_indices,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BK: tl.constexpr,
BV: tl.constexpr,
USE_G: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_t, i_bh = tl.program_id(0).to(tl.int64), tl.program_id(1).to(tl.int64)
i_b, i_h = i_bh // HV, i_bh % HV
if IS_VARLEN:
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
else:
bos, eos = i_b * T, i_b * T + T
o_t = i_t * BT + tl.arange(0, BT)
o_A = tl.arange(0, BT)
m_t = o_t < T
m_AT = (o_A[:, None] < BT) & m_t[None, :]
p_b = beta + (bos*HV + i_h) + o_t * HV
p_db = db + (bos*HV + i_h) + o_t * HV
p_A = A + (bos*HV + i_h) * BT + o_A[:, None] + o_t[None, :] * (HV*BT)
b_b = tl.load(p_b, mask=m_t, other=0.0)
b_db = tl.zeros([BT], dtype=tl.float32)
b_A = tl.load(p_A, mask=m_AT, other=0.0)
b_dA = tl.zeros([BT, BT], dtype=tl.float32)
if USE_G:
p_g = g + (bos*HV + i_h) + o_t * HV
b_g = tl.load(p_g, mask=m_t, other=0.0)
b_g_exp = exp2(b_g)
b_dg = tl.zeros([BT], dtype=tl.float32)
for i_k in range(tl.cdiv(K, BK)):
o_k = i_k * BK + tl.arange(0, BK)
m_k = m_t[:, None] & (o_k[None, :] < K)
p_k = k + (bos*H + i_h // (HV // H)) * K + o_t[:, None] * (H*K) + o_k[None, :]
p_dk = dk + (bos*HV + i_h) * K + o_t[:, None] * (HV*K) + o_k[None, :]
p_dw = dw + (bos*HV + i_h) * K + o_t[:, None] * (HV*K) + o_k[None, :]
# [BT, BK]
b_k = tl.load(p_k, mask=m_k, other=0.0)
if USE_G:
b_kbg = b_k * (b_b * b_g_exp)[:, None]
else:
b_kbg = b_k * b_b[:, None]
b_dw = tl.load(p_dw, mask=m_k, other=0.0)
b_dA += tl.dot(b_dw, tl.trans(b_kbg).to(b_dw.dtype))
b_dkbg = tl.dot(b_A, b_dw)
if USE_G:
b_dk = b_dkbg * (b_g_exp * b_b)[:, None]
b_db += tl.sum(b_dkbg * b_k * b_g_exp[:, None], 1)
b_dg += tl.sum(b_dkbg * b_kbg, 1)
else:
b_dk = b_dkbg * b_b[:, None]
b_db += tl.sum(b_dkbg * b_k, 1)
tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), mask=m_k)
for i_v in range(tl.cdiv(V, BV)):
o_v = i_v * BV + tl.arange(0, BV)
m_v = m_t[:, None] & (o_v[None, :] < V)
p_v = v + (bos*HV + i_h) * V + o_t[:, None] * (HV*V) + o_v[None, :]
p_dv = dv + (bos*HV + i_h) * V + o_t[:, None] * (HV*V) + o_v[None, :]
p_du = du + (bos*HV + i_h) * V + o_t[:, None] * (HV*V) + o_v[None, :]
b_v = tl.load(p_v, mask=m_v, other=0.0)
b_vb = (b_v * b_b[:, None]).to(b_v.dtype)
b_du = tl.load(p_du, mask=m_v, other=0.0)
b_dA += tl.dot(b_du, tl.trans(b_vb))
b_dvb = tl.dot(b_A, b_du)
b_dv = b_dvb * b_b[:, None]
b_db += tl.sum(b_dvb * b_v, 1)
tl.store(p_dv, b_dv.to(p_dv.dtype.element_ty), mask=m_v)
m_A = (o_t[:, None] > o_t[None, :]) & (m_t[:, None] & m_t)
b_dA = tl.where(m_A, b_dA, 0)
b_dA = tl.dot(b_dA.to(b_A.dtype), b_A)
b_dA = tl.dot(b_A, b_dA.to(b_A.dtype))
if USE_G:
b_dA *= exp2(b_g[:, None] - b_g[None, :])
b_A = tl.zeros([BT, BT], dtype=tl.float32)
b_dA = tl.where(m_A, -b_dA, 0).to(k.dtype.element_ty)
tl.debug_barrier()
for i_k in range(tl.cdiv(K, BK)):
o_k = i_k * BK + tl.arange(0, BK)
m_k = m_t[:, None] & (o_k[None, :] < K)
p_k = k + (bos*H + i_h // (HV // H)) * K + o_t[:, None] * (H*K) + o_k[None, :]
p_dk = dk + (bos*HV + i_h) * K + o_t[:, None] * (HV*K) + o_k[None, :]
b_k = tl.load(p_k, mask=m_k, other=0.0)
b_kt = tl.trans(b_k)
b_kb = b_k * b_b[:, None]
b_A += tl.dot(b_k, b_kt)
b_dkb = tl.dot(b_dA, b_k)
b_db += tl.sum(b_dkb * b_k, 1)
b_dk = b_dkb * b_b[:, None] + tl.trans(tl.dot(tl.trans(b_kb).to(b_dA.dtype), b_dA))
b_dk += tl.load(p_dk, mask=m_k, other=0.0)
tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), mask=m_k)
tl.store(p_db, b_db.to(p_db.dtype.element_ty), mask=m_t)
b_A *= b_b[:, None]
if USE_G:
b_AdA = b_dA * b_A
p_dg = dg + (bos*HV + i_h) + o_t * HV
b_dg += tl.sum(b_AdA, axis=1) - tl.sum(b_AdA, axis=0)
tl.store(p_dg, b_dg.to(p_dg.dtype.element_ty), mask=m_t)
@dispatch('gated_delta_rule')
def recompute_w_u_fwd(
k: torch.Tensor,
v: torch.Tensor,
beta: torch.Tensor,
A: torch.Tensor,
g: torch.Tensor | None = None,
cu_seqlens: torch.LongTensor | None = None,
chunk_indices: torch.LongTensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
B, T, H, K, V, HV = *k.shape, v.shape[-1], v.shape[2]
BT = A.shape[-1]
BK = 64
BV = 64
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
w = k.new_empty(B, T, HV, K)
u = torch.empty_like(v)
recompute_w_u_fwd_kernel[(NT, B*HV)](
k=k,
v=v,
beta=beta,
w=w,
u=u,
A=A,
g=g,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BT=BT,
BK=BK,
BV=BV,
)
return w, u
@dispatch('gated_delta_rule')
def prepare_wy_repr_bwd(
k: torch.Tensor,
v: torch.Tensor,
beta: torch.Tensor,
A: torch.Tensor,
dw: torch.Tensor,
du: torch.Tensor,
g: torch.Tensor = None,
cu_seqlens: torch.LongTensor | None = None,
chunk_indices: torch.LongTensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
B, T, H, K, V, HV = *k.shape, v.shape[-1], v.shape[2]
BT = A.shape[-1]
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
CONST_TILING = 64 if check_shared_mem() else 32
BK = min(max(triton.next_power_of_2(K), 16), CONST_TILING)
BV = min(max(triton.next_power_of_2(V), 16), CONST_TILING)
dk = k.new_empty(B, T, HV, K)
dv = torch.empty_like(v)
dg = torch.empty_like(g) if g is not None else None
db = torch.empty_like(beta)
prepare_wy_repr_bwd_kernel[(NT, B * HV)](
k=k,
v=v,
beta=beta,
g=g,
A=A,
dw=dw,
du=du,
dk=dk,
dv=dv,
db=db,
dg=dg,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BT=BT,
BK=BK,
BV=BV,
)
if H != HV:
dk = dk.view(B, T, H, HV // H, K).sum(3)
return dk, dv, db, dg
fwd_recompute_w_u = recompute_w_u_fwd
bwd_prepare_wy_repr = prepare_wy_repr_bwd