Initial K3 snapshot: 0.5B KDA/MLA/MoE train path

Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias
backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
This commit is contained in:
dela
2026-08-25 14:43:17 +08:00
commit 584f7e9e73
140 changed files with 21592 additions and 0 deletions
+14
View File
@@ -0,0 +1,14 @@
"""KDA operators, composable layers, and CausalLM."""
from .models.causal_lm import CausalLM
from .models.config import KDAConfig
from .models.k3_config import K3Config
from .ops import chunk_kda
__all__ = [
"CausalLM",
"KDAConfig",
"K3Config",
"chunk_kda",
]
__version__ = "0.0.1"
+21
View File
@@ -0,0 +1,21 @@
MIT License
Copyright (c) 2023-2026 Songlin Yang, Yu Zhang, Zhiyuan Li
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
+9
View File
@@ -0,0 +1,9 @@
# Vendored FLA KDA kernels
Subset of [flash-linear-attention](https://github.com/fla-org/flash-linear-attention)
used by `kda.ops` `backend="triton"`.
- License: MIT (see `LICENSE`)
- Upstream version tag in `__init__.py`
- Import path is `kda._fla.*`, not `fla.*`
- Not included: context parallel, Ascend, TileLang, `flash_kda`, non-KDA ops
+7
View File
@@ -0,0 +1,7 @@
"""Vendored NVIDIA-Triton KDA path from flash-linear-attention (MIT).
This package is imported as ``kda._fla``, never as the upstream ``fla``
distribution. Context-parallel, Ascend, and TileLang backends are omitted.
"""
__version__ = "0.5.2"
+1
View File
@@ -0,0 +1 @@
# Vendored FLA modules used by KDA (l2norm).
+3
View File
@@ -0,0 +1,3 @@
from kda._fla.ops.backends import BackendRegistry, BaseBackend, dispatch
__all__ = ["BackendRegistry", "BaseBackend", "dispatch"]
+299
View File
@@ -0,0 +1,299 @@
# 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 as nn
import triton
import triton.language as tl
from kda._fla.modules.backends import dispatch
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.utils import IS_AMD, autotune_cache_kwargs, input_guard
BT_LIST = [8, 16, 32, 64, 128]
NUM_WARPS_AUTOTUNE = [1, 2, 4, 8, 16] if IS_AMD else [1, 2, 4, 8, 16, 32]
@triton.autotune(
configs=[triton.Config({}, num_warps=num_warps) for num_warps in NUM_WARPS_AUTOTUNE],
key=["D"],
**autotune_cache_kwargs,
)
@triton.jit
def l2norm_fwd_kernel1(
x,
y,
rstd,
eps,
D,
BD: tl.constexpr,
):
i_t = tl.program_id(0).to(tl.int64)
x += i_t * D
y += i_t * D
# Compute mean and variance
cols = tl.arange(0, BD)
mask = cols < D
b_x = tl.load(x + cols, mask=mask, other=0.0).to(tl.float32)
b_rstd = 1 / tl.sqrt(tl.sum(b_x * b_x) + eps)
b_y = b_x * b_rstd
tl.store(y + cols, b_y, mask=mask)
tl.store(rstd + i_t, b_rstd)
@triton.autotune(
configs=[triton.Config({}, num_warps=num_warps) for num_warps in NUM_WARPS_AUTOTUNE],
key=["D"],
**autotune_cache_kwargs,
)
@triton.jit
def l2norm_bwd_kernel1(
y,
rstd,
dy,
dx,
eps,
D,
BD: tl.constexpr,
):
i_t = tl.program_id(0).to(tl.int64)
y += i_t * D
dx += i_t * D
dy += i_t * D
cols = tl.arange(0, BD)
mask = cols < D
b_y = tl.load(y + cols, mask=mask, other=0.0).to(tl.float32)
b_rstd = tl.load(rstd + i_t).to(tl.float32)
b_dy = tl.load(dy + cols, mask=mask, other=0.0).to(tl.float32)
b_dx = b_dy * b_rstd - tl.sum(b_dy * b_y) * b_y * b_rstd
tl.store(dx + cols, b_dx, mask=mask)
@fla_cache_autotune(
configs=[triton.Config({"BT": BT}, num_warps=num_warps) for num_warps in [1, 2, 4, 8, 16] for BT in BT_LIST],
key=["D", "NB"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=["T"])
def l2norm_fwd_kernel(
x,
y,
rstd,
eps,
T,
D: tl.constexpr,
BD: tl.constexpr,
NB: tl.constexpr,
BT: tl.constexpr,
):
i_t = tl.program_id(0).to(tl.int64)
o_t = i_t * BT + tl.arange(0, BT)
o_d = tl.arange(0, BD)
m_t = o_t < T
m_x = m_t[:, None] & (o_d[None, :] < D)
p_x = x + o_t[:, None] * D + o_d[None, :]
p_y = y + o_t[:, None] * D + o_d[None, :]
p_rstd = rstd + o_t
b_x = tl.load(p_x, mask=m_x, other=0.0).to(tl.float32)
b_rstd = 1 / tl.sqrt(tl.sum(b_x * b_x, 1) + eps)
b_y = b_x * b_rstd[:, None]
tl.store(p_y, b_y.to(p_y.dtype.element_ty), mask=m_x)
tl.store(p_rstd, b_rstd.to(p_rstd.dtype.element_ty), mask=m_t)
@fla_cache_autotune(
configs=[triton.Config({"BT": BT}, num_warps=num_warps) for num_warps in [1, 2, 4, 8, 16] for BT in BT_LIST],
key=["D", "NB"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=["T"])
def l2norm_bwd_kernel(
y,
rstd,
dy,
dx,
eps,
T,
D: tl.constexpr,
BD: tl.constexpr,
NB: tl.constexpr,
BT: tl.constexpr,
):
i_t = tl.program_id(0).to(tl.int64)
o_t = i_t * BT + tl.arange(0, BT)
o_d = tl.arange(0, BD)
m_t = o_t < T
m_x = m_t[:, None] & (o_d[None, :] < D)
p_y = y + o_t[:, None] * D + o_d[None, :]
p_rstd = rstd + o_t
p_dy = dy + o_t[:, None] * D + o_d[None, :]
p_dx = dx + o_t[:, None] * D + o_d[None, :]
b_y = tl.load(p_y, mask=m_x, other=0.0).to(tl.float32)
b_rstd = tl.load(p_rstd, mask=m_t, other=0.0).to(tl.float32)
b_dy = tl.load(p_dy, mask=m_x, other=0.0).to(tl.float32)
b_dx = b_dy * b_rstd[:, None] - tl.sum(b_dy * b_y, 1)[:, None] * b_y * b_rstd[:, None]
tl.store(p_dx, b_dx.to(p_dx.dtype.element_ty), mask=m_x)
@dispatch('modules')
def l2norm_fwd(
x: torch.Tensor,
eps: float = 1e-6,
output_dtype: torch.dtype | None = None,
):
x_shape_og = x.shape
x = x.view(-1, x.shape[-1])
# allocate output
if output_dtype is None:
y = torch.empty_like(x)
else:
y = torch.empty_like(x, dtype=output_dtype)
assert y.stride(-1) == 1
T, D = x.shape[0], x.shape[-1]
# Less than 64KB per feature: enqueue fused kernel
MAX_FUSED_SIZE = 65536 // x.element_size()
BD = min(MAX_FUSED_SIZE, triton.next_power_of_2(D))
if D > BD:
raise RuntimeError("This layer doesn't support feature dim >= 64KB.")
rstd = torch.empty((T,), dtype=torch.float32, device=x.device)
if D <= 512:
# NOTE(tylerr): Avoid excessive recompilation and autotuning by tolerating a larger range
# of T before recompiling the kernel.
# NB = triton.cdiv(T, 2048)
NB = triton.cdiv(T, 2048 * 32)
def grid(meta):
return (triton.cdiv(T, meta["BT"]),)
l2norm_fwd_kernel[grid](
x=x,
y=y,
rstd=rstd,
eps=eps,
T=T,
D=D,
BD=BD,
NB=NB,
)
else:
l2norm_fwd_kernel1[(T,)](
x=x,
y=y,
rstd=rstd,
eps=eps,
D=D,
BD=BD,
)
return y.view(x_shape_og), rstd.view(x_shape_og[:-1])
@dispatch('modules')
def l2norm_bwd(
y: torch.Tensor,
rstd: torch.Tensor,
dy: torch.Tensor,
eps: float = 1e-6,
):
y_shape_og = y.shape
y = y.view(-1, dy.shape[-1])
dy = dy.view(-1, dy.shape[-1])
assert dy.shape == y.shape
# allocate output
dx = torch.empty_like(y)
T, D = y.shape[0], y.shape[-1]
# Less than 64KB per feature: enqueue fused kernel
MAX_FUSED_SIZE = 65536 // y.element_size()
BD = min(MAX_FUSED_SIZE, triton.next_power_of_2(D))
if D > BD:
raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.")
if D <= 512:
# NOTE(tylerr): Avoid excessive recompilation and autotuning by tolerating a larger range
# of T before recompiling the kernel.
# NB = triton.cdiv(T, 2048)
NB = triton.cdiv(T, 2048 * 32)
def grid(meta):
return (triton.cdiv(T, meta["BT"]),)
l2norm_bwd_kernel[grid](
y=y,
rstd=rstd,
dy=dy,
dx=dx,
eps=eps,
T=T,
D=D,
BD=BD,
NB=NB,
)
else:
l2norm_bwd_kernel1[(T,)](
y=y,
rstd=rstd,
dy=dy,
dx=dx,
eps=eps,
D=D,
BD=BD,
)
return dx.view(y_shape_og)
class L2NormFunction(torch.autograd.Function):
@staticmethod
@input_guard
def forward(
ctx,
x,
eps=1e-6,
output_dtype=None,
):
y, rstd = l2norm_fwd(x, eps, output_dtype)
ctx.eps = eps
ctx.x_dtype = x.dtype
ctx.save_for_backward(y, rstd)
return y
@staticmethod
@input_guard
def backward(ctx, dy):
y, rstd = ctx.saved_tensors
dx = l2norm_bwd(y, rstd, dy, ctx.eps)
return dx, None, None
def l2norm(
x: torch.Tensor,
eps: float = 1e-6,
output_dtype: torch.dtype | None = None,
) -> torch.Tensor:
return L2NormFunction.apply(x, eps, output_dtype)
l2_norm = l2norm
class L2Norm(nn.Module):
def __init__(
self,
eps: float = 1e-6,
output_dtype: torch.dtype | None = None,
):
super().__init__()
self.eps = eps
self.output_dtype = output_dtype
def forward(self, x: torch.Tensor) -> torch.Tensor:
return l2norm(x, self.eps, self.output_dtype)
+1
View File
@@ -0,0 +1 @@
# Vendored FLA ops subset.
+34
View File
@@ -0,0 +1,34 @@
"""Identity dispatch: keep the NVIDIA Triton implementation in this tree."""
from __future__ import annotations
from collections.abc import Callable
from typing import TypeVar
F = TypeVar("F", bound=Callable)
def dispatch(operation: str):
def decorator(func: F) -> F:
return func
return decorator
class BaseBackend:
backend_type = "triton"
def is_available(self) -> bool:
return True
def is_enabled(self) -> bool:
return True
class BackendRegistry:
@classmethod
def ensure_initialized(cls, operation: str) -> None:
return None
__all__ = ["BackendRegistry", "BaseBackend", "dispatch"]
+1
View File
@@ -0,0 +1 @@
# Vendored FLA common kernels used by KDA.
+806
View File
@@ -0,0 +1,806 @@
# 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 kda._fla.ops.backends import dispatch
from kda._fla.ops.utils import prepare_chunk_indices, prepare_chunk_offsets
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.ops.utils.op import exp2
from kda._fla.utils import (
IS_INTEL,
IS_NVIDIA_BLACKWELL,
IS_NVIDIA_HOPPER,
autotune_cache_kwargs,
check_shared_mem,
)
NUM_WARPS = [2, 4] if IS_NVIDIA_HOPPER else [2, 4, 8, 16]
# TODO: Triton mainline fixes a Blackwell tl.dot recurrence race.
# Keep this kernel on num_warps=2 for Blackwell until Triton 3.8 is released
# and we re-validate the wider config space.
# Intel needs more warps than NVIDIA here: 8 warps is ~1.5x faster than the best
# config reachable under the [2, 4] cap.
if IS_NVIDIA_BLACKWELL:
GATED_DELTA_RULE_FWD_H_NUM_WARPS = [2]
elif IS_INTEL:
GATED_DELTA_RULE_FWD_H_NUM_WARPS = [2, 4, 8, 16]
else:
GATED_DELTA_RULE_FWD_H_NUM_WARPS = [2, 4]
@triton.heuristics({
'USE_G': lambda args: args['g'] is not None,
'USE_GK': lambda args: args['gk'] is not None,
'USE_INITIAL_STATE': lambda args: args['h0'] is not None,
'STORE_FINAL_STATE': lambda args: args['ht'] is not None,
'SAVE_NEW_VALUE': lambda args: args['v_new'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({'BV': BV}, num_warps=num_warps, num_stages=num_stages)
for num_warps in GATED_DELTA_RULE_FWD_H_NUM_WARPS
for num_stages in ([2, 3, 4] if check_shared_mem('ampere') else [2, 1])
for BV in ([32, 64] if check_shared_mem('ada') else [32])
],
key=['H', 'HV', 'K', 'V', 'BT', 'STATE_V_FIRST'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_gated_delta_rule_fwd_kernel_h_blockdim64(
k,
v,
w,
v_new,
g,
gk,
h,
h0,
ht,
cu_seqlens,
chunk_offsets,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BV: tl.constexpr,
USE_G: tl.constexpr,
USE_GK: tl.constexpr,
USE_INITIAL_STATE: tl.constexpr,
STORE_FINAL_STATE: tl.constexpr,
SAVE_NEW_VALUE: tl.constexpr,
STATE_V_FIRST: tl.constexpr,
IS_VARLEN: 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_h = i_nh // HV, i_nh % HV
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
NT = tl.cdiv(T, BT)
boh = tl.load(chunk_offsets + i_n).to(tl.int64)
else:
bos, eos = i_n * T, i_n * T + T
NT = tl.cdiv(T, BT)
boh = i_n * NT
if STATE_V_FIRST:
b_h1 = tl.zeros([BV, 64], dtype=tl.float32)
if K > 64:
b_h2 = tl.zeros([BV, 64], dtype=tl.float32)
if K > 128:
b_h3 = tl.zeros([BV, 64], dtype=tl.float32)
if K > 192:
b_h4 = tl.zeros([BV, 64], dtype=tl.float32)
else:
b_h1 = tl.zeros([64, BV], dtype=tl.float32)
if K > 64:
b_h2 = tl.zeros([64, BV], dtype=tl.float32)
if K > 128:
b_h3 = tl.zeros([64, BV], dtype=tl.float32)
if K > 192:
b_h4 = tl.zeros([64, BV], dtype=tl.float32)
# calculate offset
h += (boh * HV + i_h).to(tl.int64) * K*V
v += (bos * HV + i_h).to(tl.int64) * V
k += (bos * H + i_h // (HV // H)).to(tl.int64) * K
w += (bos * HV + i_h).to(tl.int64) * K
if SAVE_NEW_VALUE:
v_new += (bos * HV + i_h).to(tl.int64) * V
if USE_INITIAL_STATE:
h0 = h0 + i_nh * K*V
if STORE_FINAL_STATE:
ht = ht + i_nh * K*V
# load initial state
o_v = i_v * BV + tl.arange(0, BV)
m_v = o_v < V
o_k1 = tl.arange(0, 64)
m_k1 = o_k1 < K
o_k2 = 64 + o_k1
m_k2 = o_k2 < K
o_k3 = 128 + o_k1
m_k3 = o_k3 < K
o_k4 = 192 + o_k1
m_k4 = o_k4 < K
if USE_INITIAL_STATE:
if STATE_V_FIRST:
p_h0_1 = h0 + o_v[:, None] * K + o_k1[None, :]
m_h0_1 = m_v[:, None] & m_k1[None, :]
else:
p_h0_1 = h0 + o_k1[:, None] * V + o_v[None, :]
m_h0_1 = m_k1[:, None] & m_v[None, :]
b_h1 += tl.load(p_h0_1, mask=m_h0_1, other=0.0).to(tl.float32)
if K > 64:
if STATE_V_FIRST:
p_h0_2 = h0 + o_v[:, None] * K + o_k2[None, :]
m_h0_2 = m_v[:, None] & m_k2[None, :]
else:
p_h0_2 = h0 + o_k2[:, None] * V + o_v[None, :]
m_h0_2 = m_k2[:, None] & m_v[None, :]
b_h2 += tl.load(p_h0_2, mask=m_h0_2, other=0.0).to(tl.float32)
if K > 128:
if STATE_V_FIRST:
p_h0_3 = h0 + o_v[:, None] * K + o_k3[None, :]
m_h0_3 = m_v[:, None] & m_k3[None, :]
else:
p_h0_3 = h0 + o_k3[:, None] * V + o_v[None, :]
m_h0_3 = m_k3[:, None] & m_v[None, :]
b_h3 += tl.load(p_h0_3, mask=m_h0_3, other=0.0).to(tl.float32)
if K > 192:
if STATE_V_FIRST:
p_h0_4 = h0 + o_v[:, None] * K + o_k4[None, :]
m_h0_4 = m_v[:, None] & m_k4[None, :]
else:
p_h0_4 = h0 + o_k4[:, None] * V + o_v[None, :]
m_h0_4 = m_k4[:, None] & m_v[None, :]
b_h4 += tl.load(p_h0_4, mask=m_h0_4, other=0.0).to(tl.float32)
# main recurrence
for i_t in range(NT):
i_t_int64 = i_t.to(tl.int64)
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
if STATE_V_FIRST:
p_h1 = h + i_t_int64 * HV*K*V + o_v[:, None] * K + o_k1[None, :]
m_h1 = m_v[:, None] & m_k1[None, :]
else:
p_h1 = h + i_t_int64 * HV*K*V + o_k1[:, None] * V + o_v[None, :]
m_h1 = m_k1[:, None] & m_v[None, :]
tl.store(p_h1, b_h1.to(p_h1.dtype.element_ty), mask=m_h1)
if K > 64:
if STATE_V_FIRST:
p_h2 = h + i_t_int64 * HV*K*V + o_v[:, None] * K + o_k2[None, :]
m_h2 = m_v[:, None] & m_k2[None, :]
else:
p_h2 = h + i_t_int64 * HV*K*V + o_k2[:, None] * V + o_v[None, :]
m_h2 = m_k2[:, None] & m_v[None, :]
tl.store(p_h2, b_h2.to(p_h2.dtype.element_ty), mask=m_h2)
if K > 128:
if STATE_V_FIRST:
p_h3 = h + i_t_int64 * HV*K*V + o_v[:, None] * K + o_k3[None, :]
m_h3 = m_v[:, None] & m_k3[None, :]
else:
p_h3 = h + i_t_int64 * HV*K*V + o_k3[:, None] * V + o_v[None, :]
m_h3 = m_k3[:, None] & m_v[None, :]
tl.store(p_h3, b_h3.to(p_h3.dtype.element_ty), mask=m_h3)
if K > 192:
if STATE_V_FIRST:
p_h4 = h + i_t_int64 * HV*K*V + o_v[:, None] * K + o_k4[None, :]
m_h4 = m_v[:, None] & m_k4[None, :]
else:
p_h4 = h + i_t_int64 * HV*K*V + o_k4[:, None] * V + o_v[None, :]
m_h4 = m_k4[:, None] & m_v[None, :]
tl.store(p_h4, b_h4.to(p_h4.dtype.element_ty), mask=m_h4)
p_w = w + o_t[:, None] * (HV*K) + o_k1[None, :]
b_w = tl.load(p_w, mask=m_t[:, None] & m_k1[None, :], other=0.0)
if STATE_V_FIRST:
b_v = tl.dot(b_w, tl.trans(b_h1).to(b_w.dtype))
else:
b_v = tl.dot(b_w, b_h1.to(b_w.dtype))
if K > 64:
p_w = w + o_t[:, None] * (HV*K) + o_k2[None, :]
b_w = tl.load(p_w, mask=m_t[:, None] & m_k2[None, :], other=0.0)
if STATE_V_FIRST:
b_v += tl.dot(b_w, tl.trans(b_h2).to(b_w.dtype))
else:
b_v += tl.dot(b_w, b_h2.to(b_w.dtype))
if K > 128:
p_w = w + o_t[:, None] * (HV*K) + o_k3[None, :]
b_w = tl.load(p_w, mask=m_t[:, None] & m_k3[None, :], other=0.0)
if STATE_V_FIRST:
b_v += tl.dot(b_w, tl.trans(b_h3).to(b_w.dtype))
else:
b_v += tl.dot(b_w, b_h3.to(b_w.dtype))
if K > 192:
p_w = w + o_t[:, None] * (HV*K) + o_k4[None, :]
b_w = tl.load(p_w, mask=m_t[:, None] & m_k4[None, :], other=0.0)
if STATE_V_FIRST:
b_v += tl.dot(b_w, tl.trans(b_h4).to(b_w.dtype))
else:
b_v += tl.dot(b_w, b_h4.to(b_w.dtype))
p_v = v + o_t[:, None] * (HV*V) + o_v[None, :]
b_v = tl.load(p_v, mask=m_t[:, None] & m_v[None, :], other=0.0) - b_v
if SAVE_NEW_VALUE:
p_v = v_new + o_t[:, None] * (HV*V) + o_v[None, :]
tl.store(p_v, b_v.to(p_v.dtype.element_ty), mask=m_t[:, None] & m_v[None, :])
last_idx = min((i_t + 1) * BT, T) - 1
if USE_G:
b_g_last = tl.load(g + (bos * HV + last_idx * HV + i_h).to(tl.int64)).to(tl.float32)
p_g = g + (bos * HV + i_h).to(tl.int64) + o_t * HV
b_g = tl.load(p_g, mask=m_t, other=0.0).to(tl.float32)
b_v = b_v * tl.where(m_t, exp2(b_g_last - b_g), 0)[:, None]
b_g_last = exp2(b_g_last)
b_h1 *= b_g_last
if K > 64:
b_h2 *= b_g_last
if K > 128:
b_h3 *= b_g_last
if K > 192:
b_h4 *= b_g_last
if USE_GK:
o_k1 = tl.arange(0, 64)
b_gk_last1 = tl.load(gk + (bos + last_idx) * HV*K + i_h * K + o_k1, mask=(o_k1 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_h1 *= exp2(b_gk_last1)[None, :]
else:
b_h1 *= exp2(b_gk_last1)[:, None]
if K > 64:
o_k2 = 64 + o_k1
b_gk_last2 = tl.load(gk + (bos + last_idx) * HV*K + i_h * K + o_k2, mask=(o_k2 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_h2 *= exp2(b_gk_last2)[None, :]
else:
b_h2 *= exp2(b_gk_last2)[:, None]
if K > 128:
o_k3 = 128 + o_k1
b_gk_last3 = tl.load(gk + (bos + last_idx) * HV*K + i_h * K + o_k3, mask=(o_k3 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_h3 *= exp2(b_gk_last3)[None, :]
else:
b_h3 *= exp2(b_gk_last3)[:, None]
if K > 192:
o_k4 = 192 + o_k1
b_gk_last4 = tl.load(gk + (bos + last_idx) * HV*K + i_h * K + o_k4, mask=(o_k4 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_h4 *= exp2(b_gk_last4)[None, :]
else:
b_h4 *= exp2(b_gk_last4)[:, None]
b_v = b_v.to(k.dtype.element_ty)
p_k = k + o_k1[:, None] + o_t[None, :] * (H*K)
b_k = tl.load(p_k, mask=m_k1[:, None] & m_t[None, :], other=0.0)
if STATE_V_FIRST:
b_h1 += tl.trans(tl.dot(b_k, b_v))
else:
b_h1 += tl.dot(b_k, b_v)
if K > 64:
p_k = k + o_k2[:, None] + o_t[None, :] * (H*K)
b_k = tl.load(p_k, mask=m_k2[:, None] & m_t[None, :], other=0.0)
if STATE_V_FIRST:
b_h2 += tl.trans(tl.dot(b_k, b_v))
else:
b_h2 += tl.dot(b_k, b_v)
if K > 128:
p_k = k + o_k3[:, None] + o_t[None, :] * (H*K)
b_k = tl.load(p_k, mask=m_k3[:, None] & m_t[None, :], other=0.0)
if STATE_V_FIRST:
b_h3 += tl.trans(tl.dot(b_k, b_v))
else:
b_h3 += tl.dot(b_k, b_v)
if K > 192:
p_k = k + o_k4[:, None] + o_t[None, :] * (H*K)
b_k = tl.load(p_k, mask=m_k4[:, None] & m_t[None, :], other=0.0)
if STATE_V_FIRST:
b_h4 += tl.trans(tl.dot(b_k, b_v))
else:
b_h4 += tl.dot(b_k, b_v)
if STORE_FINAL_STATE:
if STATE_V_FIRST:
p_ht = ht + o_v[:, None] * K + o_k1[None, :]
m_ht = m_v[:, None] & m_k1[None, :]
else:
p_ht = ht + o_k1[:, None] * V + o_v[None, :]
m_ht = m_k1[:, None] & m_v[None, :]
tl.store(p_ht, b_h1.to(p_ht.dtype.element_ty), mask=m_ht)
if K > 64:
if STATE_V_FIRST:
p_ht = ht + o_v[:, None] * K + o_k2[None, :]
m_ht = m_v[:, None] & m_k2[None, :]
else:
p_ht = ht + o_k2[:, None] * V + o_v[None, :]
m_ht = m_k2[:, None] & m_v[None, :]
tl.store(p_ht, b_h2.to(p_ht.dtype.element_ty), mask=m_ht)
if K > 128:
if STATE_V_FIRST:
p_ht = ht + o_v[:, None] * K + o_k3[None, :]
m_ht = m_v[:, None] & m_k3[None, :]
else:
p_ht = ht + o_k3[:, None] * V + o_v[None, :]
m_ht = m_k3[:, None] & m_v[None, :]
tl.store(p_ht, b_h3.to(p_ht.dtype.element_ty), mask=m_ht)
if K > 192:
if STATE_V_FIRST:
p_ht = ht + o_v[:, None] * K + o_k4[None, :]
m_ht = m_v[:, None] & m_k4[None, :]
else:
p_ht = ht + o_k4[:, None] * V + o_v[None, :]
m_ht = m_k4[:, None] & m_v[None, :]
tl.store(p_ht, b_h4.to(p_ht.dtype.element_ty), mask=m_ht)
@triton.heuristics({
'USE_G': lambda args: args['g'] is not None,
'USE_GK': lambda args: args['gk'] is not None,
'USE_INITIAL_STATE': lambda args: args['dh0'] is not None,
'USE_FINAL_STATE_GRADIENT': lambda args: args['dht'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({'BV': BV}, num_warps=num_warps, num_stages=num_stages)
for num_warps in [2, 4]
for num_stages in ([2, 3, 4] if check_shared_mem('ampere') else [1])
for BV in ([32, 64] if check_shared_mem('ada') else [32])
],
key=['H', 'HV', 'K', 'V', 'BT', 'BV', 'USE_G', 'STATE_V_FIRST'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_gated_delta_rule_bwd_kernel_dhu_blockdim64(
q,
k,
w,
g,
gk,
dht,
dh0,
do,
dh,
dv,
dv2,
cu_seqlens,
chunk_offsets,
scale,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BV: tl.constexpr,
USE_G: tl.constexpr,
USE_GK: tl.constexpr,
USE_INITIAL_STATE: tl.constexpr,
USE_FINAL_STATE_GRADIENT: tl.constexpr,
STATE_V_FIRST: tl.constexpr,
IS_VARLEN: 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_h = i_nh // HV, i_nh % HV
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
NT = tl.cdiv(T, BT)
boh = tl.load(chunk_offsets + i_n).to(tl.int64)
else:
bos, eos = i_n * T, i_n * T + T
NT = tl.cdiv(T, BT)
boh = i_n * NT
if STATE_V_FIRST:
b_dh1 = tl.zeros([BV, 64], dtype=tl.float32)
if K > 64:
b_dh2 = tl.zeros([BV, 64], dtype=tl.float32)
if K > 128:
b_dh3 = tl.zeros([BV, 64], dtype=tl.float32)
if K > 192:
b_dh4 = tl.zeros([BV, 64], dtype=tl.float32)
else:
b_dh1 = tl.zeros([64, BV], dtype=tl.float32)
if K > 64:
b_dh2 = tl.zeros([64, BV], dtype=tl.float32)
if K > 128:
b_dh3 = tl.zeros([64, BV], dtype=tl.float32)
if K > 192:
b_dh4 = tl.zeros([64, BV], dtype=tl.float32)
# calculate offset
q += (bos * H + i_h // (HV // H)).to(tl.int64) * K
k += (bos * H + i_h // (HV // H)).to(tl.int64) * K
w += (bos * HV + i_h).to(tl.int64) * K
do += (bos * HV + i_h).to(tl.int64) * V
dv += (bos * HV + i_h).to(tl.int64) * V
dv2 += (bos * HV + i_h).to(tl.int64) * V
dh += (boh * HV + i_h).to(tl.int64) * K*V
if USE_GK:
gk += (bos * HV + i_h).to(tl.int64) * K
if USE_INITIAL_STATE:
dh0 += i_nh * K*V
if USE_FINAL_STATE_GRADIENT:
dht += i_nh * K*V
o_v = i_v * BV + tl.arange(0, BV)
m_v = o_v < V
o_k1 = tl.arange(0, 64)
m_k1 = o_k1 < K
o_k2 = 64 + o_k1
m_k2 = o_k2 < K
o_k3 = 128 + o_k1
m_k3 = o_k3 < K
o_k4 = 192 + o_k1
m_k4 = o_k4 < K
if USE_FINAL_STATE_GRADIENT:
if STATE_V_FIRST:
p_dht1 = dht + o_v[:, None] * K + o_k1[None, :]
m_dht1 = m_v[:, None] & m_k1[None, :]
else:
p_dht1 = dht + o_k1[:, None] * V + o_v[None, :]
m_dht1 = m_k1[:, None] & m_v[None, :]
b_dh1 += tl.load(p_dht1, mask=m_dht1, other=0.0)
if K > 64:
if STATE_V_FIRST:
p_dht2 = dht + o_v[:, None] * K + o_k2[None, :]
m_dht2 = m_v[:, None] & m_k2[None, :]
else:
p_dht2 = dht + o_k2[:, None] * V + o_v[None, :]
m_dht2 = m_k2[:, None] & m_v[None, :]
b_dh2 += tl.load(p_dht2, mask=m_dht2, other=0.0)
if K > 128:
if STATE_V_FIRST:
p_dht3 = dht + o_v[:, None] * K + o_k3[None, :]
m_dht3 = m_v[:, None] & m_k3[None, :]
else:
p_dht3 = dht + o_k3[:, None] * V + o_v[None, :]
m_dht3 = m_k3[:, None] & m_v[None, :]
b_dh3 += tl.load(p_dht3, mask=m_dht3, other=0.0)
if K > 192:
if STATE_V_FIRST:
p_dht4 = dht + o_v[:, None] * K + o_k4[None, :]
m_dht4 = m_v[:, None] & m_k4[None, :]
else:
p_dht4 = dht + o_k4[:, None] * V + o_v[None, :]
m_dht4 = m_k4[:, None] & m_v[None, :]
b_dh4 += tl.load(p_dht4, mask=m_dht4, other=0.0)
for i_t in range(NT - 1, -1, -1):
i_t_int64 = i_t.to(tl.int64)
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
if STATE_V_FIRST:
p_dh1 = dh + i_t_int64*HV*K*V + o_v[:, None] * K + o_k1[None, :]
m_dh1 = m_v[:, None] & m_k1[None, :]
else:
p_dh1 = dh + i_t_int64*HV*K*V + o_k1[:, None] * V + o_v[None, :]
m_dh1 = m_k1[:, None] & m_v[None, :]
tl.store(p_dh1, b_dh1.to(p_dh1.dtype.element_ty), mask=m_dh1)
if K > 64:
if STATE_V_FIRST:
p_dh2 = dh + i_t_int64*HV*K*V + o_v[:, None] * K + o_k2[None, :]
m_dh2 = m_v[:, None] & m_k2[None, :]
else:
p_dh2 = dh + i_t_int64*HV*K*V + o_k2[:, None] * V + o_v[None, :]
m_dh2 = m_k2[:, None] & m_v[None, :]
tl.store(p_dh2, b_dh2.to(p_dh2.dtype.element_ty), mask=m_dh2)
if K > 128:
if STATE_V_FIRST:
p_dh3 = dh + i_t_int64*HV*K*V + o_v[:, None] * K + o_k3[None, :]
m_dh3 = m_v[:, None] & m_k3[None, :]
else:
p_dh3 = dh + i_t_int64*HV*K*V + o_k3[:, None] * V + o_v[None, :]
m_dh3 = m_k3[:, None] & m_v[None, :]
tl.store(p_dh3, b_dh3.to(p_dh3.dtype.element_ty), mask=m_dh3)
if K > 192:
if STATE_V_FIRST:
p_dh4 = dh + i_t_int64*HV*K*V + o_v[:, None] * K + o_k4[None, :]
m_dh4 = m_v[:, None] & m_k4[None, :]
else:
p_dh4 = dh + i_t_int64*HV*K*V + o_k4[:, None] * V + o_v[None, :]
m_dh4 = m_k4[:, None] & m_v[None, :]
tl.store(p_dh4, b_dh4.to(p_dh4.dtype.element_ty), mask=m_dh4)
last_idx = min((i_t + 1) * BT, T) - 1
if USE_G:
bg_last = tl.load(g + (bos + last_idx) * HV + i_h).to(tl.float32)
p_g = g + bos * HV + i_h + o_t * HV
b_g = tl.load(p_g, mask=m_t, other=0.0).to(tl.float32)
bg_last_exp = exp2(bg_last)
b_g_exp = exp2(b_g)
p_dv = dv + o_t[:, None] * (HV*V) + o_v[None, :]
p_dv2 = dv2 + o_t[:, None] * (HV*V) + o_v[None, :]
p_do = do + o_t[:, None] * (HV*V) + o_v[None, :]
b_do = tl.load(p_do, mask=m_t[:, None] & m_v[None, :], other=0.0)
# Update dv
p_k = k + o_t[:, None] * (H*K) + o_k1[None, :]
b_k = tl.load(p_k, mask=m_t[:, None] & m_k1[None, :], other=0.0)
if USE_GK:
o_k1 = tl.arange(0, 64)
b_gk_last1 = tl.load(gk + last_idx * HV*K + o_k1, mask=(o_k1 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_dv = tl.dot(b_k, tl.trans(b_dh1).to(b_k.dtype))
else:
b_dv = tl.dot(b_k, b_dh1.to(b_k.dtype))
if K > 64:
p_k = k + o_t[:, None] * (H*K) + o_k2[None, :]
b_k = tl.load(p_k, mask=m_t[:, None] & m_k2[None, :], other=0.0)
if USE_GK:
b_gk_last2 = tl.load(gk + last_idx * HV*K + o_k2, mask=(o_k2 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_dv += tl.dot(b_k, tl.trans(b_dh2).to(b_k.dtype))
else:
b_dv += tl.dot(b_k, b_dh2.to(b_k.dtype))
if K > 128:
p_k = k + o_t[:, None] * (H*K) + o_k3[None, :]
b_k = tl.load(p_k, mask=m_t[:, None] & m_k3[None, :], other=0.0)
if USE_GK:
b_gk_last3 = tl.load(gk + last_idx * HV*K + o_k3, mask=(o_k3 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_dv += tl.dot(b_k, tl.trans(b_dh3).to(b_k.dtype))
else:
b_dv += tl.dot(b_k, b_dh3.to(b_k.dtype))
if K > 192:
p_k = k + o_t[:, None] * (H*K) + o_k4[None, :]
b_k = tl.load(p_k, mask=m_t[:, None] & m_k4[None, :], other=0.0)
if USE_GK:
b_gk_last4 = tl.load(gk + last_idx * HV*K + o_k4, mask=(o_k4 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_dv += tl.dot(b_k, tl.trans(b_dh4).to(b_k.dtype))
else:
b_dv += tl.dot(b_k, b_dh4.to(b_k.dtype))
if USE_G:
b_dv *= tl.where(m_t, exp2(bg_last - b_g), 0)[:, None]
b_dv += tl.load(p_dv, mask=m_t[:, None] & m_v[None, :], other=0.0)
tl.store(p_dv2, b_dv.to(p_dv.dtype.element_ty), mask=m_t[:, None] & m_v[None, :])
# Update dh
p_w = w + o_k1[:, None] + o_t[None, :] * (HV*K)
p_q = q + o_k1[:, None] + o_t[None, :] * (H*K)
b_w = tl.load(p_w, mask=m_k1[:, None] & m_t[None, :], other=0.0)
b_q = tl.load(p_q, mask=m_k1[:, None] & m_t[None, :], other=0.0)
if USE_G:
b_dh1 *= bg_last_exp
b_q = b_q * b_g_exp[None, :]
if USE_GK:
if STATE_V_FIRST:
b_dh1 *= exp2(b_gk_last1)[None, :]
else:
b_dh1 *= exp2(b_gk_last1[:, None])
if STATE_V_FIRST:
b_dh1 += tl.trans(tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype)))
else:
b_dh1 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
if K > 64:
p_q = q + o_k2[:, None] + o_t[None, :] * (H*K)
p_w = w + o_k2[:, None] + o_t[None, :] * (HV*K)
b_q = tl.load(p_q, mask=m_k2[:, None] & m_t[None, :], other=0.0)
b_w = tl.load(p_w, mask=m_k2[:, None] & m_t[None, :], other=0.0)
if USE_G:
b_dh2 *= bg_last_exp
b_q = b_q * b_g_exp[None, :]
if USE_GK:
if STATE_V_FIRST:
b_dh2 *= exp2(b_gk_last2)[None, :]
else:
b_dh2 *= exp2(b_gk_last2[:, None])
if STATE_V_FIRST:
b_dh2 += tl.trans(tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype)))
else:
b_dh2 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
if K > 128:
p_q = q + o_k3[:, None] + o_t[None, :] * (H*K)
p_w = w + o_k3[:, None] + o_t[None, :] * (HV*K)
b_q = tl.load(p_q, mask=m_k3[:, None] & m_t[None, :], other=0.0)
b_w = tl.load(p_w, mask=m_k3[:, None] & m_t[None, :], other=0.0)
if USE_G:
b_dh3 *= bg_last_exp
b_q = b_q * b_g_exp[None, :]
if USE_GK:
if STATE_V_FIRST:
b_dh3 *= exp2(b_gk_last3)[None, :]
else:
b_dh3 *= exp2(b_gk_last3[:, None])
if STATE_V_FIRST:
b_dh3 += tl.trans(tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype)))
else:
b_dh3 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
if K > 192:
p_q = q + o_k4[:, None] + o_t[None, :] * (H*K)
p_w = w + o_k4[:, None] + o_t[None, :] * (HV*K)
b_q = tl.load(p_q, mask=m_k4[:, None] & m_t[None, :], other=0.0)
b_w = tl.load(p_w, mask=m_k4[:, None] & m_t[None, :], other=0.0)
if USE_G:
b_dh4 *= bg_last_exp
b_q = b_q * b_g_exp[None, :]
if USE_GK:
if STATE_V_FIRST:
b_dh4 *= exp2(b_gk_last4)[None, :]
else:
b_dh4 *= exp2(b_gk_last4[:, None])
if STATE_V_FIRST:
b_dh4 += tl.trans(tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype)))
else:
b_dh4 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
if USE_INITIAL_STATE:
if STATE_V_FIRST:
p_dh0 = dh0 + o_v[:, None] * K + o_k1[None, :]
m_dh0 = m_v[:, None] & m_k1[None, :]
else:
p_dh0 = dh0 + o_k1[:, None] * V + o_v[None, :]
m_dh0 = m_k1[:, None] & m_v[None, :]
tl.store(p_dh0, b_dh1.to(p_dh0.dtype.element_ty), mask=m_dh0)
if K > 64:
if STATE_V_FIRST:
p_dh1 = dh0 + o_v[:, None] * K + o_k2[None, :]
m_dh1 = m_v[:, None] & m_k2[None, :]
else:
p_dh1 = dh0 + o_k2[:, None] * V + o_v[None, :]
m_dh1 = m_k2[:, None] & m_v[None, :]
tl.store(p_dh1, b_dh2.to(p_dh1.dtype.element_ty), mask=m_dh1)
if K > 128:
if STATE_V_FIRST:
p_dh2 = dh0 + o_v[:, None] * K + o_k3[None, :]
m_dh2 = m_v[:, None] & m_k3[None, :]
else:
p_dh2 = dh0 + o_k3[:, None] * V + o_v[None, :]
m_dh2 = m_k3[:, None] & m_v[None, :]
tl.store(p_dh2, b_dh3.to(p_dh2.dtype.element_ty), mask=m_dh2)
if K > 192:
if STATE_V_FIRST:
p_dh3 = dh0 + o_v[:, None] * K + o_k4[None, :]
m_dh3 = m_v[:, None] & m_k4[None, :]
else:
p_dh3 = dh0 + o_k4[:, None] * V + o_v[None, :]
m_dh3 = m_k4[:, None] & m_v[None, :]
tl.store(p_dh3, b_dh4.to(p_dh3.dtype.element_ty), mask=m_dh3)
@dispatch('common')
def chunk_gated_delta_rule_fwd_h(
k: torch.Tensor,
w: torch.Tensor,
u: torch.Tensor,
g: torch.Tensor | None = None,
gk: torch.Tensor | None = None,
initial_state: torch.Tensor | None = None,
output_final_state: bool = False,
chunk_size: int = 64,
save_new_value: bool = True,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
cu_seqlens_cpu: torch.LongTensor | None = None,
chunk_indices: torch.LongTensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None]:
B, T, H, K, V, HV = *k.shape, u.shape[-1], u.shape[2]
BT = chunk_size
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, chunk_size)
# N: the actual number of sequences in the batch with either equal or variable lengths
if cu_seqlens is None:
N, NT, chunk_offsets = B, triton.cdiv(T, BT), None
else:
N, NT, chunk_offsets = len(cu_seqlens) - 1, len(chunk_indices), prepare_chunk_offsets(cu_seqlens, BT)
assert K <= 256, "current kernel does not support head dimension larger than 256."
if state_v_first:
h = k.new_empty(B, NT, HV, V, K)
final_state = k.new_zeros(N, HV, V, K, dtype=torch.float32) if output_final_state else None
else:
h = k.new_empty(B, NT, HV, K, V)
final_state = k.new_zeros(N, HV, K, V, dtype=torch.float32) if output_final_state else None
v_new = torch.empty_like(u) if save_new_value else None
def grid(meta): return (triton.cdiv(V, meta['BV']) * N * HV, )
chunk_gated_delta_rule_fwd_kernel_h_blockdim64[grid](
k=k,
v=u,
w=w,
v_new=v_new,
g=g,
gk=gk,
h=h,
h0=initial_state,
ht=final_state,
cu_seqlens=cu_seqlens,
chunk_offsets=chunk_offsets,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BT=BT,
STATE_V_FIRST=state_v_first,
)
return h, v_new, final_state
@dispatch('common')
def chunk_gated_delta_rule_bwd_dhu(
q: torch.Tensor,
k: torch.Tensor,
w: torch.Tensor,
do: torch.Tensor,
dv: torch.Tensor,
g: torch.Tensor | None = None,
gk: torch.Tensor | None = None,
h0: torch.Tensor | None = None,
dht: torch.Tensor | None = None,
scale: float | None = None,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
chunk_size: int = 64,
chunk_indices: torch.LongTensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
B, T, H, K, V, HV = *q.shape, do.shape[-1], do.shape[2]
# N: the actual number of sequences in the batch with either equal or variable lengths
BT = chunk_size
assert K <= 256, "current kernel does not support head dimension being larger than 256."
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, chunk_size)
if cu_seqlens is None:
N, NT, chunk_offsets = B, triton.cdiv(T, BT), None
else:
N, NT, chunk_offsets = len(cu_seqlens) - 1, len(chunk_indices), prepare_chunk_offsets(cu_seqlens, BT)
if state_v_first:
dh = q.new_empty(B, NT, HV, V, K)
else:
dh = q.new_empty(B, NT, HV, K, V)
dh0 = torch.empty_like(h0, dtype=torch.float32) if h0 is not None else None
dv2 = torch.empty_like(dv)
def grid(meta): return (triton.cdiv(V, meta['BV']) * N * HV, )
chunk_gated_delta_rule_bwd_kernel_dhu_blockdim64[grid](
q=q,
k=k,
w=w,
g=g,
gk=gk,
dht=dht,
dh0=dh0,
do=do,
dh=dh,
dv=dv,
dv2=dv2,
cu_seqlens=cu_seqlens,
chunk_offsets=chunk_offsets,
scale=scale,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BT=BT,
STATE_V_FIRST=state_v_first,
)
return dh, dh0, dv2
+432
View File
@@ -0,0 +1,432 @@
# 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 kda._fla.ops.backends import dispatch
from kda._fla.ops.utils import prepare_chunk_offsets
from kda._fla.ops.utils.op import exp2
from kda._fla.utils import autotune_cache_kwargs, check_shared_mem
BKV_LIST = [32, 64] if check_shared_mem() else [16, 32]
@triton.heuristics({
'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,
})
@triton.autotune(
configs=[
triton.Config({'BK': BK, 'BV': BV}, num_warps=num_warps, num_stages=num_stages)
for BK in BKV_LIST
for BV in BKV_LIST
for num_warps in [1, 2, 4, 8]
for num_stages in [2, 3, 4]
],
key=['BT', 'USE_G', 'USE_GK', 'USE_GV', 'STATE_V_FIRST'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_fwd_kernel_h(
k,
v,
h,
g,
g_gamma,
gk,
gv,
h0,
ht,
cu_seqlens,
split_offsets,
T,
H: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BS: tl.constexpr,
BK: tl.constexpr,
BV: tl.constexpr,
USE_G: tl.constexpr,
USE_G_GAMMA: tl.constexpr,
USE_GK: tl.constexpr,
USE_GV: tl.constexpr,
USE_INITIAL_STATE: tl.constexpr,
STORE_FINAL_STATE: tl.constexpr,
IS_VARLEN: tl.constexpr,
STATE_V_FIRST: tl.constexpr,
):
i_k, i_v, i_nh = tl.program_id(0), tl.program_id(1), tl.program_id(2).to(tl.int64)
i_n, i_h = i_nh // H, i_nh % 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
NT, NS = tl.cdiv(T, BT), tl.cdiv(T, BS)
boh = tl.load(split_offsets + i_n).to(tl.int64)
else:
bos, eos = i_n * T, i_n * T + T
NT, NS = tl.cdiv(T, BT), tl.cdiv(T, BS)
boh = i_n * NS
NTS = BS // BT
if USE_G_GAMMA:
# decay rate given the head index
b_gamma = tl.load(g_gamma + i_h)
b_g = b_gamma * (tl.arange(0, BT) + 1)
# [BK, BV] accumulator; STATE_V_FIRST only flips the stored state's HBM layout to [V, K], applied at the load/store below.
b_h = tl.zeros([BK, BV], dtype=tl.float32)
o_k = i_k * BK + tl.arange(0, BK)
o_v = i_v * BV + tl.arange(0, BV)
if USE_INITIAL_STATE:
if STATE_V_FIRST:
p_h0 = h0 + i_nh * K*V + o_v[:, None] * K + o_k[None, :]
b_h = tl.trans(tl.load(p_h0, mask=(o_v[:, None] < V) & (o_k[None, :] < K), other=0.0)).to(tl.float32)
else:
p_h0 = h0 + i_nh * K*V + o_k[:, None] * V + o_v[None, :]
b_h = tl.load(p_h0, mask=(o_k[:, None] < K) & (o_v[None, :] < V), other=0.0).to(tl.float32)
for i_t in range(NT):
i_s = i_t // NTS
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
p_k = k + (bos*H + i_h) * K + o_k[:, None] + o_t[None, :] * (H*K)
p_v = v + (bos*H + i_h) * V + o_t[:, None] * (H*V) + o_v[None, :]
o_h = ((boh + i_s) * H + i_h).to(tl.int64) * K*V
if STATE_V_FIRST:
p_h = h + o_h + o_v[:, None] * K + o_k[None, :]
m_h = (o_v[:, None] < V) & (o_k[None, :] < K)
else:
p_h = h + o_h + o_k[:, None] * V + o_v[None, :]
m_h = (o_k[:, None] < K) & (o_v[None, :] < V)
if i_t % NTS == 0:
tl.store(p_h, (tl.trans(b_h) if STATE_V_FIRST else b_h).to(p_h.dtype.element_ty), mask=m_h)
# [BK, BT]
b_k = tl.load(p_k, mask=(o_k[:, None] < K) & m_t[None, :], other=0.0)
# [BT, BV]
b_v = tl.load(p_v, mask=m_t[:, None] & (o_v < V)[None, :], other=0.0)
last_idx = min((i_t + 1) * BT, T) - 1
# scalar decay
if USE_G:
b_g_last = tl.load(g + bos * H + last_idx * H + i_h)
p_g = g + bos*H + (i_t * BT + tl.arange(0, BT)) * H + i_h
b_g = tl.load(p_g, mask=(i_t * BT + tl.arange(0, BT) < T), other=0.)
b_h *= exp2(b_g_last)
b_v = (b_v * exp2(b_g_last - b_g)[:, None]).to(b_v.dtype)
if USE_G_GAMMA:
b_g_last = b_gamma * min(BT, T - i_t * BT)
b_h *= exp2(b_g_last)
b_v = (b_v * exp2(b_g_last - b_g)[:, None]).to(b_v.dtype)
# vector decay, h = Diag(gk) @ h
if USE_GK:
p_gk = gk + (bos*H + i_h) * K + o_k[:, None] + o_t[None, :] * (H*K)
p_gk_last = gk + (bos + last_idx) * H*K + i_h * K + i_k * BK + tl.arange(0, BK)
b_gk_last = tl.load(p_gk_last, mask=(i_k * BK + tl.arange(0, BK) < K), other=0.)
b_gk = tl.load(p_gk, mask=(o_k[:, None] < K) & m_t[None, :], other=0.0)
b_h *= exp2(b_gk_last)[:, None]
b_k = (b_k * exp2(b_gk_last[:, None] - b_gk)).to(b_k.dtype)
# vector decay, h = h @ Diag(gv)
if USE_GV:
p_gv = gv + (bos*H + i_h) * V + o_t[:, None] * (H*V) + o_v[None, :]
p_gv_last = gv + (bos + last_idx) * H*V + i_h * V + i_v * BV + tl.arange(0, BV)
b_gv_last = tl.load(p_gv_last, mask=(i_v * BV + tl.arange(0, BV) < V), other=0.)
b_gv = tl.load(p_gv, mask=m_t[:, None] & (o_v < V)[None, :], other=0.0)
b_h *= exp2(b_gv_last)[None, :]
b_v = (b_v * exp2(b_gv_last[None, :] - b_gv)).to(b_v.dtype)
b_h += tl.dot(b_k, b_v)
if STORE_FINAL_STATE:
if STATE_V_FIRST:
p_ht = ht + i_nh * K*V + o_v[:, None] * K + o_k[None, :]
tl.store(p_ht, tl.trans(b_h).to(p_ht.dtype.element_ty), mask=(o_v[:, None] < V) & (o_k[None, :] < K))
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=(o_k[:, None] < K) & (o_v[None, :] < V))
@triton.heuristics({
'STORE_INITIAL_STATE_GRADIENT': lambda args: args['dh0'] is not None,
'USE_FINAL_STATE_GRADIENT': lambda args: args['dht'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@triton.autotune(
configs=[
triton.Config({'BK': BK, 'BV': BV}, num_warps=num_warps, num_stages=num_stages)
for BK in BKV_LIST
for BV in BKV_LIST
for num_warps in [1, 2, 4, 8]
for num_stages in [2, 3, 4]
],
key=['BT', 'USE_G', 'USE_GK', 'USE_GV', 'STATE_V_FIRST'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_bwd_kernel_dh(
q,
g,
g_gamma,
gk,
gv,
do,
dh,
dht,
dh0,
cu_seqlens,
split_offsets,
scale,
T,
HQ: tl.constexpr,
H: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BS: tl.constexpr,
BK: tl.constexpr,
BV: tl.constexpr,
NG: tl.constexpr,
USE_G: tl.constexpr,
USE_G_GAMMA: tl.constexpr,
USE_GK: tl.constexpr,
USE_GV: tl.constexpr,
STORE_INITIAL_STATE_GRADIENT: tl.constexpr,
USE_FINAL_STATE_GRADIENT: tl.constexpr,
IS_VARLEN: tl.constexpr,
STATE_V_FIRST: tl.constexpr,
):
i_k, i_v, i_nh = tl.program_id(0), tl.program_id(1), tl.program_id(2).to(tl.int64)
i_n, i_hq = i_nh // HQ, i_nh % HQ
i_h = i_hq // NG
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
NT = tl.cdiv(T, BT)
NS = tl.cdiv(T, BS)
boh = tl.load(split_offsets + i_n).to(tl.int64)
else:
bos, eos = i_n * T, i_n * T + T
NT = tl.cdiv(T, BT)
NS = tl.cdiv(T, BS)
boh = i_n * NS
if USE_G_GAMMA:
b_gamma = tl.load(g_gamma + i_h)
b_g = b_gamma * (tl.arange(0, BT) + 1)
# [BK, BV] accumulator; STATE_V_FIRST only flips the stored state's HBM layout to [V, K], applied at the load/store below.
b_dh = tl.zeros([BK, BV], dtype=tl.float32)
o_k = i_k * BK + tl.arange(0, BK)
o_v = i_v * BV + tl.arange(0, BV)
if USE_FINAL_STATE_GRADIENT:
if STATE_V_FIRST:
p_dht = dht + i_nh * K*V + o_v[:, None] * K + o_k[None, :]
b_dh += tl.trans(tl.load(p_dht, mask=(o_v[:, None] < V) & (o_k[None, :] < K), other=0.0)).to(tl.float32)
else:
p_dht = dht + i_nh * K*V + o_k[:, None] * V + o_v[None, :]
b_dh += tl.load(p_dht, mask=(o_k[:, None] < K) & (o_v[None, :] < V), other=0.0).to(tl.float32)
for i_t in range(NT - 1, -1, -1):
i_s = i_t // (BS // BT)
o_dh = ((boh + i_s) * H + i_h).to(tl.int64) * K*V
if STATE_V_FIRST:
p_dh = dh + o_dh + o_v[:, None] * K + o_k[None, :]
m_dh = (o_v[:, None] < V) & (o_k[None, :] < K)
else:
p_dh = dh + o_dh + o_k[:, None] * V + o_v[None, :]
m_dh = (o_k[:, None] < K) & (o_v[None, :] < V)
if i_t % (BS // BT) == 0:
tl.store(p_dh, (tl.trans(b_dh) if STATE_V_FIRST else b_dh).to(p_dh.dtype.element_ty), mask=m_dh)
last_idx = min(i_t * BT + BT, T) - 1
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
# [BK, BT]
p_q = q + (bos*HQ + i_hq) * K + o_k[:, None] + o_t[None, :] * (HQ*K)
p_do = do + (bos*HQ + i_hq) * V + o_t[:, None] * (HQ*V) + o_v[None, :]
b_q = tl.load(p_q, mask=(o_k[:, None] < K) & m_t[None, :], other=0.0)
b_q = (b_q * scale).to(b_q.dtype)
# [BT, BV]
b_do = tl.load(p_do, mask=m_t[:, None] & (o_v < V)[None, :], other=0.0)
if USE_G:
p_g = g + (bos + i_t * BT + tl.arange(0, BT)) * H + i_h
b_g_last = tl.load(g + (bos + last_idx) * H + i_h)
b_g = tl.load(p_g, mask=(i_t * BT + tl.arange(0, BT) < T), other=0.)
b_q = (b_q * exp2(b_g)[None, :]).to(b_q.dtype)
b_dh *= exp2(b_g_last)
if USE_G_GAMMA:
b_g_last = b_gamma * min(BT, T - i_t * BT)
b_q = (b_q * exp2(b_g)[None, :]).to(b_q.dtype)
b_dh *= exp2(b_g_last)
if USE_GK:
p_gk = gk + (bos*H + i_h) * K + o_k[:, None] + o_t[None, :] * (H*K)
p_gk_last = gk + (bos + last_idx) * H*K + i_h * K + i_k * BK + tl.arange(0, BK)
b_gk = tl.load(p_gk, mask=(o_k[:, None] < K) & m_t[None, :], other=0.0)
b_gk_last = tl.load(p_gk_last, mask=(i_k * BK + tl.arange(0, BK) < K), other=0.)
b_q = (b_q * exp2(b_gk)).to(b_q.dtype)
b_dh *= exp2(b_gk_last)[:, None]
if USE_GV:
p_gv = gv + (bos*H + i_h) * V + o_t[:, None] * (H*V) + o_v[None, :]
p_gv_last = gv + (bos + last_idx) * H*V + i_h * V + i_v * BV + tl.arange(0, BV)
b_gv = tl.load(p_gv, mask=m_t[:, None] & (o_v < V)[None, :], other=0.0)
b_gv_last = tl.load(p_gv_last, mask=(i_v * BV + tl.arange(0, BV) < V), other=0.)
b_do = (b_do * exp2(b_gv))
b_dh *= exp2(b_gv_last)[None, :]
b_dh += tl.dot(b_q, b_do.to(b_q.dtype))
if STORE_INITIAL_STATE_GRADIENT:
if STATE_V_FIRST:
p_dh0 = dh0 + i_nh * K*V + o_v[:, None] * K + o_k[None, :]
tl.store(p_dh0, tl.trans(b_dh).to(p_dh0.dtype.element_ty), mask=(o_v[:, None] < V) & (o_k[None, :] < K))
else:
p_dh0 = dh0 + i_nh * K*V + o_k[:, None] * V + o_v[None, :]
tl.store(p_dh0, b_dh.to(p_dh0.dtype.element_ty), mask=(o_k[:, None] < K) & (o_v[None, :] < V))
@dispatch('common')
def chunk_fwd_h(
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor | None = None,
g_gamma: torch.Tensor | None = None,
gk: torch.Tensor | None = None,
gv: torch.Tensor | None = None,
h0: torch.Tensor | None = None,
output_final_state: bool = False,
state_v_first: bool = False,
cu_seqlens: torch.Tensor | None = None,
chunk_size: int = 64,
split_size: int | None = None,
states_in_fp32: bool = False,
) -> tuple[torch.Tensor, torch.Tensor]:
B, T, H, K, V = *k.shape, v.shape[-1]
BT = chunk_size
BS = BT if split_size is None else split_size
assert BS % BT == 0, f"The `split_size` (got {BS}) must be a multiple of `chunk_size` {BT}"
# N: the actual number of sequences in the batch with either equal or variable lengths
if cu_seqlens is None:
N, NS, split_offsets = B, triton.cdiv(T, BS), None
else:
split_offsets = prepare_chunk_offsets(cu_seqlens, BS)
N, NS = len(cu_seqlens) - 1, split_offsets[-1].item()
# `state_v_first` stores the states in V-first `[V, K]` layout instead of `[K, V]`
state_shape = (V, K) if state_v_first else (K, V)
h = k.new_empty(B, NS, H, *state_shape, dtype=k.dtype if not states_in_fp32 else torch.float)
ht = k.new_empty(N, H, *state_shape, dtype=torch.float) if output_final_state else None
def grid(meta): return (triton.cdiv(K, meta['BK']), triton.cdiv(V, meta['BV']), N * H)
chunk_fwd_kernel_h[grid](
k=k,
v=v,
h=h,
g=g,
g_gamma=g_gamma,
gk=gk,
gv=gv,
h0=h0,
ht=ht,
cu_seqlens=cu_seqlens,
split_offsets=split_offsets,
T=T,
H=H,
K=K,
V=V,
BT=BT,
BS=BS,
USE_G=g is not None,
USE_G_GAMMA=g_gamma is not None,
USE_GK=gk is not None,
USE_GV=gv is not None,
STATE_V_FIRST=state_v_first,
)
return h, ht
@dispatch('common')
def chunk_bwd_dh(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
do: torch.Tensor,
h0: torch.Tensor,
dht: torch.Tensor,
scale: float,
g: torch.Tensor | None = None,
g_gamma: torch.Tensor | None = None,
gk: torch.Tensor | None = None,
gv: torch.Tensor | None = None,
state_v_first: bool = False,
cu_seqlens: torch.Tensor | None = None,
chunk_size: int = 64,
split_size: int | None = None,
states_in_fp32: bool = False,
) -> tuple[torch.Tensor, torch.Tensor]:
B, T, H, K, V = *k.shape, v.shape[-1]
HQ = q.shape[2]
BT = chunk_size
BS = BT if split_size is None else split_size
assert BS % BT == 0, f"The `split_size` (got {BS}) must be a multiple of `chunk_size` {BT}"
# N: the actual number of sequences in the batch with either equal or variable lengths
# NG: number of groups in GQA
if cu_seqlens is None:
N, NS, split_offsets = B, triton.cdiv(T, BS), None
else:
split_offsets = prepare_chunk_offsets(cu_seqlens, BS)
N, NS = len(cu_seqlens) - 1, split_offsets[-1].item()
NG = HQ // H
# `state_v_first` stores the states in V-first `[V, K]` layout instead of `[K, V]`
state_shape = (V, K) if state_v_first else (K, V)
dh = k.new_empty(B, NS, HQ, *state_shape, dtype=k.dtype if not states_in_fp32 else torch.float)
dh0 = torch.empty_like(h0, dtype=torch.float) if h0 is not None else None
def grid(meta): return (triton.cdiv(K, meta['BK']), triton.cdiv(V, meta['BV']), N * H)
chunk_bwd_kernel_dh[grid](
q=q,
g=g,
g_gamma=g_gamma,
gk=gk,
gv=gv,
do=do,
dh=dh,
dht=dht,
dh0=dh0,
cu_seqlens=cu_seqlens,
split_offsets=split_offsets,
scale=scale,
T=T,
HQ=HQ,
H=H,
K=K,
V=V,
BT=BT,
BS=BS,
NG=NG,
USE_G=g is not None,
USE_G_GAMMA=g_gamma is not None,
USE_GK=gk is not None,
USE_GV=gv is not None,
STATE_V_FIRST=state_v_first,
)
return dh, dh0
+110
View File
@@ -0,0 +1,110 @@
# 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
# Shared gate helpers reused across delta-rule family ops (KDA, GDN, ...).
import torch
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.utils import autocast_custom_bwd, autocast_custom_fwd, input_guard
@triton.jit
def fused_beta_sigmoid_fwd_kernel(
x,
y,
scale,
n_elements,
BLOCK_SIZE: tl.constexpr,
):
pid = tl.program_id(0).to(tl.int64)
offs = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE).to(tl.int64)
mask = offs < n_elements
b_x = tl.load(x + offs, mask=mask, other=0).to(tl.float32)
b_y = scale * tl.sigmoid(b_x)
tl.store(y + offs, b_y.to(y.dtype.element_ty), mask=mask)
@triton.jit
def fused_beta_sigmoid_bwd_kernel(
x,
dy,
dx,
scale,
n_elements,
BLOCK_SIZE: tl.constexpr,
):
pid = tl.program_id(0).to(tl.int64)
offs = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE).to(tl.int64)
mask = offs < n_elements
b_x = tl.load(x + offs, mask=mask, other=0).to(tl.float32)
b_dy = tl.load(dy + offs, mask=mask, other=0).to(tl.float32)
b_y = tl.sigmoid(b_x)
b_dx = b_dy * scale * b_y * (1.0 - b_y)
tl.store(dx + offs, b_dx.to(dx.dtype.element_ty), mask=mask)
_BETA_SIGMOID_BLOCK_SIZE = 2048
_BETA_SIGMOID_NUM_WARPS = 8
@dispatch('common')
def fused_beta_sigmoid_fwd(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
y = torch.empty_like(x, dtype=torch.float32)
n_elements = x.numel()
grid = (triton.cdiv(n_elements, _BETA_SIGMOID_BLOCK_SIZE),)
fused_beta_sigmoid_fwd_kernel[grid](
x,
y,
scale,
n_elements,
BLOCK_SIZE=_BETA_SIGMOID_BLOCK_SIZE,
num_warps=_BETA_SIGMOID_NUM_WARPS,
)
return y
@dispatch('common')
def fused_beta_sigmoid_bwd(x: torch.Tensor, dy: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
dx = torch.empty_like(x)
n_elements = x.numel()
grid = (triton.cdiv(n_elements, _BETA_SIGMOID_BLOCK_SIZE),)
fused_beta_sigmoid_bwd_kernel[grid](
x,
dy,
dx,
scale,
n_elements,
BLOCK_SIZE=_BETA_SIGMOID_BLOCK_SIZE,
num_warps=_BETA_SIGMOID_NUM_WARPS,
)
return dx
class BetaSigmoidFunction(torch.autograd.Function):
@staticmethod
@input_guard
@autocast_custom_fwd
def forward(ctx, x: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
y = fused_beta_sigmoid_fwd(x, scale)
ctx.save_for_backward(x)
ctx.scale = scale
return y
@staticmethod
@input_guard
@autocast_custom_bwd
def backward(ctx, dy: torch.Tensor):
(x,) = ctx.saved_tensors
dx = fused_beta_sigmoid_bwd(x, dy, ctx.scale)
return dx.type_as(x), None
def fused_beta_sigmoid(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
return BetaSigmoidFunction.apply(x, scale)
+13
View File
@@ -0,0 +1,13 @@
"""Context-parallel stubs. Pass ``cp_context=None`` (the default)."""
class FLACPContext:
cu_seqlens = None
cu_seqlens_cpu = None
def build_cp_context(*args, **kwargs):
raise RuntimeError("Context parallel is not included in the vendored KDA kernels")
__all__ = ["FLACPContext", "build_cp_context"]
+11
View File
@@ -0,0 +1,11 @@
"""Context-parallel hooks referenced by KDA fwd/bwd. Not implemented here."""
def _cp_unsupported(*args, **kwargs):
raise RuntimeError("Context parallel is not included in the vendored KDA kernels")
chunk_gated_delta_rule_fwd_h_pre_process = _cp_unsupported
compress_h0 = _cp_unsupported
chunk_gated_delta_rule_bwd_dhu_pre_process = _cp_unsupported
expand_h0 = _cp_unsupported
+1
View File
@@ -0,0 +1 @@
# Vendored GLA chunk output kernel used by KDA.
File diff suppressed because it is too large Load Diff
+7
View File
@@ -0,0 +1,7 @@
from .chunk import chunk_kda
from .fused_recurrent import fused_recurrent_kda
__all__ = [
"chunk_kda",
"fused_recurrent_kda",
]
+443
View File
@@ -0,0 +1,443 @@
# 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
# Related files are modified and supported by the Moonshot AI Team
import warnings
import torch
from kda._fla.modules.l2norm import l2norm_bwd, l2norm_fwd
from kda._fla.ops.backends import dispatch
from kda._fla.ops.common.gate import fused_beta_sigmoid, fused_beta_sigmoid_bwd
from kda._fla.ops.cp import FLACPContext
from kda._fla.ops.kda.chunk_bwd import chunk_kda_bwd
from kda._fla.ops.kda.chunk_fwd import chunk_kda_fwd
from kda._fla.ops.utils.index import prepare_chunk_indices
from kda._fla.utils import autocast_custom_bwd, autocast_custom_fwd, input_guard
class ChunkKDAFunction(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,
A_log: torch.Tensor,
dt_bias: torch.Tensor,
scale: float,
initial_state: torch.Tensor,
output_final_state: bool = False,
use_qk_l2norm_in_kernel: bool = False,
use_gate_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,
safe_gate: bool = False,
lower_bound: float | None = None,
chunk_size: int = 64,
disable_recompute: bool = False,
return_intermediate_states: bool = False,
cp_context: FLACPContext | None = None,
):
# Apply l2norm
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_input = g
(o, final_state, g_cumsum, Aqk, Akk, w, u, qg, kg, v_new, h, initial_state) = chunk_kda_fwd(
q=q,
k=k,
v=v,
g=g_input,
beta=beta,
scale=scale,
initial_state=initial_state,
output_final_state=output_final_state,
cu_seqlens=cu_seqlens,
cu_seqlens_cpu=cu_seqlens_cpu,
chunk_indices=chunk_indices,
safe_gate=safe_gate,
lower_bound=lower_bound,
use_gate_in_kernel=use_gate_in_kernel,
A_log=A_log,
dt_bias=dt_bias,
chunk_size=chunk_size,
disable_recompute=disable_recompute,
return_intermediate_states=return_intermediate_states,
cp_context=cp_context,
state_v_first=state_v_first,
)
if return_intermediate_states:
assert torch.is_inference_mode_enabled(), "return_intermediate_states is only allowed in inference mode"
assert disable_recompute is False, "return_intermediate_states must be used with disable_recompute=False"
return o.type_as(q), final_state, h
ctx.save_for_backward(
q, q_rstd, k, k_rstd, v, g_cumsum, g_input, beta_raw, beta, A_log, dt_bias, Aqk, Akk,
w, u, qg, kg, v_new, h,
initial_state, cu_seqlens, chunk_indices
)
ctx.chunk_size = chunk_size
ctx.safe_gate = safe_gate
ctx.scale = scale
ctx.lower_bound = lower_bound
ctx.use_qk_l2norm_in_kernel = use_qk_l2norm_in_kernel
ctx.use_gate_in_kernel = use_gate_in_kernel
ctx.use_beta_sigmoid_in_kernel = use_beta_sigmoid_in_kernel
ctx.allow_neg_eigval = allow_neg_eigval
ctx.disable_recompute = disable_recompute
ctx.cp_context = cp_context
ctx.state_v_first = state_v_first
return o.type_as(q), 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_cumsum, g_input, beta_raw, beta, A_log, dt_bias, Aqk, Akk,
w, u, qg, kg, v_new, h,
initial_state, cu_seqlens, chunk_indices) = (
ctx.saved_tensors
)
dq, dk, dv, db, dg, dh0, dA, dbias = chunk_kda_bwd(
q=q,
k=k,
v=v,
beta=beta,
Aqk=Aqk,
Akk=Akk,
scale=ctx.scale,
initial_state=initial_state,
do=do,
dht=dht,
g=g_cumsum,
g_org=g_input if ctx.use_gate_in_kernel else None,
state_v_first=ctx.state_v_first,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
chunk_size=ctx.chunk_size,
safe_gate=ctx.safe_gate,
lower_bound=ctx.lower_bound,
use_gate_in_kernel=ctx.use_gate_in_kernel,
A_log=A_log,
dt_bias=dt_bias,
disable_recompute=ctx.disable_recompute,
cp_context=ctx.cp_context,
w=w,
u=u,
qg=qg,
kg=kg,
v_new=v_new,
h=h,
)
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_input), db.to(beta_raw), dA, dbias, None, dh0,
None, None, None, None, None, None, None, None, None, None, None, None, None, None)
@dispatch('kda')
@torch.compiler.disable
def chunk_kda(
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_gate_in_kernel: bool = False,
use_beta_sigmoid_in_kernel: bool = False,
allow_neg_eigval: bool = False,
safe_gate: bool = False,
lower_bound: float | None = None,
disable_recompute: bool = False,
return_intermediate_states: bool = False,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
cu_seqlens_cpu: torch.LongTensor | None = None,
cp_context: FLACPContext = 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 (in log space!) of shape ``[B, T, HV, K]``.
When ``use_gate_in_kernel=False`` (default), ``g`` should be the pre-computed decay value.
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 KDA 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 KDA decay internally.
- If ``True``:
The passed ``g`` acts as the raw input for ``-exp(A_log) * softplus(g + dt_bias.view(HV, K))``.
Note that as part of the input arguments,
``A_log`` (shape ``[HV]``) and the optional ``dt_bias`` (shape ``[HV * K]``) should be provided.
When ``lower_bound`` is set, ``A_log`` may be ``None``,
in which case the gate is ``lower_bound * sigmoid(g + dt_bias)``.
- If ``False``, ``g`` is expected to be the pre-computed decay value.
Default: ``False``.
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``.
safe_gate (bool):
Whether to clamp the gate to ``[lower_bound, 0)`` and enable M=16 TensorCore
acceleration for higher throughput. Requires ``lower_bound`` to be set.
Default: ``False``.
lower_bound (Optional[float]):
Lower bound for the forget gate (in log space). When set together with
``safe_gate=True``, changes the gate activation from
``-exp(A_log) * softplus(g + dt_bias)`` to
``lower_bound * sigmoid(exp(A_log) * (g + dt_bias))``,
which naturally clamps the output to ``[lower_bound, 0)``.
Recommended value: ``-5`` (i.e., per-step decay ``exp(-5) ≈ 0.0067``).
Default: ``None``.
disable_recompute (bool):
Whether to disable gradient recomputation in the kernel. When ``True``, the kernel
will save all intermediate activations for backward pass, which is beneficial
for training small models at the cost of increased memory usage. Default: ``False``.
return_intermediate_states (bool):
If True, returns intermediate state ``h`` for inference scenarios (e.g., vLLM).
Must be used within ``torch.inference_mode()`` and will return a 3-tuple instead of 2-tuple.
This is not intended for training as it bypasses autograd. 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.
cu_seqlens_cpu (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:
- Normal mode (return_intermediate_states=False): A tuple (o, final_state)
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``.
- Inference mode (return_intermediate_states=True): A tuple (o, final_state, h)
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``.
h (torch.Tensor):
Intermediate states of shape ``[B, NT, HV, K, V]`` and dtype ``bfloat16``.
- For equal-length sequences: ``NT = ceil(T / chunk_size)``
- For variable-length sequences (cu_seqlens): B is always 1 (flattened),
NT is the total number of chunks across all sequences.
Examples::
>>> import torch
>>> import torch.nn.functional as F
>>> from einops import rearrange
>>> from fla.ops.kda import chunk_kda
# inputs with equal lengths (no GVA, HV == H)
>>> B, T, H, K, V = 4, 2048, 4, 512, 512
>>> q = torch.randn(B, T, H, K, dtype=torch.bfloat16, device='cuda')
>>> k = torch.randn(B, T, H, K, dtype=torch.bfloat16, device='cuda')
>>> v = torch.randn(B, T, H, V, dtype=torch.bfloat16, device='cuda')
>>> beta = torch.rand(B, T, H, dtype=torch.bfloat16, device='cuda')
>>> g = torch.rand(B, T, H, K, dtype=torch.bfloat16, device='cuda')
>>> h0 = torch.randn(B, H, K, V, dtype=torch.bfloat16, device='cuda')
>>> A_log = torch.randn(H, dtype=torch.float32, device='cuda')
>>> dt_bias = torch.randn(H * K, dtype=torch.float32, device='cuda')
>>> o, ht = chunk_kda(
q, k, v, g, beta,
A_log=A_log,
dt_bias=dt_bias,
use_qk_l2norm_in_kernel=True,
use_gate_in_kernel=True,
initial_state=h0,
output_final_state=True
)
# GVA mode (HV > H)
>>> HV = 8 # 2x more value heads than qk heads
>>> v = torch.randn(B, T, HV, V, dtype=torch.bfloat16, device='cuda')
>>> g = torch.rand(B, T, HV, K, dtype=torch.bfloat16, device='cuda')
>>> beta = torch.rand(B, T, HV, dtype=torch.bfloat16, device='cuda')
>>> h0 = torch.randn(B, HV, K, V, dtype=torch.bfloat16, device='cuda')
>>> A_log = torch.randn(HV, dtype=torch.float32, device='cuda')
>>> dt_bias = torch.randn(HV * K, dtype=torch.float32, device='cuda')
>>> o, ht = chunk_kda(
q, k, v, g, beta,
A_log=A_log,
dt_bias=dt_bias,
use_qk_l2norm_in_kernel=True,
use_gate_in_kernel=True,
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_kda(
q, k, v, g, beta,
A_log=A_log,
dt_bias=dt_bias,
use_qk_l2norm_in_kernel=True,
use_gate_in_kernel=True,
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 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"
# Override cu_seqlens and cu_seqlens_cpu with the ones from the context
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]}.",
)
if initial_state is not None:
assert initial_state.dtype == torch.float32, "initial_state must be in float32."
A_log, dt_bias = None, None
if use_gate_in_kernel:
A_log, dt_bias = kwargs.get("A_log"), kwargs.get("dt_bias")
if A_log is None and lower_bound is None:
raise ValueError("`A_log` must be provided when `use_gate_in_kernel=True` and `lower_bound` is not set.")
chunk_size = kwargs.pop("chunk_size", 64)
if chunk_size not in (32, 64):
raise ValueError(f"`chunk_size` must be either 32 or 64 for KDA, got {chunk_size}.")
if safe_gate and use_gate_in_kernel:
if lower_bound is None:
raise ValueError("`lower_bound` must be specified when `safe_gate=True` and `use_gate_in_kernel=True`.")
if not (-5 <= lower_bound < 0):
raise ValueError(f"`lower_bound` must be in the safe range [-5, 0), got {lower_bound}.")
if allow_neg_eigval and not use_beta_sigmoid_in_kernel:
raise ValueError("`allow_neg_eigval=True` requires `use_beta_sigmoid_in_kernel=True`.")
# Validate head dimensions for GVA
B, T, H, K, HV = *q.shape, v.shape[2]
assert q.shape == k.shape, f"q and k must have the same shape, got q={q.shape} vs k={k.shape}"
assert K <= 256, f"Currently we only support key headdim <=256 for KDA, got {K}."
assert HV % H == 0, (
f"For GVA, num_v_heads (HV={HV}) must be evenly divisible by num_qk_heads (H={H}), "
f"but got HV % H = {HV % H}"
)
assert g.shape == (B, T, HV, K), f"g must have shape [B, T, HV, K]={[B, T, HV, K]}, got {list(g.shape)}"
assert beta.shape == (B, T, HV), f"beta must have shape [B, T, HV]={[B, T, HV]}, got {list(beta.shape)}"
if scale is None:
scale = K ** -0.5
return ChunkKDAFunction.apply(
q,
k,
v,
g,
beta,
A_log,
dt_bias,
scale,
initial_state,
output_final_state,
use_qk_l2norm_in_kernel,
use_gate_in_kernel,
use_beta_sigmoid_in_kernel,
allow_neg_eigval,
state_v_first,
cu_seqlens,
cu_seqlens_cpu,
safe_gate,
lower_bound,
chunk_size,
disable_recompute,
return_intermediate_states,
cp_context,
)
+651
View File
@@ -0,0 +1,651 @@
# 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 kda._fla.ops.backends import dispatch
from kda._fla.ops.common.chunk_delta_h import (
chunk_gated_delta_rule_bwd_dhu,
chunk_gated_delta_rule_fwd_h,
)
from kda._fla.ops.cp import FLACPContext
from kda._fla.ops.cp.chunk_delta_h import (
chunk_gated_delta_rule_bwd_dhu_pre_process,
expand_h0,
)
from kda._fla.ops.kda.chunk_intra import chunk_kda_bwd_intra
from kda._fla.ops.kda.gate import kda_gate_bwd, kda_gate_chunk_cumsum
from kda._fla.ops.kda.wy_fast import recompute_w_u_fwd
from kda._fla.ops.utils import chunk_local_cumsum, prepare_chunk_indices
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.ops.utils.constant import RCP_LN2
from kda._fla.ops.utils.op import exp2
from kda._fla.utils import (
IS_NVIDIA_HOPPER,
IS_NVIDIA_SM100,
autotune_cache_kwargs,
check_shared_mem,
)
BK_LIST = [32, 64] if check_shared_mem() else [16, 32]
BV_LIST = [64, 128] if check_shared_mem("ampere") else [16, 32]
NUM_WARPS = [2, 4] if IS_NVIDIA_HOPPER else [2, 4, 8]
@triton.heuristics(
{
"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 NUM_WARPS
for num_stages in [2, 3, 4]
],
key=["H", "HV", "K", "V", "BT", "BK", "BV"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=["T"])
def chunk_kda_bwd_kernel_dAv(
q,
k,
v,
A,
do,
dv,
dA,
cu_seqlens,
chunk_indices,
scale,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BK: tl.constexpr,
BV: 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_hv = i_bh // HV, i_bh % HV
i_h = i_hv // (HV // 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
# offset calculation
q += (bos * H + i_h) * K
k += (bos * H + i_h) * K
v += (bos * HV + i_hv) * V
do += (bos * HV + i_hv) * V
dv += (bos * HV + i_hv) * V
dA += (bos * HV + i_hv) * BT
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
o_A = tl.arange(0, BT)
m_AT = (o_A[:, None] < BT) & m_t[None, :]
p_A = A + (bos * HV + i_hv) * BT + o_A[:, None] + o_t[None, :] * (HV * BT)
b_A = tl.load(p_A, mask=m_AT, other=0.0)
m_A = (o_t[:, None] <= o_t[None, :]) & (m_t[:, None] & m_t)
b_A = tl.where(m_A, b_A, 0).to(do.dtype.element_ty)
b_dA = tl.zeros([BT, BT], dtype=tl.float32)
for i_v in range(tl.cdiv(V, BV)):
o_v = i_v * BV + tl.arange(0, BV)
m_v = o_v < V
m_vT = m_v[:, None] & m_t[None, :]
m_tv = m_t[:, None] & m_v[None, :]
p_v = v + o_v[:, None] + o_t[None, :] * (HV * V)
p_do = do + o_t[:, None] * (HV * V) + o_v[None, :]
p_dv = dv + o_t[:, None] * (HV * V) + o_v[None, :]
# [BV, BT]
b_v = tl.load(p_v, mask=m_vT, other=0.0)
# [BT, BV]
b_do = tl.load(p_do, mask=m_tv, other=0.0)
# [BT, BT]
b_dA += tl.dot(b_do, b_v)
# [BT, BV]
b_dv = tl.dot(b_A.to(b_do.dtype), b_do)
tl.store(p_dv, b_dv.to(p_dv.dtype.element_ty), mask=m_tv)
m_dA = m_t[:, None] & (o_A[None, :] < BT)
p_dA = dA + o_t[:, None] * (HV * BT) + o_A[None, :]
b_dA = tl.where(o_t[:, None] >= o_t, b_dA * scale, 0.0)
tl.store(p_dA, b_dA.to(p_dA.dtype.element_ty), mask=m_dA)
@triton.heuristics(
{
"IS_VARLEN": lambda args: args["cu_seqlens"] is not None,
}
)
@fla_cache_autotune(
configs=[
triton.Config({"BK": BK, "BV": BV}, num_warps=num_warps, num_stages=num_stages)
for BK in BK_LIST
for BV in BV_LIST
for num_warps in NUM_WARPS
for num_stages in [2, 3, 4]
if not (IS_NVIDIA_HOPPER and BK == 32 and num_warps == 4)
if not (IS_NVIDIA_SM100 and BK == 32 and num_warps != 2)
],
key=["BT", "HV", "STATE_V_FIRST"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=["T"])
def chunk_kda_bwd_kernel_wy_dqkg_fused(
q,
k,
v,
v_new,
g,
beta,
A,
h,
do,
dh,
dq,
dk,
dv,
dv2,
dg,
db,
dA,
cu_seqlens,
chunk_indices,
scale,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BK: tl.constexpr,
BV: tl.constexpr,
STATE_V_FIRST: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_t, i_bh = tl.program_id(0).to(tl.int64), tl.program_id(1)
i_b, i_hv = i_bh // HV, i_bh % HV
i_h = i_hv // (HV // H)
if IS_VARLEN:
i_tg = i_t.to(tl.int64)
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).to(tl.int32)
NT = tl.cdiv(T, BT)
else:
NT = tl.cdiv(T, BT)
i_tg = (i_b * NT + i_t).to(tl.int64)
bos, eos = (i_b * T).to(tl.int64), (i_b * T + T).to(tl.int64)
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
m_last = o_t == min(T, i_t * BT + BT) - 1
q += (bos * H + i_h) * K
k += (bos * H + i_h) * K
v += (bos * HV + i_hv) * V
v_new += (bos * HV + i_hv) * V
g += (bos * HV + i_hv) * K
beta += bos * HV + i_hv
A += (bos * HV + i_hv) * BT
h += (i_tg * HV + i_hv) * K * V
do += (bos * HV + i_hv) * V
dh += (i_tg * HV + i_hv) * K * V
dq += (bos * HV + i_hv) * K
dk += (bos * HV + i_hv) * K
dv += (bos * HV + i_hv) * V
dv2 += (bos * HV + i_hv) * V
dg += (bos * HV + i_hv) * K
db += bos * HV + i_hv
dA += (bos * HV + i_hv) * BT
p_beta = beta + o_t * HV
b_beta = tl.load(p_beta, mask=m_t, other=0.0)
o_A = tl.arange(0, BT)
m_AT = (o_A[:, None] < BT) & m_t[None, :]
p_A = A + o_A[:, None] + o_t[None, :] * (HV * BT)
b_A = tl.load(p_A, mask=m_AT, other=0.0)
b_dA = tl.zeros([BT, BT], dtype=tl.float32)
b_db = 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 = o_k < K
m_tk = m_t[:, None] & m_k[None, :]
p_k = k + o_t[:, None] * (H * K) + o_k[None, :]
p_g = g + o_t[:, None] * (HV * K) + o_k[None, :]
b_k = tl.load(p_k, mask=m_tk, other=0.0)
b_g = tl.load(p_g, mask=m_tk, other=0.0).to(tl.float32)
p_gn = g + (min(T, i_t * BT + BT) - 1).to(tl.int64) * HV * K + o_k
b_gn = tl.load(p_gn, mask=m_k, other=0).to(tl.float32)
b_dq = tl.zeros([BT, BK], dtype=tl.float32)
b_dk = tl.zeros([BT, BK], dtype=tl.float32)
b_dw = tl.zeros([BT, BK], dtype=tl.float32)
b_dgk = tl.zeros([BK], dtype=tl.float32)
for i_v in range(tl.cdiv(V, BV)):
o_v = i_v * BV + tl.arange(0, BV)
m_tv = m_t[:, None] & (o_v[None, :] < V)
m_h = (o_v[:, None] < V) & m_k[None, :]
p_v_new = v_new + o_t[:, None] * (HV * V) + o_v[None, :]
p_do = do + o_t[:, None] * (HV * V) + o_v[None, :]
if STATE_V_FIRST:
p_h = h + o_v[:, None] * K + o_k[None, :]
p_dh = dh + o_v[:, None] * K + o_k[None, :]
else:
p_h = h + o_v[:, None] + o_k[None, :] * V
p_dh = dh + o_v[:, None] + o_k[None, :] * V
p_dv = dv + o_t[:, None] * (HV * V) + o_v[None, :]
# [BT, BV]
b_v_new = tl.load(p_v_new, mask=m_tv, other=0.0)
b_do = tl.load(p_do, mask=m_tv, other=0.0)
# [BV, BK]
b_h = tl.load(p_h, mask=m_h, other=0.0)
b_dh = tl.load(p_dh, mask=m_h, other=0.0)
# [BT, BV]
b_dv = tl.load(p_dv, mask=m_tv, other=0.0)
b_dgk += tl.sum(b_h * b_dh, axis=0)
b_dq += tl.dot(b_do, b_h.to(b_do.dtype))
b_dk += tl.dot(b_v_new, b_dh.to(b_v_new.dtype))
b_dw += tl.dot(b_dv.to(b_v_new.dtype), b_h.to(b_v_new.dtype))
tl.debug_barrier() # DO NOT REMOVE THIS LINE!
if i_k == 0:
p_v = v + o_t[:, None] * (HV * V) + o_v[None, :]
p_dv2 = dv2 + o_t[:, None] * (HV * V) + o_v[None, :]
b_v = tl.load(p_v, mask=m_tv, other=0.0)
b_dA += tl.dot(b_dv, tl.trans(b_v))
b_dvb = tl.dot(b_A, b_dv)
b_dv2 = b_dvb * b_beta[:, None]
b_db += tl.sum(b_dvb * b_v, 1)
tl.store(p_dv2, b_dv2.to(p_dv2.dtype.element_ty), mask=m_tv)
b_gk_exp = exp2(b_g)
b_gb = b_gk_exp * b_beta[:, None]
b_dgk *= exp2(b_gn)
b_dq = b_dq * b_gk_exp * scale
b_dk = b_dk * tl.where(m_t[:, None], exp2(b_gn[None, :] - b_g), 0)
b_kg = b_k * b_gk_exp
b_dw = -b_dw.to(b_A.dtype)
b_dA += tl.dot(b_dw, tl.trans(b_kg.to(b_A.dtype)))
b_dkgb = tl.dot(b_A, b_dw)
b_db += tl.sum(b_dkgb * b_kg, 1)
p_q = q + o_t[:, None] * (H * K) + o_k[None, :]
b_q = tl.load(p_q, mask=m_tk, other=0.0)
b_kdk = b_k * b_dk
b_dgk += tl.sum(b_kdk, axis=0)
b_dg = (
b_q * b_dq
- b_kdk
+ m_last[:, None] * b_dgk
+ b_kg * b_dkgb * b_beta[:, None]
)
b_dk = b_dk + b_dkgb * b_gb
p_dq = dq + o_t[:, None] * (HV * K) + o_k[None, :]
p_dk = dk + o_t[:, None] * (HV * K) + o_k[None, :]
p_dg = dg + o_t[:, None] * (HV * K) + o_k[None, :]
tl.store(p_dq, b_dq.to(p_dq.dtype.element_ty), mask=m_tk)
tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), mask=m_tk)
tl.store(p_dg, b_dg.to(p_dg.dtype.element_ty), mask=m_tk)
m_A = (o_t[:, None] > o_t[None, :]) & (m_t[:, None] & m_t)
b_dA = tl.where(m_A, b_dA * b_beta[None, :], 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))
b_dA = tl.where(m_A, -b_dA, 0)
m_dA = m_t[:, None] & (o_A[None, :] < BT)
p_dA = dA + o_t[:, None] * (HV * BT) + o_A[None, :]
p_db = db + o_t * HV
tl.store(p_dA, b_dA.to(p_dA.dtype.element_ty), mask=m_dA)
tl.store(p_db, b_db.to(p_db.dtype.element_ty), mask=m_t)
@dispatch("kda")
def chunk_kda_bwd_dAv(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
do: torch.Tensor,
A: torch.Tensor | None = None,
scale: float = None,
cu_seqlens: torch.LongTensor | None = None,
chunk_size: int = 64,
chunk_indices: torch.LongTensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
B, T, H, K, HV, V = *k.shape, do.shape[2], do.shape[-1]
BT = chunk_size
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
# H100 can have larger block size
if check_shared_mem("hopper", k.device.index):
CONST_TILING = 128
elif check_shared_mem:
CONST_TILING = 64
else:
CONST_TILING = 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)
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
dA = v.new_empty(B, T, HV, BT, dtype=torch.float)
dv = torch.empty_like(do)
grid = (NT, B * HV)
chunk_kda_bwd_kernel_dAv[grid](
q=q,
k=k,
v=v,
A=A,
do=do,
dv=dv,
dA=dA,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
scale=scale,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BT=BT,
BK=BK,
BV=BV,
)
return dA, dv
@dispatch("kda")
def chunk_kda_bwd_wy_dqkg_fused(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
v_new: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
A: torch.Tensor,
h: torch.Tensor,
do: torch.Tensor,
dh: torch.Tensor,
dv: torch.Tensor,
scale: float | None = None,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
chunk_size: int = 64,
chunk_indices: torch.LongTensor | None = None,
):
B, T, H, K, HV, V = *k.shape, v.shape[2], v.shape[-1]
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)
# dq, dk are allocated at HV dimension; caller reduces to H if GVA
dq = g.new_empty(B, T, HV, K, dtype=torch.float)
dk = g.new_empty(B, T, HV, K, dtype=torch.float)
dv2 = torch.empty_like(v)
dg = torch.empty_like(g, dtype=torch.float)
db = torch.empty_like(beta, dtype=torch.float)
dA = torch.empty_like(A, dtype=torch.float)
grid = (NT, B * HV)
chunk_kda_bwd_kernel_wy_dqkg_fused[grid](
q=q,
k=k,
v=v,
v_new=v_new,
g=g,
beta=beta,
A=A,
h=h,
do=do,
dh=dh,
dq=dq,
dk=dk,
dv=dv,
dv2=dv2,
dg=dg,
db=db,
dA=dA,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
scale=scale,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BT=BT,
STATE_V_FIRST=state_v_first,
)
dv = dv2
return dq, dk, dv, db, dg, dA
def chunk_kda_bwd(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
beta: torch.Tensor,
Aqk: torch.Tensor,
Akk: torch.Tensor,
scale: float,
initial_state: torch.Tensor,
do: torch.Tensor,
dht: torch.Tensor,
g: torch.Tensor | None = None,
g_org: torch.Tensor | None = None,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
chunk_indices: torch.LongTensor | None = None,
chunk_size: int = 64,
safe_gate: bool = False,
lower_bound: float | None = None,
use_gate_in_kernel: bool = False,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
disable_recompute: bool = False,
cp_context: FLACPContext | None = None,
**kwargs,
):
H, HV = q.shape[2], v.shape[2]
G = HV // H
if disable_recompute is False:
if use_gate_in_kernel:
g = kda_gate_chunk_cumsum(
g=g_org,
A_log=A_log,
dt_bias=dt_bias,
scale=RCP_LN2,
chunk_size=chunk_size,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
lower_bound=lower_bound,
)
w, u, qg, kg = recompute_w_u_fwd(
q=q,
k=k,
v=v,
beta=beta,
A=Akk,
gk=g,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
)
if cp_context is not None:
# Restore the full initial_state tensor from the compressed version.
# Only the first sequence's state is non-zero as it's the only one that could be cross-rank.
initial_state = expand_h0(initial_state, context=cp_context)
h, v_new, _ = chunk_gated_delta_rule_fwd_h(
k=kg,
w=w,
u=u,
gk=g,
initial_state=initial_state,
output_final_state=False,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
chunk_size=chunk_size,
state_v_first=state_v_first,
)
else:
w, u, qg, kg, v_new, h = (
kwargs["w"],
kwargs["u"],
kwargs["qg"],
kwargs["kg"],
kwargs["v_new"],
kwargs["h"],
)
if cp_context is not None:
# Restore the full initial_state tensor from the compressed version.
# Only the first sequence's state is non-zero as it's the only one that could be cross-rank.
initial_state = expand_h0(initial_state, context=cp_context)
# dAqk = do @ v.T
# dv = A @ do
dAqk, dv = chunk_kda_bwd_dAv(
q=q,
k=k,
v=v_new,
do=do,
A=Aqk,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_size=chunk_size,
chunk_indices=chunk_indices,
)
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=qg,
k=kg,
w=w,
do=do,
dv=dv,
gk=g,
scale=scale,
cu_seqlens=cu_seqlens,
dht=dht,
initial_state=initial_state,
context=cp_context,
chunk_size=chunk_size,
state_v_first=state_v_first,
)
dh, dh0, dv = chunk_gated_delta_rule_bwd_dhu(
q=qg,
k=kg,
w=w,
gk=g,
h0=initial_state,
dht=dht,
do=do,
dv=dv,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_size=chunk_size,
chunk_indices=chunk_indices,
state_v_first=state_v_first,
)
dq, dk, dv, db, dg, dAkk = chunk_kda_bwd_wy_dqkg_fused(
q=q,
k=k,
v=v,
v_new=v_new,
g=g,
beta=beta,
A=Akk,
h=h,
do=do,
dh=dh,
dv=dv,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_size=chunk_size,
chunk_indices=chunk_indices,
state_v_first=state_v_first,
)
dq, dk, db, dg = chunk_kda_bwd_intra(
q=q,
k=k,
g=g,
beta=beta,
dAqk=dAqk,
dAkk=dAkk,
dq=dq,
dk=dk,
db=db,
dg=dg,
cu_seqlens=cu_seqlens,
chunk_size=chunk_size,
chunk_indices=chunk_indices,
safe_gate=safe_gate,
)
# For GVA, reduce dq and dk from [B, T, HV, K] back to [B, T, H, K]
if HV > H:
dq = dq.view(*dq.shape[:2], H, G, dq.shape[-1]).sum(dim=3)
dk = dk.view(*dk.shape[:2], H, G, dk.shape[-1]).sum(dim=3)
dA, dbias = None, None
dg = chunk_local_cumsum(
dg,
chunk_size=chunk_size,
reverse=True,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
)
if use_gate_in_kernel:
dg, dA, dbias = kda_gate_bwd(
g=g_org, A_log=A_log, dt_bias=dt_bias, dyg=dg, lower_bound=lower_bound
)
return dq, dk, dv, db, dg, dh0, dA, dbias
+134
View File
@@ -0,0 +1,134 @@
# 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
from kda._fla.ops.common.chunk_delta_h import chunk_gated_delta_rule_fwd_h
from kda._fla.ops.cp import FLACPContext
from kda._fla.ops.cp.chunk_delta_h import chunk_gated_delta_rule_fwd_h_pre_process, compress_h0
from kda._fla.ops.gla.chunk import chunk_gla_fwd_o_gk
from kda._fla.ops.kda.chunk_intra import chunk_kda_fwd_intra
from kda._fla.ops.kda.gate import kda_gate_chunk_cumsum
from kda._fla.ops.utils import chunk_local_cumsum
from kda._fla.ops.utils.constant import RCP_LN2
def chunk_kda_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,
cu_seqlens_cpu: torch.LongTensor | None = None,
chunk_indices: torch.LongTensor | None = None,
chunk_size: int = 64,
safe_gate: bool = False,
lower_bound: float | None = None,
use_gate_in_kernel: bool = False,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
disable_recompute: bool = False,
return_intermediate_states: bool = False,
cp_context: FLACPContext | None = None,
):
# Apply gate activation
g_org = None
if use_gate_in_kernel:
g_org = g
g = kda_gate_chunk_cumsum(
g=g_org,
A_log=A_log,
dt_bias=dt_bias,
scale=RCP_LN2,
chunk_size=chunk_size,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
lower_bound=lower_bound,
)
else:
g = chunk_local_cumsum(
g=g,
scale=RCP_LN2,
chunk_size=chunk_size,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices
)
# qg = None if disable_recompute is False
w, u, qg, kg, Aqk, Akk = chunk_kda_fwd_intra(
q=q,
k=k,
v=v,
gk=g,
beta=beta,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_size=chunk_size,
chunk_indices=chunk_indices,
safe_gate=safe_gate,
disable_recompute=disable_recompute
)
if cp_context is not None:
initial_state = chunk_gated_delta_rule_fwd_h_pre_process(
k=kg,
w=w,
u=u,
gk=g,
cu_seqlens=cu_seqlens,
initial_state=initial_state,
context=cp_context,
chunk_size=chunk_size,
state_v_first=state_v_first,
)
h, v_new, final_state = chunk_gated_delta_rule_fwd_h(
k=kg,
w=w,
u=u,
gk=g,
initial_state=initial_state,
output_final_state=output_final_state,
cu_seqlens=cu_seqlens,
cu_seqlens_cpu=cu_seqlens_cpu,
chunk_indices=chunk_indices,
chunk_size=chunk_size,
state_v_first=state_v_first,
)
if cp_context is not None:
# In Context Parallel (CP) mode, global initial states are not supported at the entry point.
# The `initial_state` here is computed internally via inter-rank communication.
# Since only the first sequence in the local batch can be a continuation of a cross-rank sequence,
# only the first state in the tensor is relevant. We compress it to optimize memory for `save_for_backward`.
initial_state = compress_h0(initial_state, context=cp_context)
o = chunk_gla_fwd_o_gk(
q=q,
v=v_new,
g=g,
A=Aqk,
h=h,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_size=chunk_size,
chunk_indices=chunk_indices,
state_v_first=state_v_first,
)
if disable_recompute is False:
# Delete to save memory
w, u, qg, kg, v_new = None, None, None, None, None
if not return_intermediate_states:
h = None
if use_gate_in_kernel:
g = None
return o, final_state, g, Aqk, Akk, w, u, qg, kg, v_new, h, initial_state
+962
View File
@@ -0,0 +1,962 @@
# 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 kda._fla.ops.backends import dispatch
from kda._fla.ops.kda.chunk_intra_token_parallel import chunk_kda_fwd_intra_token_parallel
from kda._fla.ops.kda.wy_fast import recompute_w_u_fwd
from kda._fla.ops.utils import prepare_chunk_indices
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.ops.utils.op import exp2, gather
from kda._fla.utils import IS_GATHER_SUPPORTED, 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')
################################################################################
# Fused inter + solve_tril kernel: compute off-diagonal Akk and solve in one pass
################################################################################
@triton.heuristics({
'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", "BT", "BC", "NC"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_kda_fwd_kernel_inter_solve_fused(
q,
k,
g,
beta,
Aqk,
Akkd,
Akk,
scale,
cu_seqlens,
chunk_indices,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
BT: tl.constexpr,
BC: tl.constexpr,
NC: tl.constexpr,
BK: tl.constexpr,
IS_VARLEN: tl.constexpr,
USE_SAFE_GATE: tl.constexpr,
):
"""
Fused kernel: compute inter-subchunk Akk + solve_tril in one pass.
Prerequisite: token_parallel has already computed diagonal Akk blocks in Akkd.
This kernel:
1. Computes off-diagonal Aqk blocks -> writes to global
2. Computes off-diagonal Akk blocks -> keeps in registers
3. Loads diagonal Akk blocks from Akkd (fp32)
4. Does forward substitution on diagonals
5. Computes merged Akk_inv
6. Writes Akk_inv to Akk
"""
i_t, i_bh = tl.program_id(0).to(tl.int64), tl.program_id(1).to(tl.int64)
i_b, i_hv = i_bh // HV, i_bh % HV
i_h = i_hv // (HV // 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
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
q += (bos * H + i_h) * K
k += (bos * H + i_h) * K
g += (bos * HV + i_hv) * K
Aqk += (bos * HV + i_hv) * BT
Akk += (bos * HV + i_hv) * BT
Akkd += (bos * HV + i_hv) * BC
o_i = tl.arange(0, BC)
m_tc1 = (i_tc1 + o_i) < T
m_tc2 = (i_tc2 + o_i) < T
m_tc3 = (i_tc3 + o_i) < T
o_c0 = i_tc0 + o_i
o_c1 = i_tc1 + o_i
o_c2 = i_tc2 + o_i
o_c3 = i_tc3 + o_i
m_tc0 = o_c0 < T
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)
b_Aqk10 = tl.zeros([BC, BC], dtype=tl.float32)
b_Akk10 = tl.zeros([BC, BC], dtype=tl.float32)
b_Aqk20 = tl.zeros([BC, BC], dtype=tl.float32)
b_Akk20 = tl.zeros([BC, BC], dtype=tl.float32)
b_Aqk21 = tl.zeros([BC, BC], dtype=tl.float32)
b_Akk21 = tl.zeros([BC, BC], dtype=tl.float32)
b_Aqk30 = tl.zeros([BC, BC], dtype=tl.float32)
b_Akk30 = tl.zeros([BC, BC], dtype=tl.float32)
b_Aqk31 = tl.zeros([BC, BC], dtype=tl.float32)
b_Akk31 = tl.zeros([BC, BC], dtype=tl.float32)
b_Aqk32 = tl.zeros([BC, BC], dtype=tl.float32)
b_Akk32 = tl.zeros([BC, BC], dtype=tl.float32)
################################################################################
# off-diagonal blocks
################################################################################
for i_k in range(tl.cdiv(K, BK)):
o_k = i_k * BK + tl.arange(0, BK)
m_k = o_k < K
m_ck0 = m_tc0[:, None] & m_k[None, :]
p_k0 = k + o_c0[:, None] * (H*K) + o_k[None, :]
p_g0 = g + o_c0[:, None] * (HV*K) + o_k[None, :]
b_k0 = tl.load(p_k0, mask=m_ck0, other=0.0).to(tl.float32)
b_g0 = tl.load(p_g0, mask=m_ck0, other=0.0).to(tl.float32)
if i_tc1 < T:
m_ck1 = m_tc1[:, None] & m_k[None, :]
p_q1 = q + o_c1[:, None] * (H*K) + o_k[None, :]
p_k1 = k + o_c1[:, None] * (H*K) + o_k[None, :]
p_g1 = g + o_c1[:, None] * (HV*K) + o_k[None, :]
# [BC, BK]
b_q1 = tl.load(p_q1, mask=m_ck1, other=0.0).to(tl.float32)
b_k1 = tl.load(p_k1, mask=m_ck1, other=0.0).to(tl.float32)
b_g1 = tl.load(p_g1, mask=m_ck1, other=0.0).to(tl.float32)
# [BK]
b_gn1 = tl.load(g + i_tc1 * HV*K + o_k, mask=m_k, other=0).to(tl.float32)
# [BC, BK]
b_gqn = tl.where(m_tc1[:, None], exp2(b_g1 - b_gn1[None, :]), 0)
# [BK, BC]
b_kgt = tl.trans(b_k0 * exp2(b_gn1[None, :] - b_g0))
# [BC, BC]
b_Aqk10 += tl.dot(b_q1 * b_gqn, b_kgt)
b_Akk10 += tl.dot(b_k1 * b_gqn, b_kgt)
if NC >= 3 and i_tc2 < T:
m_ck2 = m_tc2[:, None] & m_k[None, :]
p_q2 = q + o_c2[:, None] * (H*K) + o_k[None, :]
p_k2 = k + o_c2[:, None] * (H*K) + o_k[None, :]
p_g2 = g + o_c2[:, None] * (HV*K) + o_k[None, :]
# [BC, BK]
b_q2 = tl.load(p_q2, mask=m_ck2, other=0.0).to(tl.float32)
b_k2 = tl.load(p_k2, mask=m_ck2, other=0.0).to(tl.float32)
b_g2 = tl.load(p_g2, mask=m_ck2, other=0.0).to(tl.float32)
# [BK]
b_gn2 = tl.load(g + i_tc2 * HV*K + o_k, mask=m_k, other=0).to(tl.float32)
# [BC, BK]
b_gqn2 = tl.where(m_tc2[:, None], exp2(b_g2 - b_gn2[None, :]), 0)
b_qg2 = b_q2 * b_gqn2
b_kg2 = b_k2 * b_gqn2
# [BK, BC]
b_kgt = tl.trans(b_k0 * exp2(b_gn2[None, :] - b_g0))
b_Aqk20 += tl.dot(b_qg2, b_kgt)
b_Akk20 += tl.dot(b_kg2, b_kgt)
# [BC, BC]
b_kgt = tl.trans(b_k1 * exp2(b_gn2[None, :] - b_g1))
# [BC, BC]
b_Aqk21 += tl.dot(b_qg2, b_kgt)
b_Akk21 += tl.dot(b_kg2, b_kgt)
if NC >= 4 and i_tc3 < T:
m_ck3 = m_tc3[:, None] & m_k[None, :]
p_q3 = q + o_c3[:, None] * (H*K) + o_k[None, :]
p_k3 = k + o_c3[:, None] * (H*K) + o_k[None, :]
p_g3 = g + o_c3[:, None] * (HV*K) + o_k[None, :]
# [BC, BK]
b_q3 = tl.load(p_q3, mask=m_ck3, other=0.0).to(tl.float32)
b_k3 = tl.load(p_k3, mask=m_ck3, other=0.0).to(tl.float32)
b_g3 = tl.load(p_g3, mask=m_ck3, other=0.0).to(tl.float32)
# [BK]
b_gn3 = tl.load(g + i_tc3 * HV*K + o_k, mask=m_k, other=0).to(tl.float32)
# [BC, BK]
b_gqn3 = tl.where(m_tc3[:, None], exp2(b_g3 - b_gn3[None, :]), 0)
b_qg3 = b_q3 * b_gqn3
b_kg3 = b_k3 * b_gqn3
# [BK, BC]
b_kgt = tl.trans(b_k0 * exp2(b_gn3[None, :] - b_g0))
# [BC, BC]
b_Aqk30 += tl.dot(b_qg3, b_kgt)
b_Akk30 += tl.dot(b_kg3, b_kgt)
# [BK, BC]
b_kgt = tl.trans(b_k1 * exp2(b_gn3[None, :] - b_g1))
# [BC, BC]
b_Aqk31 += tl.dot(b_qg3, b_kgt)
b_Akk31 += tl.dot(b_kg3, b_kgt)
# [BK, BC]
b_kgt = tl.trans(b_k2 * exp2(b_gn3[None, :] - b_g2))
# [BC, BC]
b_Aqk32 += tl.dot(b_qg3, b_kgt)
b_Akk32 += tl.dot(b_kg3, b_kgt)
################################################################################
# save off-diagonal Aqk blocks and prepare Akk
################################################################################
if i_tc1 < T:
p_Aqk10 = Aqk + o_c1[:, None] * (HV*BT) + o_i[None, :]
tl.store(p_Aqk10, (b_Aqk10 * scale).to(Aqk.dtype.element_ty), mask=m_A1)
p_b1 = beta + bos * HV + i_hv + o_c1 * HV
b_b1 = tl.load(p_b1, mask=m_tc1, other=0.0).to(tl.float32)
b_Akk10 = b_Akk10 * b_b1[:, None]
if NC >= 3 and i_tc2 < T:
p_Aqk20 = Aqk + o_c2[:, None] * (HV*BT) + o_i[None, :]
p_Aqk21 = Aqk + o_c2[:, None] * (HV*BT) + (o_i + BC)[None, :]
tl.store(p_Aqk20, (b_Aqk20 * scale).to(Aqk.dtype.element_ty), mask=m_A2)
tl.store(p_Aqk21, (b_Aqk21 * scale).to(Aqk.dtype.element_ty), mask=m_A2)
p_b2 = beta + bos * HV + i_hv + o_c2 * HV
b_b2 = tl.load(p_b2, mask=m_tc2, other=0.0).to(tl.float32)
b_Akk20 = b_Akk20 * b_b2[:, None]
b_Akk21 = b_Akk21 * b_b2[:, None]
if NC >= 4 and i_tc3 < T:
p_Aqk30 = Aqk + o_c3[:, None] * (HV*BT) + o_i[None, :]
p_Aqk31 = Aqk + o_c3[:, None] * (HV*BT) + (o_i + BC)[None, :]
p_Aqk32 = Aqk + o_c3[:, None] * (HV*BT) + (o_i + 2*BC)[None, :]
tl.store(p_Aqk30, (b_Aqk30 * scale).to(Aqk.dtype.element_ty), mask=m_A3)
tl.store(p_Aqk31, (b_Aqk31 * scale).to(Aqk.dtype.element_ty), mask=m_A3)
tl.store(p_Aqk32, (b_Aqk32 * scale).to(Aqk.dtype.element_ty), mask=m_A3)
p_b3 = beta + bos * HV + i_hv + o_c3 * HV
b_b3 = tl.load(p_b3, mask=m_tc3, other=0.0).to(tl.float32)
b_Akk30 = b_Akk30 * b_b3[:, None]
b_Akk31 = b_Akk31 * b_b3[:, None]
b_Akk32 = b_Akk32 * b_b3[:, None]
p_Akk00 = Akkd + o_c0[:, None] * (HV*BC) + o_i[None, :]
p_Akk11 = Akkd + o_c1[:, None] * (HV*BC) + o_i[None, :]
b_Ai00 = tl.load(p_Akk00, mask=m_A0, other=0.0).to(tl.float32)
b_Ai11 = tl.load(p_Akk11, mask=m_A1, other=0.0).to(tl.float32)
if NC >= 3:
p_Akk22 = Akkd + o_c2[:, None] * (HV*BC) + o_i[None, :]
b_Ai22 = tl.load(p_Akk22, mask=m_A2, other=0.0).to(tl.float32)
if NC >= 4:
p_Akk33 = Akkd + o_c3[:, None] * (HV*BC) + o_i[None, :]
b_Ai33 = tl.load(p_Akk33, mask=m_A3, other=0.0).to(tl.float32)
################################################################################
# forward substitution on diagonals
################################################################################
if not USE_SAFE_GATE:
m_A = o_i[:, None] > o_i[None, :]
m_I = o_i[:, None] == o_i[None, :]
b_Ai00 = -tl.where(m_A, b_Ai00, 0)
b_Ai11 = -tl.where(m_A, b_Ai11, 0)
if NC >= 3:
b_Ai22 = -tl.where(m_A, b_Ai22, 0)
if NC >= 4:
b_Ai33 = -tl.where(m_A, b_Ai33, 0)
for i in range(2, min(BC, T - i_tc0)):
b_a00 = -tl.load(Akkd + (i_tc0 + i) * HV*BC + o_i)
b_a00 = tl.where(o_i < i, b_a00, 0.)
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(BC + 2, min(2*BC, T - i_tc0)):
b_a11 = -tl.load(Akkd + (i_tc0 + i) * HV*BC + o_i)
b_a11 = tl.where(o_i < i - BC, b_a11, 0.)
b_a11 += tl.sum(b_a11[:, None] * b_Ai11, 0)
b_Ai11 = tl.where((o_i == i - BC)[:, None], b_a11, b_Ai11)
if NC >= 3:
for i in range(2*BC + 2, min(3*BC, T - i_tc0)):
b_a22 = -tl.load(Akkd + (i_tc0 + i) * HV*BC + o_i)
b_a22 = tl.where(o_i < i - 2*BC, b_a22, 0.)
b_a22 += tl.sum(b_a22[:, None] * b_Ai22, 0)
b_Ai22 = tl.where((o_i == i - 2*BC)[:, None], b_a22, b_Ai22)
if NC >= 4:
for i in range(3*BC + 2, min(4*BC, T - i_tc0)):
b_a33 = -tl.load(Akkd + (i_tc0 + i) * HV*BC + o_i)
b_a33 = tl.where(o_i < i - 3*BC, b_a33, 0.)
b_a33 += tl.sum(b_a33[:, None] * b_Ai33, 0)
b_Ai33 = tl.where((o_i == i - 3*BC)[:, None], b_a33, b_Ai33)
b_Ai00 += m_I
b_Ai11 += m_I
if NC >= 3:
b_Ai22 += m_I
if NC >= 4:
b_Ai33 += m_I
################################################################################
# compute merged inverse using off-diagonals
################################################################################
# we used tf32 to maintain matrix inverse's precision whenever possible.
b_Ai10 = -tl.dot(
tl.dot(b_Ai11, b_Akk10, input_precision=SOLVE_TRIL_DOT_PRECISION),
b_Ai00,
input_precision=SOLVE_TRIL_DOT_PRECISION
)
if NC >= 3:
b_Ai21 = -tl.dot(
tl.dot(b_Ai22, b_Akk21, input_precision=SOLVE_TRIL_DOT_PRECISION),
b_Ai11,
input_precision=SOLVE_TRIL_DOT_PRECISION
)
b_Ai20 = -tl.dot(
b_Ai22,
tl.dot(b_Akk20, b_Ai00, input_precision=SOLVE_TRIL_DOT_PRECISION) +
tl.dot(b_Akk21, b_Ai10, input_precision=SOLVE_TRIL_DOT_PRECISION),
input_precision=SOLVE_TRIL_DOT_PRECISION
)
if NC >= 4:
b_Ai32 = -tl.dot(
tl.dot(b_Ai33, b_Akk32, input_precision=SOLVE_TRIL_DOT_PRECISION),
b_Ai22,
input_precision=SOLVE_TRIL_DOT_PRECISION
)
b_Ai31 = -tl.dot(
b_Ai33,
tl.dot(b_Akk31, b_Ai11, input_precision=SOLVE_TRIL_DOT_PRECISION) +
tl.dot(b_Akk32, b_Ai21, input_precision=SOLVE_TRIL_DOT_PRECISION),
input_precision=SOLVE_TRIL_DOT_PRECISION
)
b_Ai30 = -tl.dot(
b_Ai33,
tl.dot(b_Akk30, b_Ai00, input_precision=SOLVE_TRIL_DOT_PRECISION) +
tl.dot(b_Akk31, b_Ai10, input_precision=SOLVE_TRIL_DOT_PRECISION) +
tl.dot(b_Akk32, b_Ai20, input_precision=SOLVE_TRIL_DOT_PRECISION),
input_precision=SOLVE_TRIL_DOT_PRECISION
)
################################################################################
# store full Akk_inv to Akk
################################################################################
p_Akk00 = Akk + o_c0[:, None] * (HV*BT) + o_i[None, :]
p_Akk10 = Akk + o_c1[:, None] * (HV*BT) + o_i[None, :]
p_Akk11 = Akk + o_c1[:, None] * (HV*BT) + (o_i + BC)[None, :]
tl.store(p_Akk00, b_Ai00.to(Akk.dtype.element_ty), mask=m_A0)
tl.store(p_Akk10, b_Ai10.to(Akk.dtype.element_ty), mask=m_A1)
tl.store(p_Akk11, b_Ai11.to(Akk.dtype.element_ty), mask=m_A1)
if NC >= 3:
p_Akk20 = Akk + o_c2[:, None] * (HV*BT) + o_i[None, :]
p_Akk21 = Akk + o_c2[:, None] * (HV*BT) + (o_i + BC)[None, :]
p_Akk22 = Akk + o_c2[:, None] * (HV*BT) + (o_i + 2*BC)[None, :]
tl.store(p_Akk20, b_Ai20.to(Akk.dtype.element_ty), mask=m_A2)
tl.store(p_Akk21, b_Ai21.to(Akk.dtype.element_ty), mask=m_A2)
tl.store(p_Akk22, b_Ai22.to(Akk.dtype.element_ty), mask=m_A2)
if NC >= 4:
p_Akk30 = Akk + o_c3[:, None] * (HV*BT) + o_i[None, :]
p_Akk31 = Akk + o_c3[:, None] * (HV*BT) + (o_i + BC)[None, :]
p_Akk32 = Akk + o_c3[:, None] * (HV*BT) + (o_i + 2*BC)[None, :]
p_Akk33 = Akk + o_c3[:, None] * (HV*BT) + (o_i + 3*BC)[None, :]
tl.store(p_Akk30, b_Ai30.to(Akk.dtype.element_ty), mask=m_A3)
tl.store(p_Akk31, b_Ai31.to(Akk.dtype.element_ty), mask=m_A3)
tl.store(p_Akk32, b_Ai32.to(Akk.dtype.element_ty), mask=m_A3)
tl.store(p_Akk33, b_Ai33.to(Akk.dtype.element_ty), mask=m_A3)
@triton.heuristics({
'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 [1, 2, 4, 8]
for num_stages in [2, 3, 4]
],
key=['BK', 'NC', 'BT', 'HV'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['B', 'T'])
def chunk_kda_bwd_kernel_intra(
q,
k,
g,
beta,
dAqk,
dAkk,
dq,
dq2,
dk,
dk2,
dg,
dg2,
db,
cu_seqlens,
chunk_indices,
B,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
BT: tl.constexpr,
BC: tl.constexpr,
BK: tl.constexpr,
NC: tl.constexpr,
IS_VARLEN: tl.constexpr,
SAFE_GATE: tl.constexpr,
USE_GATHER: tl.constexpr,
):
i_kc, i_t, i_bh = tl.program_id(0), tl.program_id(1).to(tl.int64), tl.program_id(2).to(tl.int64)
i_b, i_hv = i_bh // HV, i_bh % HV
i_h = i_hv // (HV // H)
i_k, i_i = i_kc // NC, i_kc % NC
all = B * T
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)
else:
bos, eos = i_b * T, i_b * T + T
T = eos - bos
i_ti = i_t * BT + i_i * BC
if i_ti >= T:
return
o_k = i_k * BK + tl.arange(0, BK)
m_k = o_k < K
q += (bos * H + i_h) * K
k += (bos * H + i_h) * K
g += (bos * HV + i_hv) * K
beta += bos * HV + i_hv
dAqk += (bos * HV + i_hv) * BT
dAkk += (bos * HV + i_hv) * BT
dq += (bos * HV + i_hv) * K
dq2 += (bos * HV + i_hv) * K
dk += (bos * HV + i_hv) * K
dk2 += (bos * HV + i_hv) * K
dg += (bos * HV + i_hv) * K
dg2 += (bos * HV + i_hv) * K
db += (i_k * all + bos) * HV + i_hv
o_i = tl.arange(0, BC)
o_c = i_ti + o_i
m_c = o_c < T
m_ck = m_c[:, None] & m_k[None, :]
m_dAf = m_c[:, None] & (o_i[None, :] < BT)
m_dAt = (o_i[:, None] < BT) & m_c[None, :]
p_g = g + o_c[:, None] * (HV*K) + o_k[None, :]
b_g = tl.load(p_g, mask=m_ck, other=0.0).to(tl.float32)
p_b = beta + o_c * HV
b_b = tl.load(p_b, mask=m_c, other=0.0)
b_dq2 = tl.zeros([BC, BK], dtype=tl.float32)
b_dk2 = tl.zeros([BC, BK], dtype=tl.float32)
if i_i > 0:
p_gn = g + i_ti * HV*K + o_k
# [BK,]
b_gn = tl.load(p_gn, mask=m_k, other=0).to(tl.float32)[None, :]
for i_j in range(0, i_i):
o_j = i_t * BT + i_j * BC + o_i
m_jk = (o_j < T)[:, None] & m_k[None, :]
p_k = k + o_j[:, None] * (H*K) + o_k[None, :]
p_gk = g + o_j[:, None] * (HV*K) + o_k[None, :]
p_dAqk = dAqk + o_c[:, None] * (HV*BT) + (i_j * BC + o_i)[None, :]
p_dAkk = dAkk + o_c[:, None] * (HV*BT) + (i_j * BC + o_i)[None, :]
# [BC, BK]
b_k = tl.load(p_k, mask=m_jk, other=0.0)
b_gk = tl.load(p_gk, mask=m_jk, other=0.0)
b_kg = b_k * exp2(b_gn - b_gk)
# [BC, BC]
b_dAqk = tl.load(p_dAqk, mask=m_dAf, other=0.0)
b_dAkk = tl.load(p_dAkk, mask=m_dAf, other=0.0)
# [BC, BK]
b_dq2 += tl.dot(b_dAqk, b_kg)
b_dk2 += tl.dot(b_dAkk, b_kg)
b_gqn = exp2(b_g - b_gn)
b_dq2 *= b_gqn
b_dk2 *= b_gqn
o_i = tl.arange(0, BC)
m_dA = (i_ti + o_i) < T
o_dA = (i_ti + o_i) * HV*BT + i_i * BC
p_kj = k + i_ti * H*K + o_k
p_gkj = g + i_ti * HV*K + o_k
p_q = q + o_c[:, None] * (H*K) + o_k[None, :]
p_k = k + o_c[:, None] * (H*K) + o_k[None, :]
b_q = tl.load(p_q, mask=m_ck, other=0.0)
b_k = tl.load(p_k, mask=m_ck, other=0.0)
if SAFE_GATE:
if USE_GATHER:
b_gn = gather(b_g, tl.full([1, BK], min(BC//2, T - i_ti - 1), dtype=tl.int16), axis=0)
else:
p_gn = g + (i_ti + min(BC // 2, T - i_ti - 1)) * HV*K + o_k
b_gn = tl.load(p_gn, mask=m_k, other=0)[None, :]
p_dAqk = dAqk + o_c[:, None] * (HV*BT) + (i_i * BC + o_i)[None, :]
p_dAkk = dAkk + o_c[:, None] * (HV*BT) + (i_i * BC + o_i)[None, :]
b_dAqk_diag_qk = tl.load(p_dAqk, mask=m_dAf, other=0.0).to(tl.float32)
b_dAkk_diag_qk = tl.load(p_dAkk, mask=m_dAf, other=0.0).to(tl.float32)
m_i_diag_qk = (o_i[:, None] >= o_i[None, :]) & ((i_ti + o_i[:, None]) < T) & ((i_ti + o_i[None, :]) < T)
m_j_diag_qk = (i_ti + o_i[:, None]) < T
b_dAqk_diag_qk = tl.where(m_i_diag_qk, b_dAqk_diag_qk, 0.)
b_dAkk_diag_qk = tl.where(m_i_diag_qk, b_dAkk_diag_qk, 0.)
b_g_diag_qk = tl.where(m_j_diag_qk, b_g - b_gn, 0.)
exp_b_g_diag_qk = tl.where(m_j_diag_qk, exp2(b_g_diag_qk), 0.)
exp_neg_b_g_diag_qk = tl.where(m_j_diag_qk, exp2(-b_g_diag_qk), 0.)
b_k_exp_diag_qk = b_k * exp_neg_b_g_diag_qk
b_dq2 += tl.dot(b_dAqk_diag_qk, b_k_exp_diag_qk) * exp_b_g_diag_qk
b_dk2 += tl.dot(b_dAkk_diag_qk, b_k_exp_diag_qk) * exp_b_g_diag_qk
else:
for j in range(0, min(BC, T - i_t * BT - i_i * BC)):
# [BC]
b_dAqk = tl.load(dAqk + o_dA + j, mask=m_dA, other=0)
b_dAkk = tl.load(dAkk + o_dA + j, mask=m_dA, other=0)
# [BK]
b_kj = tl.load(p_kj, mask=m_k, other=0).to(tl.float32)
b_gkj = tl.load(p_gkj, mask=m_k, other=0).to(tl.float32)
# [BC, BK]
m_i = o_i[:, None] >= j
# [BC, BK]
b_gqk = exp2(b_g - b_gkj[None, :])
b_dq2 += tl.where(m_i, b_dAqk[:, None] * b_kj[None, :] * b_gqk, 0.)
b_dk2 += tl.where(m_i, b_dAkk[:, None] * b_kj[None, :] * b_gqk, 0.)
p_kj += H*K
p_gkj += HV*K
b_db = tl.sum(b_dk2 * b_k, 1)
b_dk2 *= b_b[:, None]
p_dq = dq + o_c[:, None] * (HV*K) + o_k[None, :]
p_dq2 = dq2 + o_c[:, None] * (HV*K) + o_k[None, :]
p_db = db + o_c * HV
b_dg2 = b_q * b_dq2
b_dq2 = b_dq2 + tl.load(p_dq, mask=m_ck, other=0.0)
tl.store(p_dq2, b_dq2.to(p_dq2.dtype.element_ty), mask=m_ck)
tl.store(p_db, b_db.to(p_db.dtype.element_ty), mask=m_c)
tl.debug_barrier()
b_dkt = tl.zeros([BC, BK], dtype=tl.float32)
NC = min(NC, tl.cdiv(T - i_t * BT, BC))
if i_i < NC - 1:
p_gn = g + (min(i_ti + BC, T) - 1) * HV*K + o_k
# [BK,]
b_gn = tl.load(p_gn, mask=m_k, other=0).to(tl.float32)[None, :]
for i_j in range(i_i + 1, NC):
o_j = i_t * BT + i_j * BC + o_i
m_j = o_j < T
m_jk = m_j[:, None] & m_k[None, :]
m_dAj = (o_i[:, None] < BT) & m_j[None, :]
p_q = q + o_j[:, None] * (H*K) + o_k[None, :]
p_k = k + o_j[:, None] * (H*K) + o_k[None, :]
p_gk = g + o_j[:, None] * (HV*K) + o_k[None, :]
p_b = beta + o_j * HV
p_dAqk = dAqk + (i_i * BC + o_i)[:, None] + o_j[None, :] * (HV*BT)
p_dAkk = dAkk + (i_i * BC + o_i)[:, None] + o_j[None, :] * (HV*BT)
# [BC]
b_b = tl.load(p_b, mask=m_j, other=0.0)
# [BC, BK]
b_q = tl.load(p_q, mask=m_jk, other=0.0)
b_kb = tl.load(p_k, mask=m_jk, other=0.0) * b_b[:, None]
b_gk = tl.load(p_gk, mask=m_jk, other=0.0).to(tl.float32)
# [BC, BC]
b_dAqk = tl.load(p_dAqk, mask=m_dAj, other=0.0)
b_dAkk = tl.load(p_dAkk, mask=m_dAj, other=0.0)
# [BC, BK]
b_gkn = exp2(b_gk - b_gn)
b_qg = b_q * tl.where(m_j[:, None], b_gkn, 0)
b_kbg = b_kb * tl.where(m_j[:, None], b_gkn, 0)
# [BC, BK]
# (SY 09/17) important to not use bf16 here to have a good precision.
b_dkt += tl.dot(b_dAqk, b_qg)
b_dkt += tl.dot(b_dAkk, b_kbg)
b_dkt *= exp2(b_gn - b_g)
o_dA = i_ti * HV*BT + i_i * BC + o_i
p_qj = q + i_ti * H*K + o_k
p_kj = k + i_ti * H*K + o_k
p_gkj = g + i_ti * HV*K + o_k
p_bj = beta + i_ti * HV
if SAFE_GATE:
if USE_GATHER:
b_gn = gather(b_g, tl.full([1, BK], min(BC//2, T - i_ti - 1), dtype=tl.int16), axis=0)
else:
p_gn = g + (i_ti + min(BC // 2, T - i_ti - 1)) * HV*K + o_k
b_gn = tl.load(p_gn, mask=m_k, other=0).to(tl.float32)[None, :]
p_q = q + o_c[:, None] * (H*K) + o_k[None, :]
b_q = tl.load(p_q, mask=m_ck, other=0.0)
p_b = beta + o_c * HV
b_b = tl.load(p_b, mask=m_c, other=0.0)
p_dAqk = dAqk + (i_i * BC + o_i)[:, None] + o_c[None, :] * (HV*BT)
p_dAkk = dAkk + (i_i * BC + o_i)[:, None] + o_c[None, :] * (HV*BT)
b_dAqk_diag_kk = tl.load(p_dAqk, mask=m_dAt, other=0.0).to(tl.float32)
b_dAkk_diag_kk = tl.load(p_dAkk, mask=m_dAt, other=0.0).to(tl.float32)
m_i_diag_kk = (o_i[:, None] <= o_i[None, :]) & ((i_ti + o_i[:, None]) < T) & ((i_ti + o_i[None, :]) < T)
m_j_diag_kk = (i_ti + o_i[:, None]) < T
b_dAqk_diag_kk = tl.where(m_i_diag_kk, b_dAqk_diag_kk, 0.)
b_dAkk_diag_kk = tl.where(m_i_diag_kk, b_dAkk_diag_kk, 0.)
# ensure numerical stability
b_g_diag_kk = tl.where(m_j_diag_kk, b_g - b_gn, 0.)
exp_b_g_diag_kk = tl.where(m_j_diag_kk, exp2(b_g_diag_kk), 0.)
exp_neg_b_g_diag_kk = tl.where(m_j_diag_kk, exp2(-b_g_diag_kk), 0.)
b_q_exp = b_q * exp_b_g_diag_kk
b_kb_exp = b_k * b_b[:, None] * exp_b_g_diag_kk
b_dkt += tl.dot(b_dAqk_diag_kk, b_q_exp) * exp_neg_b_g_diag_kk
b_dkt += tl.dot(b_dAkk_diag_kk, b_kb_exp) * exp_neg_b_g_diag_kk
else:
for j in range(0, min(BC, T - i_t * BT - i_i * BC)):
# [BC,]
b_dAqk = tl.load(dAqk + o_dA + j * HV*BT)
b_dAkk = tl.load(dAkk + o_dA + j * HV*BT)
# [BK,]
b_qj = tl.load(p_qj, mask=m_k, other=0).to(tl.float32)
b_kbj = tl.load(p_kj, mask=m_k, other=0).to(tl.float32) * tl.load(p_bj)
b_gkj = tl.load(p_gkj, mask=m_k, other=0).to(tl.float32)
# [BC, BK]
m_i = o_i[:, None] <= j
b_gkq = exp2(b_gkj[None, :] - b_g)
b_dkt += tl.where(m_i, b_dAqk[:, None] * b_qj[None, :] * b_gkq, 0.)
b_dkt += tl.where(m_i, b_dAkk[:, None] * b_kbj[None, :] * b_gkq, 0.)
p_qj += H*K
p_kj += H*K
p_gkj += HV*K
p_bj += HV
p_dk = dk + o_c[:, None] * (HV*K) + o_k[None, :]
p_dk2 = dk2 + o_c[:, None] * (HV*K) + o_k[None, :]
p_dg = dg + o_c[:, None] * (HV*K) + o_k[None, :]
p_dg2 = dg2 + o_c[:, None] * (HV*K) + o_k[None, :]
b_dg2 += (b_dk2 - b_dkt) * b_k + tl.load(p_dg, mask=m_ck, other=0.0)
b_dk2 += tl.load(p_dk, mask=m_ck, other=0.0)
b_dk2 += b_dkt
tl.store(p_dk2, b_dk2.to(p_dk2.dtype.element_ty), mask=m_ck)
tl.store(p_dg2, b_dg2.to(p_dg2.dtype.element_ty), mask=m_ck)
@triton.heuristics({
'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 [1, 2, 4, 8]
for num_stages in [2, 3, 4]
],
key=["BT", "BC", "HV"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_kda_fwd_kernel_intra_sub_chunk(
q,
k,
g,
beta,
Aqk,
Akk,
scale,
cu_seqlens,
chunk_indices,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
BT: tl.constexpr,
BC: tl.constexpr,
BK: tl.constexpr,
IS_VARLEN: tl.constexpr,
USE_GATHER: tl.constexpr,
):
i_t, i_i, i_bh = tl.program_id(0).to(tl.int64), tl.program_id(1), tl.program_id(2).to(tl.int64)
i_b, i_hv = i_bh // HV, i_bh % HV
i_h = i_hv // (HV // 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
i_ti = i_t * BT + i_i * BC
if i_ti >= T:
return
o_c = i_ti + tl.arange(0, BC)
m_c = o_c < T
q = q + (bos * H + i_h) * K
k = k + (bos * H + i_h) * K
g = g + (bos * HV + i_hv) * K
beta = beta + bos * HV + i_hv
Aqk = Aqk + (bos * HV + i_hv) * BT
Akk = Akk + (bos * HV + i_hv) * BC
o_k = tl.arange(0, BK)
m_k = o_k < K
m_ck = m_c[:, None] & m_k[None, :]
p_q = q + o_c[:, None] * (H*K) + o_k[None, :]
p_k = k + o_c[:, None] * (H*K) + o_k[None, :]
p_g = g + o_c[:, None] * (HV*K) + o_k[None, :]
p_beta = beta + o_c * HV
b_q = tl.load(p_q, mask=m_ck, other=0.0)
b_k = tl.load(p_k, mask=m_ck, other=0.0)
b_g = tl.load(p_g, mask=m_ck, other=0.0)
b_beta = tl.load(p_beta, mask=m_c, other=0.0)
if USE_GATHER:
b_gn = gather(b_g, tl.full([1, BK], min(BC//2, T - i_ti - 1), dtype=tl.int16), axis=0)
else:
# caculate offset
p_gn = g + (i_ti + min(BC // 2, T - i_ti - 1)) * HV*K + tl.arange(0, BK)
b_gn = tl.load(p_gn, mask=tl.arange(0, BK) < K, other=0.0)
b_gn = b_gn[None, :]
# current block, keep numerical stability by subtracting the left boundary
# less than 85 to avoid overflow in exp2
b_gm = (b_g - b_gn).to(tl.float32)
b_gq = tl.where(m_c[:, None], exp2(b_gm), 0.)
b_gk = tl.where(m_c[:, None], exp2(-b_gm), 0.)
b_kgt = tl.trans(b_k * b_gk)
b_Aqk = tl.dot(b_q * b_gq, b_kgt) * scale
b_Akk = tl.dot(b_k * b_gq, b_kgt) * b_beta[:, None]
o_i = tl.arange(0, BC)
m_Aqk = o_i[:, None] >= o_i[None, :]
m_Akk = o_i[:, None] > o_i[None, :]
m_I = o_i[:, None] == o_i[None, :]
b_Aqk = tl.where(m_Aqk, b_Aqk, 0.0)
b_Akk = tl.where(m_Akk, b_Akk, 0.0)
m_Aqk_st = m_c[:, None] & (o_i[None, :] < BT)
m_Akk_st = m_c[:, None] & (o_i[None, :] < BC)
p_Aqk = Aqk + o_c[:, None] * (HV*BT) + (i_i * BC + o_i)[None, :]
p_Akk = Akk + o_c[:, None] * (HV*BC) + o_i[None, :]
tl.store(p_Aqk, b_Aqk.to(Aqk.dtype.element_ty), mask=m_Aqk_st)
tl.store(p_Akk, b_Akk.to(Akk.dtype.element_ty), mask=m_Akk_st)
tl.debug_barrier()
################################################################################
# forward substitution
################################################################################
b_Ai = -b_Akk
for i in range(2, min(BC, T - i_ti)):
b_a = -tl.load(Akk + (i_ti + i) * HV*BC + o_i)
b_a = tl.where(o_i < i, b_a, 0.)
b_a += tl.sum(b_a[:, None] * b_Ai, 0)
b_Ai = tl.where((o_i == i)[:, None], b_a, b_Ai)
b_Ai += m_I
tl.store(p_Akk, b_Ai.to(Akk.dtype.element_ty), mask=m_Akk_st)
@dispatch('kda')
def chunk_kda_fwd_intra(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
gk: torch.Tensor | None = None,
beta: torch.Tensor | None = None,
scale: float | None = None,
cu_seqlens: torch.LongTensor | None = None,
chunk_size: int = 64,
chunk_indices: torch.LongTensor | None = None,
safe_gate: bool = False,
disable_recompute: bool = False,
):
B, T, H, K, HV = *k.shape, gk.shape[2]
BT = chunk_size
if BT not in (32, 64):
raise ValueError(f"KDA intra chunk kernel only supports chunk_size 32 or 64, got {BT}.")
BC = 16
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)
NC = triton.cdiv(BT, BC)
Aqk = torch.empty(B, T, HV, BT, device=k.device, dtype=k.dtype)
# Akk must be zero-initialized - kernel only writes lower triangular
Akk = torch.zeros(B, T, HV, BT, device=k.device, dtype=k.dtype)
# Separate fp32 buffer for diagonal 16x16 blocks (for precision in solve_tril)
Akkd = torch.empty(B, T, HV, BC, device=k.device, dtype=torch.float32)
# Step 1: Run token_parallel first to compute diagonal blocks into Akkd (fp32)
# Step 1: compute diagonal blocks into Akk_diag (fp32)
if safe_gate:
grid = (NT, NC, B * HV)
BK = triton.next_power_of_2(K)
chunk_kda_fwd_kernel_intra_sub_chunk[grid](
q=q,
k=k,
g=gk,
beta=beta,
Aqk=Aqk,
Akk=Akkd,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
T=T,
H=H,
HV=HV,
K=K,
BT=BT,
BC=BC,
BK=BK,
USE_GATHER=IS_GATHER_SUPPORTED,
)
else:
Aqk, Akkd = chunk_kda_fwd_intra_token_parallel(
q=q,
k=k,
gk=gk,
beta=beta,
Aqk=Aqk,
Akk=Akkd,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_size=BT,
sub_chunk_size=BC,
)
# Step 2: Fused inter + solve_tril (works for both fixed-len and varlen)
grid = (NT, B * HV)
chunk_kda_fwd_kernel_inter_solve_fused[grid](
q=q,
k=k,
g=gk,
beta=beta,
Aqk=Aqk,
Akkd=Akkd,
Akk=Akk,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
T=T,
H=H,
HV=HV,
K=K,
BT=BT,
BC=BC,
NC=NC,
USE_SAFE_GATE=safe_gate,
)
w, u, qg, kg = recompute_w_u_fwd(
k=k,
v=v,
beta=beta,
A=Akk,
q=q if disable_recompute else None,
gk=gk,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
)
return w, u, qg, kg, Aqk, Akk
@dispatch('kda')
def chunk_kda_bwd_intra(
q: torch.Tensor,
k: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
dAqk: torch.Tensor,
dAkk: torch.Tensor,
dq: torch.Tensor,
dk: torch.Tensor,
db: torch.Tensor,
dg: torch.Tensor,
cu_seqlens: torch.LongTensor | None = None,
chunk_indices: torch.LongTensor | None = None,
chunk_size: int = 64,
safe_gate: bool = False,
):
B, T, H, K, HV = *k.shape, g.shape[2]
BT = chunk_size
BC = min(16, BT)
BK = min(32, triton.next_power_of_2(K))
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)
NC = triton.cdiv(BT, BC)
NK = triton.cdiv(K, BK)
dq2 = torch.empty_like(dq)
dk2 = torch.empty_like(dk)
db2 = beta.new_empty(NK, *beta.shape, dtype=torch.float)
dg2 = torch.empty_like(dg, dtype=torch.float)
grid = (NK * NC, NT, B * HV)
chunk_kda_bwd_kernel_intra[grid](
q=q,
k=k,
g=g,
beta=beta,
dAqk=dAqk,
dAkk=dAkk,
dq=dq,
dq2=dq2,
dk=dk,
dk2=dk2,
dg=dg,
dg2=dg2,
db=db2,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
B=B,
T=T,
H=H,
HV=HV,
K=K,
BT=BT,
BC=BC,
BK=BK,
NC=NC,
SAFE_GATE=safe_gate,
USE_GATHER=IS_GATHER_SUPPORTED,
)
dq = dq2
dk = dk2
db = db2.sum(0).add_(db)
dg = dg2
return dq, dk, db, dg
@@ -0,0 +1,182 @@
# 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
# Token-parallel implementation of KDA intra chunk kernel
import torch
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.ops.utils.op import exp2
from kda._fla.utils import autotune_cache_kwargs
@triton.heuristics({
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({'BH': BH}, num_warps=num_warps)
for BH in [1, 2, 4, 8]
for num_warps in [1, 2, 4, 8]
],
key=["K", "H", "HV"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T', 'N'])
def chunk_kda_fwd_kernel_intra_token_parallel(
q,
k,
g,
beta,
Aqk,
Akk,
scale,
cu_seqlens,
N,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
BK: tl.constexpr,
BT: tl.constexpr,
BC: tl.constexpr,
BH: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_tg, i_hg = tl.program_id(0).to(tl.int64), tl.program_id(1)
if IS_VARLEN:
i_n = 0
left, right = 0, N
# Unrolled binary search (max B=2^32)
# We can limit iterations based on expected max batch size if needed
# 20 iterations covers B=1M, usually enough
for _ in range(20):
if left < right:
mid = (left + right) // 2
if i_tg < tl.load(cu_seqlens + mid + 1).to(tl.int32):
right = mid
else:
left = mid + 1
i_n = left
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
i_t = i_tg - bos
else:
bos = (i_tg // T) * T
i_t = i_tg % T
if i_t >= T:
return
i_c = i_t // BT
i_s = (i_t % BT) // BC
i_tc = i_c * BT
i_ts = i_tc + i_s * BC
G: tl.constexpr = HV // H
q += bos * H*K
k += bos * H*K
g += bos * HV*K
Aqk += bos * HV*BT
Akk += bos * HV*BC
beta += bos * HV
o_hv = i_hg * BH + tl.arange(0, BH)
o_h = o_hv // G
o_k = tl.arange(0, BK)
m_hv = o_hv < HV
m_k = o_k < K
m_hk = m_hv[:, None] & m_k[None, :]
# q/k: [B, T, H, K], manual load via mapped qk head index
p_qk = o_h[:, None] * K + o_k[None, :]
b_q = tl.load(q + i_t * H * K + p_qk, mask=m_hk, other=0).to(tl.float32)
b_k = tl.load(k + i_t * H * K + p_qk, mask=m_hk, other=0).to(tl.float32)
# g: [B, T, HV, K], beta: [B, T, HV]
p_g = g + i_t * HV * K + o_hv[:, None] * K + o_k[None, :]
p_beta = beta + i_t * HV + o_hv
b_g = tl.load(p_g, mask=m_hk, other=0.0).to(tl.float32)
b_k = b_k * tl.load(p_beta, mask=m_hv, other=0.0).to(tl.float32)[:, None]
for j in range(i_ts, min(i_t + 1, min(T, i_ts + BC))):
b_kj = tl.load(k + j * H * K + p_qk, mask=m_hk, other=0).to(tl.float32)
p_gj = g + j * HV * K + o_hv[:, None] * K + o_k[None, :]
b_gj = tl.load(p_gj, mask=m_hk, other=0.0).to(tl.float32)
b_kgj = tl.where(m_k[None, :], b_kj * exp2(b_g - b_gj), 0.0)
b_Aqk = tl.sum(b_q * b_kgj, axis=1) * scale
b_Akk = tl.sum(b_k * b_kgj, axis=1) * tl.where(j < i_t, 1.0, 0.0)
tl.store(Aqk + i_t * HV * BT + o_hv * BT + j % BT, b_Aqk.to(Aqk.dtype.element_ty), mask=m_hv)
tl.store(Akk + i_t * HV * BC + o_hv * BC + j - i_ts, b_Akk.to(Akk.dtype.element_ty), mask=m_hv)
@dispatch('kda')
def chunk_kda_fwd_intra_token_parallel(
q: torch.Tensor,
k: torch.Tensor,
gk: torch.Tensor,
beta: torch.Tensor,
Aqk: torch.Tensor,
Akk: torch.Tensor,
scale: float,
cu_seqlens: torch.LongTensor | None = None,
chunk_size: int = 64,
sub_chunk_size: int = 16,
) -> None:
"""
Token-parallel implementation: each token gets its own thread block.
Supports both fixed-length and variable-length sequences.
Reduces wasted computation on padding.
Writes directly to Aqk and Akk tensors (in-place).
Args:
q: [B, T, H, K]
k: [B, T, H, K]
gk: [B, T, HV, K] cumsum of gates (HV >= H for GVA)
beta: [B, T, HV]
Aqk: [B, T, HV, BT] output tensor to write to
Akk: [B, T, HV, BC] output tensor for diagonal blocks (fp32)
scale: attention scale
chunk_size: BT (default 64)
sub_chunk_size: BC (default 16)
"""
B, T, H, K, HV = *q.shape, gk.shape[2]
N = len(cu_seqlens) - 1 if cu_seqlens is not None else B
BT = chunk_size
BC = sub_chunk_size
BK = triton.next_power_of_2(K)
def grid(meta): return (B * T, triton.cdiv(HV, meta['BH']))
chunk_kda_fwd_kernel_intra_token_parallel[grid](
q=q,
k=k,
g=gk,
beta=beta,
Aqk=Aqk,
Akk=Akk,
scale=scale,
cu_seqlens=cu_seqlens,
N=N,
T=T,
H=H,
HV=HV,
K=K,
BK=BK,
BT=BT,
BC=BC,
)
return Aqk, Akk
+491
View File
@@ -0,0 +1,491 @@
# 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
# This kernel is modified from the Decode kernel of the vllm gdn/kda model.
import warnings
import torch
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.ops.utils.op import exp
from kda._fla.ops.utils.softplus import softplus
from kda._fla.utils import input_guard
@triton.heuristics(
{
"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,
"IS_CONTINUOUS_BATCHING": lambda args: args["ssm_state_indices"] is not None,
"IS_SPEC_DECODING": lambda args: args["num_accepted_tokens"] is not None,
"HAS_A": lambda args: args["A_log"] is not None,
"HAS_BIAS": lambda args: args["dt_bias"] is not None,
"USE_LOWER_BOUND": lambda args: args["lower_bound"] is not None,
}
)
@triton.jit(do_not_specialize=["N", "T"])
def fused_recurrent_kda_fwd_kernel(
q,
k,
v,
g,
beta,
A_log,
dt_bias,
o,
h0,
ht,
cu_seqlens,
ssm_state_indices,
num_accepted_tokens,
lower_bound,
scale: tl.constexpr,
N: tl.int64, # num of sequences
T: tl.int64, # num of tokens
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BK: tl.constexpr,
BV: tl.constexpr,
stride_init_state_token: tl.constexpr,
stride_final_state_token: tl.constexpr,
stride_indices_seq: tl.constexpr,
stride_indices_tok: tl.constexpr,
USE_INITIAL_STATE: tl.constexpr, # whether to use initial state
INPLACE_FINAL_STATE: tl.constexpr, # whether to store final state inplace
IS_BETA_HEADWISE: tl.constexpr, # whether beta is headwise vector or scalar,
USE_QK_L2NORM_IN_KERNEL: tl.constexpr,
IS_VARLEN: tl.constexpr,
IS_CONTINUOUS_BATCHING: tl.constexpr,
IS_SPEC_DECODING: tl.constexpr,
STORE_FINAL_STATE: tl.constexpr,
HAS_A: tl.constexpr,
HAS_BIAS: tl.constexpr,
USE_GATE_IN_KERNEL: tl.constexpr,
USE_LOWER_BOUND: tl.constexpr,
APPLY_BETA_SIGMOID: tl.constexpr,
ALLOW_NEG_EIGVAL: tl.constexpr,
STATE_V_FIRST: tl.constexpr,
num_stages: tl.constexpr,
):
pid = tl.program_id(0).to(tl.int64)
NV = tl.cdiv(V, BV)
NK = tl.cdiv(K, BK)
i_k = pid % NK
pid_rest = pid // NK
i_v = pid_rest % NV
i_nh = pid_rest // NV
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
if T == 0:
# no tokens to process for this sequence
return
o_k = i_k * BK + 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 IS_BETA_HEADWISE:
p_beta = beta + (bos * HV + i_hv) * V + o_v
else:
p_beta = beta + bos * HV + i_hv
p_g = g + (bos * HV + i_hv) * K + o_k
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 IS_CONTINUOUS_BATCHING:
if IS_SPEC_DECODING:
i_t = tl.load(num_accepted_tokens + i_n).to(tl.int64) - 1
else:
i_t = 0
p_h0 = (
h0
+ tl.load(ssm_state_indices + i_n * stride_indices_seq + i_t).to(
tl.int64
)
* stride_init_state_token
)
if STATE_V_FIRST:
p_h0 = p_h0 + i_hv * K * V + o_v[:, None] * K + o_k[None, :]
else:
p_h0 = p_h0 + i_hv * K * V + o_k[:, None] * V + o_v[None, :]
else:
if STATE_V_FIRST:
p_h0 = h0 + (i_n * HV + i_hv) * K * V + o_v[:, None] * K + o_k[None, :]
else:
p_h0 = h0 + (i_n * HV + i_hv) * K * V + o_k[:, None] * V + o_v[None, :]
b_h += tl.load(p_h0, mask=mask_h, other=0).to(tl.float32)
for i_t in tl.range(0, T, num_stages=num_stages):
b_q = tl.load(p_q, mask=mask_k, other=0, eviction_policy='evict_last').to(tl.float32)
b_k = tl.load(p_k, mask=mask_k, other=0, eviction_policy='evict_last').to(tl.float32)
b_v = tl.load(p_v, mask=mask_v, other=0, eviction_policy='evict_first').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
b_g = tl.load(p_g, mask=mask_k, other=0, eviction_policy='evict_last').to(tl.float32)
if USE_GATE_IN_KERNEL:
b_A = tl.load(A_log + i_hv).to(tl.float32) if HAS_A else 1.0
if HAS_BIAS:
b_bias = tl.load(dt_bias + i_hv * K + o_k, mask=mask_k, other=0).to(tl.float32)
b_g = b_g + b_bias
if USE_LOWER_BOUND:
b_gk = lower_bound * tl.sigmoid((exp(b_A) if HAS_A else b_A) * b_g)
else:
b_gk = -exp(b_A) * softplus(b_g)
else:
b_gk = b_g
if STATE_V_FIRST:
b_h *= exp(b_gk[None, :])
else:
b_h *= exp(b_gk[:, None])
if STATE_V_FIRST:
b_v -= tl.sum(b_h * b_k[None, :], 1)
else:
b_v -= tl.sum(b_h * b_k[:, None], 0)
if IS_BETA_HEADWISE:
b_beta = tl.load(p_beta, mask=mask_v, other=0, eviction_policy='evict_first').to(tl.float32)
else:
b_beta = tl.load(p_beta, eviction_policy='evict_last').to(tl.float32)
if APPLY_BETA_SIGMOID:
b_beta = tl.sigmoid(b_beta)
if ALLOW_NEG_EIGVAL:
b_beta = b_beta * 2
b_v *= b_beta
if STATE_V_FIRST:
b_h += b_v[:, None] * b_k[None, :]
b_o = tl.sum(b_h * b_q[None, :], 1)
else:
b_h += b_k[:, None] * b_v[None, :]
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, eviction_policy='evict_first')
if IS_CONTINUOUS_BATCHING:
if INPLACE_FINAL_STATE:
p_ht = (
ht
+ tl.load(ssm_state_indices + i_n * stride_indices_seq + i_t).to(
tl.int64
)
* stride_final_state_token
)
else:
p_ht = ht + (bos + i_t) * stride_final_state_token
if STATE_V_FIRST:
p_ht = p_ht + i_hv * K * V + o_v[:, None] * K + o_k[None, :]
else:
p_ht = p_ht + i_hv * K * V + o_k[:, None] * V + o_v[None, :]
tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), mask=mask_h)
p_q += H * K
p_k += H * K
p_o += HV * V
p_v += HV * V
p_g += HV * K
p_beta += HV * (V if IS_BETA_HEADWISE else 1)
if not IS_CONTINUOUS_BATCHING:
if STORE_FINAL_STATE:
if STATE_V_FIRST:
p_ht = ht + (i_n * HV + i_hv) * K * V + o_v[:, None] * K + o_k[None, :]
else:
p_ht = ht + (i_n * HV + i_hv) * K * V + o_k[:, None] * V + o_v[None, :]
tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), mask=mask_h)
@dispatch("kda")
def fused_recurrent_kda_fwd(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
initial_state: torch.Tensor | None = None,
scale: float | None = None,
output_final_state: bool = False,
inplace_final_state: bool = True,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
ssm_state_indices: torch.Tensor | None = None,
num_accepted_tokens: torch.Tensor | None = None,
use_qk_l2norm_in_kernel: bool = False,
use_gate_in_kernel: bool = False,
use_beta_sigmoid_in_kernel: bool = False,
allow_neg_eigval: bool = False,
lower_bound: float | None = None,
out: torch.Tensor | None = None,
**kwargs,
) -> tuple[torch.Tensor, torch.Tensor]:
if scale is None:
scale = k.shape[-1] ** -0.5
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 = 32
if out is None:
out = torch.zeros_like(v)
else:
assert out.shape == v.shape
if inplace_final_state:
assert initial_state is not None
final_state = initial_state
elif 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
stride_init_state_token = initial_state.stride(0) if initial_state is not None else 1
stride_final_state_token = final_state.stride(0) if final_state is not None else 1
if ssm_state_indices is None:
stride_indices_seq, stride_indices_tok = 1, 1
elif ssm_state_indices.ndim == 1:
stride_indices_seq, stride_indices_tok = ssm_state_indices.stride(0), 1
else:
stride_indices_seq, stride_indices_tok = ssm_state_indices.stride()
grid = (triton.cdiv(V, BV) * N * HV, )
fused_recurrent_kda_fwd_kernel[grid](
q=q,
k=k,
v=v,
g=g,
beta=beta,
A_log=A_log,
dt_bias=dt_bias,
o=out,
h0=initial_state,
ht=final_state,
cu_seqlens=cu_seqlens,
ssm_state_indices=ssm_state_indices,
num_accepted_tokens=num_accepted_tokens,
lower_bound=lower_bound,
scale=scale,
N=N,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BK=BK,
BV=BV,
stride_init_state_token=stride_init_state_token,
stride_final_state_token=stride_final_state_token,
stride_indices_seq=stride_indices_seq,
stride_indices_tok=stride_indices_tok,
IS_BETA_HEADWISE=beta.ndim == v.ndim,
USE_QK_L2NORM_IN_KERNEL=use_qk_l2norm_in_kernel,
INPLACE_FINAL_STATE=inplace_final_state,
USE_GATE_IN_KERNEL=use_gate_in_kernel,
APPLY_BETA_SIGMOID=use_beta_sigmoid_in_kernel,
ALLOW_NEG_EIGVAL=allow_neg_eigval,
STATE_V_FIRST=state_v_first,
num_warps=4,
num_stages=2,
)
return out, final_state
@input_guard
def fused_recurrent_kda(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
scale: float | None = None,
initial_state: torch.Tensor = None,
output_final_state: bool = False,
use_qk_l2norm_in_kernel: bool = False,
use_gate_in_kernel: bool = False,
use_beta_sigmoid_in_kernel: bool = False,
allow_neg_eigval: bool = False,
lower_bound: float | None = None,
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 is applied if `HV > H`.
g (torch.Tensor):
g (decays) of shape `[B, T, HV, K]`.
beta (torch.Tensor):
betas of shape `[B, T, HV]`.
A_log (Optional[torch.Tensor]):
Decay parameter of shape `[HV]`.
When `use_gate_in_kernel=True` together with `lower_bound`,
may be `None` to use `lower_bound * sigmoid(g + dt_bias)`.
dt_bias (Optional[torch.Tensor]):
Bias added to `g` before activation, of shape `[HV]`. Only used when `use_gate_in_kernel=True`.
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 (Optional[bool]):
Whether to compute the log-space KDA decay internally.
When `True`, `g` is the raw input and the kernel fuses gate activation into the recurrence.
Default: `False`.
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`.
lower_bound (Optional[float]):
Lower bound for the forget gate (in log space). Only used when `use_gate_in_kernel=True`. Default: `None`.
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.kda import fused_recurrent_kda
# 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, K, device='cuda'))
>>> beta = torch.rand(B, T, HV, device='cuda').sigmoid()
>>> h0 = torch.randn(B, HV, K, V, device='cuda')
>>> o, ht = fused_recurrent_kda(
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_var, ht_var = fused_recurrent_kda(
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 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 = fused_recurrent_kda_fwd(
q=q,
k=k,
v=v,
g=g,
beta=beta,
A_log=A_log,
dt_bias=dt_bias,
scale=scale,
initial_state=initial_state,
inplace_final_state=False,
output_final_state=output_final_state,
use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel,
use_gate_in_kernel=use_gate_in_kernel,
use_beta_sigmoid_in_kernel=use_beta_sigmoid_in_kernel,
allow_neg_eigval=allow_neg_eigval,
lower_bound=lower_bound,
cu_seqlens=cu_seqlens,
state_v_first=state_v_first,
)
return o, final_state
+514
View File
@@ -0,0 +1,514 @@
# 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
# This file is modified and supported by the Moonshot AI Team
import torch
import torch.nn.functional as F
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.ops.utils.index import prepare_chunk_indices
from kda._fla.ops.utils.op import exp
from kda._fla.ops.utils.softplus import softplus
from kda._fla.utils import IS_AMD, autocast_custom_bwd, autocast_custom_fwd, autotune_cache_kwargs, check_shared_mem, input_guard
BS_LIST = [32, 64] if check_shared_mem() else [16, 32]
BT_LIST_AUTOTUNE = [32, 64, 128]
NUM_WARPS_AUTOTUNE = [2, 4, 8, 16] if IS_AMD else [4, 8, 16, 32]
def naive_kda_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 KDA gate computation.
Computes: g = -A_log.exp().unsqueeze(-1) * softplus(g + dt_bias.view(g.shape[-2:]))
Args:
g (torch.Tensor):
Input tensor of shape `[..., H, K]`.
A_log (torch.Tensor):
Parameter tensor with `H` elements.
dt_bias (torch.Tensor | None):
Optional bias tensor added to `g` before activation, shape `[H * K]`.
Returns:
Output tensor of shape `[..., H, K]` .
"""
H, _ = g.shape[-2:]
g = g.float()
if dt_bias is not None:
g = g + dt_bias.view(H, -1)
g = (-A_log.view(H, 1).float().exp() * F.softplus(g.float())).to(output_dtype)
return g
def naive_kda_lowerbound_gate(
g: torch.Tensor,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
lower_bound: float = -5.0,
output_dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
"""
Torch reference implementation for KDA lowerbound gate computation.
Computes: ``g = lower_bound * sigmoid(exp(A_log) * (g + dt_bias))``.
When ``A_log`` is ``None``: ``g = lower_bound * sigmoid(g + dt_bias)``.
Args:
g (torch.Tensor):
Input tensor of shape `[..., H, K]`.
A_log (torch.Tensor | None):
Optional parameter tensor with `H` elements.
dt_bias (torch.Tensor | None):
Optional bias tensor added to `g` before activation, shape `[H * K]`.
lower_bound (float):
Lower bound for the gate output. Default: `-5.0`.
output_dtype (torch.dtype):
The dtype of the output tensor. Default: `torch.float32`.
Returns:
Output tensor of shape `[..., H, K]`.
"""
H, _ = g.shape[-2:]
g = g.float()
if dt_bias is not None:
g = g + dt_bias.view(H, -1)
if A_log is not None:
g = A_log.view(H, 1).float().exp() * g
g = lower_bound * F.sigmoid(g)
return g.to(output_dtype)
@triton.heuristics({
"HAS_A": lambda args: args["A_log"] is not None,
"HAS_BIAS": lambda args: args["dt_bias"] is not None,
"HAS_BETA": lambda args: args["beta"] is not None,
'USE_LOWER_BOUND': lambda args: args['lower_bound'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({"BT": BT}, num_warps=num_warps, num_stages=num_stages)
for BT in BT_LIST_AUTOTUNE
for num_warps in NUM_WARPS_AUTOTUNE
for num_stages in [2, 3]
],
key=["H", "D"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def kda_gate_fwd_kernel(
g,
A_log,
dt_bias,
beta,
yg,
yb,
lower_bound,
T,
H: tl.constexpr,
D: tl.constexpr,
BT: tl.constexpr,
BD: tl.constexpr,
HAS_A: tl.constexpr,
HAS_BIAS: tl.constexpr,
HAS_BETA: tl.constexpr,
USE_LOWER_BOUND: 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) if HAS_A else 1.0
o_t = i_t * BT + tl.arange(0, BT)
o_d = tl.arange(0, BD)
m_t = o_t < T
m_g = m_t[:, None] & (o_d[None, :] < D)
p_g = g + i_h * D + o_t[:, None] * (H * D) + o_d[None, :]
p_yg = yg + i_h * D + o_t[:, None] * (H * D) + o_d[None, :]
# [BT, BD]
b_g = tl.load(p_g, mask=m_g, other=0.0).to(tl.float32)
if HAS_BIAS:
o_b = i_h * D + tl.arange(0, BD)
b_g = b_g + tl.load(dt_bias + o_b, mask=o_b < H * D, other=0.0).to(tl.float32)
if not USE_LOWER_BOUND:
b_yg = -exp(b_A) * softplus(b_g)
else:
b_yg = lower_bound * tl.sigmoid((exp(b_A) if HAS_A else b_A) * b_g)
tl.store(p_yg, b_yg.to(p_yg.dtype.element_ty), mask=m_g)
if HAS_BETA:
p_b = beta + i_h + o_t * H
p_yb = yb + i_h + o_t * H
b_yb = tl.sigmoid(tl.load(p_b, mask=m_t, other=0.0).to(tl.float32))
tl.store(p_yb, b_yb.to(p_yb.dtype.element_ty), mask=m_t)
@triton.heuristics({
"HAS_A": lambda args: args["A_log"] is not None,
"HAS_BIAS": lambda args: args["dt_bias"] is not None,
"HAS_BETA": lambda args: args["beta"] is not None,
'USE_LOWER_BOUND': lambda args: args['lower_bound'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({}, num_warps=num_warps, num_stages=num_stages)
for num_warps in NUM_WARPS_AUTOTUNE
for num_stages in [2, 3]
],
key=["H", "D"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def kda_gate_bwd_kernel(
g,
A_log,
dt_bias,
beta,
dyg,
dyb,
dg,
dA,
dbeta,
lower_bound,
T,
H: tl.constexpr,
D: tl.constexpr,
BT: tl.constexpr,
BD: tl.constexpr,
HAS_A: tl.constexpr,
HAS_BIAS: tl.constexpr,
HAS_BETA: tl.constexpr,
USE_LOWER_BOUND: 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) if HAS_A else 1.0
o_t = i_t * BT + tl.arange(0, BT)
o_d = tl.arange(0, BD)
m_t = o_t < T
m_g = m_t[:, None] & (o_d[None, :] < D)
p_g = g + i_h * D + o_t[:, None] * (H * D) + o_d[None, :]
p_dg = dg + i_h * D + o_t[:, None] * (H * D) + o_d[None, :]
p_dyg = dyg + i_h * D + o_t[:, None] * (H * D) + o_d[None, :]
# [BT, BD]
b_g = tl.load(p_g, mask=m_g, other=0.0).to(tl.float32)
b_dyg = tl.load(p_dyg, mask=m_g, other=0.0).to(tl.float32)
if HAS_BIAS:
o_b = i_h * D + tl.arange(0, BD)
b_g = b_g + tl.load(dt_bias + o_b, mask=o_b < H * D, other=0.0).to(tl.float32)
# [BT, BD]
if not USE_LOWER_BOUND:
b_A = -exp(b_A)
b_yg = b_A * softplus(b_g)
b_dg = b_A * (b_dyg * tl.sigmoid(b_g))
b_dA = tl.sum(tl.sum(b_dyg * b_yg, 1), 0)
else:
b_A = exp(b_A) if HAS_A else b_A
b_inner = b_A * b_g
b_sig = tl.sigmoid(b_inner)
b_dsig = b_sig * (1.0 - b_sig)
# Common term: dy * (LB * dsig)
b_d_inner_term = b_dyg * (lower_bound * b_dsig)
# dg = d_inner_term * A
b_dg = b_d_inner_term * b_A
b_dA = tl.sum(tl.sum(b_dg * b_g, 1), 0) if HAS_A else 0.0
tl.store(p_dg, b_dg.to(p_dg.dtype.element_ty), mask=m_g)
if HAS_A:
tl.store(dA + i_t * H + i_h, b_dA)
if HAS_BETA:
p_b = beta + i_h + o_t * H
p_db = dbeta + i_h + o_t * H
p_dyb = dyb + i_h + o_t * H
b_b = tl.load(p_b, mask=m_t, other=0.0).to(tl.float32)
b_db = tl.load(p_dyb, mask=m_t, other=0.0).to(tl.float32) * b_b * (1.0 - b_b)
tl.store(p_db, b_db.to(p_db.dtype.element_ty), mask=m_t)
@dispatch('kda')
def kda_gate_fwd(
g: torch.Tensor,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
lower_bound: float | None = None,
output_dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
H, K = g.shape[-2:]
T = g.numel() // (H * K)
yg = torch.empty_like(g, dtype=output_dtype)
def grid(meta):
return (triton.cdiv(T, meta["BT"]), H)
kda_gate_fwd_kernel[grid](
g=g,
A_log=A_log,
dt_bias=dt_bias,
beta=None,
yg=yg,
yb=None,
T=T,
H=H,
D=K,
BD=triton.next_power_of_2(K),
lower_bound=lower_bound,
)
return yg
@dispatch('kda')
def kda_gate_bwd(
g: torch.Tensor,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
dyg: torch.Tensor | None = None,
lower_bound: float | None = None,
) -> tuple[torch.Tensor, torch.Tensor | None, torch.Tensor | None]:
H, K = g.shape[-2:]
T = g.numel() // (H * K)
BT = 32
NT = triton.cdiv(T, BT)
dg = torch.empty_like(g, dtype=torch.float32)
dA = g.new_empty(NT, H, dtype=torch.float32) if A_log is not None else None
grid = (triton.cdiv(T, BT), H)
kda_gate_bwd_kernel[grid](
g=g,
A_log=A_log,
dt_bias=dt_bias,
beta=None,
dyg=dyg,
dyb=None,
dg=dg,
dA=dA,
dbeta=None,
T=T,
H=H,
D=K,
BT=BT,
BD=triton.next_power_of_2(K),
lower_bound=lower_bound,
)
dg = dg.view_as(g).type_as(g)
dA = dA.sum(0).view_as(A_log).type_as(A_log) if A_log is not None else None
# dt_bias is [HV, K] in KDAAttention and [HV*K] in some FLA call sites.
dbias = (
dg.view(-1, H * K).sum(0).reshape_as(dt_bias).type_as(dt_bias)
if dt_bias is not None
else None
)
return dg, dA, dbias
class KDAGateFunction(torch.autograd.Function):
@staticmethod
@input_guard
@autocast_custom_fwd
def forward(
ctx,
g: torch.Tensor,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
lower_bound: float | None = None,
output_dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
yg = kda_gate_fwd(
g=g,
A_log=A_log,
dt_bias=dt_bias,
lower_bound=lower_bound,
output_dtype=output_dtype
)
ctx.save_for_backward(g, A_log, dt_bias)
ctx.lower_bound = lower_bound
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 = kda_gate_bwd(
g=g,
A_log=A_log,
dt_bias=dt_bias,
dyg=dyg,
lower_bound=ctx.lower_bound
)
return dg, dA, dbias, None, None
@dispatch('kda')
@torch.compiler.disable
def fused_kda_gate(
g: torch.Tensor,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
lower_bound: float | None = None,
output_dtype: torch.dtype = torch.float32,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
"""
Fused KDA gate computation with autograd support.
Computes: g = -A_log.exp().unsqueeze(-1) * softplus(g + dt_bias.view(g.shape[-2:]))
When ``lower_bound`` is set: g = lower_bound * sigmoid(exp(A_log) * (g + dt_bias)).
When ``A_log`` is ``None`` (requires ``lower_bound``): g = lower_bound * sigmoid(g + dt_bias).
Args:
g (torch.Tensor):
Input tensor of shape `[..., H, K]`.
A_log (torch.Tensor | None):
Optional parameter tensor with `H` elements.
When ``None``, the gate reduces to ``lower_bound * sigmoid(g + dt_bias)`` (requires ``lower_bound``).
dt_bias (torch.Tensor | None):
Optional bias tensor added to `g` before activation, shape `[H * K]`.
Returns:
Output tensor of shape `[..., H, K]`.
"""
return KDAGateFunction.apply(g, A_log, dt_bias, lower_bound, output_dtype)
@triton.heuristics({
"HAS_A": lambda args: args["A_log"] is not None,
"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,
'USE_LOWER_BOUND': lambda args: args['lower_bound'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({'BS': BS}, num_warps=num_warps)
for BS in BS_LIST
for num_warps in [2, 4, 8]
],
key=['H', 'S', 'BT', 'IS_VARLEN', 'REVERSE'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def kda_gate_chunk_cumsum_vector_kernel(
s,
A_log,
dt_bias,
o,
scale,
cu_seqlens,
chunk_indices,
lower_bound,
T,
H: tl.constexpr,
S: tl.constexpr,
BT: tl.constexpr,
BS: tl.constexpr,
REVERSE: tl.constexpr,
HAS_A: tl.constexpr,
HAS_BIAS: tl.constexpr,
HAS_SCALE: tl.constexpr,
IS_VARLEN: tl.constexpr,
USE_LOWER_BOUND: tl.constexpr,
):
i_s, i_t, i_bh = tl.program_id(0), tl.program_id(1).to(tl.int64), tl.program_id(2).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)
o_s = i_s * BS + tl.arange(0, BS)
m_s = (o_t[:, None] < T) & (o_s[None, :] < S)
p_s = s + (bos * H + i_h) * S + o_t[:, None] * (H*S) + o_s[None, :]
p_o = o + (bos * H + i_h) * S + o_t[:, None] * (H*S) + o_s[None, :]
# [BT, BS]
b_s = tl.load(p_s, mask=m_s, other=0.0).to(tl.float32)
# Apply dt_bias if exists
if HAS_BIAS:
b_bias = tl.load(dt_bias + i_h * S + o_s, mask=o_s < S, other=0.0).to(tl.float32)
b_s = b_s + b_bias[None, :]
b_A = tl.load(A_log + i_h).to(tl.float32) if HAS_A else 1.0
if not USE_LOWER_BOUND:
# Apply gate: -exp(A_log) * softplus(g + bias)
b_gate = -exp(b_A) * softplus(b_s)
else:
b_gate = lower_bound * tl.sigmoid((exp(b_A) if HAS_A else b_A) * b_s)
# Apply chunk local cumsum
if REVERSE:
b_o = tl.cumsum(b_gate, axis=0, reverse=True)
else:
b_o = tl.cumsum(b_gate, axis=0)
if HAS_SCALE:
b_o *= scale
tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=m_s)
@input_guard
@dispatch('kda')
def kda_gate_chunk_cumsum(
g: torch.Tensor,
A_log: torch.Tensor | None,
chunk_size: int,
scale: float = None,
dt_bias: torch.Tensor | None = None,
cu_seqlens: torch.Tensor | None = None,
output_dtype: torch.dtype | None = torch.float,
chunk_indices: torch.LongTensor | None = None,
lower_bound: float | None = None,
**kwargs,
) -> torch.Tensor:
if cu_seqlens is not None:
assert g.shape[0] == 1, "Only batch size 1 is supported when cu_seqlens are provided"
assert len(g.shape) == 4
B, T, H, S = 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)
assert chunk_size == 2**(chunk_size.bit_length()-1), "chunk_size must be a power of 2"
g_org, g = g, torch.empty_like(g, dtype=output_dtype or g.dtype)
def grid(meta): return (triton.cdiv(meta['S'], meta['BS']), NT, B * H)
kda_gate_chunk_cumsum_vector_kernel[grid](
s=g_org,
A_log=A_log,
dt_bias=dt_bias,
o=g,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
lower_bound=lower_bound,
T=T,
H=H,
S=S,
BT=BT,
REVERSE=False,
)
return g
+369
View File
@@ -0,0 +1,369 @@
# 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 kda._fla.ops.backends import dispatch
from kda._fla.ops.utils import prepare_chunk_indices
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.ops.utils.op import exp2
from kda._fla.utils import autotune_cache_kwargs, check_shared_mem
@triton.heuristics({
'STORE_QG': lambda args: args['qg'] is not None,
'STORE_KG': lambda args: args['kg'] 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 [2, 4, 8]
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_kda_kernel(
q,
k,
qg,
kg,
v,
beta,
w,
u,
A,
gk,
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,
STORE_QG: tl.constexpr,
STORE_KG: 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_hv = i_bh // HV, i_bh % HV
i_h = i_hv // (HV // 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
k += (bos * H + i_h) * K
v += (bos * HV + i_hv) * V
u += (bos * HV + i_hv) * V
w += (bos * HV + i_hv) * K
gk += (bos * HV + i_hv) * K
beta += bos * HV + i_hv
A += (bos * HV + i_hv) * BT
if STORE_QG:
q += (bos * H + i_h) * K
qg += (bos * HV + i_hv) * K
if STORE_KG:
kg += (bos * HV + i_hv) * K
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
p_b = beta + o_t * HV
b_b = tl.load(p_b, mask=m_t, other=0.0)
o_A = tl.arange(0, BT)
m_A = m_t[:, None] & (o_A[None, :] < BT)
p_A = A + 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 + o_t[:, None] * (HV*V) + o_v[None, :]
p_u = u + 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)
tl.store(p_u, b_u.to(p_u.dtype.element_ty), mask=m_v)
for i_k in range(tl.cdiv(K, BK)):
o_k = i_k * BK + tl.arange(0, BK)
m_k = o_k < K
m_tk = m_t[:, None] & m_k[None, :]
p_w = w + o_t[:, None] * (HV*K) + o_k[None, :]
p_k = k + o_t[:, None] * (H*K) + o_k[None, :]
b_k = tl.load(p_k, mask=m_tk, other=0.0)
b_kb = b_k * b_b[:, None]
p_gk = gk + o_t[:, None] * (HV*K) + o_k[None, :]
b_gk = tl.load(p_gk, mask=m_tk, other=0.0).to(tl.float32)
b_kb *= exp2(b_gk)
if STORE_QG:
p_q = q + o_t[:, None] * (H*K) + o_k[None, :]
p_qg = qg + o_t[:, None] * (HV*K) + o_k[None, :]
b_q = tl.load(p_q, mask=m_tk, other=0.0)
b_qg = b_q * exp2(b_gk)
tl.store(p_qg, b_qg.to(p_qg.dtype.element_ty), mask=m_tk)
if STORE_KG:
last_idx = min(i_t * BT + BT, T) - 1
b_gn = tl.load(gk + last_idx * HV*K + o_k, mask=m_k, other=0.).to(tl.float32)
b_kg = b_k * tl.where((i_t * BT + tl.arange(0, BT) < T)[:, None], exp2(b_gn[None, :] - b_gk), 0)
p_kg = kg + o_t[:, None] * (HV*K) + o_k[None, :]
tl.store(p_kg, b_kg.to(p_kg.dtype.element_ty), mask=m_tk)
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_tk)
@triton.heuristics({
'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 [2, 4]
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 prepare_wy_repr_bwd_kda_kernel(
k,
v,
beta,
gk,
A,
dA,
dw,
du,
dk,
dk2,
dv,
db,
dg,
dg2,
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,
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_hv = i_bh // HV, i_bh % HV
i_h = i_hv // (HV // 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
k += (bos * H + i_h) * K
v += (bos * HV + i_hv) * V
beta += bos * HV + i_hv
gk += (bos * HV + i_hv) * K
A += (bos * HV + i_hv) * BT
dA += (bos * HV + i_hv) * BT
dk += (bos * HV + i_hv) * K
dk2 += (bos * HV + i_hv) * K
dw += (bos * HV + i_hv) * K
du += (bos * HV + i_hv) * V
dv += (bos * HV + i_hv) * V
db += bos * HV + i_hv
dg += (bos * HV + i_hv) * K
dg2 += (bos * HV + i_hv) * K
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
p_b = beta + o_t * HV
p_db = db + o_t * HV
o_A = tl.arange(0, BT)
m_AT = (o_A[:, None] < BT) & m_t[None, :]
p_A = A + 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)
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 + o_t[:, None] * (H*K) + o_k[None, :]
p_dk = dk + o_t[:, None] * (HV*K) + o_k[None, :]
p_dk2 = dk2 + o_t[:, None] * (HV*K) + o_k[None, :]
p_dw = dw + o_t[:, None] * (HV*K) + o_k[None, :]
p_dg = dg + o_t[:, None] * (HV*K) + o_k[None, :]
p_dg2 = dg2 + o_t[:, None] * (HV*K) + o_k[None, :]
# [BT, BK]
b_k = tl.load(p_k, mask=m_k, other=0.0)
p_gk = gk + o_t[:, None] * (HV*K) + o_k[None, :]
b_gk_exp = exp2(tl.load(p_gk, mask=m_k, other=0.0))
b_kbg = b_k * b_b[:, None] * b_gk_exp
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)
b_dk = b_dkbg * b_gk_exp * b_b[:, None] + tl.load(p_dk, mask=m_k, other=0.0)
b_db += tl.sum(b_dkbg * b_k * b_gk_exp, 1)
b_dg = b_kbg * b_dkbg + tl.load(p_dg, mask=m_k, other=0.0)
tl.store(p_dk2, b_dk.to(p_dk2.dtype.element_ty), mask=m_k)
tl.store(p_dg2, b_dg.to(p_dg2.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 + o_t[:, None] * (HV*V) + o_v[None, :]
p_dv = dv + o_t[:, None] * (HV*V) + o_v[None, :]
p_du = du + 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))
b_dA = tl.where(m_A, -b_dA, 0)
m_dA = m_t[:, None] & (o_A[None, :] < BT)
p_dA = dA + o_t[:, None] * (HV*BT) + o_A[None, :]
tl.store(p_dA, b_dA.to(p_dA.dtype.element_ty), mask=m_dA)
tl.store(p_db, b_db.to(p_db.dtype.element_ty), mask=m_t)
@dispatch('kda')
def recompute_w_u_fwd(
k: torch.Tensor,
v: torch.Tensor,
beta: torch.Tensor,
A: torch.Tensor,
gk: torch.Tensor,
q: torch.Tensor | None = None,
cu_seqlens: torch.LongTensor | None = None,
chunk_indices: torch.LongTensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None, torch.Tensor | None]:
B, T, H, K, V = *k.shape, v.shape[-1]
HV = 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 = torch.empty(B, T, HV, K, device=k.device, dtype=k.dtype)
u = torch.empty_like(v)
qg = torch.empty(B, T, HV, K, device=k.device, dtype=k.dtype) if q is not None else None
kg = torch.empty(B, T, HV, K, device=k.device, dtype=k.dtype)
recompute_w_u_fwd_kda_kernel[(NT, B*HV)](
q=q,
k=k,
qg=qg,
kg=kg,
v=v,
beta=beta,
w=w,
u=u,
A=A,
gk=gk,
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, qg, kg
def prepare_wy_repr_bwd(
k: torch.Tensor,
v: torch.Tensor,
beta: torch.Tensor,
gk: torch.Tensor,
A: torch.Tensor,
dk: torch.Tensor,
dw: torch.Tensor,
du: torch.Tensor,
dg: torch.Tensor,
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 = *k.shape, v.shape[-1]
HV = 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)
dk2 = torch.empty_like(dk, dtype=torch.float)
dv = torch.empty_like(v)
dg2 = torch.empty_like(gk, dtype=torch.float)
dA = torch.empty_like(A, dtype=torch.float)
db = torch.empty_like(beta, dtype=torch.float)
prepare_wy_repr_bwd_kda_kernel[(NT, B * HV)](
k=k,
v=v,
beta=beta,
gk=gk,
A=A,
dA=dA,
dw=dw,
du=du,
dk=dk,
dk2=dk2,
dv=dv,
db=db,
dg=dg,
dg2=dg2,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BT=BT,
BK=BK,
BV=BV,
)
dk = dk2
dg = dg2
return dk, dv, db, dg, dA
+14
View File
@@ -0,0 +1,14 @@
from .cumsum import (
chunk_local_cumsum,
chunk_local_cumsum_scalar,
chunk_local_cumsum_vector,
)
from .index import prepare_chunk_indices, prepare_chunk_offsets
__all__ = [
"chunk_local_cumsum",
"chunk_local_cumsum_scalar",
"chunk_local_cumsum_vector",
"prepare_chunk_indices",
"prepare_chunk_offsets",
]
+449
View File
@@ -0,0 +1,449 @@
# 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 dataclasses
import enum
import json
import logging
import os
import re
from functools import cache, lru_cache
from pathlib import Path
from typing import Any
import torch
import triton
from packaging import version
from triton.runtime.autotuner import Autotuner
TRITON_ABOVE_3_5_1 = version.parse(triton.__version__) >= version.parse("3.5.1")
TRITON_ABOVE_3_4_0 = version.parse(triton.__version__) >= version.parse("3.4.0")
class FlaCacheMode(enum.Enum):
"""Controls how FLA loads kernel configs from its config cache (FLA_CACHE_MODE env var).
DISABLED — skip all cache lookups, always fall back to Triton autotune (default when FLA_CACHE_MODE is unset)
STRICT — exact key match only; falls back to Triton autotune if no match
FUZZY — exact key match → fuzzy key match; falls back to Triton autotune if no match
FULL — exact key match → fuzzy key match → default_config fallback
DEFAULT — use only the top-level default_config field, skip key-based lookup
ALWAYS — like DEFAULT, but re-reads config files on every kernel call;
useful for debugging: edit default_config in a JSON file and the next
kernel call picks it up without restarting the process
"""
DISABLED = "disabled"
STRICT = "strict"
FUZZY = "fuzzy"
FULL = "full"
DEFAULT = "default"
ALWAYS = "always"
def uses_default_config(self) -> bool:
"""Return True for modes that may fall back to default_config (FULL, DEFAULT, ALWAYS)."""
return self in (FlaCacheMode.FULL, FlaCacheMode.DEFAULT, FlaCacheMode.ALWAYS)
@classmethod
def from_env(cls) -> "FlaCacheMode":
mode_str = os.environ.get("FLA_CACHE_MODE", cls.DISABLED.value)
try:
return cls(mode_str)
except ValueError:
valid = [m.value for m in cls]
raise ValueError(
f"Invalid FLA_CACHE_MODE={mode_str!r}. Valid values: {valid}"
) from None
FLA_CACHE_MODE: FlaCacheMode = FlaCacheMode.from_env()
logger = logging.getLogger(__name__)
def sanitize_gpu_name(gpu_name: str) -> str:
sanitized = re.sub(r"[^0-9A-Za-z]+", "_", gpu_name)
sanitized = sanitized.strip("_")
return sanitized or "unknown_gpu"
@lru_cache(maxsize=1)
def get_gpu_info():
"""Get GPU model information.
This function detects the GPU model and returns a sanitized string identifier.
It prioritizes FLA_GPU_NAME environment variable if set, then detects from
available hardware (CUDA, ROCm, Intel GPU, or CPU).
"""
# Check if GPU name is overridden via environment variable
gpu_name = None
# Check if GPU name is overridden via environment variable
if "FLA_GPU_NAME" in os.environ:
gpu_name = os.environ["FLA_GPU_NAME"]
# Try to get device name based on availability
elif torch.cuda.is_available():
# Works for both NVIDIA and AMD GPUs (ROCm)
gpu_name = torch.cuda.get_device_name(0)
elif hasattr(torch, 'xpu') and torch.xpu.is_available():
gpu_name = torch.xpu.get_device_name(0)
if gpu_name:
return sanitize_gpu_name(gpu_name)
# Default to CPU if no GPU available
return "cpu"
def get_fla_config_dir() -> Path:
"""Get FLA's configs directory.
The directory can be overridden by setting the FLA_CONFIG_DIR environment variable.
If set, configs will be loaded directly from $FLA_CONFIG_DIR/. Otherwise FLA
falls back to the default fla/configs/{GPU}/ directory in the project.
"""
# Check if custom config dir is set via environment variable
if "FLA_CONFIG_DIR" in os.environ:
return Path(os.environ["FLA_CONFIG_DIR"])
# Default: project_dir/fla/configs/{GPU}/
project_dir = Path(__file__).parent.parent.parent
return project_dir / "configs" / get_gpu_info()
@dataclasses.dataclass(frozen=True)
class AutotuneKey:
"""Autotune key with exact/fuzzy matching, serialization, and construction helpers."""
autotune_key: tuple[Any, ...]
@staticmethod
def normalize_autotune_key(value: Any) -> Any:
if isinstance(value, (list, tuple)):
return [AutotuneKey.normalize_autotune_key(v) for v in value]
if isinstance(value, dict):
return {k: AutotuneKey.normalize_autotune_key(v) for k, v in value.items()}
return value
@staticmethod
def serialize(key: Any) -> str:
return json.dumps(AutotuneKey.normalize_autotune_key(key), separators=(",", ":"), sort_keys=True)
@staticmethod
def key_hash(key: Any) -> str:
import hashlib
return hashlib.md5(AutotuneKey.serialize(key).encode()).hexdigest()
@staticmethod
def is_numeric(value: Any) -> bool:
return isinstance(value, (int, float)) and not isinstance(value, bool)
@staticmethod
def keys_fuzzy_match(cached_key: Any, requested_key: Any) -> bool:
# Fuzzy match: numeric leaves are compatible regardless of their actual numeric values
# (e.g. a config tuned for seq_len=1024 can apply to seq_len=2048).
# Structure (type, length, dict keys) must still match exactly.
if AutotuneKey.is_numeric(cached_key) and AutotuneKey.is_numeric(requested_key):
return True
if isinstance(cached_key, (list, tuple)) and isinstance(requested_key, (list, tuple)):
return len(cached_key) == len(requested_key) and all(
AutotuneKey.keys_fuzzy_match(c, r) for c, r in zip(cached_key, requested_key)
)
if isinstance(cached_key, dict) and isinstance(requested_key, dict):
return cached_key.keys() == requested_key.keys() and all(
AutotuneKey.keys_fuzzy_match(cached_key[k], requested_key[k]) for k in cached_key
)
return cached_key == requested_key
@classmethod
def build(
cls,
arg_names: list[str],
key_names: list[str],
positional_args: tuple[Any, ...],
runtime_kwargs: dict[str, Any],
) -> "AutotuneKey":
named_args = dict(zip(arg_names, positional_args))
all_args = {**named_args, **runtime_kwargs}
tracked_args = {k: v for (k, v) in all_args.items() if k in arg_names}
tuning_key = [tracked_args[name] for name in key_names if name in tracked_args]
for arg in tracked_args.values():
if hasattr(arg, "dtype"):
tuning_key.append(str(arg.dtype))
return cls(autotune_key=tuple(tuning_key))
def exact_matches(self, entry_key: Any) -> bool:
return self.serialize(self.autotune_key) == self.serialize(entry_key)
def fuzzy_matches(self, entry_key: Any) -> bool:
self_normalized = self.normalize_autotune_key(self.autotune_key)
entry_normalized = self.normalize_autotune_key(entry_key)
return (
isinstance(self_normalized, list)
and isinstance(entry_normalized, list)
and len(self_normalized) == len(entry_normalized)
and AutotuneKey.keys_fuzzy_match(self_normalized, entry_normalized)
)
@dataclasses.dataclass(frozen=True)
class KernelConfigFile:
"""Validated in-memory representation of a {kernel_name}.json config file."""
kernel_name: str | None
triton_version: str | None
autotune_entries: dict[str, dict[str, Any]] | None
default_config: dict[str, Any] | None
@classmethod
def from_dict(cls, config_file: Path, data: Any) -> "KernelConfigFile | None":
"""Parse and validate a raw JSON dict. Returns None (with a warning) if malformed."""
def fail(msg, *args):
logger.warning(msg, *args)
raise ValueError
try:
if not isinstance(data, dict):
fail("Malformed config %s: root is %s, expected dict", config_file, type(data).__name__)
raw_entries = data.get("autotune_entries")
entries: dict[str, dict[str, Any]] | None = None
if raw_entries is not None:
if not isinstance(raw_entries, dict):
fail("Malformed config %s: 'autotune_entries' is %s, expected dict",
config_file, type(raw_entries).__name__)
for h, entry in raw_entries.items():
if not isinstance(entry, dict):
fail("Malformed config %s: autotune_entries[%r] is %s, expected dict",
config_file, h, type(entry).__name__)
if not isinstance(entry.get("config"), dict):
fail("Malformed config %s: autotune_entries[%r] missing valid 'config' field", config_file, h)
entries = raw_entries
default_config = data.get("default_config")
if default_config is not None and not isinstance(default_config, dict):
fail("Malformed config %s: 'default_config' is %s, expected dict", config_file, type(default_config).__name__)
return cls(
kernel_name=data.get("kernel_name"),
triton_version=data.get("triton_version"),
autotune_entries=entries,
default_config=default_config,
)
except ValueError:
return None
@classmethod
def from_file(cls, config_file: Path) -> "KernelConfigFile | None":
"""Read and validate a config file. Returns None if the file is missing or malformed."""
config_data = read_config_file(config_file)
if config_data is None:
return None
return cls.from_dict(config_file, config_data)
def lookup_exact(self, key: AutotuneKey) -> dict[str, Any] | None:
if self.autotune_entries is None:
return None
return self.autotune_entries.get(AutotuneKey.key_hash(key.autotune_key))
def lookup_fuzzy(self, key: AutotuneKey) -> dict[str, Any] | None:
if self.autotune_entries is None:
return None
for entry in self.autotune_entries.values():
if key.fuzzy_matches(entry.get("autotune_key")):
return entry
return None
@cache
def load_config_file(config_file: Path) -> dict[str, Any] | None:
try:
with open(config_file) as f:
return json.load(f)
except Exception as e:
logger.warning("Error reading config file %s: %s", config_file, e)
return None
def read_config_file(config_file: Path) -> dict[str, Any] | None:
"""Read a config file, bypassing the in-process cache in ALWAYS mode."""
if FLA_CACHE_MODE is FlaCacheMode.ALWAYS:
return load_config_file.__wrapped__(config_file)
return load_config_file(config_file)
def load_cached_config(kernel_name: str, autotune_key: AutotuneKey | None = None) -> dict[str, Any] | None:
"""
Load cached best config for a kernel from FLA configs directory.
This function loads the cached best configuration for a given kernel name
from get_fla_config_dir()/{kernel_name}.json.
Cache files may contain multiple autotune entries keyed by Triton's
runtime tuning key plus a top-level default config.
If the config file is not found or cannot be loaded, a warning is printed
and None is returned, allowing fallback to Triton's autotune.
The lookup mode is controlled by the FLA_CACHE_MODE environment variable (see FlaCacheMode).
Args:
kernel_name: Name of the kernel (e.g., "causal_conv1d_fwd_kernel")
autotune_key: Triton autotune key for the current invocation
Returns:
Best config dictionary or None if not found or disabled
"""
if FLA_CACHE_MODE is FlaCacheMode.DISABLED:
return None
config_dir = get_fla_config_dir()
config_file = config_dir / f"{kernel_name}.json"
if not config_file.exists():
return None
config_data = read_config_file(config_file)
if config_data is None:
return None
config = KernelConfigFile.from_dict(config_file, config_data)
if config is None:
return None
if FLA_CACHE_MODE is FlaCacheMode.DEFAULT or FLA_CACHE_MODE is FlaCacheMode.ALWAYS:
return config.default_config
# STRICT mode: exact match only, no fuzzy fallback
if FLA_CACHE_MODE is FlaCacheMode.STRICT:
if autotune_key is not None:
entry = config.lookup_exact(autotune_key)
if entry is not None:
return entry["config"]
return None
# FULL and FUZZY modes: try exact key match first, then fuzzy match
if autotune_key is not None:
entry = config.lookup_exact(autotune_key) or config.lookup_fuzzy(autotune_key)
if entry is not None:
return entry["config"]
if FLA_CACHE_MODE is FlaCacheMode.FUZZY:
return None
# FULL mode: fall back to default_config, then legacy raw config (no autotune_entries)
if config.default_config is not None:
return config.default_config
if config.autotune_entries is not None:
return None
return config_data
class CachedAutotuner(Autotuner):
"""
A modified autotuner that loads best config from FLA's config directory.
This class extends Triton's Autotuner but overrides the run method to
try loading cached configuration first before falling back to autotune.
"""
def __init__(self, fn, arg_names, configs, key, reset_to_zero, restore_value, **kwargs):
super().__init__(fn, arg_names, configs, key, reset_to_zero, restore_value, **kwargs)
self.kernel_name = fn.fn.__name__ if hasattr(fn, 'fn') else fn.__name__
# None-safe pre/post hooks: Triton's defaults crash when a restore_value / reset_to_zero arg
# is None (idiomatic for optional pointers gated by a tl.constexpr flag).
# Fixed upstream in triton-lang/triton#10295 — remove this override once FLA's minimum Triton version has it.
if not self.user_defined_pre_hook and (self.reset_to_zero or self.restore_value):
def _pre_hook(kw, reset_only=False):
for n in self.reset_to_zero:
if kw[n] is not None:
kw[n].zero_()
if not reset_only:
self.restore_copies = {n: kw[n].clone() for n in self.restore_value if kw[n] is not None}
self.pre_hook = _pre_hook
if not self.user_defined_post_hook and self.restore_value:
def _post_hook(kw, exception):
for n, copy in self.restore_copies.items():
kw[n].copy_(copy)
self.restore_copies = {}
self.post_hook = _post_hook
def should_check_fla_cache(self, key: AutotuneKey) -> bool:
if FLA_CACHE_MODE is FlaCacheMode.DISABLED:
return False
if FLA_CACHE_MODE is FlaCacheMode.ALWAYS:
return True
return key.autotune_key not in self.cache
def run(self, *args, **kwargs):
key = AutotuneKey.build(self.arg_names, self.keys, args, kwargs)
if self.should_check_fla_cache(key):
self.maybe_load_cached_config(key)
return super().run(*args, **kwargs)
def maybe_load_cached_config(self, key: AutotuneKey):
best_config = load_cached_config(self.kernel_name, key)
if best_config is not None:
kw = best_config["kwargs"]
num_warps = best_config["num_warps"]
num_stages = best_config["num_stages"]
extra = {
"num_ctas": best_config["num_ctas"],
"maxnreg": best_config.get("maxnreg"),
"pre_hook": None,
"ir_override": best_config.get("ir_override"),
} if TRITON_ABOVE_3_5_1 else {}
cfg = triton.Config(kw, num_warps=num_warps, num_stages=num_stages, **extra)
self.cache[key.autotune_key] = cfg
else:
logger.debug(
"No cached config found for kernel %s and key %s; falling back to Triton autotune",
self.kernel_name,
list(key.autotune_key),
)
def fla_cache_autotune(configs, key=None, prune_configs_by=None, reset_to_zero=None, restore_value=None,
pre_hook=None, post_hook=None, warmup=None, rep=None, use_cuda_graph=False,
do_bench=None, cache_results=False):
"""
Decorator for auto-tuning a :code:`triton.jit`'d function with FLA config support.
Extends Triton's autotune to load best configurations from FLA's config directory
(default: fla/configs/{GPU}/, or FLA_CONFIG_DIR/ when overridden), keyed by kernel
name from {kernel_name}.json. Lookup behaviour is controlled by FLA_CACHE_MODE.
Falls back to normal Triton autotuning when no cached config is found.
"""
# key can be None when we want to use cache only (no fallback autotune)
if key is None:
key = []
def decorator(fn):
kwargs = {}
if TRITON_ABOVE_3_4_0:
kwargs = {"cache_results": cache_results}
return CachedAutotuner(fn, fn.arg_names, configs, key, reset_to_zero, restore_value,
pre_hook=pre_hook, post_hook=post_hook,
prune_configs_by=prune_configs_by, warmup=warmup, rep=rep,
use_cuda_graph=use_cuda_graph, do_bench=do_bench,
**kwargs,
)
return decorator
def configure_fla_cache_autotune():
triton.autotune = fla_cache_autotune
logger.info(
"configure_fla_cache_autotune() is enabling FLA fla_cache_autotune; "
"triton.autotune will be replaced with fla_cache_autotune."
)
def restore_autotune_backend():
from triton.runtime.autotuner import autotune as original_autotune
triton.autotune = original_autotune
logger.info(
"restore_autotune_backend() is restoring Triton's original autotune; "
"triton.autotune will be replaced with triton.runtime.autotuner.autotune."
)
+10
View File
@@ -0,0 +1,10 @@
# 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
# Approximate value of 1/ln(2), used for log/exp base conversion
# Best FP32 approximation: 1.4426950216 (hex 0x3FB8AA3B)
RCP_LN2 = 1.4426950216
+468
View File
@@ -0,0 +1,468 @@
# 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 kda._fla.ops.backends import dispatch
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.ops.utils.index import prepare_chunk_indices
from kda._fla.utils import autotune_cache_kwargs, check_shared_mem, input_guard
BS_LIST = [32, 64] if check_shared_mem() else [16, 32]
@triton.heuristics({
'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=['B', 'H', 'BT', 'IS_VARLEN', 'REVERSE'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_local_cumsum_scalar_kernel(
s,
o,
scale,
cu_seqlens,
chunk_indices,
T,
B: tl.constexpr,
H: tl.constexpr,
BT: tl.constexpr,
REVERSE: 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_s = s + bos*H + i_h + o_t * H
p_o = o + bos*H + i_h + o_t * H
# [BT]
b_s = tl.load(p_s, mask=m_t, other=0.0).to(tl.float32)
if REVERSE:
b_o = tl.cumsum(b_s, axis=0, reverse=True)
else:
b_o = tl.cumsum(b_s, axis=0)
if HAS_SCALE:
b_o *= scale
tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=m_t)
@triton.heuristics({
'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({'BS': BS}, num_warps=num_warps)
for BS in BS_LIST
for num_warps in [2, 4, 8]
],
key=['B', 'H', 'S', 'BT', 'IS_VARLEN', 'REVERSE'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_local_cumsum_vector_kernel(
s,
o,
scale,
cu_seqlens,
chunk_indices,
T,
B: tl.constexpr,
H: tl.constexpr,
S: tl.constexpr,
BT: tl.constexpr,
BS: tl.constexpr,
REVERSE: tl.constexpr,
HAS_SCALE: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_s, i_t, i_bh = tl.program_id(0), tl.program_id(1).to(tl.int64), tl.program_id(2).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)
o_s = i_s * BS + tl.arange(0, BS)
m_s = (o_t[:, None] < T) & (o_s[None, :] < S)
p_s = s + (bos * H + i_h) * S + o_t[:, None] * (H*S) + o_s[None, :]
p_o = o + (bos * H + i_h) * S + o_t[:, None] * (H*S) + o_s[None, :]
# [BT, BS]
b_s = tl.load(p_s, mask=m_s, other=0.0).to(tl.float32)
if REVERSE:
b_o = tl.cumsum(b_s, axis=0, reverse=True)
else:
b_o = tl.cumsum(b_s, axis=0)
if HAS_SCALE:
b_o *= scale
tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=m_s)
@triton.heuristics({
'HAS_SCALE': lambda args: args['scale'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@triton.autotune(
configs=[
triton.Config({'BT': BT}, num_warps=num_warps, num_stages=num_stages)
for BT in [32, 64, 128, 256]
for num_warps in [2, 4, 8]
for num_stages in [1, 2, 3, 4]
],
key=['B', 'H', 'IS_VARLEN', 'REVERSE'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_global_cumsum_scalar_kernel(
s,
o,
scale,
cu_seqlens,
T,
B: tl.constexpr,
H: tl.constexpr,
BT: tl.constexpr,
REVERSE: tl.constexpr,
HAS_SCALE: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_nh = tl.program_id(0).to(tl.int64)
i_n, i_h = i_nh // H, i_nh % 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)
else:
bos, eos = i_n * T, i_n * T + T
T = eos - bos
b_z = tl.zeros([], dtype=tl.float32)
NT = tl.cdiv(T, BT)
for i_c in range(NT):
i_t = NT - 1 - i_c if REVERSE else i_c
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
p_s = s + bos*H + i_h + o_t * H
p_o = o + bos*H + i_h + o_t * H
b_s = tl.load(p_s, mask=m_t, other=0.0).to(tl.float32)
if REVERSE:
b_o = tl.cumsum(b_s, axis=0, reverse=True)
else:
b_o = tl.cumsum(b_s, axis=0)
b_ss = tl.sum(b_s, 0)
b_o += b_z
if i_c >= 0:
b_z += b_ss
if HAS_SCALE:
b_o *= scale
tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=m_t)
@triton.heuristics({
'HAS_SCALE': lambda args: args['scale'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@triton.autotune(
configs=[
triton.Config({'BT': BT}, num_warps=num_warps, num_stages=num_stages)
for BT in [16, 32, 64, 128]
for num_warps in [2, 4, 8]
for num_stages in [1, 2, 3, 4]
],
key=['B', 'H', 'S', 'IS_VARLEN', 'REVERSE'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_global_cumsum_vector_kernel(
s,
o,
scale,
cu_seqlens,
T,
B: tl.constexpr,
H: tl.constexpr,
S: tl.constexpr,
BT: tl.constexpr,
BS: tl.constexpr,
REVERSE: tl.constexpr,
HAS_SCALE: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_s, i_nh = tl.program_id(0), tl.program_id(1).to(tl.int64)
i_n, i_h = i_nh // H, i_nh % 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)
else:
bos, eos = i_n * T, i_n * T + T
T = eos - bos
b_z = tl.zeros([BS], dtype=tl.float32)
NT = tl.cdiv(T, BT)
for i_c in range(NT):
i_t = NT - 1 - i_c if REVERSE else i_c
o_t = i_t * BT + tl.arange(0, BT)
o_s = i_s * BS + tl.arange(0, BS)
m_s = (o_t[:, None] < T) & (o_s[None, :] < S)
p_s = s + (bos * H + i_h) * S + o_t[:, None] * (H*S) + o_s[None, :]
p_o = o + (bos * H + i_h) * S + o_t[:, None] * (H*S) + o_s[None, :]
# [BT, BS]
b_s = tl.load(p_s, mask=m_s, other=0.0).to(tl.float32)
if REVERSE:
b_c = b_z[None, :] + tl.cumsum(b_s, axis=0, reverse=True)
else:
b_c = b_z[None, :] + tl.cumsum(b_s, axis=0)
if HAS_SCALE:
b_c *= scale
tl.store(p_o, b_c.to(p_o.dtype.element_ty), mask=m_s)
b_z += tl.sum(b_s, 0)
def chunk_local_cumsum_scalar(
g: torch.Tensor,
chunk_size: int,
reverse: bool = False,
scale: float = None,
cu_seqlens: torch.Tensor | None = None,
output_dtype: torch.dtype | None = torch.float,
chunk_indices: torch.LongTensor | None = None,
**kwargs,
) -> torch.Tensor:
if 'head_first' in kwargs:
raise DeprecationWarning(
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
)
B, T, H = g.shape
assert chunk_size == 2**(chunk_size.bit_length()-1), "chunk_size must be a power of 2"
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)
g_org, g = g, torch.empty_like(g, dtype=output_dtype or g.dtype)
grid = (NT, B * H)
chunk_local_cumsum_scalar_kernel[grid](
s=g_org,
o=g,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
T=T,
B=B,
H=H,
BT=BT,
REVERSE=reverse,
)
return g
def chunk_local_cumsum_vector(
g: torch.Tensor,
chunk_size: int,
reverse: bool = False,
scale: float = None,
cu_seqlens: torch.Tensor | None = None,
output_dtype: torch.dtype | None = torch.float,
chunk_indices: torch.LongTensor | None = None,
**kwargs,
) -> torch.Tensor:
if 'head_first' in kwargs:
raise DeprecationWarning(
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
)
B, T, H, S = 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)
assert chunk_size == 2**(chunk_size.bit_length()-1), "chunk_size must be a power of 2"
g_org, g = g, torch.empty_like(g, dtype=output_dtype or g.dtype)
def grid(meta): return (triton.cdiv(meta['S'], meta['BS']), NT, B * H)
# keep cummulative normalizer in fp32
# this kernel is equivalent to
# g = g.view(B, H, NT, BT, -1).cumsum(-2).view(B, H, T, -1)
chunk_local_cumsum_vector_kernel[grid](
s=g_org,
o=g,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
T=T,
B=B,
H=H,
S=S,
BT=BT,
REVERSE=reverse,
)
return g
@input_guard
def chunk_global_cumsum_scalar(
s: torch.Tensor,
reverse: bool = False,
cu_seqlens: torch.Tensor | None = None,
scale: float = None,
output_dtype: torch.dtype | None = torch.float,
**kwargs,
) -> torch.Tensor:
if 'head_first' in kwargs:
raise DeprecationWarning(
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
)
B, T, H = s.shape
N = len(cu_seqlens) - 1 if cu_seqlens is not None else B
z = torch.empty_like(s, dtype=output_dtype or s.dtype)
grid = (N * H,)
chunk_global_cumsum_scalar_kernel[grid](
s=s,
o=z,
scale=scale,
cu_seqlens=cu_seqlens,
T=T,
B=B,
H=H,
REVERSE=reverse,
)
return z
@input_guard
def chunk_global_cumsum_vector(
s: torch.Tensor,
reverse: bool = False,
cu_seqlens: torch.Tensor | None = None,
scale: float = None,
output_dtype: torch.dtype | None = torch.float,
**kwargs,
) -> torch.Tensor:
if 'head_first' in kwargs:
raise DeprecationWarning(
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
)
B, T, H, S = s.shape
N = len(cu_seqlens) - 1 if cu_seqlens is not None else B
BS = min(32, triton.next_power_of_2(S))
z = torch.empty_like(s, dtype=output_dtype or s.dtype)
grid = (triton.cdiv(S, BS), N * H)
chunk_global_cumsum_vector_kernel[grid](
s=s,
o=z,
scale=scale,
cu_seqlens=cu_seqlens,
T=T,
B=B,
H=H,
S=S,
BS=BS,
REVERSE=reverse,
)
return z
@input_guard
@dispatch('utils')
def chunk_global_cumsum(
s: torch.Tensor,
reverse: bool = False,
cu_seqlens: torch.Tensor | None = None,
scale: float = None,
output_dtype: torch.dtype | None = torch.float,
**kwargs,
) -> torch.Tensor:
if 'head_first' in kwargs:
raise DeprecationWarning(
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
)
if cu_seqlens is not None:
assert s.shape[0] == 1, "Only batch size 1 is supported when cu_seqlens are provided"
if len(s.shape) == 3:
return chunk_global_cumsum_scalar(
s=s,
reverse=reverse,
cu_seqlens=cu_seqlens,
scale=scale,
output_dtype=output_dtype,
)
elif len(s.shape) == 4:
return chunk_global_cumsum_vector(
s=s,
reverse=reverse,
cu_seqlens=cu_seqlens,
scale=scale,
output_dtype=output_dtype,
)
else:
raise ValueError(
f"Unsupported input shape {s.shape}, "
f"which should be [B, T, H] or [B, T, H, D]",
)
@input_guard
@dispatch('utils')
def chunk_local_cumsum(
g: torch.Tensor,
chunk_size: int,
reverse: bool = False,
scale: float = None,
cu_seqlens: torch.Tensor | None = None,
output_dtype: torch.dtype | None = torch.float,
chunk_indices: torch.LongTensor | None = None,
**kwargs,
) -> torch.Tensor:
if 'head_first' in kwargs:
raise DeprecationWarning(
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
)
if cu_seqlens is not None:
assert g.shape[0] == 1, "Only batch size 1 is supported when cu_seqlens are provided"
if len(g.shape) == 3:
return chunk_local_cumsum_scalar(
g=g,
chunk_size=chunk_size,
reverse=reverse,
scale=scale,
cu_seqlens=cu_seqlens,
output_dtype=output_dtype,
chunk_indices=chunk_indices,
)
elif len(g.shape) == 4:
return chunk_local_cumsum_vector(
g=g,
chunk_size=chunk_size,
reverse=reverse,
scale=scale,
cu_seqlens=cu_seqlens,
output_dtype=output_dtype,
chunk_indices=chunk_indices,
)
else:
raise ValueError(
f"Unsupported input shape {g.shape}, "
f"which should be (B, T, H) or (B, T, H, D)",
)
+183
View File
@@ -0,0 +1,183 @@
# 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 kda._fla.utils import autotune_cache_kwargs, tensor_cache
@triton.autotune(
configs=[
triton.Config({}, num_warps=num_warps)
for num_warps in [4, 8, 16, 32]
],
key=['B'],
**autotune_cache_kwargs,
)
@triton.jit
def prepare_position_ids_kernel(
y,
cu_seqlens,
B: tl.constexpr,
):
i_n = tl.program_id(0)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
T = eos - bos
o = tl.arange(0, B)
for i in range(0, tl.cdiv(T, B) * B, B):
o_i = o + i
tl.store(y + bos + o_i, o_i, o_i < T)
@tensor_cache
def prepare_lens(cu_seqlens: torch.LongTensor) -> torch.LongTensor:
return torch.diff(cu_seqlens)
@tensor_cache
def prepare_lens_from_mask(mask: torch.BoolTensor) -> torch.LongTensor:
return mask.sum(dim=-1, dtype=torch.int32)
@tensor_cache
def prepare_cu_seqlens_from_lens(
lens: torch.LongTensor,
dtype: torch.dtype | None = torch.int32,
) -> torch.LongTensor:
return F.pad(lens.cumsum(dim=0, dtype=dtype), (1, 0))
@tensor_cache
def prepare_cu_seqlens_from_mask(
mask: torch.BoolTensor,
dtype: torch.dtype | None = torch.int32,
) -> torch.LongTensor:
return prepare_cu_seqlens_from_lens(prepare_lens_from_mask(mask), dtype)
@tensor_cache
def prepare_split_cu_seqlens(
batch_size: int | None = None,
seq_len: int | None = None,
split_size: int | None = None,
cu_seqlens: torch.LongTensor | None = None,
dtype: torch.dtype | None = torch.int32,
device: torch.device | None = torch.device('cpu'),
) -> torch.LongTensor:
"""Sub-split a (optionally packed) batch along the token axis.
Two calling modes:
- **Rectangular batch**: pass `batch_size` and `seq_len`, leave
`cu_seqlens=None`. Internally synthesizes `[0, L, 2L, ..., B*L]`.
- **Packed varlen**: pass `cu_seqlens`. `batch_size` and `seq_len` are
ignored (kept as optional kwargs for backward-compat with callers
that used to pass dummies).
`split_size` is always required.
The legacy positional signature `(batch_size, seq_len, split_size, ...)`
continues to work — the first two args retain their position but may now
be omitted when `cu_seqlens` is supplied.
"""
if split_size is None:
raise TypeError("prepare_split_cu_seqlens() requires `split_size`")
if cu_seqlens is None:
if batch_size is None or seq_len is None:
raise TypeError(
"prepare_split_cu_seqlens() requires either `cu_seqlens`, "
"or both `batch_size` and `seq_len`"
)
total_tokens = batch_size * seq_len
cu_seqlens = list(range(0, total_tokens, seq_len)) + [total_tokens]
else:
cu_seqlens = cu_seqlens.tolist()
return torch.tensor(
[
i
for bos, eos in zip(cu_seqlens[:-1], cu_seqlens[1:], strict=False)
for i in range(bos, eos, split_size)
] + [cu_seqlens[-1]],
dtype=dtype,
device=device,
)
def _segmented_arange(counts: torch.LongTensor) -> tuple[torch.LongTensor, torch.LongTensor]:
"""Expand per-segment counts into flat per-slot index tensors.
Given segment sizes ``counts = [c0, c1, ...]``, return two 1-D tensors of
length ``counts.sum()`` that together label every slot with its segment and
its position within that segment.
Example -- ``counts = [2, 3]`` (segment 0 spans 2 slots, segment 1 spans 3)::
seg_id = [0, 0, 1, 1, 1] # which segment each slot belongs to
intra_idx = [0, 1, 0, 1, 2] # running index within that segment
With CUDA ``counts``, ``repeat_interleave`` reads ``counts.sum()`` on the
host (one device sync). Pass host-side counts to avoid it.
"""
seg_id = torch.repeat_interleave(
torch.arange(counts.numel(), device=counts.device, dtype=counts.dtype),
counts,
)
seg_start = F.pad(counts.cumsum(0), (1, 0))[:-1]
intra_idx = torch.arange(seg_id.shape[0], device=counts.device, dtype=counts.dtype) - seg_start[seg_id]
return seg_id, intra_idx
@tensor_cache
def prepare_position_ids(cu_seqlens: torch.LongTensor, cu_seqlens_cpu: torch.LongTensor | None = None) -> torch.LongTensor:
src = cu_seqlens_cpu if cu_seqlens_cpu is not None else cu_seqlens
_, position_ids = _segmented_arange(prepare_lens(src))
return position_ids.to(cu_seqlens)
@tensor_cache
def prepare_sequence_ids(cu_seqlens: torch.LongTensor, cu_seqlens_cpu: torch.LongTensor | None = None) -> torch.LongTensor:
return prepare_position_ids(cu_seqlens, cu_seqlens_cpu).eq(0).cumsum(0) - 1
@tensor_cache
def prepare_token_indices(cu_seqlens: torch.LongTensor, cu_seqlens_cpu: torch.LongTensor | None = None) -> torch.LongTensor:
position_ids = prepare_position_ids(cu_seqlens, cu_seqlens_cpu)
return torch.stack([prepare_sequence_ids(cu_seqlens, cu_seqlens_cpu), position_ids], 1).to(cu_seqlens)
@tensor_cache
def prepare_chunk_indices(
cu_seqlens: torch.LongTensor,
chunk_size: int,
cu_seqlens_cpu: torch.LongTensor | None = None,
) -> torch.LongTensor:
src = cu_seqlens_cpu if cu_seqlens_cpu is not None else cu_seqlens
chunk_counts = (prepare_lens(src) + (chunk_size - 1)).div(chunk_size, rounding_mode='floor')
seg_id, intra_chunk_idx = _segmented_arange(chunk_counts)
return torch.stack([seg_id, intra_chunk_idx], 1).to(cu_seqlens)
@tensor_cache
def prepare_chunk_offsets(
cu_seqlens: torch.LongTensor,
chunk_size: int,
) -> torch.LongTensor:
return F.pad(triton.cdiv(prepare_lens(cu_seqlens), chunk_size), (1, 0), value=0).cumsum(-1)
@tensor_cache
def get_max_num_splits(
cu_seqlens: torch.LongTensor,
chunk_size: int,
cu_seqlens_cpu: torch.LongTensor | None = None
) -> int:
if cu_seqlens_cpu is not None:
return triton.cdiv(int(max(prepare_lens(cu_seqlens_cpu))), chunk_size)
return triton.cdiv(int(max(prepare_lens(cu_seqlens))), chunk_size)
+101
View File
@@ -0,0 +1,101 @@
# 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 os
import triton
import triton.language as tl
import triton.language.extra.libdevice as tldevice
from kda._fla.utils import IS_GATHER_SUPPORTED, IS_NVIDIA_BLACKWELL
if os.environ.get('FLA_USE_FAST_OPS', '0') == '1':
@triton.jit
def exp(x): return tldevice.fast_expf(x.to(tl.float32))
@triton.jit
def exp2(x): return tldevice.exp2(x.to(tl.float32))
@triton.jit
def log(x): return tldevice.fast_logf(x.to(tl.float32))
@triton.jit
def log2(x): return tldevice.fast_log2f(x.to(tl.float32))
@triton.jit
def tanh(x): return tldevice.fast_tanhf(x.to(tl.float32))
else:
@triton.jit
def exp(x): return tl.exp(x.to(tl.float32))
@triton.jit
def exp2(x): return tl.math.exp2(x.to(tl.float32))
@triton.jit
def log(x): return tl.log(x.to(tl.float32))
@triton.jit
def log2(x): return tl.log2(x.to(tl.float32))
@triton.jit
def tanh(x): return tldevice.tanh(x.to(tl.float32))
if IS_NVIDIA_BLACKWELL:
"""
Compute tl.dot with Blackwell workaround.
On SM100 datacenter and SM120 consumer Blackwell GPUs, wraps the result in
inline assembly to prevent the TritonGPUHoistTMEMAlloc pass from incorrectly
fusing add and dot operations.
See: https://github.com/fla-org/flash-linear-attention/issues/638
TODO: Remove this workaround once the Triton compiler bug is fixed.
Track upstream issue at: https://github.com/triton-lang/triton/issues/8695
"""
@triton.jit
def safe_dot(a, b, allow_tf32: tl.constexpr = None):
return tl.inline_asm_elementwise(
asm="mov.f32 $0, $1;",
constraints="=r,r",
args=[tl.dot(a, b, allow_tf32=allow_tf32)],
dtype=tl.float32,
is_pure=True,
pack=1,
)
else:
@triton.jit
def safe_dot(a, b, allow_tf32: tl.constexpr = None):
return tl.dot(a, b, allow_tf32=allow_tf32)
if not IS_GATHER_SUPPORTED:
@triton.jit
def gather(src, index, axis, _builder=None):
"""
Gather operation that works when tl.gather is not supported.
This is a fallback implementation that returns None.
Just to make triton compiler happy.
"""
return None
else:
gather = tl.gather
if hasattr(triton.language, '_experimental_make_tensor_descriptor'):
# For Triton 3.3.x
make_tensor_descriptor = triton.language._experimental_make_tensor_descriptor
elif hasattr(triton.language, 'make_tensor_descriptor'):
# For Triton 3.4.x and later
make_tensor_descriptor = triton.language.make_tensor_descriptor
else:
"""
Fallback implementation when TMA is not supported.
Returns None to indicate TMA descriptors are unavailable.
Just make triton compiler happy.
"""
@triton.jit
def make_tensor_descriptor(
base,
shape,
strides,
block_shape,
_builder=None,
):
return None
+115
View File
@@ -0,0 +1,115 @@
# 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
# REVISED FROM
# https://github.com/shawntan/stickbreaking-attention/blob/main/stickbreaking_attention/sb_varlen/softplus.py
import triton
from triton import language as tl
from kda._fla.utils import IS_NVIDIA
def _generate_softplus(num_pack):
template = """
.reg .pred p;
setp.gt.f32 p, ${in_reg}, 20.;
@p mov.f32 ${out_reg}, ${in_reg};
@!p mul.f32 ${out_reg}, ${in_reg}, 1.4426950408889634;
@!p ex2.approx.ftz.f32 ${out_reg}, ${out_reg};
@!p add.f32 ${out_reg}, ${out_reg}, 1.0;
@!p lg2.approx.ftz.f32 ${out_reg}, ${out_reg};
@!p mul.f32 ${out_reg}, ${out_reg}, 0.6931471805599453;
"""
out_str = ""
for i in range(num_pack):
inner_str = template.format(out_reg=i, in_reg=i + num_pack)
out_str += "{" + inner_str + "}\n"
# flatten out because torch.compile doesn't like newlines
out_str = " ".join(out_str.split("\n"))
return out_str
def _generate_softplus2(num_pack):
template = """
.reg .pred p;
setp.gt.f32 p, ${in_reg}, 15.;
@p mov.f32 ${out_reg}, ${in_reg};
@!p ex2.approx.ftz.f32 ${out_reg}, ${in_reg};
@!p add.f32 ${out_reg}, ${out_reg}, 1.0;
@!p lg2.approx.ftz.f32 ${out_reg}, ${out_reg};
"""
out_str = ""
for i in range(num_pack):
inner_str = template.format(out_reg=i, in_reg=i + num_pack)
out_str += "{" + inner_str + "}\n"
# flatten out because torch.compile doesn't like newlines
out_str = " ".join(out_str.split("\n"))
return out_str
def _generate_constraints(num_pack):
return ",".join("=r" for i in range(num_pack)) + "," + ",".join("r" for i in range(num_pack))
_NUM_REG = 1
s_softplus: tl.constexpr = tl.constexpr(_generate_softplus(_NUM_REG))
s_softplus2: tl.constexpr = tl.constexpr(_generate_softplus2(_NUM_REG))
s_constraints: tl.constexpr = tl.constexpr(_generate_constraints(_NUM_REG))
NUM_REG: tl.constexpr = tl.constexpr(_NUM_REG)
@triton.jit
def softplus_nv(x):
# equivalent to:
# return tl.where(x < 20.0, tl.math.log(1 + tl.math.exp(x)), x)
return tl.inline_asm_elementwise(
asm=s_softplus,
constraints=s_constraints,
pack=NUM_REG,
args=[
x,
],
dtype=tl.float32,
is_pure=True,
)
@triton.jit
def softplus_triton(x):
return tl.where(x < 20.0, tl.math.log(1 + tl.math.exp(x)), x)
@triton.jit
def softplus2_nv(x):
# equivalent to:
# return tl.where(x < 15.0, tl.math.log2(1 + tl.math.exp2(x)), x)
return tl.inline_asm_elementwise(
asm=s_softplus2,
constraints=s_constraints,
pack=NUM_REG,
args=[
x,
],
dtype=tl.float32,
is_pure=True,
)
@triton.jit
def softplus2_triton(x):
return tl.where(x < 15.0, tl.math.log2(1 + tl.math.exp2(x)), x)
if IS_NVIDIA:
softplus = softplus_nv
softplus2 = softplus2_nv
else:
softplus = softplus_triton
softplus2 = softplus2_triton
+92
View File
@@ -0,0 +1,92 @@
# 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 sys
from ._compat import ( # noqa: F401
SUPPORTS_AUTOTUNE_CACHE,
TRITON_ABOVE_3_4_0,
TRITON_ABOVE_3_5_1,
TRITON_ABOVE_3_7_1,
autotune_cache_kwargs,
find_spec_cached,
has_usable_nvcc,
)
from ._config import ( # noqa: F401
FLA_CACHE_RESULTS,
FLA_CI_ENV,
FLA_DISABLE_TENSOR_CACHE,
FLA_TENSOR_CACHE_SIZE,
)
from ._decorators import ( # noqa: F401
Action,
checkpoint,
contiguous,
deprecate_kwarg,
input_guard,
require_version,
tensor_cache,
)
from ._device import ( # noqa: F401
IS_AMD,
IS_ARM,
IS_GATHER_SUPPORTED,
IS_INTEL,
IS_INTEL_ALCHEMIST,
IS_NPU,
IS_NVIDIA,
IS_NVIDIA_BLACKWELL,
IS_NVIDIA_HOPPER,
IS_NVIDIA_SM100,
IS_NVIDIA_SM120,
IS_TF32_SUPPORTED,
IS_TMA_SUPPORTED,
Backend,
autocast_custom_bwd,
autocast_custom_fwd,
check_environments,
check_pytorch_version,
check_shared_mem,
custom_device_ctx,
device,
device_name,
device_platform,
device_torch_lib,
get_all_max_shared_mem,
get_available_device,
get_device_capability,
get_device_smem_optin,
get_multiprocessor_count,
map_triton_backend_to_torch_device,
)
from ._testing import assert_close, get_abs_err, get_err_ratio # noqa: F401
def _register_aliases():
current_module = sys.modules[__name__]
for key in (
'IS_AMD',
'IS_ARM',
'IS_INTEL',
'IS_INTEL_ALCHEMIST',
'IS_NVIDIA',
'IS_NPU',
'IS_NVIDIA_BLACKWELL',
'IS_NVIDIA_HOPPER',
'IS_NVIDIA_SM100',
'IS_NVIDIA_SM120',
'IS_TF32_SUPPORTED',
'IS_GATHER_SUPPORTED',
'IS_TMA_SUPPORTED',
):
if hasattr(current_module, key):
setattr(current_module, key.lower(), getattr(current_module, key))
_register_aliases()
del _register_aliases
+65
View File
@@ -0,0 +1,65 @@
# 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 functools
import importlib.metadata
import inspect
import logging
import os
import shutil
from importlib.util import find_spec
from pathlib import Path
import triton
from packaging import version as package_version
from ._config import FLA_CACHE_RESULTS
logger = logging.getLogger(__name__)
TRITON_ABOVE_3_4_0 = package_version.parse(triton.__version__) >= package_version.parse("3.4.0")
TRITON_ABOVE_3_5_1 = package_version.parse(triton.__version__) >= package_version.parse("3.5.1")
TRITON_ABOVE_3_7_1 = package_version.parse(triton.__version__) >= package_version.parse("3.7.1")
SUPPORTS_AUTOTUNE_CACHE = "cache_results" in inspect.signature(triton.autotune).parameters
autotune_cache_kwargs = {"cache_results": FLA_CACHE_RESULTS} if SUPPORTS_AUTOTUNE_CACHE else {}
@functools.cache
def find_spec_cached(name):
return find_spec(name)
@functools.cache
def has_usable_nvcc() -> bool:
"""Whether a usable nvcc compiler is available for TileLang's JIT.
Mirrors the guesses in ``tilelang.env._find_cuda_home`` (env
CUDA_HOME/CUDA_PATH, nvcc on PATH, the ``nvidia-cuda-nvcc`` wheel,
/usr/local/cuda), but verifies the nvcc binary actually exists —
only ``nvidia-cuda-nvcc`` >= 13.0 ships it, the ``-cu12`` variant
installs just ptxas.
"""
cuda_home = os.environ.get("CUDA_HOME") or os.environ.get("CUDA_PATH")
if cuda_home is not None and (Path(cuda_home) / "bin" / "nvcc").exists():
return True
if shutil.which("nvcc") is not None:
return True
try:
files = importlib.metadata.files("nvidia-cuda-nvcc") or []
except importlib.metadata.PackageNotFoundError:
files = []
if any(f.name in ("nvcc", "nvcc.exe") for f in files):
return True
if (Path("/usr/local/cuda") / "bin" / "nvcc").exists():
return True
logger.info(
"[FLA Backend] TileLang is installed but no usable nvcc compiler was found; falling back to Triton. "
"Install a CUDA toolkit or nvidia-cuda-nvcc, or set FLA_TILELANG=0 to disable TileLang explicitly."
)
return False
+17
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
import os
FLA_CI_ENV = os.getenv("FLA_CI_ENV") == "1"
FLA_CACHE_RESULTS = os.getenv('FLA_CACHE_RESULTS', '1') == '1'
FLA_DISABLE_TENSOR_CACHE = os.getenv('FLA_DISABLE_TENSOR_CACHE', '0') == '1'
try:
FLA_TENSOR_CACHE_SIZE = int(os.getenv('FLA_TENSOR_CACHE_SIZE', "4"))
except ValueError:
FLA_TENSOR_CACHE_SIZE = 4
+336
View File
@@ -0,0 +1,336 @@
# 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 contextlib
import functools
import inspect
import sys
import warnings
from collections import deque
from collections.abc import Callable
from enum import Enum
from typing import Any
import torch
from packaging import version as package_version
from .. import __version__
from ._config import FLA_DISABLE_TENSOR_CACHE, FLA_TENSOR_CACHE_SIZE
from ._device import custom_device_ctx
class Action(Enum):
NONE = "none"
NOTIFY = "notify"
NOTIFY_ALWAYS = "notify_always"
RAISE = "raise"
def tensor_cache(
fn: Callable[..., torch.Tensor],
) -> Callable[..., torch.Tensor]:
"""
A decorator that memoizes the most recent results of a function call by argument identity.
The decorator keeps a bounded queue of up to ``FLA_TENSOR_CACHE_SIZE`` (default 4)
recent ``(args, kwargs, result)`` triples. On each call, every cached entry is checked
in order; an entry is considered a hit when the positional arg count and kwarg key set
match and every argument is the *same object* (``is`` identity) as the cached one. On a
hit the cached result is returned and ``fn`` is skipped; on a miss ``fn`` is invoked and
the new triple is appended (evicting the oldest when the queue is full).
Caching is fully bypassed when the ``FLA_DISABLE_TENSOR_CACHE`` environment variable is
set to ``'1'``.
Args:
fn (Callable[..., torch.Tensor]):
The function to be decorated. Intended for functions whose inputs are tensors
(or other objects compared by identity) and whose output is a tensor.
Returns:
Callable[..., torch.Tensor]:
A wrapped version of ``fn`` backed by an identity-based bounded cache.
"""
cached: deque = deque(maxlen=FLA_TENSOR_CACHE_SIZE)
def cache_disabled() -> bool:
utils_module = sys.modules.get('kda._fla.utils')
return getattr(utils_module, 'FLA_DISABLE_TENSOR_CACHE', FLA_DISABLE_TENSOR_CACHE)
@functools.wraps(fn)
def wrapper(*args: Any, **kwargs: Any) -> Any:
if cache_disabled():
return fn(*args, **kwargs)
for cached_args, cached_kwargs, cached_result in cached:
if len(args) != len(cached_args) or len(kwargs) != len(cached_kwargs):
continue
if all(a is b for a, b in zip(args, cached_args, strict=False)) and \
all(k in cached_kwargs and v is cached_kwargs[k] for k, v in kwargs.items()):
return cached_result
result = fn(*args, **kwargs)
cached.append((args, kwargs, result))
return result
return wrapper
def _skip_contiguous(
no_guard_contiguous: bool | list[str] | tuple[str, ...] | set[str],
param_name: str,
skip_params: set[str],
) -> bool:
return no_guard_contiguous is True or param_name in skip_params
def _contiguous_if_needed(arg: Any, skip: bool) -> Any:
if isinstance(arg, torch.Tensor) and not skip:
return arg.contiguous()
return arg
def input_guard(
fn: Callable[..., torch.Tensor] | None = None,
*,
no_guard_contiguous: bool | list[str] | tuple[str, ...] | set[str] = False,
) -> Callable[[Callable[..., torch.Tensor]], Callable[..., torch.Tensor]] | Callable[..., torch.Tensor]:
"""
A decorator to make sure all input tensors are contiguous and set the device based on input tensors.
Args:
no_guard_contiguous (bool | list[str] | tuple[str, ...] | set[str]):
If True, skip all contiguous checks. If a list/tuple/set of parameter names, skip contiguous check for those parameters.
"""
def decorator(fn: Callable[..., torch.Tensor]) -> Callable[..., torch.Tensor]:
# Get function signature for parameter name mapping
sig = inspect.signature(fn)
param_names = list(sig.parameters.keys())
skip_params = set(no_guard_contiguous) if isinstance(no_guard_contiguous, (list, tuple, set)) else set()
@functools.wraps(fn)
def wrapper(*args, **kwargs):
# Process args with parameter name mapping
processed_args = []
for i, arg in enumerate(args):
if i < len(param_names):
param_name = param_names[i]
else:
# For *args beyond signature, use position as name
param_name = f"__arg_{i}"
processed_args.append(_contiguous_if_needed(
arg, _skip_contiguous(no_guard_contiguous, param_name, skip_params)))
# Process kwargs
processed_kwargs = {}
for k, v in kwargs.items():
processed_kwargs[k] = _contiguous_if_needed(v, _skip_contiguous(no_guard_contiguous, k, skip_params))
tensor = None
for arg in args:
if isinstance(arg, torch.Tensor):
tensor = arg
break
if tensor is None:
for value in kwargs.values():
if isinstance(value, torch.Tensor):
tensor = value
break
if tensor is not None:
ctx = custom_device_ctx(tensor.device.index)
else:
ctx = contextlib.nullcontext()
with ctx:
return fn(*processed_args, **processed_kwargs)
return wrapper
# Handle direct usage without parentheses: @input_guard
if fn is not None:
return decorator(fn)
return decorator
def contiguous(fn: Callable[..., torch.Tensor]) -> Callable[..., torch.Tensor]:
"""Alias for input_guard() without parameters."""
return input_guard(fn)
def require_version(version, hint):
"""
Perform a runtime check of the dependency versions, using the exact same syntax used by pip.
"""
def decorator(fn):
@functools.wraps(fn)
def wrapper(ctx, *args, **kwargs):
from transformers.utils.versions import require_version
require_version(version, hint)
return fn(
ctx,
*(i if not isinstance(i, torch.Tensor) else i.contiguous() for i in args),
**{k: (v if not isinstance(v, torch.Tensor) else v.contiguous()) for k, v in kwargs.items()},
)
return wrapper
return decorator
def deprecate_kwarg(
old_name: str,
version: str,
new_name: str | None = None,
warn_if_greater_or_equal_version: bool = False,
raise_if_greater_or_equal_version: bool = False,
raise_if_both_names: bool = False,
additional_message: str | None = None,
):
"""
Decorator to notify users about deprecated keyword arguments, replacing them with a new name if specified.
This decorator allows you to:
- Notify users when a keyword argument is deprecated.
- Automatically replace deprecated keyword arguments with new ones.
- Raise an error if deprecated arguments are used, depending on the specified conditions.
By default, the decorator notifies the user about the deprecated argument while the `fla.__version__` < specified `version`
in the decorator. To keep notifications with any version `warn_if_greater_or_equal_version=True` can be set.
Args:
old_name (`str`):
Name of the deprecated keyword argument.
version (`str`):
The version in which the keyword argument was (or will be) deprecated.
new_name (`Optional[str]`, *optional*):
The new name for the deprecated keyword argument.
If specified, the deprecated keyword argument will be replaced with this new name.
warn_if_greater_or_equal_version (`bool`, *optional*, defaults to `False`):
Whether to show warning if current `fla` version is greater or equal to the deprecated version.
raise_if_greater_or_equal_version (`bool`, *optional*, defaults to `False`):
Whether to raise `ValueError` if current `fla` version is greater or equal to the deprecated version.
raise_if_both_names (`bool`, *optional*, defaults to `False`):
Whether to raise `ValueError` if both deprecated and new keyword arguments are set.
additional_message (`Optional[str]`, *optional*):
An additional message to append to the default deprecation message.
Raises:
ValueError:
If `raise_if_greater_or_equal_version` is `True` and the current version >= the deprecated one,
or if `raise_if_both_names` is `True` and both old and new keyword arguments are provided.
Returns:
Callable:
A wrapped function that handles the deprecated keyword arguments according to the specified parameters.
Example usage with renaming argument:
```python
@deprecate_kwarg("reduce_labels", new_name="do_reduce_labels", version="6.0.0")
def my_function(do_reduce_labels):
print(do_reduce_labels)
my_function(reduce_labels=True) # Will show a deprecation warning and use do_reduce_labels=True
```
Example usage without renaming argument:
```python
@deprecate_kwarg("max_size", version="6.0.0")
def my_function(max_size):
print(max_size)
my_function(max_size=1333) # Will show a deprecation warning
```
"""
deprecated_version = package_version.parse(version)
current_version = package_version.parse(__version__)
is_greater_or_equal_version = current_version >= deprecated_version
if is_greater_or_equal_version:
version_message = f"and removed starting from version {version}"
else:
version_message = f"and will be removed in version {version}"
def wrapper(func):
# Required for better warning message
sig = inspect.signature(func)
function_named_args = set(sig.parameters.keys())
is_instance_method = "self" in function_named_args
is_class_method = "cls" in function_named_args
@functools.wraps(func)
def wrapped_func(*args, **kwargs):
# Get class + function name (just for better warning message)
func_name = func.__name__
if is_instance_method:
func_name = f"{args[0].__class__.__name__}.{func_name}"
elif is_class_method:
func_name = f"{args[0].__name__}.{func_name}"
minimum_action = Action.NONE
message = None
# deprecated kwarg and its new version are set for function call -> replace it with new name
if old_name in kwargs and new_name in kwargs:
minimum_action = Action.RAISE if raise_if_both_names else Action.NOTIFY_ALWAYS
message = (
f"Both `{old_name}` and `{new_name}` are set for `{func_name}`. "
f"Using `{new_name}={kwargs[new_name]}` and ignoring deprecated `{old_name}={kwargs[old_name]}`."
)
kwargs.pop(old_name)
# only deprecated kwarg is set for function call -> replace it with new name
elif old_name in kwargs and new_name is not None and new_name not in kwargs:
minimum_action = Action.NOTIFY
message = (
f"`{old_name}` is deprecated {version_message} for `{func_name}`. "
f"Use `{new_name}` instead."
)
kwargs[new_name] = kwargs.pop(old_name)
# deprecated kwarg is not set for function call and new name is not specified -> just notify
elif old_name in kwargs:
minimum_action = Action.NOTIFY
message = f"`{old_name}` is deprecated {version_message} for `{func_name}`."
if message is not None and additional_message is not None:
message = f"{message} {additional_message}"
# update minimum_action if argument is ALREADY deprecated (current version >= deprecated version)
if is_greater_or_equal_version:
# change to (NOTIFY, NOTIFY_ALWAYS) -> RAISE if specified
# in case we want to raise error for already deprecated arguments
if raise_if_greater_or_equal_version and minimum_action != Action.NONE:
minimum_action = Action.RAISE
# change to NOTIFY -> NONE if specified (NOTIFY_ALWAYS can't be changed to NONE)
elif not warn_if_greater_or_equal_version and minimum_action == Action.NOTIFY:
minimum_action = Action.NONE
# raise error or notify user
if minimum_action == Action.RAISE:
raise ValueError(message)
elif minimum_action in (Action.NOTIFY, Action.NOTIFY_ALWAYS):
# DeprecationWarning is ignored by default, so we use FutureWarning instead
warnings.warn(message, FutureWarning, stacklevel=2)
return func(*args, **kwargs)
return wrapped_func
return wrapper
def checkpoint(fn):
@functools.wraps(fn)
def wrapper(*args, **kwargs):
return torch.utils.checkpoint.checkpoint(fn, *args, **kwargs)
return wrapper
+245
View File
@@ -0,0 +1,245 @@
# 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 contextlib
import functools
import logging
import os
import platform
import sys
import warnings
from enum import Enum
from functools import cache, lru_cache
import torch
import triton
from packaging import version as package_version
logger = logging.getLogger(__name__)
@lru_cache(maxsize=1)
def check_environments():
"""
Checks the current operating system, Triton version, and Python version,
issuing warnings if they don't meet recommendations.
This function's body only runs once due to lru_cache.
"""
# Check Operating System
if sys.platform == 'win32':
# Check if triton-windows is installed
try:
from importlib.metadata import PackageNotFoundError, metadata
metadata('triton-windows')
# triton-windows is installed, no warning needed
except PackageNotFoundError:
logger.warning(
"Detected Windows operating system. Consider installing triton-windows "
"(https://github.com/triton-lang/triton-windows) for better compatibility. "
"Without it, some features may not work correctly.",
)
triton_version = package_version.parse(triton.__version__)
required_triton_version = package_version.parse("3.3.0")
if triton_version < required_triton_version:
logger.warning(
f"Current Triton version {triton_version} is below the recommended 3.3.0 version. "
"Errors may occur and these issues will not be fixed. "
"Please consider upgrading Triton.",
)
# Check Python version
py_version = package_version.parse(f"{sys.version_info.major}.{sys.version_info.minor}")
required_py_version = package_version.parse("3.11")
if py_version < required_py_version:
logger.warning(
f"Current Python version {py_version} is below the recommended 3.11 version. "
"It is recommended to upgrade to Python 3.11 or higher for the best experience.",
)
return None
check_environments()
def _cpu_device_warning():
warnings.warn(('Triton is not supported on current platform, roll back to CPU.'), stacklevel=2)
@cache
def check_pytorch_version(version_s: str = '2.4') -> bool:
return package_version.parse(torch.__version__) >= package_version.parse(version_s)
@cache
def get_multiprocessor_count(tensor_idx: int = 0, *, use_aicore: bool = False) -> int:
try:
return triton.runtime.driver.active.utils.get_device_properties(tensor_idx)['multiprocessor_count']
except Exception:
# Maybe we use a NPU device.
try:
if triton.runtime.driver.active.get_current_target().backend == 'npu':
props = triton.runtime.driver.active.utils.get_device_properties(tensor_idx)
return props['num_aicore'] if use_aicore else props['num_vectorcore']
except Exception:
logger.debug('Failed to get NPU multiprocessor count, falling back to 1.', exc_info=True)
return 1
@cache
def get_device_capability(device_index: int = 0) -> tuple[int, int]:
major, minor = torch.cuda.get_device_capability(device_index)
return int(major), int(minor)
@cache
def get_device_smem_optin(device_index: int = 0) -> int:
props = torch.cuda.get_device_properties(device_index)
return int(getattr(props, 'shared_memory_per_block_optin', props.shared_memory_per_block))
@cache
def get_available_device() -> str:
try:
return triton.runtime.driver.active.get_current_target().backend
except Exception:
_cpu_device_warning()
return 'cpu'
def map_triton_backend_to_torch_device() -> str:
backend = get_available_device() # 'cuda' | 'hip' | 'xpu' | 'cpu' | ...
return {'cuda': 'cuda', 'hip': 'cuda', 'xpu': 'xpu'}.get(backend, backend)
# For AMD GPUs, the triton backend is 'hip', while for Nvidia GPUs, the triton backend is 'cuda'.
# However, the torch backend is 'cuda' for both Nvidia and AMD GPUs.
# Therefore, we need to check the triton backend to determine the actual GPU vendor.
device = get_available_device() if get_available_device() != 'hip' else 'cuda'
device_torch_lib = getattr(torch, device)
device_platform = get_available_device()
device_name = map_triton_backend_to_torch_device()
IS_AMD = (device_platform == 'hip')
IS_ARM = platform.machine().lower() in ('aarch64', 'arm64')
IS_INTEL = (device_platform == 'xpu')
IS_INTEL_ALCHEMIST = (IS_INTEL and 'Intel(R) Arc(TM) A' in torch.xpu.get_device_name(0))
IS_NPU = (device_platform == 'npu')
IS_NVIDIA = (device_platform == 'cuda')
IS_NVIDIA_HOPPER = (
IS_NVIDIA and (
'NVIDIA H' in torch.cuda.get_device_name(0)
or torch.cuda.get_device_capability()[0] == 9
)
)
IS_NVIDIA_SM100 = (IS_NVIDIA and torch.cuda.get_device_capability()[0] == 10)
# NOTE: exactly 12.0 — 12.1 (GB10) is a different target that FlashQLA rejects at import time.
IS_NVIDIA_SM120 = (IS_NVIDIA and torch.cuda.get_device_capability() == (12, 0))
IS_NVIDIA_BLACKWELL = (IS_NVIDIA and torch.cuda.get_device_capability()[0] in (10, 12))
# Nvidia Ampere or newer, haven't check AMD and intel yet.
IS_TF32_SUPPORTED = (IS_NVIDIA and torch.cuda.get_device_capability(0)[0] >= 8)
IS_GATHER_SUPPORTED = hasattr(triton.language, 'gather')
IS_TMA_SUPPORTED = (
IS_NVIDIA
and torch.cuda.get_device_capability(0)[0] >= 9
and os.environ.get('FLA_USE_TMA', '0') == '1'
and (
hasattr(triton.language, '_experimental_make_tensor_descriptor')
or hasattr(triton.language, 'make_tensor_descriptor')
)
)
if IS_NVIDIA and not IS_TF32_SUPPORTED:
# Make old card happy, since triton will use tf32 by default.
# This is a workaround for old nvidia card.
os.environ['TRITON_F32_DEFAULT'] = 'ieee'
def _default_alloc_fn(size: int, alignment: int, stream: int | None):
return torch.empty(size, device=torch.device(device_name, device_torch_lib.current_device()), dtype=torch.int8)
if IS_TMA_SUPPORTED:
logger.info('TMA is supported, using TMA by default.')
triton.set_allocator(_default_alloc_fn)
elif IS_NVIDIA_BLACKWELL:
# Blackwell (SM100 datacenter / SM120 consumer): Triton compiler may emit global_scratch for
# autotuned kernels even without TMA. Register a default allocator to
# prevent NullAllocator crashes. See triton-lang/triton#10002.
logger.info('Blackwell detected: registering default global_scratch allocator.')
triton.set_allocator(_default_alloc_fn)
def get_all_max_shared_mem():
try:
return [
triton.runtime.driver.active.utils.get_device_properties(i)['max_shared_mem']
for i in range(device_torch_lib.device_count())
]
except Exception:
_cpu_device_warning()
return [-1]
class Backend(Enum):
ADA = 101376 # RTX 4090
AMPERE = 166912 # A100
HOPPER = 232448 # H100
DEFAULT = 102400 # Default
@classmethod
def get_shared_memory(cls, arch: str) -> int:
try:
return cls[arch.upper()].value
except KeyError:
return cls.DEFAULT.value
@cache
def check_shared_mem(arch: str = "none", tensor_idx: int = 0) -> bool:
try:
device_shared_mem_list = get_all_max_shared_mem()
max_shared_memory = device_shared_mem_list[tensor_idx]
return max_shared_memory >= Backend.get_shared_memory(arch)
except Exception:
return False
if check_pytorch_version('2.4'):
if device == 'cpu':
device = 'cuda'
device_torch_lib = getattr(torch, device)
autocast_custom_fwd = functools.partial(torch.amp.custom_fwd, device_type=device)
autocast_custom_bwd = functools.partial(torch.amp.custom_bwd, device_type=device)
def custom_device_ctx(index: int):
if index is None:
return contextlib.nullcontext()
try:
return device_torch_lib.device(index)
except (AttributeError, AssertionError, RuntimeError):
return contextlib.nullcontext()
else:
assert device == 'cuda', 'Only cuda device is supported for PyTorch version < 2.4.0.'
autocast_custom_fwd = device_torch_lib.amp.custom_fwd
autocast_custom_bwd = device_torch_lib.amp.custom_bwd
def custom_device_ctx(index: int):
if index is None:
return contextlib.nullcontext()
try:
return torch.cuda.device(index)
except (AttributeError, AssertionError, RuntimeError):
return contextlib.nullcontext()
+41
View File
@@ -0,0 +1,41 @@
# 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 logging
import warnings
import torch
from ._config import FLA_CI_ENV
logger = logging.getLogger(__name__)
def get_abs_err(x, y):
return (x.detach() - y.detach()).flatten().abs().max().item()
def get_err_ratio(x, y):
err = (x.detach() - y.detach()).flatten().square().mean().sqrt().item()
base = (x.detach()).flatten().square().mean().sqrt().item()
return err / (base + 1e-8)
def assert_close(prefix, ref, tri, ratio, warning=False, err_atol=1e-6):
abs_atol = get_abs_err(ref, tri)
error_rate = get_err_ratio(ref, tri)
msg = f"{prefix:>16} diff: {abs_atol:.6f} ratio: {error_rate:.6f}"
logger.info(msg)
if abs_atol <= err_atol:
return
assert not torch.isnan(ref).any(), f"{prefix}: NaN detected in ref"
assert not torch.isnan(tri).any(), f"{prefix}: NaN detected in tri"
if warning or (FLA_CI_ENV and (error_rate < 0.01 or abs_atol <= 0.3)):
if error_rate > ratio:
warnings.warn(msg)
else:
assert error_rate < ratio, msg
+23
View File
@@ -0,0 +1,23 @@
"""Composable mixing layers: attn and ffn both map [B,T,D] -> [B,T,D].
Depth mixing (AttnRes) is not a layer_specs kind. CausalLM reads
``config.attnres`` (off | full | block) and wraps DecoderBlock sublayers.
"""
from .block import DecoderBlock, build_attn, build_ffn
from .kda_attn import KDAAttention
from .latent_moe import LatentMoE
from .mla import GatedMLA
from .rmsnorm import RMSNorm
from .swiglu import SwiGLUMLP
__all__ = [
"DecoderBlock",
"GatedMLA",
"KDAAttention",
"LatentMoE",
"RMSNorm",
"SwiGLUMLP",
"build_attn",
"build_ffn",
]
+519
View File
@@ -0,0 +1,519 @@
"""
Attention Residual in one file
Reference:
Kimi Team, Guangyu Chen, Yu Zhang, Jianlin Su, Weixin Xu, Siyuan Pan,
Yaoyu Wang, Yucheng Wang, Guanduo Chen, et al.
"Attention Residuals." arXiv:2603.15031, 2026.
https://arxiv.org/abs/2603.15031
This module is a compact PyTorch reference implementation of:
- Full AttnRes
- Block AttnRes
- two-phase inter/intra-block computation from the paper
CausalLM wires Full/Block stacks when ``config.attnres`` is ``full`` or
``block``. Standard residual (``x += attn; x += ffn``) is ``attnres="off"``.
"""
import torch
import torch.nn.functional as F
from einops import rearrange
from torch import Tensor, nn
ATTNRES_MODES = ("off", "full", "block")
def exists(x):
return x is not None
def validate_attnres(mode: str, block_size: int | None) -> None:
if mode not in ATTNRES_MODES:
raise ValueError(f"attnres must be one of {ATTNRES_MODES}, got {mode!r}")
if block_size is not None and block_size < 1:
raise ValueError(f"attnres_block_size must be >= 1, got {block_size}")
def atomic_block_size(num_hidden_layers: int, attnres_block_size: int | None) -> int:
"""DecoderBlocks per AttnRes block, converted to attn|ffn atomic layers.
``None`` targets about 8 blocks: ``max(1, ceil(L / 8))`` DecoderBlocks.
"""
layers_per_block = (
attnres_block_size
if attnres_block_size is not None
else max(1, (num_hidden_layers + 7) // 8)
)
if layers_per_block < 1:
raise ValueError(f"attnres_block_size must be >= 1, got {layers_per_block}")
return layers_per_block * 2
class BorrowedSubLayer(nn.Module):
"""``fn(norm(x))`` without registering ``norm``/``fn`` (owned by DecoderBlock)."""
def __init__(self, norm: nn.Module, fn: nn.Module):
super().__init__()
self._borrowed = (norm, fn)
def forward(self, x: Tensor) -> Tensor:
norm, fn = self._borrowed
return fn(norm(x))
def rms(x: Tensor, eps: float):
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + eps)
class RMSNorm(nn.Module):
def __init__(self, dim: int, eps: float):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x: Tensor) -> Tensor:
return rms(x, self.eps) * self.weight
class DepthResidual(nn.Module):
"""
h_l = sum_i softmax_i(w_l^T RMSNorm(v_i))*v_i
Keep query and RMSNorm gain separate
Since q^T (gamma * RMS(v)) == (q * gamma)^T RMS(v),
we can fold gamma into q for scoring.
"""
def __init__(self, dim: int, eps: float = 1e-8, zero_init: bool = True):
super().__init__()
self.query = nn.Parameter(torch.zeros(dim))
self.norm = RMSNorm(dim, eps=eps)
if not zero_init:
nn.init.normal_(self.query, std=0.02)
def effective_query(self) -> Tensor:
return (self.query * self.norm.weight).float()
def logits(self, sources: Tensor | list[Tensor] | tuple[Tensor, ...]) -> Tensor:
sources = stack_layers(sources) # [n, b, t, d]
q = self.effective_query() # [d]
k = rms(sources.float(), self.norm.eps) # [n, b, t, d]
return torch.einsum("d, n b t d -> n b t", q, k)
def forward(self, sources: Tensor | list[Tensor] | tuple[Tensor, ...]) -> Tensor:
sources = stack_layers(sources)
weights = self.logits(sources).softmax(dim=0)
out = torch.einsum("n b t, n b t d -> b t d", weights, sources.float())
return out.to(sources.dtype)
class DepthResidualList(nn.Module):
def __init__(self, dim: int, depth: int, eps: float, zero_init: bool = True):
super().__init__()
# for L layers (depth), create depth residual modules
self.layers = nn.ModuleList(
[DepthResidual(dim, eps=eps, zero_init=zero_init) for _ in range(depth)]
)
def __getitem__(self, idx: int) -> DepthResidual:
return self.layers[idx]
def __iter__(self):
return iter(self.layers)
def __len__(self):
return len(self.layers)
# attnres stacks
class FullAttnResStack(nn.Module):
"""
Full AttnRes over atomic layers
eg: f_1,...,f_L
Each entry in `layers` should already be a full atomic layer fxn
x -> f_l(x)
"""
def __init__(
self,
dim: int,
layers,
*,
eps: float = 1e-8,
zero_init_queries: bool = True,
is_final_aggregate: bool = True,
):
super().__init__()
self.layers = nn.ModuleList(list(layers))
self.eps = eps
depth = len(self.layers)
self.residuals = DepthResidualList(dim, depth, eps, zero_init_queries)
self.final_residual = (
DepthResidual(dim, eps, zero_init_queries) if is_final_aggregate else None
)
def forward_naive(self, x: Tensor) -> Tensor:
sources = [x]
for layer, residual in zip(self.layers, self.residuals):
h = residual(sources)
out = layer(h)
sources.append(out)
return (
self.final_residual(sources) if exists(self.final_residual) else sources[-1]
)
def forward_two_phase(self, x: Tensor, schedule_block_size: int) -> Tensor:
assert schedule_block_size > 0
sources = [x]
depth = len(self.layers)
start = 0
while start < depth:
end = min(start + schedule_block_size, depth)
queries = torch.stack(
[self.residuals[i].effective_query() for i in range(start, end)], dim=0
)
inter_sources = stack_layers(sources)
inter_stats = attn_with_stats(queries, inter_sources, self.eps)
local_outputs = [] # outputs of intra-block
for local_idx, layer_idx in enumerate(range(start, end)):
stats = inter_stats.select(local_idx)
if len(local_outputs) > 0:
intra_sources = stack_layers(local_outputs)
intra = attn_with_stats(
queries[local_idx : local_idx + 1], intra_sources, self.eps
).select(0)
stats = merge_attn_stats(stats, intra)
h = stats.normalized()
out = self.layers[layer_idx](h)
local_outputs.append(out)
sources.append(out)
start = end
return (
self.final_residual(sources) if exists(self.final_residual) else sources[-1]
)
def forward(self, x: Tensor, schedule_block_size: int | None = None) -> Tensor:
if schedule_block_size is None:
return self.forward_naive(x)
return self.forward_two_phase(x, schedule_block_size)
class BlockAttnResStack(nn.Module):
"""
Block AttnRes over atomic layers
`block_size` is in atomic layers, not Transformer blocks.
Eg: block_size=8 -> 4 transformer blocks when layers alternate attn/MLP
The default forward path is the two-phase algorithm from the paper:
phase 1: batch inter-block attn from all queries in the block
phase 2: merge the evolving intra-block partial sum with online softmax
"""
def __init__(
self,
dim: int,
layers,
*,
block_size: int,
eps: float = 1e-8,
zero_init_queries: bool = True,
is_final_aggregate: bool = True,
):
super().__init__()
self.layers = nn.ModuleList(list(layers))
assert len(self.layers) > 0
assert block_size > 0
self.block_size = block_size
self.eps = eps
depth = len(self.layers)
self.residuals = DepthResidualList(
dim, depth, eps=eps, zero_init=zero_init_queries
)
self.final_residual = (
DepthResidual(dim, eps=eps, zero_init=zero_init_queries)
if is_final_aggregate
else None
)
def forward_naive(self, x: Tensor) -> Tensor:
blocks = [x] # b_0=embedding/input representation
partial = None
for layer_idx, (layer, residual) in enumerate(
zip(self.layers, self.residuals), start=1
):
sources = blocks if partial is None else blocks + [partial]
h = residual(sources)
out = layer(h)
partial = out if partial is None else (partial + out)
if (layer_idx % self.block_size == 0) or (layer_idx == len(self.layers)):
blocks.append(partial)
partial = None
return (
self.final_residual(blocks) if exists(self.final_residual) else blocks[-1]
)
def _run_block_two_phase(
self, blocks: list[Tensor], start: int, end: int
) -> Tensor:
queries = torch.stack(
[self.residuals[i].effective_query() for i in range(start, end)], dim=0
)
inter_sources = stack_layers(blocks)
inter = attn_with_stats(queries, inter_sources, self.eps)
partial = None
for local_idx, layer_idx in enumerate(range(start, end)):
stats = inter.select(local_idx)
if partial is not None:
intra = single_source_stats(queries[local_idx], partial, self.eps)
stats = merge_attn_stats(stats, intra)
h = stats.normalized()
out = self.layers[layer_idx](h)
partial = out if partial is None else (partial + out)
return partial
def forward(self, x: Tensor) -> Tensor:
blocks = [x]
depth = len(self.layers)
start = 0
while start < depth:
end = min(start + self.block_size, depth)
blocks.append(self._run_block_two_phase(blocks, start, end))
start = end
return (
self.final_residual(blocks) if exists(self.final_residual) else blocks[-1]
)
# helpers
def stack_layers(sources: Tensor | list[Tensor] | tuple[Tensor, ...]) -> Tensor:
if isinstance(sources, Tensor):
assert sources.ndim == 4, f"expected [n, b, t, d] got {tuple(sources.shape)}"
return sources
assert len(sources) > 0, "needs at least one source"
return torch.stack(tuple(sources), dim=0)
class SingleAttnStats:
def __init__(self, numer: Tensor, denom: Tensor, max: Tensor):
self.numer = numer # [b,t,d]
self.max = max # [b,t]
self.denom = denom # [b,t]
def normalized(self) -> Tensor:
return self.numer / self.denom[..., None]
class AttnStats:
# store the numerator => e^{s_{j}-m} * v_j where m is the max score so far
# store the max m = max(s_j)
# store the denominator sum_j e^{s_{j}-m}
def __init__(self, numer: Tensor, denom: Tensor, max: Tensor):
self.numer = numer # [q,b,t,d]
self.max = max # [q,b,t]
self.denom = denom # [q,b,t]
def select(self, idx: int) -> "SingleAttnStats":
return SingleAttnStats(self.numer[idx], self.denom[idx], self.max[idx])
def attn_with_stats(queries: Tensor, sources: Tensor, eps: float = 1e-8) -> AttnStats:
"""
queries: [q, d]
sources: [n, b, t, d]
Returns the following for online softmax:
numer = sum_i exp(logit_i - m)*v_i
m = max_i logit_i
denom = sum_i exp(logit_i - m)
"""
normed = rms(sources, eps)
logits = torch.einsum("q d, n b t d -> q n b t", queries, normed)
m = logits.amax(dim=1)
weights = torch.exp(logits - m[:, None])
numer = torch.einsum("q n b t, n b t d -> q b t d", weights, sources)
denom = weights.sum(dim=1)
return AttnStats(numer, denom, m)
def single_source_stats(
query: Tensor, source: Tensor, eps: float = 1e-8
) -> SingleAttnStats:
score = torch.einsum("d, b t d -> b t", query, rms(source, eps))
denom = torch.ones_like(score)
return SingleAttnStats(source, denom, score)
def merge_attn_stats(a: SingleAttnStats, b: SingleAttnStats) -> SingleAttnStats:
m = torch.maximum(a.max, b.max)
wa = torch.exp(a.max - m)
wb = torch.exp(b.max - m)
numer = wa[..., None] * a.numer + wb[..., None] * b.numer
denom = wa * a.denom + wb * b.denom
return SingleAttnStats(numer, denom, m)
# transformer
class PreNorm(nn.Module):
def __init__(self, dim: int, fn: nn.Module, eps: float = 1e-8):
super().__init__()
self.norm = RMSNorm(dim, eps=eps)
self.fn = fn
def forward(self, x: Tensor) -> Tensor:
return self.fn(self.norm(x))
class CausalAttention(nn.Module):
def __init__(
self, dim: int, heads: int = 8, dim_head: int = 64, dropout: float = 0.0
):
super().__init__()
inner_dim = heads * dim_head
self.heads = heads
self.dim_head = dim_head
self.dropout = dropout
self.to_qkv = nn.Linear(dim, inner_dim * 3, bias=False)
self.to_out = nn.Linear(inner_dim, dim, bias=False)
def forward(self, x: Tensor) -> Tensor:
q, k, v = self.to_qkv(x).chunk(3, dim=-1)
def split_heads(y: Tensor) -> Tensor:
return rearrange(y, "b t (h d) -> b h t d", h=self.heads)
q, k, v = map(split_heads, (q, k, v))
out = F.scaled_dot_product_attention(
q, k, v, is_causal=True, dropout_p=self.dropout if self.training else 0.0
)
out = rearrange(out, "b h t d -> b t (h d)")
return self.to_out(out)
class SwiGLU(nn.Module):
def __init__(self, dim: int, mult: int = 4, dropout: float = 0.0):
# dropout not needed unless training on a smaller training data
super().__init__()
inner_dim = dim * mult
self.to_hidden = nn.Linear(dim, inner_dim * 2, bias=False)
self.to_out = nn.Linear(inner_dim, dim, bias=False)
self.dropout = nn.Dropout(dropout)
def forward(self, x: Tensor) -> Tensor:
gate, value = self.to_hidden(x).chunk(2, dim=-1)
x = F.silu(gate) * value
x = self.dropout(x)
return self.to_out(x)
class AttnResTransformer(nn.Module):
"""
Small GPT-style reference model using AttnRes
Using plain PyTorch: tok/pos embedding, alternating
causal attn, SwiGLU MLP layers, final norm, output head.
"""
def __init__(
self,
*,
num_tokens: int,
dim: int,
depth: int,
max_seq_len: int,
heads: int = 8,
dim_head: int = 64,
ff_mult: int = 4,
attn_dropout: float = 0.0,
ff_dropout: float = 0.0,
attnres: str = "block", # full or block
block_size: int = 8,
zero_init_queries: bool = True,
is_final_aggregate: bool = True,
eps: float = 1e-8,
):
super().__init__()
assert attnres in {"full", "block"}
self.max_seq_len = max_seq_len
self.attnres = attnres
self.token_emb = nn.Embedding(num_tokens, dim)
self.pos_emb = nn.Embedding(max_seq_len, dim)
atomic_layers = []
for _ in range(depth):
atomic_layers.append(
PreNorm(dim, CausalAttention(dim, heads, dim_head, attn_dropout), eps)
)
atomic_layers.append(PreNorm(dim, SwiGLU(dim, ff_mult, ff_dropout), eps))
if attnres == "full":
self.backbone = FullAttnResStack(
dim,
atomic_layers,
eps=eps,
zero_init_queries=zero_init_queries,
is_final_aggregate=is_final_aggregate,
)
else:
self.backbone = BlockAttnResStack(
dim,
atomic_layers,
block_size=block_size,
eps=eps,
zero_init_queries=zero_init_queries,
is_final_aggregate=is_final_aggregate,
)
self.final_norm = RMSNorm(dim, eps)
self.to_logits = nn.Linear(dim, num_tokens, bias=False)
def forward(self, ids: Tensor, schedule_block_size: int | None = None) -> Tensor:
b, t = ids.shape
assert t <= self.max_seq_len
pos = torch.arange(t, device=ids.device)
x = self.token_emb(ids) + self.pos_emb(pos)[None, :, :]
if self.attnres == "full":
x = self.backbone(x, schedule_block_size=schedule_block_size)
else:
x = self.backbone(x)
x = self.final_norm(x)
return self.to_logits(x)
__all__ = [
"ATTNRES_MODES",
"RMSNorm",
"DepthResidual",
"DepthResidualList",
"FullAttnResStack",
"BlockAttnResStack",
"BorrowedSubLayer",
"PreNorm",
"CausalAttention",
"SwiGLU",
"AttnResTransformer",
"atomic_block_size",
"validate_attnres",
]
+51
View File
@@ -0,0 +1,51 @@
"""Decoder block: x += attn(norm(x)); x += ffn(norm(x)).
attn/ffn are any modules with forward: [B,T,D] -> [B,T,D].
"""
from __future__ import annotations
from torch import nn
from .kda_attn import KDAAttention
from .latent_moe import LatentMoE
from .mla import GatedMLA
from .rmsnorm import RMSNorm
from .swiglu import SwiGLUMLP
def build_attn(config, kind: str) -> nn.Module:
if kind == "kda":
return KDAAttention.from_config(config)
if kind == "mla":
return GatedMLA.from_config(config)
raise ValueError(f"unknown attn kind: {kind}")
def build_ffn(config, kind: str) -> nn.Module:
if kind == "swiglu":
return SwiGLUMLP.from_config(config)
if kind == "moe":
return LatentMoE.from_config(config)
raise ValueError(f"unknown ffn kind: {kind}")
class DecoderBlock(nn.Module):
def __init__(self, hidden_size: int, norm_eps: float, attn: nn.Module, ffn: nn.Module):
super().__init__()
self.attn_norm = RMSNorm(hidden_size, norm_eps)
self.attn = attn
self.ffn_norm = RMSNorm(hidden_size, norm_eps)
self.ffn = ffn
@classmethod
def from_spec(cls, config, attn_kind: str, ffn_kind: str) -> DecoderBlock:
return cls(
config.hidden_size,
config.norm_eps,
build_attn(config, attn_kind),
build_ffn(config, ffn_kind),
)
def forward(self, x):
x = x + self.attn(self.attn_norm(x))
return x + self.ffn(self.ffn_norm(x))
+99
View File
@@ -0,0 +1,99 @@
"""KDA attention: project q/k/v/g/beta, run chunk_kda, project back to D."""
from __future__ import annotations
import torch
from torch import nn
from ..ops.api import chunk_kda
class KDAAttention(nn.Module):
"""Mixing module: x [B,T,D] -> y [B,T,D]."""
def __init__(
self,
hidden_size: int,
num_heads: int,
num_value_heads: int,
head_dim: int,
*,
chunk_size: int = 16,
initializer_range: float = 0.02,
use_gate_in_kernel: bool = True,
use_qk_l2norm_in_kernel: bool = True,
use_beta_sigmoid_in_kernel: bool = True,
lower_bound: float | None = -5.0,
kda_backend: str = "reference",
):
super().__init__()
if num_value_heads % num_heads:
raise ValueError("num_value_heads must be divisible by num_heads")
self.hidden_size = hidden_size
self.num_heads = num_heads
self.num_value_heads = num_value_heads
self.head_dim = head_dim
self.chunk_size = chunk_size
self.initializer_range = initializer_range
self.use_gate_in_kernel = use_gate_in_kernel
self.use_qk_l2norm_in_kernel = use_qk_l2norm_in_kernel
self.use_beta_sigmoid_in_kernel = use_beta_sigmoid_in_kernel
self.lower_bound = lower_bound
self.kda_backend = kda_backend
H, HV, K, V = num_heads, num_value_heads, head_dim, head_dim
self.q_proj = nn.Linear(hidden_size, H * K, bias=False)
self.k_proj = nn.Linear(hidden_size, H * K, bias=False)
self.v_proj = nn.Linear(hidden_size, HV * V, bias=False)
self.g_proj = nn.Linear(hidden_size, HV * K, bias=False)
self.beta_proj = nn.Linear(hidden_size, HV, bias=False)
self.o_proj = nn.Linear(HV * V, hidden_size, bias=False)
self.A_log = nn.Parameter(torch.zeros(HV))
# With safe_gate=-5, bias=-4 starts at g≈-0.09 (about 91% state retention).
self.dt_bias = nn.Parameter(torch.full((HV, K), -4.0))
self.apply(self._init_weights)
@classmethod
def from_config(cls, config) -> KDAAttention:
return cls(
hidden_size=config.hidden_size,
num_heads=config.num_heads,
num_value_heads=getattr(config, "num_value_heads", config.num_heads),
head_dim=config.head_dim,
chunk_size=config.chunk_size,
initializer_range=config.initializer_range,
use_gate_in_kernel=config.use_gate_in_kernel,
use_qk_l2norm_in_kernel=config.use_qk_l2norm_in_kernel,
use_beta_sigmoid_in_kernel=config.use_beta_sigmoid_in_kernel,
lower_bound=config.lower_bound,
kda_backend=config.kda_backend,
)
def _init_weights(self, module):
if isinstance(module, nn.Linear):
nn.init.normal_(module.weight, std=self.initializer_range)
def forward(self, x: torch.Tensor):
B, T, _ = x.shape
H, HV, K, V = self.num_heads, self.num_value_heads, self.head_dim, self.head_dim
q = self.q_proj(x).view(B, T, H, K)
k = self.k_proj(x).view(B, T, H, K)
v = self.v_proj(x).view(B, T, HV, V)
g_raw = self.g_proj(x).view(B, T, HV, K)
beta_raw = self.beta_proj(x).view(B, T, HV)
o, _ = chunk_kda(
q,
k,
v,
g_raw,
beta_raw,
A_log=self.A_log,
dt_bias=self.dt_bias,
use_qk_l2norm_in_kernel=self.use_qk_l2norm_in_kernel,
use_gate_in_kernel=self.use_gate_in_kernel,
use_beta_sigmoid_in_kernel=self.use_beta_sigmoid_in_kernel,
safe_gate=self.lower_bound is not None,
lower_bound=self.lower_bound,
chunk_size=self.chunk_size,
backend=self.kda_backend,
)
return self.o_proj(o.reshape(B, T, HV * V))
+118
View File
@@ -0,0 +1,118 @@
"""Stable LatentMoE (K3): shared 全宽 + routed 半宽专家 + SiTU-GLU + Top-k.
对照 learning/kimi-k3-notes §Stable LatentMoE:
z = W_down(x) [B, T, ℓ] ℓ = d/2 latent 接口宽
u = Σ_{i∈Top-k(x)} p_i E_i^rt(z) [B, T, ℓ] routed 专家只在 ℓ 上算
y = Σ_j E_j^sh(x) + W_up RMSNorm(u) [B, T, d] shared 全宽
SiTU-GLU: gate = β1·tanh(W_g x/β1)⊙σ(W_g x); up = β2·tanh(W_u x/β2)
||SiTU-GLU||_∞ ≤ β1·β2 (=100), 原点附近≈SwiGLU, 远端软饱和防低精度溢出.
E: R^in → R^in (内部中间维 d_ff).
Router: Top-k logits 基于全宽 x (笔记 Topk(x)); 归一化权重取 softmax(topk).
"""
from __future__ import annotations
import torch
import torch.nn.functional as F
from torch import nn
from .rmsnorm import RMSNorm
class SiTU(nn.Module):
"""SiTU-GLU expert: gate 支软上限 β1, up 支软上限 β2, 输出回到输入维."""
def __init__(self, dim_in: int, dim_ff: int, beta1: float = 4.0, beta2: float = 25.0):
super().__init__()
self.beta1, self.beta2 = beta1, beta2
self.w_g = nn.Linear(dim_in, dim_ff, bias=False)
self.w_u = nn.Linear(dim_in, dim_ff, bias=False)
self.w_o = nn.Linear(dim_ff, dim_in, bias=False)
def forward(self, x: torch.Tensor):
wg = self.w_g(x)
g = self.beta1 * torch.tanh(wg / self.beta1) * torch.sigmoid(wg)
u = self.beta2 * torch.tanh(self.w_u(x) / self.beta2)
return self.w_o(g * u)
class LatentMoE(nn.Module):
def __init__(
self,
hidden_size: int,
latent_size: int,
n_routed: int,
top_k: int,
n_shared: int,
d_ff: int,
beta1: float = 4.0,
beta2: float = 25.0,
):
super().__init__()
self.latent_size = latent_size
self.n_routed = n_routed
self.top_k = top_k
self.down = nn.Linear(hidden_size, latent_size, bias=False) # W↓
self.router = nn.Linear(hidden_size, n_routed, bias=False) # Top-k logits
self.shared = nn.ModuleList(
[SiTU(hidden_size, d_ff, beta1, beta2) for _ in range(n_shared)]
)
self.experts = nn.ModuleList(
[SiTU(latent_size, d_ff, beta1, beta2) for _ in range(n_routed)]
)
self.norm = RMSNorm(latent_size)
self.up = nn.Linear(latent_size, hidden_size, bias=False) # W↑
self.last_route_ids: torch.Tensor | None = None
@classmethod
def from_config(cls, config) -> LatentMoE:
return cls(
config.hidden_size,
config.moe_latent_size,
config.n_routed,
config.top_k,
config.n_shared,
config.moe_d_ff,
config.situ_beta1,
config.situ_beta2,
)
def forward(self, x: torch.Tensor):
B, T, _ = x.shape
z = self.down(x) # [B, T, ℓ]
logits = self.router(x) # [B, T, n_routed]
topk = torch.topk(logits, self.top_k, dim=-1)
ids = topk.indices # [B, T, k]
self.last_route_ids = ids.detach()
probs = F.softmax(topk.values, dim=-1) # [B, T, k]
# 向量化 routed: 预计算全部专家输出, 按 token 的 Top-k id 取
all_out = torch.stack([e(z) for e in self.experts]) # [R, B, T, ℓ]
all_out = all_out.permute(1, 2, 0, 3).reshape(B * T, self.n_routed, self.latent_size)
u = torch.zeros(B, T, self.latent_size, device=x.device, dtype=x.dtype)
for i in range(self.top_k):
idx = ids[:, :, i].reshape(B * T) # [B*T]
sel = all_out[torch.arange(B * T, device=x.device), idx] # [B*T, ℓ]
u += probs[:, :, i : i + 1] * sel.reshape(B, T, self.latent_size)
shared_out = torch.stack([e(x) for e in self.shared]).sum(0) # [B, T, d]
return shared_out + self.up(self.norm(u))
def moe_route_frac(model: nn.Module) -> torch.Tensor | None:
"""Mean expert occupancy over LatentMoE layers from the last forward."""
hists: list[torch.Tensor] = []
n_routed: int | None = None
for module in model.modules():
if not isinstance(module, LatentMoE) or module.last_route_ids is None:
continue
n_routed = module.n_routed
ids = module.last_route_ids.reshape(-1)
hists.append(torch.bincount(ids, minlength=n_routed).float())
if not hists or n_routed is None:
return None
stacked = torch.stack(hists).sum(0)
return stacked / stacked.sum().clamp_min(1.0)
+97
View File
@@ -0,0 +1,97 @@
"""Gated MLA (K3): NoPE, latent KV compression, matrix absorption, full-rank output gate.
K3 相对 DeepSeek MLA 的三个改动 (对照 learning/kimi-k3-notes):
1. NoPE — 不显式 RoPE; 位置感交给夹层 KDA 的 decay/gate。
2. 矩阵吸收 — 训练/推理都不解压 K/V: q 吸收 W_UK 后直接与 latent c 内积,
输出先在 latent 加权再乘 W_UV 还原 (v2 吸收版)。
3. Full-rank 输出门 — y = W_o[ σ(W_g x) ⊙ õ ]。
形状 (小规模 toy, d 为 hidden):
c = RMSNorm(kv_down(x)) [B, T, r] latent
q = q_up(RMSNorm(q_down(x))) [B, T, H, d_q] d_q = d_nope (NoPE)
W_UK = kv_up[.., :H*d_q].view(H,d_q,r) W_UV = kv_up[.., H*d_q:].view(H,d_v,r)
score = (q @ W_UK^T) @ c^T [B, H, T, T] causal
õ = (softmax(score) @ c) @ W_UV^T [B, T, H, d_v]
y = o_proj( σ(W_g x) ⊙ õ_head ) [B, T, d]
"""
from __future__ import annotations
import torch
import torch.nn.functional as F
from torch import nn
from .rmsnorm import RMSNorm
class GatedMLA(nn.Module):
def __init__(
self,
hidden_size: int,
num_heads: int,
kv_lora_rank: int,
q_lora_rank: int,
qk_nope_head_dim: int,
v_head_dim: int,
):
super().__init__()
self.hidden_size = hidden_size
self.num_heads = num_heads
self.qk_nope_head_dim = qk_nope_head_dim
self.v_head_dim = v_head_dim
# Q 低秩路径 (NoPE, 只有 nope 段)
self.q_down = nn.Linear(hidden_size, q_lora_rank, bias=False)
self.q_norm = RMSNorm(q_lora_rank)
self.q_up = nn.Linear(q_lora_rank, num_heads * qk_nope_head_dim, bias=False)
# KV latent 压缩 + 解压 (W_UK | W_UV 拼接在同一矩阵里)
self.kv_down = nn.Linear(hidden_size, kv_lora_rank, bias=False)
self.kv_norm = RMSNorm(kv_lora_rank)
self.kv_up = nn.Linear(
kv_lora_rank, num_heads * (qk_nope_head_dim + v_head_dim), bias=False
)
# Full-rank 输出门: σ(W_g x) 与 õ (H*d_v 维) 逐元素相乘
self.gate = nn.Linear(hidden_size, num_heads * v_head_dim, bias=False)
self.o_proj = nn.Linear(num_heads * v_head_dim, hidden_size, bias=False)
@classmethod
def from_config(cls, config) -> GatedMLA:
return cls(
config.hidden_size,
config.num_heads,
config.kv_lora_rank,
config.q_lora_rank,
config.qk_nope_head_dim,
config.v_head_dim,
)
def forward(self, x: torch.Tensor):
B, T, _ = x.shape
H, r = self.num_heads, self.kv_up.in_features
c = self.kv_norm(self.kv_down(x)) # [B, T, r]
q = self.q_up(self.q_norm(self.q_down(x))) # [B, T, H*d_q]
q = q.view(B, T, H, self.qk_nope_head_dim) # [B, T, H, d_q]
w = self.kv_up.weight # [H*(d_q+d_v), r]
w_uk = w[: H * self.qk_nope_head_dim].view(H, self.qk_nope_head_dim, r)
w_uv = w[H * self.qk_nope_head_dim :].view(H, self.v_head_dim, r)
# 吸收 W_UK 进 query: score = (q @ W_UK^T) @ c^T
q_absorb = torch.einsum("bthd,hdj->bthj", q, w_uk) # [B, T, H, r]
scores = torch.einsum("bthj,bsj->bhts", q_absorb, c) # [B, H, T, T]
mask = torch.triu(
torch.ones(T, T, dtype=torch.bool, device=x.device), diagonal=1
)
scores = scores.masked_fill(mask, float("-inf"))
attn = F.softmax(scores, dim=-1) # [B, H, T, T]
# 先在 latent 加权, 再乘 W_UV^T 还原 v —— 永不解压
latent_out = torch.einsum("bhts,bsj->bhtj", attn, c) # [B, H, T, r]
o_heads = torch.einsum("bhtj,hvj->bhtv", latent_out, w_uv) # [B, H, T, d_v]
o_heads = o_heads.transpose(1, 2).reshape(B, T, H * self.v_head_dim)
gate = torch.sigmoid(self.gate(x)) # [B, T, H*d_v]
return self.o_proj(gate * o_heads) # [B, T, d]
+17
View File
@@ -0,0 +1,17 @@
"""RMSNorm used by attention, FFN, and the final LM stem."""
from __future__ import annotations
import torch
from torch import nn
class RMSNorm(nn.Module):
def __init__(self, dim: int, eps: float = 1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(dim))
self.eps = eps
def forward(self, x: torch.Tensor):
dtype = x.dtype
x = x.float()
return (x * torch.rsqrt(x.square().mean(-1, keepdim=True) + self.eps)).to(dtype) * self.weight
+20
View File
@@ -0,0 +1,20 @@
"""SwiGLU FFN: x [B,T,D] -> y [B,T,D]."""
from __future__ import annotations
import torch.nn.functional as F
from torch import nn
class SwiGLUMLP(nn.Module):
def __init__(self, hidden_size: int, intermediate_size: int):
super().__init__()
self.w1 = nn.Linear(hidden_size, intermediate_size, bias=False)
self.w3 = nn.Linear(hidden_size, intermediate_size, bias=False)
self.w2 = nn.Linear(intermediate_size, hidden_size, bias=False)
@classmethod
def from_config(cls, config) -> SwiGLUMLP:
return cls(config.hidden_size, config.intermediate_size)
def forward(self, x):
return self.w2(F.silu(self.w1(x)) * self.w3(x))
+7
View File
@@ -0,0 +1,7 @@
"""Configs and the single CausalLM entry."""
from .causal_lm import CausalLM
from .config import KDAConfig
from .k3_config import K3Config
__all__ = ["CausalLM", "K3Config", "KDAConfig"]
+123
View File
@@ -0,0 +1,123 @@
"""Causal LM stem: embed -> DecoderBlock* -> norm -> lm_head.
KDA-only and K3-like both use this class. Config.layer_specs() chooses
attn/ffn per layer: ("kda"|"mla", "swiglu"|"moe").
``config.attnres`` selects the depth mixer:
off — standard residual inside each DecoderBlock (default)
full — Full AttnRes over attn|ffn sublayers
block — Block AttnRes (K3); block size from ``attnres_block_size``
"""
from __future__ import annotations
import torch
import torch.nn.functional as F
from torch import nn
from torch.utils.checkpoint import checkpoint as activation_checkpoint
from ..layers.attn_res import (
BlockAttnResStack,
BorrowedSubLayer,
FullAttnResStack,
atomic_block_size,
)
from ..layers.block import DecoderBlock
from ..layers.rmsnorm import RMSNorm
def _build_mixer(config, blocks: nn.ModuleList):
mode = getattr(config, "attnres", "off")
if mode == "off":
return None
atomics = []
for block in blocks:
atomics.append(BorrowedSubLayer(block.attn_norm, block.attn))
atomics.append(BorrowedSubLayer(block.ffn_norm, block.ffn))
kwargs = dict(
eps=config.norm_eps,
zero_init_queries=getattr(config, "attnres_zero_init_queries", True),
is_final_aggregate=getattr(config, "attnres_final_aggregate", True),
)
if mode == "full":
return FullAttnResStack(config.hidden_size, atomics, **kwargs)
if mode == "block":
return BlockAttnResStack(
config.hidden_size,
atomics,
block_size=atomic_block_size(
config.num_hidden_layers, getattr(config, "attnres_block_size", None)
),
**kwargs,
)
raise ValueError(f"unknown attnres mode: {mode!r}")
class CausalLM(nn.Module):
def __init__(self, config):
super().__init__()
self.config = config
self.attnres = getattr(config, "attnres", "off")
self.embedding = nn.Embedding(config.vocab_size, config.hidden_size)
self.blocks = nn.ModuleList(
[
DecoderBlock.from_spec(config, attn, ffn)
for attn, ffn in config.layer_specs()
]
)
self.mixer = _build_mixer(config, self.blocks)
self.gradient_checkpointing = bool(
getattr(config, "gradient_checkpointing", False)
)
self.norm = RMSNorm(config.hidden_size, config.norm_eps)
self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
nn.init.normal_(self.embedding.weight, std=config.initializer_range)
nn.init.normal_(self.lm_head.weight, std=config.initializer_range)
if config.tie_word_embeddings:
self.lm_head.weight = self.embedding.weight
def forward(
self,
input_ids: torch.Tensor,
labels: torch.Tensor | None = None,
ignore_index: int = -100,
):
x = self.embedding(input_ids)
if self.mixer is None:
for block in self.blocks:
if self.gradient_checkpointing and self.training:
x = activation_checkpoint(block, x, use_reentrant=False)
else:
x = block(x)
elif self.gradient_checkpointing and self.training:
x = activation_checkpoint(self.mixer, x, use_reentrant=False)
else:
x = self.mixer(x)
logits = self.lm_head(self.norm(x))
if labels is None:
return logits
return F.cross_entropy(
logits[:, :-1].reshape(-1, logits.size(-1)),
labels[:, 1:].reshape(-1),
ignore_index=ignore_index,
)
@torch.inference_mode()
def generate(
self,
input_ids: torch.Tensor,
max_new_tokens: int,
temperature: float = 0.0,
eos_token_id: int | None = None,
):
for _ in range(max_new_tokens):
logits = self(input_ids)[:, -1]
if temperature > 0:
probs = F.softmax(logits / temperature, dim=-1)
next_token = torch.multinomial(probs, 1)
else:
next_token = logits.argmax(-1, keepdim=True)
input_ids = torch.cat((input_ids, next_token), dim=1)
if eos_token_id is not None and (next_token.squeeze(-1) == eos_token_id).all():
break
return input_ids
+62
View File
@@ -0,0 +1,62 @@
"""KDAConfig — toy Causal LM hyperparameters.
Defaults match the working reference-backend model: GVA with G=2,
safe gate (lower_bound=-5), q/k L2-norm and beta sigmoid inside the op.
"""
from __future__ import annotations
from dataclasses import dataclass
@dataclass
class KDAConfig:
hidden_size: int = 64
num_hidden_layers: int = 2
num_heads: int = 4
num_value_heads: int = 8 # G = num_value_heads // num_heads
head_dim: int = 16
chunk_size: int = 16
vocab_size: int = 256
intermediate_size: int = 128
max_position_embeddings: int = 128
initializer_range: float = 0.02
norm_eps: float = 1e-6
use_gate_in_kernel: bool = True
use_qk_l2norm_in_kernel: bool = True
use_beta_sigmoid_in_kernel: bool = True
lower_bound: float | None = -5.0
tie_word_embeddings: bool = False
kda_backend: str = "reference" # reference | triton | fla
attnres: str = "off" # off | full | block
attnres_block_size: int | None = None # DecoderBlocks / block; None ≈ L/8
attnres_zero_init_queries: bool = True
attnres_final_aggregate: bool = True
gradient_checkpointing: bool = False
@property
def H(self) -> int: return self.num_heads
@property
def G(self) -> int: return self.num_value_heads // self.num_heads
@property
def HV(self) -> int: return self.num_value_heads
@property
def K(self) -> int: return self.head_dim
@property
def V(self) -> int: return self.head_dim
def __post_init__(self):
from ..layers.attn_res import validate_attnres
if self.num_value_heads % self.num_heads:
raise ValueError("num_value_heads must be divisible by num_heads")
supported = {"reference", "triton", "fla", "torch", "auto"}
if self.kda_backend not in supported:
raise ValueError(f"kda_backend must be one of {sorted(supported)}")
validate_attnres(self.attnres, self.attnres_block_size)
def layer_specs(self) -> list[tuple[str, str]]:
return [("kda", "swiglu")] * self.num_hidden_layers
+128
View File
@@ -0,0 +1,128 @@
"""K3Config — Kimi K3 架构的小规模复现配置 (KDA + Gated MLA + Stable LatentMoE).
对照 learning/kimi-k3-notes §尺寸速查 (真实 K3 → 本 toy 缩比):
hidden 7168 → 256; L 93 → 4; H=HV 96 → 8; K=V 128 → 16;
MLA kv_lora 512 → 32, q_lora 1536 → 64, nope/v 128 → 16;
MoE ℓ=d/2=3584 → 128, 896/16 → 16/2, shared 2, d_ff 3072 → 96.
Hybrid Attention (K3): 每 4 层 1 次 Gated MLA, 末层强制 MLA.
Presets:
toy — ~8M, 自训 8k SP, 本地过拟合
0.5b — ~482M, Qwen3 词表, 32–40GB bf16;默认 step 是冒烟,翻译前置用 --max-tokens
"""
from __future__ import annotations
from dataclasses import dataclass
# Qwen3 config.json; train_k3 overrides with len(tokenizer).
QWEN3_VOCAB_SIZE = 151936
@dataclass
class K3Config:
# 主干
hidden_size: int = 256
num_hidden_layers: int = 4
vocab_size: int = 8192 # toy: data/spm_4k; 0.5b: Qwen3
initializer_range: float = 0.02
norm_eps: float = 1e-6
tie_word_embeddings: bool = False
max_position_embeddings: int = 2048 # NoPE, 仅语义保留
# KDA (K3: H = HV = 96, 无 GVA)
num_heads: int = 8
head_dim: int = 16
chunk_size: int = 16
lower_bound: float | None = -5.0
use_gate_in_kernel: bool = True
use_qk_l2norm_in_kernel: bool = True
use_beta_sigmoid_in_kernel: bool = True
# Gated MLA (NoPE)
kv_lora_rank: int = 32
q_lora_rank: int = 64
qk_nope_head_dim: int = 16
v_head_dim: int = 16
# Stable LatentMoE
moe_latent_size: int = 128 # ℓ = d/2
n_routed: int = 16
top_k: int = 2
n_shared: int = 2
moe_d_ff: int = 96
situ_beta1: float = 4.0
situ_beta2: float = 25.0
kda_backend: str = "reference"
# Depth mixer. off = DecoderBlock residual; block matches K3.
attnres: str = "off" # off | full | block
attnres_block_size: int | None = None # DecoderBlocks / AttnRes block; None ≈ L/8
attnres_zero_init_queries: bool = True
attnres_final_aggregate: bool = True
gradient_checkpointing: bool = False
def __post_init__(self):
from ..layers.attn_res import validate_attnres
validate_attnres(self.attnres, self.attnres_block_size)
@classmethod
def preset(cls, name: str) -> K3Config:
if name == "toy":
return cls()
if name in {"0.5b", "500m"}:
# H * head_dim == hidden. Routed 16: LatentMoE still runs every expert.
# ~482M with tied Qwen3 embeddings. 6×(3 KDA + 1 MLA).
return cls(
hidden_size=768,
num_hidden_layers=24,
vocab_size=QWEN3_VOCAB_SIZE,
tie_word_embeddings=True,
max_position_embeddings=2048,
num_heads=12,
head_dim=64,
chunk_size=64,
kv_lora_rank=192,
q_lora_rank=512,
qk_nope_head_dim=64,
v_head_dim=64,
moe_latent_size=384,
n_routed=16,
top_k=2,
n_shared=2,
moe_d_ff=512,
# The pure-PyTorch reference is far too slow at this size.
kda_backend="triton",
gradient_checkpointing=True,
)
raise ValueError(f"unknown preset: {name}")
@property
def H(self) -> int:
return self.num_heads
@property
def HV(self) -> int:
return self.num_heads
@property
def K(self) -> int:
return self.head_dim
@property
def V(self) -> int:
return self.head_dim
def layer_types(self) -> list[str]:
"""Hybrid pattern: 每 4 层 1 次 MLA (0-based 层 3,7,...), 末层强制 MLA."""
types = ["kda"] * self.num_hidden_layers
for i in range(self.num_hidden_layers):
if i % 4 == 3:
types[i] = "mla"
types[-1] = "mla"
return types
def layer_specs(self) -> list[tuple[str, str]]:
return [(kind, "moe") for kind in self.layer_types()]
+5
View File
@@ -0,0 +1,5 @@
"""KDA operator API and implementation backends."""
from .api import chunk_kda
__all__ = ["chunk_kda"]
+167
View File
@@ -0,0 +1,167 @@
"""Training-facing KDA operator with the same boundary as FLA's ``chunk_kda``."""
from __future__ import annotations
import warnings
from functools import lru_cache
import torch
import torch.nn.functional as F
from .reference.chunkwise import DECAY_BLOCK, _EXP_LIMIT, naive_chunk_kda
@lru_cache(maxsize=1)
def _fla_chunk_kda():
try:
from fla.ops.kda import chunk_kda
except ImportError:
return None
return chunk_kda
def _reference_chunk_size(T: int, requested: int) -> int:
size = min(T, requested)
while T % size:
size -= 1
return size
def chunk_kda(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
*,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
scale: float | None = None,
initial_state: torch.Tensor | None = None,
output_final_state: bool = False,
use_qk_l2norm_in_kernel: bool = False,
use_gate_in_kernel: bool = False,
use_beta_sigmoid_in_kernel: bool = False,
safe_gate: bool = False,
lower_bound: float | None = None,
chunk_size: int = 64,
backend: str = "reference",
):
"""Run an explicitly selected KDA implementation.
``reference`` and its legacy alias ``torch`` use this repository's
differentiable PyTorch implementation. ``triton`` uses the vendored
FLA NVIDIA Triton kernels in ``kda._fla`` (chunk_size 32 or 64,
CUDA). ``fla`` is reserved for explicit upstream parity runs.
"""
supported = {"reference", "triton", "fla", "torch", "auto"}
if backend not in supported:
raise ValueError(f"backend must be one of {sorted(supported)}")
if not use_qk_l2norm_in_kernel:
# Backend-independent: this is a property of the recurrence, not of any
# one implementation.
warnings.warn(
"use_qk_l2norm_in_kernel=False: KDA's chunkwise form assumes "
"||k||=1 so that I + tril(A_kk*beta) has a convergent Neumann "
"series. Unnormalised k makes the exact output grow like "
"||k||^chunk_size and can reach inf on any backend.",
RuntimeWarning,
stacklevel=2,
)
if backend == "auto":
warnings.warn(
"backend='auto' is deprecated and now selects the local reference backend; "
"use backend='fla' explicitly for upstream FLA",
DeprecationWarning,
stacklevel=2,
)
backend = "reference"
if backend == "torch":
backend = "reference"
if backend == "triton":
from .triton.chunk import chunk_kda as triton_chunk_kda
fla_chunk = 32 if chunk_size <= 32 else 64
return triton_chunk_kda(
q,
k,
v,
g,
beta,
scale=scale,
initial_state=initial_state,
output_final_state=output_final_state,
chunk_size=fla_chunk,
use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel,
use_gate_in_kernel=use_gate_in_kernel,
use_beta_sigmoid_in_kernel=use_beta_sigmoid_in_kernel,
A_log=A_log,
dt_bias=dt_bias,
safe_gate=safe_gate,
lower_bound=lower_bound,
)
if backend == "fla":
fused_op = _fla_chunk_kda()
if fused_op is None:
raise RuntimeError(
"backend='fla' requires a complete flash-linear-attention installation"
)
return fused_op(
q,
k,
v,
g,
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_gate_in_kernel=use_gate_in_kernel,
use_beta_sigmoid_in_kernel=use_beta_sigmoid_in_kernel,
safe_gate=safe_gate,
lower_bound=lower_bound,
chunk_size=32 if chunk_size <= 32 else 64,
)
if safe_gate and lower_bound is not None:
# _decayed_dot exponentiates at most DECAY_BLOCK steps of gate decay,
# and safe_gate bounds each step by |lower_bound|.
budget = DECAY_BLOCK * abs(lower_bound)
if budget > _EXP_LIMIT:
raise ValueError(
f"lower_bound={lower_bound} allows a gate span of {budget:.1f} "
f"per {DECAY_BLOCK}-row block, which overflows exp() "
f"(limit {_EXP_LIMIT:.1f}) and yields NaN. Use "
f"|lower_bound| < {_EXP_LIMIT / DECAY_BLOCK:.2f} or "
"backend='triton'."
)
if use_qk_l2norm_in_kernel:
q, k = F.normalize(q, dim=-1), F.normalize(k, dim=-1)
if use_beta_sigmoid_in_kernel:
beta = beta.sigmoid()
if use_gate_in_kernel:
if A_log is None:
raise ValueError("A_log is required when use_gate_in_kernel=True")
bias = 0 if dt_bias is None else dt_bias.view(g.shape[-2:])
gate_input = g + bias
rate = A_log.exp().view(1, 1, -1, 1)
if safe_gate:
if lower_bound is None:
raise ValueError("lower_bound is required when safe_gate=True")
g = lower_bound * torch.sigmoid(rate * gate_input)
else:
g = -rate * F.softplus(gate_input)
return naive_chunk_kda(
q,
k,
v,
g,
beta,
scale=scale,
initial_state=initial_state,
output_final_state=output_final_state,
chunk_size=_reference_chunk_size(q.shape[1], chunk_size),
)
+5
View File
@@ -0,0 +1,5 @@
"""Incremental recurrent KDA implementations and state containers."""
from .fused import KDAState, fused_recurrent_kda, fused_recurrent_kda_step
__all__ = ["KDAState", "fused_recurrent_kda", "fused_recurrent_kda_step"]
+73
View File
@@ -0,0 +1,73 @@
"""L6: FLA fused recurrent KDA decode with optional step cache."""
from __future__ import annotations
from dataclasses import dataclass
import torch
from kda._fla.ops.kda.fused_recurrent import fused_recurrent_kda as _fused_recurrent_kda
@dataclass
class KDAState:
"""Mutable recurrent state cache: ``S`` is ``[B, HV, K, V]``."""
S: torch.Tensor
pos: int = 0
def reset(self):
self.S.zero_()
self.pos = 0
def fused_recurrent_kda_step(
state: KDAState,
q_t: torch.Tensor,
k_t: torch.Tensor,
v_t: torch.Tensor,
g_t: torch.Tensor,
beta_t: torch.Tensor,
scale: float | None = None,
):
"""Single-token step. Inputs are ``[B, H|HV, ...]`` (no time dim)."""
o, ht = _fused_recurrent_kda(
q_t.unsqueeze(1),
k_t.unsqueeze(1),
v_t.unsqueeze(1),
g_t.unsqueeze(1),
beta_t.unsqueeze(1),
scale=scale,
initial_state=state.S,
output_final_state=True,
)
state.S = ht
state.pos += 1
return o.squeeze(1)
def fused_recurrent_kda(
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,
**kwargs,
):
return _fused_recurrent_kda(
q,
k,
v,
g,
beta,
scale=scale,
initial_state=initial_state,
output_final_state=output_final_state,
**kwargs,
)
__all__ = ["KDAState", "fused_recurrent_kda", "fused_recurrent_kda_step"]
+13
View File
@@ -0,0 +1,13 @@
"""Readable PyTorch implementations used as correctness references."""
from .chunkwise import naive_chunk_kda
from .gate import kda_gate_naive, kda_gate_reference
from .recurrent import naive_kda, naive_kda_fwd
__all__ = [
"kda_gate_naive",
"kda_gate_reference",
"naive_chunk_kda",
"naive_kda",
"naive_kda_fwd",
]
+155
View File
@@ -0,0 +1,155 @@
"""Pure-PyTorch chunked reference implementation of KDA."""
from __future__ import annotations
import math
import warnings
import torch
from einops import rearrange
#: ``exp`` overflows past this exponent in fp32 and bf16 (both top out at 3.4e38).
_EXP_LIMIT = math.log(torch.finfo(torch.float32).max)
#: Row-block size for :func:`_decayed_dot`.
#:
#: The g_ref GEMM exponentiates the gate span between the reference row and the
#: rows/columns it covers, so the block size caps that exponent at
#: ``DECAY_BLOCK * max|g|``. With the default ``lower_bound=-5`` gate that is
#: ``16 * 5 = 80 < ln(3.4e38) = 88.7``, i.e. fp32/bf16-safe for any chunk size.
#: Referencing a whole 64-row chunk instead would allow ``64 * 5 = 320`` and
#: overflow to NaN once the gate saturates.
DECAY_BLOCK = 16
def _decayed_dot(x: torch.Tensor, k: torch.Tensor, g: torch.Tensor) -> torch.Tensor:
"""Return ``A[..., i, j] = <x_i, exp(g_i-g_j) * k_j>`` (FLA g_ref GEMM).
Only the causal part (``j <= i``) is exact; callers mask the rest, which is
left at zero. Rows are processed in blocks of :data:`DECAY_BLOCK` against
the block's own first row, which is what bounds the exponent: for a row
block starting at ``r``, ``exp(g_i - g_ref)`` spans at most ``DECAY_BLOCK``
steps, and ``exp(g_ref - g_j)`` is ``<= 1`` for ``j < r`` and likewise spans
at most ``DECAY_BLOCK`` steps for ``j >= r``.
"""
C = g.shape[-2]
out = g.new_zeros(*g.shape[:-1], C)
for r in range(0, C, DECAY_BLOCK):
end = min(r + DECAY_BLOCK, C)
g_ref = g[..., r : r + 1, :]
rows = x[..., r:end, :] * (g[..., r:end, :] - g_ref).exp()
cols = k[..., :end, :] * (g_ref - g[..., :end, :]).exp()
out[..., r:end, :end] = rows @ cols.transpose(-1, -2)
return out
#: Whether :func:`naive_chunk_kda` checks the gate span against the ``exp``
#: budget. The check costs one device sync per call; set it to ``False`` if that
#: matters more than diagnosing a NaN.
CHECK_DECAY_SPAN = True
def _max_decay_span(g_cumsum: torch.Tensor) -> torch.Tensor:
"""Largest ``|g_ref - g_j|`` any row block will exponentiate."""
C = g_cumsum.shape[-2]
if C % DECAY_BLOCK == 0:
blocks = g_cumsum.unflatten(-2, (C // DECAY_BLOCK, DECAY_BLOCK))
return (blocks[..., :1, :] - blocks).abs().amax()
return torch.stack(
[
(g_cumsum[..., r : r + 1, :] - g_cumsum[..., r : r + DECAY_BLOCK, :])
.abs()
.amax()
for r in range(0, C, DECAY_BLOCK)
]
).amax()
def _warn_if_decay_span_overflows(g_cumsum: torch.Tensor) -> None:
"""Warn when a row block's gate span is about to overflow ``exp``.
``DECAY_BLOCK`` bounds this for the default ``safe_gate`` path, but an
unbounded gate (``-A.exp() * softplus(x)``) can still exceed it.
"""
span = _max_decay_span(g_cumsum).item()
if span > _EXP_LIMIT:
warnings.warn(
f"gate span within a {DECAY_BLOCK}-row block is {span:.1f} > "
f"{_EXP_LIMIT:.1f}; exp() will overflow to inf and the output will "
"be NaN. Reduce the gate magnitude (e.g. safe_gate with a smaller "
"|lower_bound|) or use backend='triton'.",
RuntimeWarning,
stacklevel=3,
)
def naive_chunk_kda(
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,
chunk_size: int = 64,
):
"""Chunk-parallel, inter-chunk recurrent KDA reference.
Shapes are ``q/k: [B,T,H,K]``, ``v: [B,T,HV,V]``,
``g: [B,T,HV,K]`` and ``beta: [B,T,HV]``.
"""
dtype = v.dtype
B, T, H, K = q.shape
HV, V = v.shape[2], v.shape[-1]
C = chunk_size
assert HV % H == 0, f"HV={HV} must be divisible by H={H}"
assert T % C == 0, f"T={T} must be divisible by chunk_size={C}"
scale = K**-0.5 if scale is None else scale
q, k = [
rearrange(x, "b (n c) h d -> b h n c d", c=C)
.repeat_interleave(HV // H, dim=1)
for x in (q, k)
]
v, g = [rearrange(x, "b (n c) h d -> b h n c d", c=C) for x in (v, g)]
beta = rearrange(beta, "b (n c) h -> b h n c", c=C)
q = q * scale
g = g.cumsum(dim=-2)
if CHECK_DECAY_SPAN:
_warn_if_decay_span_overflows(g)
# r_i + sum_{j<i} beta_j <k_i, exp(g_i-g_j)k_j> r_j
# = v_i - <exp(g_i)k_i, S_start>.
mask_upper = torch.triu(torch.ones(C, C, dtype=torch.bool, device=q.device))
mask_strict_upper = torch.triu(mask_upper, diagonal=1)
eye = torch.eye(C, dtype=q.dtype, device=q.device)
A_kk = _decayed_dot(k, k, g)
M = eye + (A_kk * beta[..., None, :]).masked_fill(mask_upper, 0)
W = torch.linalg.solve_triangular(M, g.exp() * k, upper=False)
U = torch.linalg.solve_triangular(M, v, upper=False)
# Output includes the current token, hence the diagonal is retained.
A_qk = (_decayed_dot(q, k, g) * beta[..., None, :]).masked_fill(mask_strict_upper, 0)
S = q.new_zeros(B, HV, K, V)
if initial_state is not None:
S = S + initial_state
o = v.new_empty(B, HV, T // C, C, V)
for n in range(T // C):
q_n, k_n, g_n = q[:, :, n], k[:, :, n], g[:, :, n]
r = U[:, :, n] - W[:, :, n] @ S
o[:, :, n] = (q_n * g_n.exp()) @ S + A_qk[:, :, n] @ r
decay = (g_n[:, :, -1:, :] - g_n).exp()
S = S * g_n[:, :, -1, :, None].exp()
S = S + (decay * k_n).transpose(-1, -2) @ (r * beta[:, :, n, :, None])
if not output_final_state:
S = None
return rearrange(o, "b h n c d -> b (n c) h d").to(dtype), S
# Backward-compatible name used by earlier notes/scripts.
naive_chunk_kda_fwd = naive_chunk_kda
+47
View File
@@ -0,0 +1,47 @@
"""PyTorch references for the two KDA gate activations."""
from __future__ import annotations
import torch
import torch.nn.functional as F
def kda_gate_reference(
g: torch.Tensor,
A_log: torch.Tensor,
dt_bias: torch.Tensor | None = None,
*,
safe_gate: bool = False,
lower_bound: float | None = None,
) -> torch.Tensor:
"""Compute the official KDA gate semantics in PyTorch.
``A_log`` is head-wise with shape ``[HV]`` and ``dt_bias`` is
per-dimension with shape ``[HV, K]`` (or flattened to ``[HV*K]``).
"""
HV, K = g.shape[-2:]
gate_input = g if dt_bias is None else g + dt_bias.view(HV, K)
rate = A_log.view(HV, 1).exp()
if safe_gate:
if lower_bound is None:
raise ValueError("lower_bound is required when safe_gate=True")
return lower_bound * torch.sigmoid(rate * gate_input)
return -rate * F.softplus(gate_input)
def kda_gate_naive(
g: torch.Tensor,
A_log: torch.Tensor,
dt_bias: torch.Tensor | None = None,
lower_bound: float | None = None,
) -> torch.Tensor:
"""Compatibility name matching FLA's reference gate convention."""
return kda_gate_reference(
g,
A_log,
dt_bias,
safe_gate=lower_bound is not None,
lower_bound=lower_bound,
)
__all__ = ["kda_gate_naive", "kda_gate_reference"]
+298
View File
@@ -0,0 +1,298 @@
"""L1: Naive recurrent KDA fwd+bwd (torch only).
公式 (per timestep t, log-space gate; q/k 入口 H 维, 内部 repeat_interleave 到 HV):
S_t = exp(g_t) * S_{t-1} + (beta_t * k_t) outer (v_t - k_t . (exp(g_t) * S_{t-1}))
o_t = (q_t * scale) . S_t
backward (BPTT, T -> 0):
设 dS_t 为进入 t 步累积的反传梯度 (含 o_t 反传).
1. o_t = q_t . S_t -> dS_t += q_t outer do_t (i.e. dS = dS + q_t·do_t)
dq_t = do_t . S_t^T -> einsum('bhv,bhkv->bhk')
2. S_t = S_decay + a_t outer r_t, a_t = b_t k_t, r_t = v_t - k_t . S_decay
其中 S_decay = exp(g_t) * S_{t-1}
dS_{t-1} = exp(g_t) * (dS_t - r_t outer da_t - a_t outer dr_t) via residual 反传
更具体:
dS_decay = dS_t - (a_t outer dr_t) - (da_t outer r_t)
dS_{t-1} += exp(g_t) * dS_decay
其中 dr_t = -dv_t + dS_t . a_t^T (因为 r_t = v - k·S_dec, dr 来自 -dv - k·dS_decay)
da_t = -r_t outer dS_t? 让我直接推导下面.
推导 (设 G1 = S_t, 走 a = r 反向链 通过 autograd):
o_t = q_t . G1
dq_t = do_t . G1^T -> [B,HV,K]
dG1 = q_t outer do_t -> [B,HV,K,V] = dS_t (上游)
G1 = Sdec + a outer r -> Sdec = G1[...] (跳过)
dSdec = dG1
da_t = r_t outer dG1 -> [B,HV,K] (因为 a outer r 是 K-V, d(a outer r) = r outer d[...,V])
但在 einsum 表示: dA_t.grad = einsum('bhkv,bhv->bhk', dS_t, r_t)
dr_t = a_t outer dG1 -> [B,HV,V] = einsum('bhkv,bhk->bhv', dS_t, a_t)
plus: S_t = Sdec + a outer r -> a outer r - outer product 形状是 [B,HV,K,V] = einsum('bhk,bhv->bhkv')
d(a outer r) 的雅可比: let G1_m = a_t ⊗ r_t (rank-1 matrix per (b,h))
dG1_m[i,j] = da_t[i] * r_t[j] + a_t[i] * dr_t[j]
在外积形式, 即 dG1_m = a_outer r 的张量积正交分解:
da_t = sum_j r_t[j] dG1_m[i,j] = einsum('bhkv,bhv->bhk', dG1_m, r_t)
dr_t = sum_i a_t[i] dG1_m[i,j] = einsum('bhkv,bhk->bhv', dG1_m, a_t)
因为 a_t = b_t k_t -> da_t = db_t k_t + b_t dk_t (b_t 是 ...)
db_t = einsum('bhk,bhk->bh', da_t, k_t)
dk_t_a = b_t * da_t (来自 a_t 路径, 还有来自 r_t 路径和 S_dec 路径)
因为 r_t = v_t - k_t . S_dec -> 注 rk_t grad via dg,S_dec 和 dv_t
dv_t = -dr_t (实际 dr 的负梯度) 即 dv_t = -dr_t
这里 r_t = v_t - k_t · S_dec, 写作矩阵乘 r = v - einsum('bhk,bhkv->bhv', k, S_dec)
dr = -dv - einsum('bhk,bhkv->bhv', dk_from_r, S_dec) + eigengrad via S_dec
更精确的反向: r_t = v_t - k_t . S_dec
dv_t += -dr_t -> dv_t = -dr_t
dk_t_r_path = -S_dec outer dr_t (即 -dS_dec 传递来自 k_t 的部分)
具体: d(k·S) = dk·S + k·dS -> dS_dec 这层, dk 的贡献: -S_dec outer dr_t
即 dk_t_r = einsum('bhv,bhkv->bhk', -dr_t, S_dec)
dS_dec_r = -k_t outer dr_t = -einsum('bhv,bhk->bhkv', dr_t, k_t)
合并: dS_dec 合总 = dG1 + (-k_t outer dr_t)
= dS_t - k_t outer dr_t
(相加过的 dv, dk_r, dS_dec_r 都上面项)
Sdec = exp(g_t) * S_{t-1}:
dS_{t-1} = exp(g_t) ⊙ dS_dec (因为 Sdec = exp_g * S_prev, 微分后 exp_g 直接相乘)
dg_t = exp(g_t) * S_prev * dS_dec (微分时对 g_t (log-space) 求偏导数)
即 dg_t = exp(g_t) * (S_{t-1} ⊙ dS_dec) -> 沿 K 维求和
in einsum: dg_t = sum over v of (exp(g_t) * S_{t-1}) ⊙ dS_dec ...\n
= einsum('bhk, bhk, bhkv -> bhk', exp_g, S_prev, dS_dec)
更简洁: Sdec = exp_g * S_prev (per-(b,h,k)/v), 故 dSdec/dg_t = S_prev * exp_g
所以 dg_t = sum_v S_prev_sub_k_dim * exp_g * dS_dec -> [B, HV, K]
einsum: dg_t = einsum('bhkv,bhkv->bhk', Sdec, dS_dec)
(因为 Sdec = S_prev * exp_g, sum_v Sdec[:, :, :, v] * dS_dec[:, :, :, v] = sum_v Sdec_eachK * dSdec_eachK)
einsum上是 einsum('bhkv,bhkv->bhk', Sdec, dSdec)
dS_{t-1} = exp_g ⊙ dSdec (per (b,h,k,v) entrywise multiply exp_g with dSdec)
GVA 反归约:
q,k 入口 [B, T, H, K] --repeat_interleave(G, dim=2)--> [B, T, HV, K]
内部计算后, dq/dk 在 HV 维上 -> dV 拿 shape [B,T,HV,K]
bwd 通过 sum 回 H: dq_H = dq_HV.view(B,T,H,G,K).sum(dim=3) -> [B,T,H,K]
(因为 repeat_interleave 是复制, 反传是 sum 路径相同意义)
记号对照:
a_t = b_t * k_t (a = beta * k) [B, HV, K]
r_t = v_t - k_t . S_dec (residual) [B, HV, V]
S_dec = exp(g_t) * S_{t-1} [B, HV, K, V]
S_t = S_dec + a_t outer r_t [B, HV, K, V]
o_t = q_t . S_t = (q_t_eff * scale) . S_t [B, HV, V]
"""
from __future__ import annotations
import math
import torch
def naive_kda_fwd(
q: torch.Tensor, # [B, T, H, K]
k: torch.Tensor, # [B, T, H, K]
v: torch.Tensor, # [B, T, HV, V]
g: torch.Tensor, # [B, T, HV, K]
beta: torch.Tensor, # [B, T, HV]
scale: float | None = None,
initial_state: torch.Tensor | None = None, # [B, HV, K, V]
output_final_state: bool = False,
*,
force_float32: bool = False,
):
"""纯 forward, 不带 autograd. 与上游 naive_recurrent_kda 数值等价.
force_float32=True 时强制 fp32 计算 (与上游对拍时用);
默认保持输入 dtype (gradcheck 用 fp64).
"""
dtype = v.dtype
B, T, H, K = q.shape
HV, V = v.shape[2], v.shape[-1]
G = HV // H
if scale is None:
scale = 1.0 / math.sqrt(K)
# 上游强制 fp32; 本实现默认保留输入 dtype 以便 gradcheck 适用 fp64
# force_float32=True 时与上游逐位对齐
work_dtype = torch.float if force_float32 else q.dtype
q = q.to(work_dtype)
k = k.to(work_dtype)
v = v.to(work_dtype)
g = g.to(work_dtype)
beta = beta.to(work_dtype)
# GVA: expand q/k from H to HV
qe = q.repeat_interleave(G, dim=2) * scale # [B, T, HV, K]
ke = k.repeat_interleave(G, dim=2) # [B, T, HV, K]
S = torch.zeros(B, HV, K, V, dtype=work_dtype, device=q.device)
if initial_state is not None:
S = S + initial_state.to(work_dtype)
o = torch.empty(B, T, HV, V, dtype=work_dtype, device=q.device)
for t in range(T):
q_t = qe[:, t] # [B, HV, K]
k_t = ke[:, t] # [B, HV, K]
v_t = v[:, t] # [B, HV, V]
g_t = g[:, t] # [B, HV, K]
b_t = beta[:, t] # [B, HV]
S_dec = S * g_t.exp().unsqueeze(-1) # [B, HV, K, V]
p_t = torch.einsum('b h k, b h k v -> b h v', k_t, S_dec) # [B, HV, V]
r_t = v_t - p_t # [B, HV, V]
a_t = b_t.unsqueeze(-1) * k_t # [B, HV, K]
S = S_dec + torch.einsum('b h k, b h v -> b h k v', a_t, r_t)
o[:, t] = torch.einsum('b h k, b h k v -> b h v', q_t, S)
if not output_final_state:
S = None
return o.to(dtype), S
class KDAFunction(torch.autograd.Function):
"""autograd Function (forward + backward).
forward 入参顺序 (q, k, v, g, beta, scale, initial_state, output_final_state)
backward 必须返回一致: (dq, dk, dv, dg, dbeta, None, dinit_state, None)
"""
@staticmethod
def forward(ctx, q, k, v, g, beta, scale, initial_state, output_final_state):
dtype = v.dtype
B, T, H, K = q.shape
HV, V = v.shape[2], v.shape[-1]
G = HV // H
if scale is None:
scale = 1.0 / math.sqrt(K)
work_dtype = q.dtype
qf = q.to(work_dtype).contiguous()
kf = k.to(work_dtype).contiguous()
vf = v.to(work_dtype).contiguous()
gf = g.to(work_dtype).contiguous()
bf = beta.to(work_dtype).contiguous()
# GVA: expand q/k from H to HV
qe = qf.repeat_interleave(G, dim=2) * scale # [B, T, HV, K]
ke = kf.repeat_interleave(G, dim=2) # [B, T, HV, K]
S = torch.zeros(B, HV, K, V, dtype=work_dtype, device=q.device)
if initial_state is not None:
S = S + initial_state.to(work_dtype)
o = torch.empty(B, T, HV, V, dtype=work_dtype, device=q.device)
q_ts, k_ts, b_ts, S_decs, r_ts, a_ts, exp_g_ts = [], [], [], [], [], [], []
for t in range(T):
q_t = qe[:, t]
k_t = ke[:, t]
v_t = vf[:, t]
g_t = gf[:, t]
b_t = bf[:, t]
exp_g_t = g_t.exp()
S_dec = S * exp_g_t.unsqueeze(-1)
p_t = torch.einsum('b h k, b h k v -> b h v', k_t, S_dec)
r_t = v_t - p_t
a_t = b_t.unsqueeze(-1) * k_t
S = S_dec + torch.einsum('b h k, b h v -> b h k v', a_t, r_t)
o[:, t] = torch.einsum('b h k, b h k v -> b h v', q_t, S)
q_ts.append(q_t)
k_ts.append(k_t)
b_ts.append(b_t)
S_decs.append(S_dec)
r_ts.append(r_t)
a_ts.append(a_t)
exp_g_ts.append(exp_g_t)
ctx.save_for_backward(
torch.stack(q_ts, dim=1),
torch.stack(k_ts, dim=1),
torch.stack(b_ts, dim=1),
torch.stack(S_decs, dim=1),
torch.stack(r_ts, dim=1),
torch.stack(a_ts, dim=1),
torch.stack(exp_g_ts, dim=1),
)
ctx.G = G
ctx.H = H
ctx.HV = HV
ctx.K = K
ctx.V = V
ctx.T = T
ctx.B = B
ctx.scale = scale
ctx.dtype = dtype
ctx.has_initial_state = initial_state is not None
ctx.output_final_state = output_final_state
final_S = S if output_final_state else None
return o.to(dtype), final_S
@staticmethod
def backward(ctx, do, dS):
q_ts, k_ts, b_ts, S_decs, r_ts, a_ts, exp_g_ts = ctx.saved_tensors
B, T, H, HV, K, V, G = ctx.B, ctx.T, ctx.H, ctx.HV, ctx.K, ctx.V, ctx.G
work_dtype = q_ts.dtype
device = q_ts.device
dq_e = torch.zeros(B, T, HV, K, dtype=work_dtype, device=device)
dk_e = torch.zeros(B, T, HV, K, dtype=work_dtype, device=device)
dv = torch.zeros(B, T, HV, V, dtype=work_dtype, device=device)
dg = torch.zeros(B, T, HV, K, dtype=work_dtype, device=device)
dbeta= torch.zeros(B, T, HV, dtype=work_dtype, device=device)
if dS is None:
dS_acc = torch.zeros(B, HV, K, V, dtype=work_dtype, device=device)
else:
dS_acc = dS.to(work_dtype).clone()
for t in range(T - 1, -1, -1):
q_t = q_ts[:, t]
k_t = k_ts[:, t]
b_t = b_ts[:, t]
S_dec = S_decs[:, t]
r_t = r_ts[:, t]
a_t = a_ts[:, t]
exp_g_t = exp_g_ts[:, t]
do_t = do[:, t].to(work_dtype)
S_t = S_dec + torch.einsum('b h k, b h v -> b h k v', a_t, r_t)
dS_acc = dS_acc + torch.einsum('b h k, b h v -> b h k v', q_t, do_t)
dq_e[:, t] = torch.einsum('b h v, b h k v -> b h k', do_t, S_t)
da_t = torch.einsum('b h v, b h k v -> b h k', r_t, dS_acc)
dr_t = torch.einsum('b h k, b h k v -> b h v', a_t, dS_acc)
dbeta[:, t] = torch.einsum('b h k, b h k -> b h', k_t, da_t)
dk_t_a = b_t.unsqueeze(-1) * da_t
dv[:, t] = dr_t
dS_dec_from_r = -torch.einsum('b h v, b h k -> b h k v', dr_t, k_t)
dk_t_r = -torch.einsum('b h v, b h k v -> b h k', dr_t, S_dec)
dS_dec_total = dS_acc + dS_dec_from_r
dk_e[:, t] = dk_t_a + dk_t_r
dg[:, t] = torch.einsum('b h k v, b h k v -> b h k', S_dec, dS_dec_total)
dS_acc = exp_g_t.unsqueeze(-1) * dS_dec_total
if HV > H:
dq_H = dq_e.view(B, T, H, G, K).sum(dim=3)
dk_H = dk_e.view(B, T, H, G, K).sum(dim=3)
else:
dq_H = dq_e
dk_H = dk_e
# q 在 forward 内被乘过 scale (qe = q * scale), chain rule: dq_orig = dq_e * scale
dq_H = dq_H * ctx.scale
return (dq_H.to(ctx.dtype), dk_H.to(ctx.dtype), dv.to(ctx.dtype),
dg.to(ctx.dtype), dbeta.to(ctx.dtype), None, None, None)
def naive_kda(
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,
):
"""对外入口: 调 KDAFunction.apply."""
return KDAFunction.apply(q, k, v, g, beta, scale, initial_state, output_final_state)
+7
View File
@@ -0,0 +1,7 @@
"""Local Triton KDA kernels vendored from FLA chunk_{fwd,intra,bwd,wy,gate}."""
from .chunk import ChunkKDAFunction, chunk_kda
from .chunk_fwd import chunk_kda_fwd
from .gate import kda_gate_fwd
__all__ = ["ChunkKDAFunction", "chunk_kda", "chunk_kda_fwd", "kda_gate_fwd"]
+5
View File
@@ -0,0 +1,5 @@
"""FLA ``chunk_kda`` surface used by ``ops.api`` backend='triton'."""
from kda._fla.ops.kda.chunk import ChunkKDAFunction, chunk_kda
__all__ = ["ChunkKDAFunction", "chunk_kda"]
+5
View File
@@ -0,0 +1,5 @@
"""Vendored FLA chunk KDA backward."""
from kda._fla.ops.kda.chunk_bwd import chunk_kda_bwd
__all__ = ["chunk_kda_bwd"]
+37
View File
@@ -0,0 +1,37 @@
"""Vendored FLA chunk KDA forward, returning ``(o, ht)`` like the public op."""
from __future__ import annotations
import torch
from kda._fla.ops.kda.chunk import chunk_kda
from kda._fla.ops.kda.chunk_fwd import chunk_kda_fwd as fla_chunk_kda_fwd
__all__ = ["chunk_kda_fwd", "fla_chunk_kda_fwd"]
def chunk_kda_fwd(
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,
chunk_size: int = 64,
**kwargs,
):
"""Chunked KDA forward with FLA kernels. Returns ``(o, ht)``."""
return chunk_kda(
q,
k,
v,
g,
beta,
scale=scale,
initial_state=initial_state,
output_final_state=output_final_state,
chunk_size=chunk_size,
**kwargs,
)
+36
View File
@@ -0,0 +1,36 @@
"""Vendored FLA KDA gate fusion (standard + safe gate + chunk cumsum)."""
from __future__ import annotations
import torch
from kda._fla.ops.kda.gate import (
kda_gate_bwd,
kda_gate_chunk_cumsum,
kda_gate_fwd as _kda_gate_fwd,
)
DEFAULT_LOWER_BOUND = -5.0
def kda_gate_fwd(
g: torch.Tensor,
A_log: torch.Tensor,
dt_bias: torch.Tensor | None = None,
lower_bound: float | None = DEFAULT_LOWER_BOUND,
):
return _kda_gate_fwd(
g,
A_log=A_log,
dt_bias=dt_bias,
lower_bound=lower_bound,
output_dtype=g.dtype,
)
__all__ = [
"DEFAULT_LOWER_BOUND",
"kda_gate_bwd",
"kda_gate_chunk_cumsum",
"kda_gate_fwd",
]
+5
View File
@@ -0,0 +1,5 @@
"""Vendored FLA WY recompute used by the chunk KDA backward."""
from kda._fla.ops.kda.wy_fast import recompute_w_u_fwd
__all__ = ["recompute_w_u_fwd"]
+5
View File
@@ -0,0 +1,5 @@
"""Training and checkpoint helpers."""
from .toy import load_ckpt, make_toy_data, save_ckpt, train_one_batch
__all__ = ["load_ckpt", "make_toy_data", "save_ckpt", "train_one_batch"]
+355
View File
@@ -0,0 +1,355 @@
"""Pretrain / SFT sample construction.
Pretrain: Wikipedia parquet → tokenize → pack (B, T). Languages mix 1:1 by
token via seq_len-sized blocks so each training chunk is monolingual.
SFT: instruction-parallel rows → prompt-masked labels. Template lives in
``prompts.instruction_prompt`` (same string as eval_mt).
"""
from __future__ import annotations
import json
import os
from dataclasses import dataclass
from pathlib import Path
from typing import Iterable, Protocol
import torch
from .prompts import instruction_prompt
WIKI_SHARD_TOTAL = {"zh": 6, "en": 41}
WIKI_BASE = (
"https://huggingface.co/datasets/wikimedia/wikipedia/resolve/main/20231101.{lang}"
)
IGNORE_INDEX = -100
class Tokenizer(Protocol):
vocab_size: int
def encode(self, text: str) -> list[int]: ...
def decode(self, ids: list[int]) -> str: ...
@dataclass
class SentencePieceTokenizer:
_sp: object
@property
def vocab_size(self) -> int:
return int(self._sp.vocab_size())
def encode(self, text: str) -> list[int]:
return list(self._sp.encode(text, out_type=int))
def decode(self, ids: list[int]) -> str:
return str(self._sp.decode(ids))
@dataclass
class HuggingFaceTokenizer:
_tok: object
@property
def vocab_size(self) -> int:
return int(len(self._tok))
def encode(self, text: str) -> list[int]:
return list(self._tok.encode(text, add_special_tokens=False))
def decode(self, ids: list[int]) -> str:
return str(self._tok.decode(ids, skip_special_tokens=True))
def load_tokenizer(source: str) -> Tokenizer:
"""`.model` 走 SentencePiece, 其它当作 HuggingFace 名或本地目录."""
if source.endswith(".model"):
from sentencepiece import SentencePieceProcessor
return SentencePieceTokenizer(SentencePieceProcessor(model_file=source))
from transformers import AutoTokenizer
tok = AutoTokenizer.from_pretrained(source, trust_remote_code=True)
return HuggingFaceTokenizer(tok)
def pretrain_dir() -> Path:
for candidate in (
os.environ.get("KDA_PRETRAIN_DIR"),
"/data/pretrain",
"data/pretrain",
):
if candidate and Path(candidate).is_dir():
return Path(candidate)
return Path("data/pretrain")
def _wiki_files(lang: str, n_shards: int) -> list[str]:
if lang not in WIKI_SHARD_TOTAL:
raise ValueError(f"unsupported wiki lang {lang!r}; expected zh or en")
total = WIKI_SHARD_TOTAL[lang]
n = min(max(n_shards, 1), total)
base = WIKI_BASE.format(lang=lang)
return [f"{base}/train-{i:05d}-of-{total:05d}.parquet" for i in range(n)]
def _cache_path(cache_dir: Path, lang: str, n_shards: int, limit: int) -> Path:
return cache_dir / f"wiki-{lang}-n{n_shards}-limit{limit}.jsonl"
def fetch_wiki_texts(
limit: int,
lang: str = "zh",
n_shards: int = 2,
cache_dir: str | Path | None = None,
) -> list[str]:
"""Load up to ``limit`` article bodies, caching jsonl under pretrain_dir."""
cache = Path(cache_dir) if cache_dir is not None else pretrain_dir()
cache.mkdir(parents=True, exist_ok=True)
path = _cache_path(cache, lang, n_shards, limit)
if path.exists():
texts: list[str] = []
with path.open(encoding="utf-8") as fh:
for line in fh:
line = line.strip()
if not line:
continue
texts.append(json.loads(line)["text"])
if len(texts) >= limit:
break
if texts:
return texts
from datasets import load_dataset
files = _wiki_files(lang, n_shards)
ds = load_dataset("parquet", data_files=files, split="train", streaming=True)
texts = []
for i, row in enumerate(ds):
if i >= limit:
break
texts.append(row["text"])
tmp = path.with_suffix(path.suffix + ".tmp")
with tmp.open("w", encoding="utf-8") as fh:
for text in texts:
fh.write(json.dumps({"text": text}, ensure_ascii=False) + "\n")
tmp.replace(path)
return texts
def tokenize_corpus(texts: list[str], tok: Tokenizer) -> list[int]:
ids: list[int] = []
for text in texts:
ids.extend(tok.encode(text))
return ids
def interleave_balanced(ids_a: list[int], ids_b: list[int], block: int) -> list[int]:
"""1:1 by token: seq_len-sized monolingual blocks, drop the longer tail."""
if block < 1:
raise ValueError(f"block must be >= 1, got {block}")
n = min(len(ids_a), len(ids_b))
n = (n // block) * block
out: list[int] = []
a, b = ids_a, ids_b
for i in range(0, n, block):
out.extend(a[i : i + block])
out.extend(b[i : i + block])
return out
def chunk_ids(ids: list[int], batch: int, seq_len: int) -> torch.Tensor:
"""切成 (num_chunks, B, T); 末尾不足部分丢弃."""
n = (len(ids) // (batch * seq_len)) * (batch * seq_len)
t = torch.tensor(ids[:n], dtype=torch.long)
if n == 0:
return t.view(0, batch, seq_len)
return t.view(batch, -1, seq_len).transpose(0, 1)
def split_heldout(
chunks: torch.Tensor,
frac: float = 0.01,
min_heldout: int = 1,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Last ``frac`` of packed chunks for CE only. Empty held-out if too few."""
n = int(chunks.size(0))
if n <= 1 or frac <= 0:
return chunks, chunks[:0]
h = max(min_heldout, int(n * frac))
h = min(h, n - 1)
return chunks[:-h], chunks[-h:]
def iter_chunks(chunks: torch.Tensor):
"""逐块产出 (input_ids, labels), labels 右移 (模型内 CE shift)."""
for chunk in chunks:
yield chunk, chunk.clone()
def iter_indexed(chunks: torch.Tensor, start: int = 0):
"""Infinite cycle with a global index (for --resume)."""
n = int(chunks.size(0))
if n == 0:
raise ValueError("no training chunks")
i = start
while True:
x = chunks[i % n]
yield i, x, x.clone()
i += 1
def load_pretrain_chunks(
tok: Tokenizer,
*,
langs: Iterable[str],
limit: int,
batch: int,
seq_len: int,
heldout_frac: float = 0.01,
n_shards: int = 2,
cache_dir: str | Path | None = None,
) -> tuple[torch.Tensor, torch.Tensor, int]:
"""Fetch / cache / tokenize / pack. Returns train chunks, held-out, token count."""
lang_list = [lang.strip() for lang in langs if lang.strip()]
if not lang_list:
raise ValueError("langs must contain at least one of zh, en")
streams: list[list[int]] = []
for lang in lang_list:
print(f"loading {limit} wiki articles ({lang}) ...")
texts = fetch_wiki_texts(limit, lang=lang, n_shards=n_shards, cache_dir=cache_dir)
streams.append(tokenize_corpus(texts, tok))
print(f" {lang}: {len(streams[-1]):,} tokens from {len(texts)} articles")
if len(streams) == 1:
ids = streams[0]
else:
ids = streams[0]
for extra in streams[1:]:
ids = interleave_balanced(ids, extra, seq_len)
chunks = chunk_ids(ids, batch, seq_len)
train, held = split_heldout(chunks, heldout_frac)
return train, held, len(ids)
def pad_id(tok: Tokenizer) -> int:
inner = getattr(tok, "_tok", None)
if inner is not None:
pid = getattr(inner, "pad_token_id", None)
if pid is not None:
return int(pid)
eid = getattr(inner, "eos_token_id", None)
if eid is not None:
return int(eid)
return 0
def eos_id(tok: Tokenizer) -> int | None:
inner = getattr(tok, "_tok", None)
if inner is not None:
eid = getattr(inner, "eos_token_id", None)
if eid is not None:
return int(eid)
convert = getattr(inner, "convert_tokens_to_ids", None)
if convert is not None:
tid = convert("<|im_end|>")
if isinstance(tid, int) and tid >= 0:
return tid
return None
def encode_sft_row(
tok: Tokenizer,
src: str,
tgt: str,
target_lang: str,
max_len: int,
eos: int | None = None,
) -> tuple[list[int], list[int]]:
prompt_ids = tok.encode(instruction_prompt(src, target_lang))
tgt_ids = tok.encode(tgt)
if eos is not None:
tgt_ids = tgt_ids + [eos]
ids = prompt_ids + tgt_ids
labels = [IGNORE_INDEX] * len(prompt_ids) + list(tgt_ids)
if len(ids) > max_len:
overflow = len(ids) - max_len
cut = min(overflow, max(len(prompt_ids) - 1, 0))
ids = ids[cut:]
labels = labels[cut:]
if len(ids) > max_len:
ids = ids[:max_len]
labels = labels[:max_len]
return ids, labels
def load_sft_rows(path: str | Path) -> list[dict]:
"""jsonl ``{src,tgt,target_lang}`` or TSV ``src\\ttgt\\ttarget_lang``."""
p = Path(path)
rows: list[dict] = []
text = p.read_text(encoding="utf-8")
if p.suffix == ".jsonl" or p.suffix == ".json":
for line in text.splitlines():
line = line.strip()
if not line:
continue
obj = json.loads(line)
rows.append(
{
"src": obj["src"],
"tgt": obj["tgt"],
"target_lang": obj.get("target_lang", "en"),
}
)
return rows
for line in text.splitlines():
line = line.strip()
if not line or line.startswith("#"):
continue
parts = line.split("\t")
if len(parts) < 2:
raise ValueError(f"SFT TSV needs src, tgt [, target_lang]: {line[:80]!r}")
lang = parts[2] if len(parts) > 2 else "en"
rows.append({"src": parts[0], "tgt": parts[1], "target_lang": lang})
return rows
def collate_sft(
rows: list[dict],
tok: Tokenizer,
max_len: int,
) -> tuple[torch.Tensor, torch.Tensor]:
pad = pad_id(tok)
eos = eos_id(tok)
encoded = [
encode_sft_row(tok, r["src"], r["tgt"], r["target_lang"], max_len, eos)
for r in rows
]
width = min(max(len(ids) for ids, _ in encoded), max_len)
width = max(width, 2)
bsz = len(encoded)
input_ids = torch.full((bsz, width), pad, dtype=torch.long)
labels = torch.full((bsz, width), IGNORE_INDEX, dtype=torch.long)
for i, (ids, lab) in enumerate(encoded):
n = min(len(ids), width)
input_ids[i, :n] = torch.tensor(ids[:n], dtype=torch.long)
labels[i, :n] = torch.tensor(lab[:n], dtype=torch.long)
return input_ids, labels
def iter_sft_batches(
rows: list[dict],
tok: Tokenizer,
batch: int,
max_len: int,
start: int = 0,
):
n = len(rows)
if n == 0:
raise ValueError("no SFT rows")
i = start
while True:
sl = [rows[j % n] for j in range(i, i + batch)]
yield i, *collate_sft(sl, tok, max_len)
i += batch
+139
View File
@@ -0,0 +1,139 @@
"""Greedy translation eval on line-aligned src/ref files.
python -m kda.training.eval_mt \\
--ckpt ckpts/k3_wiki.pt --src /data/eval/zh2en.src.txt \\
--ref /data/eval/zh2en.ref.txt --target-lang en
"""
from __future__ import annotations
import argparse
import json
from pathlib import Path
import torch
from kda.training.data import eos_id, load_tokenizer
from kda.training.prompts import instruction_prompt
from kda.training.success import _chrf, _detect_lang, translation_success
from kda.training.toy import load_ckpt
def _read_lines(path: str) -> list[str]:
return [ln.strip() for ln in Path(path).read_text(encoding="utf-8").splitlines() if ln.strip()]
def _instruction(src: str, target_lang: str) -> str:
return instruction_prompt(src, target_lang)
@torch.inference_mode()
def decode_one(model, tok, prompt: str, device: str, max_new: int) -> str:
ids = tok.encode(prompt)
if not ids:
return ""
inp = torch.tensor([ids], dtype=torch.long, device=device)
out = model.generate(inp, max_new, eos_token_id=eos_id(tok))
gen = out[0, inp.size(1) :].tolist()
return tok.decode(gen).strip()
def evaluate_pairs(
model,
tok,
srcs: list[str],
refs: list[str],
*,
target_lang: str,
device: str,
max_new: int,
limit: int | None,
) -> dict:
n = len(srcs)
if limit is not None:
n = min(n, limit)
hyps: list[str] = []
wins = 0
copies = 0
lang_ok = 0
chrf_sum = 0.0
for i in range(n):
src, ref = srcs[i], refs[i]
hyp = decode_one(model, tok, _instruction(src, target_lang), device, max_new)
hyps.append(hyp)
ok = translation_success(src, hyp, ref, target_lang=target_lang)
wins += int(ok)
copies += int(_chrf(hyp, src) >= 80.0 or hyp == src)
want = "zh" if target_lang.startswith("zh") else "en"
lang_ok += int(_detect_lang(hyp) == want)
chrf_sum += _chrf(hyp, ref)
corpus = {}
try:
from sacrebleu.metrics import BLEU, CHRF
corpus["chrf"] = float(CHRF(word_order=2).corpus_score(hyps, [refs[:n]]).score)
corpus["bleu"] = float(BLEU().corpus_score(hyps, [refs[:n]]).score)
except Exception:
corpus["chrf"] = chrf_sum / max(n, 1)
corpus["bleu"] = None
return {
"n": n,
"success_rate": wins / max(n, 1),
"copy_rate": copies / max(n, 1),
"lang_ok": lang_ok / max(n, 1),
"chrf": corpus["chrf"],
"bleu": corpus["bleu"],
"hyps": hyps,
}
def main() -> None:
p = argparse.ArgumentParser(description=__doc__)
p.add_argument("--ckpt", required=True)
p.add_argument("--tokenizer", default=None, help="override ckpt tokenizer field")
p.add_argument("--src", default=None, help="one source sentence per line")
p.add_argument("--ref", default=None, help="one reference sentence per line")
p.add_argument("--target-lang", default="en", choices=["en", "zh"])
p.add_argument("--max-new", type=int, default=64)
p.add_argument("--limit", type=int, default=None)
p.add_argument("--prefix", default=None, help="single-prompt smoke decode")
p.add_argument("--device", default="auto")
args = p.parse_args()
device = args.device
if device == "auto":
device = "cuda" if torch.cuda.is_available() else "cpu"
model, _config = load_ckpt(args.ckpt)
model.to(device).eval()
payload = torch.load(args.ckpt, map_location="cpu", weights_only=False)
tok_src = args.tokenizer or payload.get("tokenizer")
if not tok_src:
raise SystemExit("need --tokenizer or a 'tokenizer' field in the checkpoint")
tok = load_tokenizer(tok_src)
if args.prefix:
print(decode_one(model, tok, args.prefix, device, args.max_new))
if args.src and args.ref:
srcs, refs = _read_lines(args.src), _read_lines(args.ref)
if len(srcs) != len(refs):
raise SystemExit(f"src/ref length mismatch: {len(srcs)} vs {len(refs)}")
out = evaluate_pairs(
model,
tok,
srcs,
refs,
target_lang=args.target_lang,
device=device,
max_new=args.max_new,
limit=args.limit,
)
printable = {k: v for k, v in out.items() if k != "hyps"}
print(json.dumps(printable, ensure_ascii=False, indent=2))
elif not args.prefix:
raise SystemExit("pass --prefix and/or --src + --ref")
if __name__ == "__main__":
main()
+7
View File
@@ -0,0 +1,7 @@
"""Instruction strings shared by SFT and eval. Do not drift."""
def instruction_prompt(src: str, target_lang: str) -> str:
if target_lang.startswith("zh"):
return f"Translate to Chinese:\n{src}"
return f"Translate to English:\n{src}"
+47
View File
@@ -0,0 +1,47 @@
"""LR scale and token-horizon helpers for train_k3 / train_sft."""
from __future__ import annotations
import math
def lr_scale(
opt_step: int,
warmup: int,
total_opt: int,
min_ratio: float = 0.1,
) -> float:
"""Linear warmup (optimizer steps) then cosine down to ``min_ratio``.
``opt_step`` is 0-indexed at the optimizer update that is about to run.
"""
if warmup > 0 and opt_step < warmup:
return (opt_step + 1) / warmup
denom = max(total_opt - warmup - 1, 1)
progress = min(max(opt_step - warmup, 0) / denom, 1.0)
cosine = 0.5 * (1.0 + math.cos(math.pi * progress))
return min_ratio + (1.0 - min_ratio) * cosine
def tokens_per_micro(batch: int, seq_len: int) -> int:
return batch * seq_len
def total_opt_steps(
*,
max_tokens: int | None,
max_micro: int | None,
batch: int,
seq_len: int,
grad_acc: int,
) -> int:
"""Optimizer-step horizon used by cosine. At least 1."""
acc = max(grad_acc, 1)
candidates: list[int] = []
if max_tokens is not None and max_tokens > 0:
tpm = max(tokens_per_micro(batch, seq_len), 1)
candidates.append(math.ceil(max_tokens / (tpm * acc)))
if max_micro is not None and max_micro > 0:
candidates.append(math.ceil(max_micro / acc))
if not candidates:
return 1
return max(min(candidates), 1)
+86
View File
@@ -0,0 +1,86 @@
"""Frozen translation success() — SFT eval and RL reward must call this."""
from __future__ import annotations
import re
CHRF_MIN = 40.0
COPY_CHRF_MAX = 80.0
_CJK = re.compile(r"[\u4e00-\u9fff]")
def _detect_lang(text: str) -> str | None:
sample = text.strip()
if not sample:
return None
try:
from langdetect import detect
tag = detect(sample)
except Exception:
if _CJK.search(sample):
return "zh"
if any(c.isascii() and c.isalpha() for c in sample):
return "en"
return None
if tag.startswith("zh"):
return "zh"
return tag[:2]
def _chrf(hyp: str, ref: str) -> float:
"""chrF++ in 0–100. Falls back to char unigram F if sacrebleu is missing."""
if not hyp or not ref:
return 0.0
try:
from sacrebleu.metrics import CHRF
return float(CHRF(word_order=2).sentence_score(hyp, [ref]).score)
except Exception:
hyp_c, ref_c = list(hyp), list(ref)
if not hyp_c:
return 0.0
ref_set = set(ref_c)
overlap = sum(1 for c in hyp_c if c in ref_set)
prec = overlap / len(hyp_c)
rec = overlap / max(len(ref_c), 1)
if prec + rec == 0:
return 0.0
return 100.0 * 2 * prec * rec / (prec + rec)
def translation_success(
src: str,
hyp: str,
ref: str | None = None,
*,
target_lang: str,
chrf_min: float = CHRF_MIN,
copy_chrf_max: float = COPY_CHRF_MAX,
) -> bool:
"""Binary task success for zh↔en instruction translation.
1. non-empty hyp, no instruction leak prefix
2. langid(hyp) matches target_lang (zh / en)
3. hyp is not a copy of src
4. if ref is given, chrF(hyp, ref) >= chrf_min
"""
hyp = hyp.strip()
src = src.strip()
if not hyp:
return False
leak = ("翻译如下", "translate to", "translation:", "译文:")
head = hyp[:40].lower()
if any(p in head or p in hyp[:20] for p in leak):
return False
want = "zh" if target_lang.startswith("zh") else "en"
got = _detect_lang(hyp)
if got != want:
return False
if src and _chrf(hyp, src) >= copy_chrf_max:
return False
if hyp == src:
return False
if ref is not None and _chrf(hyp, ref.strip()) < chrf_min:
return False
return True
+110
View File
@@ -0,0 +1,110 @@
"""L7: toy training loop — overfit 起步.
target:
端到端验证模型 + 数据流 + optimizer + ckpt + generate.
toy data:
建一份 256-token vocab 的小数据集: e.g. 1000 个长度 32 随机 token 序列
起步只取 batch=4, 看能否在 ~320 steps 内把 loss 压到 < 0.1 (overfit 单 batch).
step:
optimizer = AdamW(lr=1e-3, wd=0.01)
loss.backward(); optimizer.step(); optimizer.zero_grad()
every N steps: 打印 loss
end: 保存 ckpt to ckpts/kda_toy.pt
ckpt:
save:
torch.save({"model_state": model.state_dict(), "config": asdict(config)}, path)
load:
torch.load -> model.load_state_dict
"""
from __future__ import annotations
import os
from dataclasses import asdict, fields
import torch
from ..models.causal_lm import CausalLM
from ..models.config import KDAConfig
from ..models.k3_config import K3Config
def make_toy_data(batch: int = 4, seq_len: int = 32, vocab: int = 256, seed: int = 42):
"""单 batch overfit 数据: 同一组序列循环."""
torch.manual_seed(seed)
seq = torch.randint(0, vocab, (batch, seq_len), dtype=torch.long)
return seq # 用作 input_ids 和 labels (shift one inside forward)
def train_one_batch(model, optimizer, input_ids, labels):
optimizer.zero_grad(set_to_none=True)
loss = model(input_ids, labels=labels)
loss.backward()
optimizer.step()
return loss.detach()
def save_ckpt(model, config, path: str):
os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
torch.save({"model_state": model.state_dict(), "config": asdict(config)}, path)
#: The feed-forward submodule was named after its contents (``mlp`` in the
#: dense config, ``moe`` in K3) before both were unified under ``ffn``.
#: Checkpoints saved before that rename still carry the old prefixes.
_LEGACY_PREFIXES = {
".mlp.": ".ffn.",
".mlp_norm.": ".ffn_norm.",
".moe.": ".ffn.",
".moe_norm.": ".ffn_norm.",
}
def _rename_legacy_keys(state: dict) -> dict:
def fix(key: str) -> str:
for old, new in _LEGACY_PREFIXES.items():
if old in key:
return key.replace(old, new)
return key
return {fix(k): v for k, v in state.items()}
def _config_from(payload_config: dict) -> K3Config | KDAConfig:
"""Pick the config class the checkpoint was written with.
``moe_latent_size`` is a K3-only field, so its presence identifies the
hybrid K3 architecture; anything else is the dense KDA config.
"""
cls = K3Config if "moe_latent_size" in payload_config else KDAConfig
known = {item.name for item in fields(cls)}
return cls(**{k: v for k, v in payload_config.items() if k in known})
def load_ckpt(path: str, model: CausalLM | None = None) -> tuple[CausalLM, K3Config | KDAConfig]:
payload = torch.load(path, map_location="cpu", weights_only=False)
config = _config_from(payload["config"])
if model is None:
model = CausalLM(config)
model.load_state_dict(_rename_legacy_keys(payload["model_state"]))
return model, config
def main():
"""主入口: overfit 起步. 320 steps 期望 loss < 0.1."""
device = "cuda" if torch.cuda.is_available() else "cpu"
config = KDAConfig()
model = CausalLM(config).to(device)
tokens = make_toy_data(seq_len=32, vocab=config.vocab_size).to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=0.01)
for step in range(320):
loss = train_one_batch(model, optimizer, tokens, tokens)
if step % 64 == 0 or step == 319:
print(f"step {step:3d} loss {loss.item():.4f}")
save_ckpt(model, config, "ckpts/kda_toy.pt")
if __name__ == "__main__":
main()
+51
View File
@@ -0,0 +1,51 @@
"""Train a SentencePiece tokenizer on a Chinese Wikipedia subset.
用法:
uv run python kda/training/train_tokenizer.py \
--out data/spm_4k --vocab-size 4096 --limit 20000
产出:
data/spm_4k.model / data/spm_4k.vocab (BPE/unigram, 中文小语料)
"""
from __future__ import annotations
import argparse
import sentencepiece as spm
from .data import fetch_wiki_texts
def main() -> None:
p = argparse.ArgumentParser(description=__doc__)
p.add_argument("--out", default="data/spm_4k", help="输出前缀 (model/vocab 文件)")
p.add_argument("--vocab-size", type=int, default=8192)
p.add_argument("--limit", type=int, default=20000, help="用于训练的 wiki 文章数")
p.add_argument("--model-type", default="unigram", choices=["unigram", "bpe"])
p.add_argument("--character-coverage", type=float, default=0.9995)
args = p.parse_args()
texts = fetch_wiki_texts(args.limit)
corpus = "".join(texts)
tmp = args.out + ".corpus.txt"
with open(tmp, "w", encoding="utf-8") as f:
f.write(corpus)
print(f"corpus: {len(corpus):,} chars from {len(texts)} articles")
spm.SentencePieceTrainer.train(
input=tmp,
model_prefix=args.out,
vocab_size=args.vocab_size,
model_type=args.model_type,
character_coverage=args.character_coverage,
unk_id=0,
pad_id=1,
bos_id=-1,
eos_id=-1,
num_threads=4,
)
print(f"tokenizer saved: {args.out}.model / {args.out}.vocab")
if __name__ == "__main__":
main()