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,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
|
||||
Reference in New Issue
Block a user