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:
@@ -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"
|
||||
@@ -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.
|
||||
@@ -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
|
||||
@@ -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"
|
||||
@@ -0,0 +1 @@
|
||||
# Vendored FLA modules used by KDA (l2norm).
|
||||
@@ -0,0 +1,3 @@
|
||||
from kda._fla.ops.backends import BackendRegistry, BaseBackend, dispatch
|
||||
|
||||
__all__ = ["BackendRegistry", "BaseBackend", "dispatch"]
|
||||
@@ -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)
|
||||
@@ -0,0 +1 @@
|
||||
# Vendored FLA ops subset.
|
||||
@@ -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"]
|
||||
@@ -0,0 +1 @@
|
||||
# Vendored FLA common kernels used by KDA.
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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"]
|
||||
@@ -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
|
||||
@@ -0,0 +1 @@
|
||||
# Vendored GLA chunk output kernel used by KDA.
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,7 @@
|
||||
from .chunk import chunk_kda
|
||||
from .fused_recurrent import fused_recurrent_kda
|
||||
|
||||
__all__ = [
|
||||
"chunk_kda",
|
||||
"fused_recurrent_kda",
|
||||
]
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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."
|
||||
)
|
||||
@@ -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
|
||||
@@ -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)",
|
||||
)
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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))
|
||||
@@ -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))
|
||||
@@ -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)
|
||||
@@ -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]
|
||||
@@ -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
|
||||
@@ -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))
|
||||
@@ -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"]
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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()]
|
||||
@@ -0,0 +1,5 @@
|
||||
"""KDA operator API and implementation backends."""
|
||||
|
||||
from .api import chunk_kda
|
||||
|
||||
__all__ = ["chunk_kda"]
|
||||
+167
@@ -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),
|
||||
)
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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
|
||||
@@ -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"]
|
||||
@@ -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)
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -0,0 +1,5 @@
|
||||
"""Vendored FLA chunk KDA backward."""
|
||||
|
||||
from kda._fla.ops.kda.chunk_bwd import chunk_kda_bwd
|
||||
|
||||
__all__ = ["chunk_kda_bwd"]
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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}"
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user