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

Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias
backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
This commit is contained in:
dela
2026-08-25 14:43:17 +08:00
commit 584f7e9e73
140 changed files with 21592 additions and 0 deletions
+1
View File
@@ -0,0 +1 @@
# Vendored FLA ops subset.
+34
View File
@@ -0,0 +1,34 @@
"""Identity dispatch: keep the NVIDIA Triton implementation in this tree."""
from __future__ import annotations
from collections.abc import Callable
from typing import TypeVar
F = TypeVar("F", bound=Callable)
def dispatch(operation: str):
def decorator(func: F) -> F:
return func
return decorator
class BaseBackend:
backend_type = "triton"
def is_available(self) -> bool:
return True
def is_enabled(self) -> bool:
return True
class BackendRegistry:
@classmethod
def ensure_initialized(cls, operation: str) -> None:
return None
__all__ = ["BackendRegistry", "BaseBackend", "dispatch"]
+1
View File
@@ -0,0 +1 @@
# Vendored FLA common kernels used by KDA.
+806
View File
@@ -0,0 +1,806 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import torch
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.ops.utils import prepare_chunk_indices, prepare_chunk_offsets
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.ops.utils.op import exp2
from kda._fla.utils import (
IS_INTEL,
IS_NVIDIA_BLACKWELL,
IS_NVIDIA_HOPPER,
autotune_cache_kwargs,
check_shared_mem,
)
NUM_WARPS = [2, 4] if IS_NVIDIA_HOPPER else [2, 4, 8, 16]
# TODO: Triton mainline fixes a Blackwell tl.dot recurrence race.
# Keep this kernel on num_warps=2 for Blackwell until Triton 3.8 is released
# and we re-validate the wider config space.
# Intel needs more warps than NVIDIA here: 8 warps is ~1.5x faster than the best
# config reachable under the [2, 4] cap.
if IS_NVIDIA_BLACKWELL:
GATED_DELTA_RULE_FWD_H_NUM_WARPS = [2]
elif IS_INTEL:
GATED_DELTA_RULE_FWD_H_NUM_WARPS = [2, 4, 8, 16]
else:
GATED_DELTA_RULE_FWD_H_NUM_WARPS = [2, 4]
@triton.heuristics({
'USE_G': lambda args: args['g'] is not None,
'USE_GK': lambda args: args['gk'] is not None,
'USE_INITIAL_STATE': lambda args: args['h0'] is not None,
'STORE_FINAL_STATE': lambda args: args['ht'] is not None,
'SAVE_NEW_VALUE': lambda args: args['v_new'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({'BV': BV}, num_warps=num_warps, num_stages=num_stages)
for num_warps in GATED_DELTA_RULE_FWD_H_NUM_WARPS
for num_stages in ([2, 3, 4] if check_shared_mem('ampere') else [2, 1])
for BV in ([32, 64] if check_shared_mem('ada') else [32])
],
key=['H', 'HV', 'K', 'V', 'BT', 'STATE_V_FIRST'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_gated_delta_rule_fwd_kernel_h_blockdim64(
k,
v,
w,
v_new,
g,
gk,
h,
h0,
ht,
cu_seqlens,
chunk_offsets,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BV: tl.constexpr,
USE_G: tl.constexpr,
USE_GK: tl.constexpr,
USE_INITIAL_STATE: tl.constexpr,
STORE_FINAL_STATE: tl.constexpr,
SAVE_NEW_VALUE: tl.constexpr,
STATE_V_FIRST: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
pid = tl.program_id(0)
NV = tl.cdiv(V, BV)
i_v, i_nh = pid % NV, (pid // NV).to(tl.int64)
i_n, i_h = i_nh // HV, i_nh % HV
if IS_VARLEN:
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
NT = tl.cdiv(T, BT)
boh = tl.load(chunk_offsets + i_n).to(tl.int64)
else:
bos, eos = i_n * T, i_n * T + T
NT = tl.cdiv(T, BT)
boh = i_n * NT
if STATE_V_FIRST:
b_h1 = tl.zeros([BV, 64], dtype=tl.float32)
if K > 64:
b_h2 = tl.zeros([BV, 64], dtype=tl.float32)
if K > 128:
b_h3 = tl.zeros([BV, 64], dtype=tl.float32)
if K > 192:
b_h4 = tl.zeros([BV, 64], dtype=tl.float32)
else:
b_h1 = tl.zeros([64, BV], dtype=tl.float32)
if K > 64:
b_h2 = tl.zeros([64, BV], dtype=tl.float32)
if K > 128:
b_h3 = tl.zeros([64, BV], dtype=tl.float32)
if K > 192:
b_h4 = tl.zeros([64, BV], dtype=tl.float32)
# calculate offset
h += (boh * HV + i_h).to(tl.int64) * K*V
v += (bos * HV + i_h).to(tl.int64) * V
k += (bos * H + i_h // (HV // H)).to(tl.int64) * K
w += (bos * HV + i_h).to(tl.int64) * K
if SAVE_NEW_VALUE:
v_new += (bos * HV + i_h).to(tl.int64) * V
if USE_INITIAL_STATE:
h0 = h0 + i_nh * K*V
if STORE_FINAL_STATE:
ht = ht + i_nh * K*V
# load initial state
o_v = i_v * BV + tl.arange(0, BV)
m_v = o_v < V
o_k1 = tl.arange(0, 64)
m_k1 = o_k1 < K
o_k2 = 64 + o_k1
m_k2 = o_k2 < K
o_k3 = 128 + o_k1
m_k3 = o_k3 < K
o_k4 = 192 + o_k1
m_k4 = o_k4 < K
if USE_INITIAL_STATE:
if STATE_V_FIRST:
p_h0_1 = h0 + o_v[:, None] * K + o_k1[None, :]
m_h0_1 = m_v[:, None] & m_k1[None, :]
else:
p_h0_1 = h0 + o_k1[:, None] * V + o_v[None, :]
m_h0_1 = m_k1[:, None] & m_v[None, :]
b_h1 += tl.load(p_h0_1, mask=m_h0_1, other=0.0).to(tl.float32)
if K > 64:
if STATE_V_FIRST:
p_h0_2 = h0 + o_v[:, None] * K + o_k2[None, :]
m_h0_2 = m_v[:, None] & m_k2[None, :]
else:
p_h0_2 = h0 + o_k2[:, None] * V + o_v[None, :]
m_h0_2 = m_k2[:, None] & m_v[None, :]
b_h2 += tl.load(p_h0_2, mask=m_h0_2, other=0.0).to(tl.float32)
if K > 128:
if STATE_V_FIRST:
p_h0_3 = h0 + o_v[:, None] * K + o_k3[None, :]
m_h0_3 = m_v[:, None] & m_k3[None, :]
else:
p_h0_3 = h0 + o_k3[:, None] * V + o_v[None, :]
m_h0_3 = m_k3[:, None] & m_v[None, :]
b_h3 += tl.load(p_h0_3, mask=m_h0_3, other=0.0).to(tl.float32)
if K > 192:
if STATE_V_FIRST:
p_h0_4 = h0 + o_v[:, None] * K + o_k4[None, :]
m_h0_4 = m_v[:, None] & m_k4[None, :]
else:
p_h0_4 = h0 + o_k4[:, None] * V + o_v[None, :]
m_h0_4 = m_k4[:, None] & m_v[None, :]
b_h4 += tl.load(p_h0_4, mask=m_h0_4, other=0.0).to(tl.float32)
# main recurrence
for i_t in range(NT):
i_t_int64 = i_t.to(tl.int64)
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
if STATE_V_FIRST:
p_h1 = h + i_t_int64 * HV*K*V + o_v[:, None] * K + o_k1[None, :]
m_h1 = m_v[:, None] & m_k1[None, :]
else:
p_h1 = h + i_t_int64 * HV*K*V + o_k1[:, None] * V + o_v[None, :]
m_h1 = m_k1[:, None] & m_v[None, :]
tl.store(p_h1, b_h1.to(p_h1.dtype.element_ty), mask=m_h1)
if K > 64:
if STATE_V_FIRST:
p_h2 = h + i_t_int64 * HV*K*V + o_v[:, None] * K + o_k2[None, :]
m_h2 = m_v[:, None] & m_k2[None, :]
else:
p_h2 = h + i_t_int64 * HV*K*V + o_k2[:, None] * V + o_v[None, :]
m_h2 = m_k2[:, None] & m_v[None, :]
tl.store(p_h2, b_h2.to(p_h2.dtype.element_ty), mask=m_h2)
if K > 128:
if STATE_V_FIRST:
p_h3 = h + i_t_int64 * HV*K*V + o_v[:, None] * K + o_k3[None, :]
m_h3 = m_v[:, None] & m_k3[None, :]
else:
p_h3 = h + i_t_int64 * HV*K*V + o_k3[:, None] * V + o_v[None, :]
m_h3 = m_k3[:, None] & m_v[None, :]
tl.store(p_h3, b_h3.to(p_h3.dtype.element_ty), mask=m_h3)
if K > 192:
if STATE_V_FIRST:
p_h4 = h + i_t_int64 * HV*K*V + o_v[:, None] * K + o_k4[None, :]
m_h4 = m_v[:, None] & m_k4[None, :]
else:
p_h4 = h + i_t_int64 * HV*K*V + o_k4[:, None] * V + o_v[None, :]
m_h4 = m_k4[:, None] & m_v[None, :]
tl.store(p_h4, b_h4.to(p_h4.dtype.element_ty), mask=m_h4)
p_w = w + o_t[:, None] * (HV*K) + o_k1[None, :]
b_w = tl.load(p_w, mask=m_t[:, None] & m_k1[None, :], other=0.0)
if STATE_V_FIRST:
b_v = tl.dot(b_w, tl.trans(b_h1).to(b_w.dtype))
else:
b_v = tl.dot(b_w, b_h1.to(b_w.dtype))
if K > 64:
p_w = w + o_t[:, None] * (HV*K) + o_k2[None, :]
b_w = tl.load(p_w, mask=m_t[:, None] & m_k2[None, :], other=0.0)
if STATE_V_FIRST:
b_v += tl.dot(b_w, tl.trans(b_h2).to(b_w.dtype))
else:
b_v += tl.dot(b_w, b_h2.to(b_w.dtype))
if K > 128:
p_w = w + o_t[:, None] * (HV*K) + o_k3[None, :]
b_w = tl.load(p_w, mask=m_t[:, None] & m_k3[None, :], other=0.0)
if STATE_V_FIRST:
b_v += tl.dot(b_w, tl.trans(b_h3).to(b_w.dtype))
else:
b_v += tl.dot(b_w, b_h3.to(b_w.dtype))
if K > 192:
p_w = w + o_t[:, None] * (HV*K) + o_k4[None, :]
b_w = tl.load(p_w, mask=m_t[:, None] & m_k4[None, :], other=0.0)
if STATE_V_FIRST:
b_v += tl.dot(b_w, tl.trans(b_h4).to(b_w.dtype))
else:
b_v += tl.dot(b_w, b_h4.to(b_w.dtype))
p_v = v + o_t[:, None] * (HV*V) + o_v[None, :]
b_v = tl.load(p_v, mask=m_t[:, None] & m_v[None, :], other=0.0) - b_v
if SAVE_NEW_VALUE:
p_v = v_new + o_t[:, None] * (HV*V) + o_v[None, :]
tl.store(p_v, b_v.to(p_v.dtype.element_ty), mask=m_t[:, None] & m_v[None, :])
last_idx = min((i_t + 1) * BT, T) - 1
if USE_G:
b_g_last = tl.load(g + (bos * HV + last_idx * HV + i_h).to(tl.int64)).to(tl.float32)
p_g = g + (bos * HV + i_h).to(tl.int64) + o_t * HV
b_g = tl.load(p_g, mask=m_t, other=0.0).to(tl.float32)
b_v = b_v * tl.where(m_t, exp2(b_g_last - b_g), 0)[:, None]
b_g_last = exp2(b_g_last)
b_h1 *= b_g_last
if K > 64:
b_h2 *= b_g_last
if K > 128:
b_h3 *= b_g_last
if K > 192:
b_h4 *= b_g_last
if USE_GK:
o_k1 = tl.arange(0, 64)
b_gk_last1 = tl.load(gk + (bos + last_idx) * HV*K + i_h * K + o_k1, mask=(o_k1 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_h1 *= exp2(b_gk_last1)[None, :]
else:
b_h1 *= exp2(b_gk_last1)[:, None]
if K > 64:
o_k2 = 64 + o_k1
b_gk_last2 = tl.load(gk + (bos + last_idx) * HV*K + i_h * K + o_k2, mask=(o_k2 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_h2 *= exp2(b_gk_last2)[None, :]
else:
b_h2 *= exp2(b_gk_last2)[:, None]
if K > 128:
o_k3 = 128 + o_k1
b_gk_last3 = tl.load(gk + (bos + last_idx) * HV*K + i_h * K + o_k3, mask=(o_k3 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_h3 *= exp2(b_gk_last3)[None, :]
else:
b_h3 *= exp2(b_gk_last3)[:, None]
if K > 192:
o_k4 = 192 + o_k1
b_gk_last4 = tl.load(gk + (bos + last_idx) * HV*K + i_h * K + o_k4, mask=(o_k4 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_h4 *= exp2(b_gk_last4)[None, :]
else:
b_h4 *= exp2(b_gk_last4)[:, None]
b_v = b_v.to(k.dtype.element_ty)
p_k = k + o_k1[:, None] + o_t[None, :] * (H*K)
b_k = tl.load(p_k, mask=m_k1[:, None] & m_t[None, :], other=0.0)
if STATE_V_FIRST:
b_h1 += tl.trans(tl.dot(b_k, b_v))
else:
b_h1 += tl.dot(b_k, b_v)
if K > 64:
p_k = k + o_k2[:, None] + o_t[None, :] * (H*K)
b_k = tl.load(p_k, mask=m_k2[:, None] & m_t[None, :], other=0.0)
if STATE_V_FIRST:
b_h2 += tl.trans(tl.dot(b_k, b_v))
else:
b_h2 += tl.dot(b_k, b_v)
if K > 128:
p_k = k + o_k3[:, None] + o_t[None, :] * (H*K)
b_k = tl.load(p_k, mask=m_k3[:, None] & m_t[None, :], other=0.0)
if STATE_V_FIRST:
b_h3 += tl.trans(tl.dot(b_k, b_v))
else:
b_h3 += tl.dot(b_k, b_v)
if K > 192:
p_k = k + o_k4[:, None] + o_t[None, :] * (H*K)
b_k = tl.load(p_k, mask=m_k4[:, None] & m_t[None, :], other=0.0)
if STATE_V_FIRST:
b_h4 += tl.trans(tl.dot(b_k, b_v))
else:
b_h4 += tl.dot(b_k, b_v)
if STORE_FINAL_STATE:
if STATE_V_FIRST:
p_ht = ht + o_v[:, None] * K + o_k1[None, :]
m_ht = m_v[:, None] & m_k1[None, :]
else:
p_ht = ht + o_k1[:, None] * V + o_v[None, :]
m_ht = m_k1[:, None] & m_v[None, :]
tl.store(p_ht, b_h1.to(p_ht.dtype.element_ty), mask=m_ht)
if K > 64:
if STATE_V_FIRST:
p_ht = ht + o_v[:, None] * K + o_k2[None, :]
m_ht = m_v[:, None] & m_k2[None, :]
else:
p_ht = ht + o_k2[:, None] * V + o_v[None, :]
m_ht = m_k2[:, None] & m_v[None, :]
tl.store(p_ht, b_h2.to(p_ht.dtype.element_ty), mask=m_ht)
if K > 128:
if STATE_V_FIRST:
p_ht = ht + o_v[:, None] * K + o_k3[None, :]
m_ht = m_v[:, None] & m_k3[None, :]
else:
p_ht = ht + o_k3[:, None] * V + o_v[None, :]
m_ht = m_k3[:, None] & m_v[None, :]
tl.store(p_ht, b_h3.to(p_ht.dtype.element_ty), mask=m_ht)
if K > 192:
if STATE_V_FIRST:
p_ht = ht + o_v[:, None] * K + o_k4[None, :]
m_ht = m_v[:, None] & m_k4[None, :]
else:
p_ht = ht + o_k4[:, None] * V + o_v[None, :]
m_ht = m_k4[:, None] & m_v[None, :]
tl.store(p_ht, b_h4.to(p_ht.dtype.element_ty), mask=m_ht)
@triton.heuristics({
'USE_G': lambda args: args['g'] is not None,
'USE_GK': lambda args: args['gk'] is not None,
'USE_INITIAL_STATE': lambda args: args['dh0'] is not None,
'USE_FINAL_STATE_GRADIENT': lambda args: args['dht'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({'BV': BV}, num_warps=num_warps, num_stages=num_stages)
for num_warps in [2, 4]
for num_stages in ([2, 3, 4] if check_shared_mem('ampere') else [1])
for BV in ([32, 64] if check_shared_mem('ada') else [32])
],
key=['H', 'HV', 'K', 'V', 'BT', 'BV', 'USE_G', 'STATE_V_FIRST'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_gated_delta_rule_bwd_kernel_dhu_blockdim64(
q,
k,
w,
g,
gk,
dht,
dh0,
do,
dh,
dv,
dv2,
cu_seqlens,
chunk_offsets,
scale,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BV: tl.constexpr,
USE_G: tl.constexpr,
USE_GK: tl.constexpr,
USE_INITIAL_STATE: tl.constexpr,
USE_FINAL_STATE_GRADIENT: tl.constexpr,
STATE_V_FIRST: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
pid = tl.program_id(0)
NV = tl.cdiv(V, BV)
i_v, i_nh = pid % NV, (pid // NV).to(tl.int64)
i_n, i_h = i_nh // HV, i_nh % HV
if IS_VARLEN:
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
NT = tl.cdiv(T, BT)
boh = tl.load(chunk_offsets + i_n).to(tl.int64)
else:
bos, eos = i_n * T, i_n * T + T
NT = tl.cdiv(T, BT)
boh = i_n * NT
if STATE_V_FIRST:
b_dh1 = tl.zeros([BV, 64], dtype=tl.float32)
if K > 64:
b_dh2 = tl.zeros([BV, 64], dtype=tl.float32)
if K > 128:
b_dh3 = tl.zeros([BV, 64], dtype=tl.float32)
if K > 192:
b_dh4 = tl.zeros([BV, 64], dtype=tl.float32)
else:
b_dh1 = tl.zeros([64, BV], dtype=tl.float32)
if K > 64:
b_dh2 = tl.zeros([64, BV], dtype=tl.float32)
if K > 128:
b_dh3 = tl.zeros([64, BV], dtype=tl.float32)
if K > 192:
b_dh4 = tl.zeros([64, BV], dtype=tl.float32)
# calculate offset
q += (bos * H + i_h // (HV // H)).to(tl.int64) * K
k += (bos * H + i_h // (HV // H)).to(tl.int64) * K
w += (bos * HV + i_h).to(tl.int64) * K
do += (bos * HV + i_h).to(tl.int64) * V
dv += (bos * HV + i_h).to(tl.int64) * V
dv2 += (bos * HV + i_h).to(tl.int64) * V
dh += (boh * HV + i_h).to(tl.int64) * K*V
if USE_GK:
gk += (bos * HV + i_h).to(tl.int64) * K
if USE_INITIAL_STATE:
dh0 += i_nh * K*V
if USE_FINAL_STATE_GRADIENT:
dht += i_nh * K*V
o_v = i_v * BV + tl.arange(0, BV)
m_v = o_v < V
o_k1 = tl.arange(0, 64)
m_k1 = o_k1 < K
o_k2 = 64 + o_k1
m_k2 = o_k2 < K
o_k3 = 128 + o_k1
m_k3 = o_k3 < K
o_k4 = 192 + o_k1
m_k4 = o_k4 < K
if USE_FINAL_STATE_GRADIENT:
if STATE_V_FIRST:
p_dht1 = dht + o_v[:, None] * K + o_k1[None, :]
m_dht1 = m_v[:, None] & m_k1[None, :]
else:
p_dht1 = dht + o_k1[:, None] * V + o_v[None, :]
m_dht1 = m_k1[:, None] & m_v[None, :]
b_dh1 += tl.load(p_dht1, mask=m_dht1, other=0.0)
if K > 64:
if STATE_V_FIRST:
p_dht2 = dht + o_v[:, None] * K + o_k2[None, :]
m_dht2 = m_v[:, None] & m_k2[None, :]
else:
p_dht2 = dht + o_k2[:, None] * V + o_v[None, :]
m_dht2 = m_k2[:, None] & m_v[None, :]
b_dh2 += tl.load(p_dht2, mask=m_dht2, other=0.0)
if K > 128:
if STATE_V_FIRST:
p_dht3 = dht + o_v[:, None] * K + o_k3[None, :]
m_dht3 = m_v[:, None] & m_k3[None, :]
else:
p_dht3 = dht + o_k3[:, None] * V + o_v[None, :]
m_dht3 = m_k3[:, None] & m_v[None, :]
b_dh3 += tl.load(p_dht3, mask=m_dht3, other=0.0)
if K > 192:
if STATE_V_FIRST:
p_dht4 = dht + o_v[:, None] * K + o_k4[None, :]
m_dht4 = m_v[:, None] & m_k4[None, :]
else:
p_dht4 = dht + o_k4[:, None] * V + o_v[None, :]
m_dht4 = m_k4[:, None] & m_v[None, :]
b_dh4 += tl.load(p_dht4, mask=m_dht4, other=0.0)
for i_t in range(NT - 1, -1, -1):
i_t_int64 = i_t.to(tl.int64)
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
if STATE_V_FIRST:
p_dh1 = dh + i_t_int64*HV*K*V + o_v[:, None] * K + o_k1[None, :]
m_dh1 = m_v[:, None] & m_k1[None, :]
else:
p_dh1 = dh + i_t_int64*HV*K*V + o_k1[:, None] * V + o_v[None, :]
m_dh1 = m_k1[:, None] & m_v[None, :]
tl.store(p_dh1, b_dh1.to(p_dh1.dtype.element_ty), mask=m_dh1)
if K > 64:
if STATE_V_FIRST:
p_dh2 = dh + i_t_int64*HV*K*V + o_v[:, None] * K + o_k2[None, :]
m_dh2 = m_v[:, None] & m_k2[None, :]
else:
p_dh2 = dh + i_t_int64*HV*K*V + o_k2[:, None] * V + o_v[None, :]
m_dh2 = m_k2[:, None] & m_v[None, :]
tl.store(p_dh2, b_dh2.to(p_dh2.dtype.element_ty), mask=m_dh2)
if K > 128:
if STATE_V_FIRST:
p_dh3 = dh + i_t_int64*HV*K*V + o_v[:, None] * K + o_k3[None, :]
m_dh3 = m_v[:, None] & m_k3[None, :]
else:
p_dh3 = dh + i_t_int64*HV*K*V + o_k3[:, None] * V + o_v[None, :]
m_dh3 = m_k3[:, None] & m_v[None, :]
tl.store(p_dh3, b_dh3.to(p_dh3.dtype.element_ty), mask=m_dh3)
if K > 192:
if STATE_V_FIRST:
p_dh4 = dh + i_t_int64*HV*K*V + o_v[:, None] * K + o_k4[None, :]
m_dh4 = m_v[:, None] & m_k4[None, :]
else:
p_dh4 = dh + i_t_int64*HV*K*V + o_k4[:, None] * V + o_v[None, :]
m_dh4 = m_k4[:, None] & m_v[None, :]
tl.store(p_dh4, b_dh4.to(p_dh4.dtype.element_ty), mask=m_dh4)
last_idx = min((i_t + 1) * BT, T) - 1
if USE_G:
bg_last = tl.load(g + (bos + last_idx) * HV + i_h).to(tl.float32)
p_g = g + bos * HV + i_h + o_t * HV
b_g = tl.load(p_g, mask=m_t, other=0.0).to(tl.float32)
bg_last_exp = exp2(bg_last)
b_g_exp = exp2(b_g)
p_dv = dv + o_t[:, None] * (HV*V) + o_v[None, :]
p_dv2 = dv2 + o_t[:, None] * (HV*V) + o_v[None, :]
p_do = do + o_t[:, None] * (HV*V) + o_v[None, :]
b_do = tl.load(p_do, mask=m_t[:, None] & m_v[None, :], other=0.0)
# Update dv
p_k = k + o_t[:, None] * (H*K) + o_k1[None, :]
b_k = tl.load(p_k, mask=m_t[:, None] & m_k1[None, :], other=0.0)
if USE_GK:
o_k1 = tl.arange(0, 64)
b_gk_last1 = tl.load(gk + last_idx * HV*K + o_k1, mask=(o_k1 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_dv = tl.dot(b_k, tl.trans(b_dh1).to(b_k.dtype))
else:
b_dv = tl.dot(b_k, b_dh1.to(b_k.dtype))
if K > 64:
p_k = k + o_t[:, None] * (H*K) + o_k2[None, :]
b_k = tl.load(p_k, mask=m_t[:, None] & m_k2[None, :], other=0.0)
if USE_GK:
b_gk_last2 = tl.load(gk + last_idx * HV*K + o_k2, mask=(o_k2 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_dv += tl.dot(b_k, tl.trans(b_dh2).to(b_k.dtype))
else:
b_dv += tl.dot(b_k, b_dh2.to(b_k.dtype))
if K > 128:
p_k = k + o_t[:, None] * (H*K) + o_k3[None, :]
b_k = tl.load(p_k, mask=m_t[:, None] & m_k3[None, :], other=0.0)
if USE_GK:
b_gk_last3 = tl.load(gk + last_idx * HV*K + o_k3, mask=(o_k3 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_dv += tl.dot(b_k, tl.trans(b_dh3).to(b_k.dtype))
else:
b_dv += tl.dot(b_k, b_dh3.to(b_k.dtype))
if K > 192:
p_k = k + o_t[:, None] * (H*K) + o_k4[None, :]
b_k = tl.load(p_k, mask=m_t[:, None] & m_k4[None, :], other=0.0)
if USE_GK:
b_gk_last4 = tl.load(gk + last_idx * HV*K + o_k4, mask=(o_k4 < K), other=0.).to(tl.float32)
if STATE_V_FIRST:
b_dv += tl.dot(b_k, tl.trans(b_dh4).to(b_k.dtype))
else:
b_dv += tl.dot(b_k, b_dh4.to(b_k.dtype))
if USE_G:
b_dv *= tl.where(m_t, exp2(bg_last - b_g), 0)[:, None]
b_dv += tl.load(p_dv, mask=m_t[:, None] & m_v[None, :], other=0.0)
tl.store(p_dv2, b_dv.to(p_dv.dtype.element_ty), mask=m_t[:, None] & m_v[None, :])
# Update dh
p_w = w + o_k1[:, None] + o_t[None, :] * (HV*K)
p_q = q + o_k1[:, None] + o_t[None, :] * (H*K)
b_w = tl.load(p_w, mask=m_k1[:, None] & m_t[None, :], other=0.0)
b_q = tl.load(p_q, mask=m_k1[:, None] & m_t[None, :], other=0.0)
if USE_G:
b_dh1 *= bg_last_exp
b_q = b_q * b_g_exp[None, :]
if USE_GK:
if STATE_V_FIRST:
b_dh1 *= exp2(b_gk_last1)[None, :]
else:
b_dh1 *= exp2(b_gk_last1[:, None])
if STATE_V_FIRST:
b_dh1 += tl.trans(tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype)))
else:
b_dh1 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
if K > 64:
p_q = q + o_k2[:, None] + o_t[None, :] * (H*K)
p_w = w + o_k2[:, None] + o_t[None, :] * (HV*K)
b_q = tl.load(p_q, mask=m_k2[:, None] & m_t[None, :], other=0.0)
b_w = tl.load(p_w, mask=m_k2[:, None] & m_t[None, :], other=0.0)
if USE_G:
b_dh2 *= bg_last_exp
b_q = b_q * b_g_exp[None, :]
if USE_GK:
if STATE_V_FIRST:
b_dh2 *= exp2(b_gk_last2)[None, :]
else:
b_dh2 *= exp2(b_gk_last2[:, None])
if STATE_V_FIRST:
b_dh2 += tl.trans(tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype)))
else:
b_dh2 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
if K > 128:
p_q = q + o_k3[:, None] + o_t[None, :] * (H*K)
p_w = w + o_k3[:, None] + o_t[None, :] * (HV*K)
b_q = tl.load(p_q, mask=m_k3[:, None] & m_t[None, :], other=0.0)
b_w = tl.load(p_w, mask=m_k3[:, None] & m_t[None, :], other=0.0)
if USE_G:
b_dh3 *= bg_last_exp
b_q = b_q * b_g_exp[None, :]
if USE_GK:
if STATE_V_FIRST:
b_dh3 *= exp2(b_gk_last3)[None, :]
else:
b_dh3 *= exp2(b_gk_last3[:, None])
if STATE_V_FIRST:
b_dh3 += tl.trans(tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype)))
else:
b_dh3 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
if K > 192:
p_q = q + o_k4[:, None] + o_t[None, :] * (H*K)
p_w = w + o_k4[:, None] + o_t[None, :] * (HV*K)
b_q = tl.load(p_q, mask=m_k4[:, None] & m_t[None, :], other=0.0)
b_w = tl.load(p_w, mask=m_k4[:, None] & m_t[None, :], other=0.0)
if USE_G:
b_dh4 *= bg_last_exp
b_q = b_q * b_g_exp[None, :]
if USE_GK:
if STATE_V_FIRST:
b_dh4 *= exp2(b_gk_last4)[None, :]
else:
b_dh4 *= exp2(b_gk_last4[:, None])
if STATE_V_FIRST:
b_dh4 += tl.trans(tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype)))
else:
b_dh4 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
if USE_INITIAL_STATE:
if STATE_V_FIRST:
p_dh0 = dh0 + o_v[:, None] * K + o_k1[None, :]
m_dh0 = m_v[:, None] & m_k1[None, :]
else:
p_dh0 = dh0 + o_k1[:, None] * V + o_v[None, :]
m_dh0 = m_k1[:, None] & m_v[None, :]
tl.store(p_dh0, b_dh1.to(p_dh0.dtype.element_ty), mask=m_dh0)
if K > 64:
if STATE_V_FIRST:
p_dh1 = dh0 + o_v[:, None] * K + o_k2[None, :]
m_dh1 = m_v[:, None] & m_k2[None, :]
else:
p_dh1 = dh0 + o_k2[:, None] * V + o_v[None, :]
m_dh1 = m_k2[:, None] & m_v[None, :]
tl.store(p_dh1, b_dh2.to(p_dh1.dtype.element_ty), mask=m_dh1)
if K > 128:
if STATE_V_FIRST:
p_dh2 = dh0 + o_v[:, None] * K + o_k3[None, :]
m_dh2 = m_v[:, None] & m_k3[None, :]
else:
p_dh2 = dh0 + o_k3[:, None] * V + o_v[None, :]
m_dh2 = m_k3[:, None] & m_v[None, :]
tl.store(p_dh2, b_dh3.to(p_dh2.dtype.element_ty), mask=m_dh2)
if K > 192:
if STATE_V_FIRST:
p_dh3 = dh0 + o_v[:, None] * K + o_k4[None, :]
m_dh3 = m_v[:, None] & m_k4[None, :]
else:
p_dh3 = dh0 + o_k4[:, None] * V + o_v[None, :]
m_dh3 = m_k4[:, None] & m_v[None, :]
tl.store(p_dh3, b_dh4.to(p_dh3.dtype.element_ty), mask=m_dh3)
@dispatch('common')
def chunk_gated_delta_rule_fwd_h(
k: torch.Tensor,
w: torch.Tensor,
u: torch.Tensor,
g: torch.Tensor | None = None,
gk: torch.Tensor | None = None,
initial_state: torch.Tensor | None = None,
output_final_state: bool = False,
chunk_size: int = 64,
save_new_value: bool = True,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
cu_seqlens_cpu: torch.LongTensor | None = None,
chunk_indices: torch.LongTensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None]:
B, T, H, K, V, HV = *k.shape, u.shape[-1], u.shape[2]
BT = chunk_size
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, chunk_size)
# N: the actual number of sequences in the batch with either equal or variable lengths
if cu_seqlens is None:
N, NT, chunk_offsets = B, triton.cdiv(T, BT), None
else:
N, NT, chunk_offsets = len(cu_seqlens) - 1, len(chunk_indices), prepare_chunk_offsets(cu_seqlens, BT)
assert K <= 256, "current kernel does not support head dimension larger than 256."
if state_v_first:
h = k.new_empty(B, NT, HV, V, K)
final_state = k.new_zeros(N, HV, V, K, dtype=torch.float32) if output_final_state else None
else:
h = k.new_empty(B, NT, HV, K, V)
final_state = k.new_zeros(N, HV, K, V, dtype=torch.float32) if output_final_state else None
v_new = torch.empty_like(u) if save_new_value else None
def grid(meta): return (triton.cdiv(V, meta['BV']) * N * HV, )
chunk_gated_delta_rule_fwd_kernel_h_blockdim64[grid](
k=k,
v=u,
w=w,
v_new=v_new,
g=g,
gk=gk,
h=h,
h0=initial_state,
ht=final_state,
cu_seqlens=cu_seqlens,
chunk_offsets=chunk_offsets,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BT=BT,
STATE_V_FIRST=state_v_first,
)
return h, v_new, final_state
@dispatch('common')
def chunk_gated_delta_rule_bwd_dhu(
q: torch.Tensor,
k: torch.Tensor,
w: torch.Tensor,
do: torch.Tensor,
dv: torch.Tensor,
g: torch.Tensor | None = None,
gk: torch.Tensor | None = None,
h0: torch.Tensor | None = None,
dht: torch.Tensor | None = None,
scale: float | None = None,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
chunk_size: int = 64,
chunk_indices: torch.LongTensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
B, T, H, K, V, HV = *q.shape, do.shape[-1], do.shape[2]
# N: the actual number of sequences in the batch with either equal or variable lengths
BT = chunk_size
assert K <= 256, "current kernel does not support head dimension being larger than 256."
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, chunk_size)
if cu_seqlens is None:
N, NT, chunk_offsets = B, triton.cdiv(T, BT), None
else:
N, NT, chunk_offsets = len(cu_seqlens) - 1, len(chunk_indices), prepare_chunk_offsets(cu_seqlens, BT)
if state_v_first:
dh = q.new_empty(B, NT, HV, V, K)
else:
dh = q.new_empty(B, NT, HV, K, V)
dh0 = torch.empty_like(h0, dtype=torch.float32) if h0 is not None else None
dv2 = torch.empty_like(dv)
def grid(meta): return (triton.cdiv(V, meta['BV']) * N * HV, )
chunk_gated_delta_rule_bwd_kernel_dhu_blockdim64[grid](
q=q,
k=k,
w=w,
g=g,
gk=gk,
dht=dht,
dh0=dh0,
do=do,
dh=dh,
dv=dv,
dv2=dv2,
cu_seqlens=cu_seqlens,
chunk_offsets=chunk_offsets,
scale=scale,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BT=BT,
STATE_V_FIRST=state_v_first,
)
return dh, dh0, dv2
+432
View File
@@ -0,0 +1,432 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import torch
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.ops.utils import prepare_chunk_offsets
from kda._fla.ops.utils.op import exp2
from kda._fla.utils import autotune_cache_kwargs, check_shared_mem
BKV_LIST = [32, 64] if check_shared_mem() else [16, 32]
@triton.heuristics({
'USE_INITIAL_STATE': lambda args: args['h0'] is not None,
'STORE_FINAL_STATE': lambda args: args['ht'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@triton.autotune(
configs=[
triton.Config({'BK': BK, 'BV': BV}, num_warps=num_warps, num_stages=num_stages)
for BK in BKV_LIST
for BV in BKV_LIST
for num_warps in [1, 2, 4, 8]
for num_stages in [2, 3, 4]
],
key=['BT', 'USE_G', 'USE_GK', 'USE_GV', 'STATE_V_FIRST'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_fwd_kernel_h(
k,
v,
h,
g,
g_gamma,
gk,
gv,
h0,
ht,
cu_seqlens,
split_offsets,
T,
H: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BS: tl.constexpr,
BK: tl.constexpr,
BV: tl.constexpr,
USE_G: tl.constexpr,
USE_G_GAMMA: tl.constexpr,
USE_GK: tl.constexpr,
USE_GV: tl.constexpr,
USE_INITIAL_STATE: tl.constexpr,
STORE_FINAL_STATE: tl.constexpr,
IS_VARLEN: tl.constexpr,
STATE_V_FIRST: tl.constexpr,
):
i_k, i_v, i_nh = tl.program_id(0), tl.program_id(1), tl.program_id(2).to(tl.int64)
i_n, i_h = i_nh // H, i_nh % H
if IS_VARLEN:
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
NT, NS = tl.cdiv(T, BT), tl.cdiv(T, BS)
boh = tl.load(split_offsets + i_n).to(tl.int64)
else:
bos, eos = i_n * T, i_n * T + T
NT, NS = tl.cdiv(T, BT), tl.cdiv(T, BS)
boh = i_n * NS
NTS = BS // BT
if USE_G_GAMMA:
# decay rate given the head index
b_gamma = tl.load(g_gamma + i_h)
b_g = b_gamma * (tl.arange(0, BT) + 1)
# [BK, BV] accumulator; STATE_V_FIRST only flips the stored state's HBM layout to [V, K], applied at the load/store below.
b_h = tl.zeros([BK, BV], dtype=tl.float32)
o_k = i_k * BK + tl.arange(0, BK)
o_v = i_v * BV + tl.arange(0, BV)
if USE_INITIAL_STATE:
if STATE_V_FIRST:
p_h0 = h0 + i_nh * K*V + o_v[:, None] * K + o_k[None, :]
b_h = tl.trans(tl.load(p_h0, mask=(o_v[:, None] < V) & (o_k[None, :] < K), other=0.0)).to(tl.float32)
else:
p_h0 = h0 + i_nh * K*V + o_k[:, None] * V + o_v[None, :]
b_h = tl.load(p_h0, mask=(o_k[:, None] < K) & (o_v[None, :] < V), other=0.0).to(tl.float32)
for i_t in range(NT):
i_s = i_t // NTS
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
p_k = k + (bos*H + i_h) * K + o_k[:, None] + o_t[None, :] * (H*K)
p_v = v + (bos*H + i_h) * V + o_t[:, None] * (H*V) + o_v[None, :]
o_h = ((boh + i_s) * H + i_h).to(tl.int64) * K*V
if STATE_V_FIRST:
p_h = h + o_h + o_v[:, None] * K + o_k[None, :]
m_h = (o_v[:, None] < V) & (o_k[None, :] < K)
else:
p_h = h + o_h + o_k[:, None] * V + o_v[None, :]
m_h = (o_k[:, None] < K) & (o_v[None, :] < V)
if i_t % NTS == 0:
tl.store(p_h, (tl.trans(b_h) if STATE_V_FIRST else b_h).to(p_h.dtype.element_ty), mask=m_h)
# [BK, BT]
b_k = tl.load(p_k, mask=(o_k[:, None] < K) & m_t[None, :], other=0.0)
# [BT, BV]
b_v = tl.load(p_v, mask=m_t[:, None] & (o_v < V)[None, :], other=0.0)
last_idx = min((i_t + 1) * BT, T) - 1
# scalar decay
if USE_G:
b_g_last = tl.load(g + bos * H + last_idx * H + i_h)
p_g = g + bos*H + (i_t * BT + tl.arange(0, BT)) * H + i_h
b_g = tl.load(p_g, mask=(i_t * BT + tl.arange(0, BT) < T), other=0.)
b_h *= exp2(b_g_last)
b_v = (b_v * exp2(b_g_last - b_g)[:, None]).to(b_v.dtype)
if USE_G_GAMMA:
b_g_last = b_gamma * min(BT, T - i_t * BT)
b_h *= exp2(b_g_last)
b_v = (b_v * exp2(b_g_last - b_g)[:, None]).to(b_v.dtype)
# vector decay, h = Diag(gk) @ h
if USE_GK:
p_gk = gk + (bos*H + i_h) * K + o_k[:, None] + o_t[None, :] * (H*K)
p_gk_last = gk + (bos + last_idx) * H*K + i_h * K + i_k * BK + tl.arange(0, BK)
b_gk_last = tl.load(p_gk_last, mask=(i_k * BK + tl.arange(0, BK) < K), other=0.)
b_gk = tl.load(p_gk, mask=(o_k[:, None] < K) & m_t[None, :], other=0.0)
b_h *= exp2(b_gk_last)[:, None]
b_k = (b_k * exp2(b_gk_last[:, None] - b_gk)).to(b_k.dtype)
# vector decay, h = h @ Diag(gv)
if USE_GV:
p_gv = gv + (bos*H + i_h) * V + o_t[:, None] * (H*V) + o_v[None, :]
p_gv_last = gv + (bos + last_idx) * H*V + i_h * V + i_v * BV + tl.arange(0, BV)
b_gv_last = tl.load(p_gv_last, mask=(i_v * BV + tl.arange(0, BV) < V), other=0.)
b_gv = tl.load(p_gv, mask=m_t[:, None] & (o_v < V)[None, :], other=0.0)
b_h *= exp2(b_gv_last)[None, :]
b_v = (b_v * exp2(b_gv_last[None, :] - b_gv)).to(b_v.dtype)
b_h += tl.dot(b_k, b_v)
if STORE_FINAL_STATE:
if STATE_V_FIRST:
p_ht = ht + i_nh * K*V + o_v[:, None] * K + o_k[None, :]
tl.store(p_ht, tl.trans(b_h).to(p_ht.dtype.element_ty), mask=(o_v[:, None] < V) & (o_k[None, :] < K))
else:
p_ht = ht + i_nh * K*V + o_k[:, None] * V + o_v[None, :]
tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), mask=(o_k[:, None] < K) & (o_v[None, :] < V))
@triton.heuristics({
'STORE_INITIAL_STATE_GRADIENT': lambda args: args['dh0'] is not None,
'USE_FINAL_STATE_GRADIENT': lambda args: args['dht'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@triton.autotune(
configs=[
triton.Config({'BK': BK, 'BV': BV}, num_warps=num_warps, num_stages=num_stages)
for BK in BKV_LIST
for BV in BKV_LIST
for num_warps in [1, 2, 4, 8]
for num_stages in [2, 3, 4]
],
key=['BT', 'USE_G', 'USE_GK', 'USE_GV', 'STATE_V_FIRST'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_bwd_kernel_dh(
q,
g,
g_gamma,
gk,
gv,
do,
dh,
dht,
dh0,
cu_seqlens,
split_offsets,
scale,
T,
HQ: tl.constexpr,
H: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BS: tl.constexpr,
BK: tl.constexpr,
BV: tl.constexpr,
NG: tl.constexpr,
USE_G: tl.constexpr,
USE_G_GAMMA: tl.constexpr,
USE_GK: tl.constexpr,
USE_GV: tl.constexpr,
STORE_INITIAL_STATE_GRADIENT: tl.constexpr,
USE_FINAL_STATE_GRADIENT: tl.constexpr,
IS_VARLEN: tl.constexpr,
STATE_V_FIRST: tl.constexpr,
):
i_k, i_v, i_nh = tl.program_id(0), tl.program_id(1), tl.program_id(2).to(tl.int64)
i_n, i_hq = i_nh // HQ, i_nh % HQ
i_h = i_hq // NG
if IS_VARLEN:
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
NT = tl.cdiv(T, BT)
NS = tl.cdiv(T, BS)
boh = tl.load(split_offsets + i_n).to(tl.int64)
else:
bos, eos = i_n * T, i_n * T + T
NT = tl.cdiv(T, BT)
NS = tl.cdiv(T, BS)
boh = i_n * NS
if USE_G_GAMMA:
b_gamma = tl.load(g_gamma + i_h)
b_g = b_gamma * (tl.arange(0, BT) + 1)
# [BK, BV] accumulator; STATE_V_FIRST only flips the stored state's HBM layout to [V, K], applied at the load/store below.
b_dh = tl.zeros([BK, BV], dtype=tl.float32)
o_k = i_k * BK + tl.arange(0, BK)
o_v = i_v * BV + tl.arange(0, BV)
if USE_FINAL_STATE_GRADIENT:
if STATE_V_FIRST:
p_dht = dht + i_nh * K*V + o_v[:, None] * K + o_k[None, :]
b_dh += tl.trans(tl.load(p_dht, mask=(o_v[:, None] < V) & (o_k[None, :] < K), other=0.0)).to(tl.float32)
else:
p_dht = dht + i_nh * K*V + o_k[:, None] * V + o_v[None, :]
b_dh += tl.load(p_dht, mask=(o_k[:, None] < K) & (o_v[None, :] < V), other=0.0).to(tl.float32)
for i_t in range(NT - 1, -1, -1):
i_s = i_t // (BS // BT)
o_dh = ((boh + i_s) * H + i_h).to(tl.int64) * K*V
if STATE_V_FIRST:
p_dh = dh + o_dh + o_v[:, None] * K + o_k[None, :]
m_dh = (o_v[:, None] < V) & (o_k[None, :] < K)
else:
p_dh = dh + o_dh + o_k[:, None] * V + o_v[None, :]
m_dh = (o_k[:, None] < K) & (o_v[None, :] < V)
if i_t % (BS // BT) == 0:
tl.store(p_dh, (tl.trans(b_dh) if STATE_V_FIRST else b_dh).to(p_dh.dtype.element_ty), mask=m_dh)
last_idx = min(i_t * BT + BT, T) - 1
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
# [BK, BT]
p_q = q + (bos*HQ + i_hq) * K + o_k[:, None] + o_t[None, :] * (HQ*K)
p_do = do + (bos*HQ + i_hq) * V + o_t[:, None] * (HQ*V) + o_v[None, :]
b_q = tl.load(p_q, mask=(o_k[:, None] < K) & m_t[None, :], other=0.0)
b_q = (b_q * scale).to(b_q.dtype)
# [BT, BV]
b_do = tl.load(p_do, mask=m_t[:, None] & (o_v < V)[None, :], other=0.0)
if USE_G:
p_g = g + (bos + i_t * BT + tl.arange(0, BT)) * H + i_h
b_g_last = tl.load(g + (bos + last_idx) * H + i_h)
b_g = tl.load(p_g, mask=(i_t * BT + tl.arange(0, BT) < T), other=0.)
b_q = (b_q * exp2(b_g)[None, :]).to(b_q.dtype)
b_dh *= exp2(b_g_last)
if USE_G_GAMMA:
b_g_last = b_gamma * min(BT, T - i_t * BT)
b_q = (b_q * exp2(b_g)[None, :]).to(b_q.dtype)
b_dh *= exp2(b_g_last)
if USE_GK:
p_gk = gk + (bos*H + i_h) * K + o_k[:, None] + o_t[None, :] * (H*K)
p_gk_last = gk + (bos + last_idx) * H*K + i_h * K + i_k * BK + tl.arange(0, BK)
b_gk = tl.load(p_gk, mask=(o_k[:, None] < K) & m_t[None, :], other=0.0)
b_gk_last = tl.load(p_gk_last, mask=(i_k * BK + tl.arange(0, BK) < K), other=0.)
b_q = (b_q * exp2(b_gk)).to(b_q.dtype)
b_dh *= exp2(b_gk_last)[:, None]
if USE_GV:
p_gv = gv + (bos*H + i_h) * V + o_t[:, None] * (H*V) + o_v[None, :]
p_gv_last = gv + (bos + last_idx) * H*V + i_h * V + i_v * BV + tl.arange(0, BV)
b_gv = tl.load(p_gv, mask=m_t[:, None] & (o_v < V)[None, :], other=0.0)
b_gv_last = tl.load(p_gv_last, mask=(i_v * BV + tl.arange(0, BV) < V), other=0.)
b_do = (b_do * exp2(b_gv))
b_dh *= exp2(b_gv_last)[None, :]
b_dh += tl.dot(b_q, b_do.to(b_q.dtype))
if STORE_INITIAL_STATE_GRADIENT:
if STATE_V_FIRST:
p_dh0 = dh0 + i_nh * K*V + o_v[:, None] * K + o_k[None, :]
tl.store(p_dh0, tl.trans(b_dh).to(p_dh0.dtype.element_ty), mask=(o_v[:, None] < V) & (o_k[None, :] < K))
else:
p_dh0 = dh0 + i_nh * K*V + o_k[:, None] * V + o_v[None, :]
tl.store(p_dh0, b_dh.to(p_dh0.dtype.element_ty), mask=(o_k[:, None] < K) & (o_v[None, :] < V))
@dispatch('common')
def chunk_fwd_h(
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor | None = None,
g_gamma: torch.Tensor | None = None,
gk: torch.Tensor | None = None,
gv: torch.Tensor | None = None,
h0: torch.Tensor | None = None,
output_final_state: bool = False,
state_v_first: bool = False,
cu_seqlens: torch.Tensor | None = None,
chunk_size: int = 64,
split_size: int | None = None,
states_in_fp32: bool = False,
) -> tuple[torch.Tensor, torch.Tensor]:
B, T, H, K, V = *k.shape, v.shape[-1]
BT = chunk_size
BS = BT if split_size is None else split_size
assert BS % BT == 0, f"The `split_size` (got {BS}) must be a multiple of `chunk_size` {BT}"
# N: the actual number of sequences in the batch with either equal or variable lengths
if cu_seqlens is None:
N, NS, split_offsets = B, triton.cdiv(T, BS), None
else:
split_offsets = prepare_chunk_offsets(cu_seqlens, BS)
N, NS = len(cu_seqlens) - 1, split_offsets[-1].item()
# `state_v_first` stores the states in V-first `[V, K]` layout instead of `[K, V]`
state_shape = (V, K) if state_v_first else (K, V)
h = k.new_empty(B, NS, H, *state_shape, dtype=k.dtype if not states_in_fp32 else torch.float)
ht = k.new_empty(N, H, *state_shape, dtype=torch.float) if output_final_state else None
def grid(meta): return (triton.cdiv(K, meta['BK']), triton.cdiv(V, meta['BV']), N * H)
chunk_fwd_kernel_h[grid](
k=k,
v=v,
h=h,
g=g,
g_gamma=g_gamma,
gk=gk,
gv=gv,
h0=h0,
ht=ht,
cu_seqlens=cu_seqlens,
split_offsets=split_offsets,
T=T,
H=H,
K=K,
V=V,
BT=BT,
BS=BS,
USE_G=g is not None,
USE_G_GAMMA=g_gamma is not None,
USE_GK=gk is not None,
USE_GV=gv is not None,
STATE_V_FIRST=state_v_first,
)
return h, ht
@dispatch('common')
def chunk_bwd_dh(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
do: torch.Tensor,
h0: torch.Tensor,
dht: torch.Tensor,
scale: float,
g: torch.Tensor | None = None,
g_gamma: torch.Tensor | None = None,
gk: torch.Tensor | None = None,
gv: torch.Tensor | None = None,
state_v_first: bool = False,
cu_seqlens: torch.Tensor | None = None,
chunk_size: int = 64,
split_size: int | None = None,
states_in_fp32: bool = False,
) -> tuple[torch.Tensor, torch.Tensor]:
B, T, H, K, V = *k.shape, v.shape[-1]
HQ = q.shape[2]
BT = chunk_size
BS = BT if split_size is None else split_size
assert BS % BT == 0, f"The `split_size` (got {BS}) must be a multiple of `chunk_size` {BT}"
# N: the actual number of sequences in the batch with either equal or variable lengths
# NG: number of groups in GQA
if cu_seqlens is None:
N, NS, split_offsets = B, triton.cdiv(T, BS), None
else:
split_offsets = prepare_chunk_offsets(cu_seqlens, BS)
N, NS = len(cu_seqlens) - 1, split_offsets[-1].item()
NG = HQ // H
# `state_v_first` stores the states in V-first `[V, K]` layout instead of `[K, V]`
state_shape = (V, K) if state_v_first else (K, V)
dh = k.new_empty(B, NS, HQ, *state_shape, dtype=k.dtype if not states_in_fp32 else torch.float)
dh0 = torch.empty_like(h0, dtype=torch.float) if h0 is not None else None
def grid(meta): return (triton.cdiv(K, meta['BK']), triton.cdiv(V, meta['BV']), N * H)
chunk_bwd_kernel_dh[grid](
q=q,
g=g,
g_gamma=g_gamma,
gk=gk,
gv=gv,
do=do,
dh=dh,
dht=dht,
dh0=dh0,
cu_seqlens=cu_seqlens,
split_offsets=split_offsets,
scale=scale,
T=T,
HQ=HQ,
H=H,
K=K,
V=V,
BT=BT,
BS=BS,
NG=NG,
USE_G=g is not None,
USE_G_GAMMA=g_gamma is not None,
USE_GK=gk is not None,
USE_GV=gv is not None,
STATE_V_FIRST=state_v_first,
)
return dh, dh0
+110
View File
@@ -0,0 +1,110 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
# Shared gate helpers reused across delta-rule family ops (KDA, GDN, ...).
import torch
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.utils import autocast_custom_bwd, autocast_custom_fwd, input_guard
@triton.jit
def fused_beta_sigmoid_fwd_kernel(
x,
y,
scale,
n_elements,
BLOCK_SIZE: tl.constexpr,
):
pid = tl.program_id(0).to(tl.int64)
offs = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE).to(tl.int64)
mask = offs < n_elements
b_x = tl.load(x + offs, mask=mask, other=0).to(tl.float32)
b_y = scale * tl.sigmoid(b_x)
tl.store(y + offs, b_y.to(y.dtype.element_ty), mask=mask)
@triton.jit
def fused_beta_sigmoid_bwd_kernel(
x,
dy,
dx,
scale,
n_elements,
BLOCK_SIZE: tl.constexpr,
):
pid = tl.program_id(0).to(tl.int64)
offs = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE).to(tl.int64)
mask = offs < n_elements
b_x = tl.load(x + offs, mask=mask, other=0).to(tl.float32)
b_dy = tl.load(dy + offs, mask=mask, other=0).to(tl.float32)
b_y = tl.sigmoid(b_x)
b_dx = b_dy * scale * b_y * (1.0 - b_y)
tl.store(dx + offs, b_dx.to(dx.dtype.element_ty), mask=mask)
_BETA_SIGMOID_BLOCK_SIZE = 2048
_BETA_SIGMOID_NUM_WARPS = 8
@dispatch('common')
def fused_beta_sigmoid_fwd(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
y = torch.empty_like(x, dtype=torch.float32)
n_elements = x.numel()
grid = (triton.cdiv(n_elements, _BETA_SIGMOID_BLOCK_SIZE),)
fused_beta_sigmoid_fwd_kernel[grid](
x,
y,
scale,
n_elements,
BLOCK_SIZE=_BETA_SIGMOID_BLOCK_SIZE,
num_warps=_BETA_SIGMOID_NUM_WARPS,
)
return y
@dispatch('common')
def fused_beta_sigmoid_bwd(x: torch.Tensor, dy: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
dx = torch.empty_like(x)
n_elements = x.numel()
grid = (triton.cdiv(n_elements, _BETA_SIGMOID_BLOCK_SIZE),)
fused_beta_sigmoid_bwd_kernel[grid](
x,
dy,
dx,
scale,
n_elements,
BLOCK_SIZE=_BETA_SIGMOID_BLOCK_SIZE,
num_warps=_BETA_SIGMOID_NUM_WARPS,
)
return dx
class BetaSigmoidFunction(torch.autograd.Function):
@staticmethod
@input_guard
@autocast_custom_fwd
def forward(ctx, x: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
y = fused_beta_sigmoid_fwd(x, scale)
ctx.save_for_backward(x)
ctx.scale = scale
return y
@staticmethod
@input_guard
@autocast_custom_bwd
def backward(ctx, dy: torch.Tensor):
(x,) = ctx.saved_tensors
dx = fused_beta_sigmoid_bwd(x, dy, ctx.scale)
return dx.type_as(x), None
def fused_beta_sigmoid(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
return BetaSigmoidFunction.apply(x, scale)
+13
View File
@@ -0,0 +1,13 @@
"""Context-parallel stubs. Pass ``cp_context=None`` (the default)."""
class FLACPContext:
cu_seqlens = None
cu_seqlens_cpu = None
def build_cp_context(*args, **kwargs):
raise RuntimeError("Context parallel is not included in the vendored KDA kernels")
__all__ = ["FLACPContext", "build_cp_context"]
+11
View File
@@ -0,0 +1,11 @@
"""Context-parallel hooks referenced by KDA fwd/bwd. Not implemented here."""
def _cp_unsupported(*args, **kwargs):
raise RuntimeError("Context parallel is not included in the vendored KDA kernels")
chunk_gated_delta_rule_fwd_h_pre_process = _cp_unsupported
compress_h0 = _cp_unsupported
chunk_gated_delta_rule_bwd_dhu_pre_process = _cp_unsupported
expand_h0 = _cp_unsupported
+1
View File
@@ -0,0 +1 @@
# Vendored GLA chunk output kernel used by KDA.
File diff suppressed because it is too large Load Diff
+7
View File
@@ -0,0 +1,7 @@
from .chunk import chunk_kda
from .fused_recurrent import fused_recurrent_kda
__all__ = [
"chunk_kda",
"fused_recurrent_kda",
]
+443
View File
@@ -0,0 +1,443 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
# Related files are modified and supported by the Moonshot AI Team
import warnings
import torch
from kda._fla.modules.l2norm import l2norm_bwd, l2norm_fwd
from kda._fla.ops.backends import dispatch
from kda._fla.ops.common.gate import fused_beta_sigmoid, fused_beta_sigmoid_bwd
from kda._fla.ops.cp import FLACPContext
from kda._fla.ops.kda.chunk_bwd import chunk_kda_bwd
from kda._fla.ops.kda.chunk_fwd import chunk_kda_fwd
from kda._fla.ops.utils.index import prepare_chunk_indices
from kda._fla.utils import autocast_custom_bwd, autocast_custom_fwd, input_guard
class ChunkKDAFunction(torch.autograd.Function):
@staticmethod
@input_guard
@autocast_custom_fwd
def forward(
ctx,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
A_log: torch.Tensor,
dt_bias: torch.Tensor,
scale: float,
initial_state: torch.Tensor,
output_final_state: bool = False,
use_qk_l2norm_in_kernel: bool = False,
use_gate_in_kernel: bool = False,
use_beta_sigmoid_in_kernel: bool = False,
allow_neg_eigval: bool = False,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
cu_seqlens_cpu: torch.LongTensor | None = None,
safe_gate: bool = False,
lower_bound: float | None = None,
chunk_size: int = 64,
disable_recompute: bool = False,
return_intermediate_states: bool = False,
cp_context: FLACPContext | None = None,
):
# Apply l2norm
q_rstd, k_rstd = None, None
if use_qk_l2norm_in_kernel:
q, q_rstd = l2norm_fwd(q)
k, k_rstd = l2norm_fwd(k)
beta_raw = beta
if use_beta_sigmoid_in_kernel:
beta = fused_beta_sigmoid(beta_raw, scale=2.0 if allow_neg_eigval else 1.0)
chunk_indices = None
if cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(
cu_seqlens,
chunk_size,
cu_seqlens_cpu=cu_seqlens_cpu,
)
g_input = g
(o, final_state, g_cumsum, Aqk, Akk, w, u, qg, kg, v_new, h, initial_state) = chunk_kda_fwd(
q=q,
k=k,
v=v,
g=g_input,
beta=beta,
scale=scale,
initial_state=initial_state,
output_final_state=output_final_state,
cu_seqlens=cu_seqlens,
cu_seqlens_cpu=cu_seqlens_cpu,
chunk_indices=chunk_indices,
safe_gate=safe_gate,
lower_bound=lower_bound,
use_gate_in_kernel=use_gate_in_kernel,
A_log=A_log,
dt_bias=dt_bias,
chunk_size=chunk_size,
disable_recompute=disable_recompute,
return_intermediate_states=return_intermediate_states,
cp_context=cp_context,
state_v_first=state_v_first,
)
if return_intermediate_states:
assert torch.is_inference_mode_enabled(), "return_intermediate_states is only allowed in inference mode"
assert disable_recompute is False, "return_intermediate_states must be used with disable_recompute=False"
return o.type_as(q), final_state, h
ctx.save_for_backward(
q, q_rstd, k, k_rstd, v, g_cumsum, g_input, beta_raw, beta, A_log, dt_bias, Aqk, Akk,
w, u, qg, kg, v_new, h,
initial_state, cu_seqlens, chunk_indices
)
ctx.chunk_size = chunk_size
ctx.safe_gate = safe_gate
ctx.scale = scale
ctx.lower_bound = lower_bound
ctx.use_qk_l2norm_in_kernel = use_qk_l2norm_in_kernel
ctx.use_gate_in_kernel = use_gate_in_kernel
ctx.use_beta_sigmoid_in_kernel = use_beta_sigmoid_in_kernel
ctx.allow_neg_eigval = allow_neg_eigval
ctx.disable_recompute = disable_recompute
ctx.cp_context = cp_context
ctx.state_v_first = state_v_first
return o.type_as(q), final_state
@staticmethod
@input_guard
@autocast_custom_bwd
def backward(
ctx,
do: torch.Tensor,
dht: torch.Tensor,
):
(q, q_rstd, k, k_rstd, v, g_cumsum, g_input, beta_raw, beta, A_log, dt_bias, Aqk, Akk,
w, u, qg, kg, v_new, h,
initial_state, cu_seqlens, chunk_indices) = (
ctx.saved_tensors
)
dq, dk, dv, db, dg, dh0, dA, dbias = chunk_kda_bwd(
q=q,
k=k,
v=v,
beta=beta,
Aqk=Aqk,
Akk=Akk,
scale=ctx.scale,
initial_state=initial_state,
do=do,
dht=dht,
g=g_cumsum,
g_org=g_input if ctx.use_gate_in_kernel else None,
state_v_first=ctx.state_v_first,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
chunk_size=ctx.chunk_size,
safe_gate=ctx.safe_gate,
lower_bound=ctx.lower_bound,
use_gate_in_kernel=ctx.use_gate_in_kernel,
A_log=A_log,
dt_bias=dt_bias,
disable_recompute=ctx.disable_recompute,
cp_context=ctx.cp_context,
w=w,
u=u,
qg=qg,
kg=kg,
v_new=v_new,
h=h,
)
if ctx.use_qk_l2norm_in_kernel:
dq = l2norm_bwd(q, q_rstd, dq)
dk = l2norm_bwd(k, k_rstd, dk)
if ctx.use_beta_sigmoid_in_kernel:
db = fused_beta_sigmoid_bwd(beta_raw, db, scale=2.0 if ctx.allow_neg_eigval else 1.0)
return (dq.to(q), dk.to(k), dv.to(v), dg.to(g_input), db.to(beta_raw), dA, dbias, None, dh0,
None, None, None, None, None, None, None, None, None, None, None, None, None, None)
@dispatch('kda')
@torch.compiler.disable
def chunk_kda(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
scale: float | None = None,
initial_state: torch.Tensor | None = None,
output_final_state: bool = False,
use_qk_l2norm_in_kernel: bool = False,
use_gate_in_kernel: bool = False,
use_beta_sigmoid_in_kernel: bool = False,
allow_neg_eigval: bool = False,
safe_gate: bool = False,
lower_bound: float | None = None,
disable_recompute: bool = False,
return_intermediate_states: bool = False,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
cu_seqlens_cpu: torch.LongTensor | None = None,
cp_context: FLACPContext = None,
**kwargs,
):
r"""
Args:
q (torch.Tensor):
queries of shape ``[B, T, H, K]``.
k (torch.Tensor):
keys of shape ``[B, T, H, K]``.
v (torch.Tensor):
values of shape ``[B, T, HV, V]``.
GVA (Grouped Value Attention) is applied if ``HV > H``, where ``HV`` must be divisible by ``H``.
g (torch.Tensor):
(forget) gating tensor (in log space!) of shape ``[B, T, HV, K]``.
When ``use_gate_in_kernel=False`` (default), ``g`` should be the pre-computed decay value.
When ``use_gate_in_kernel=True``, ``g`` is the raw input before gate activation;
the kernel fuses ``-exp(A_log) * softplus(g + dt_bias)`` + chunk cumsum internally.
beta (torch.Tensor):
betas of shape ``[B, T, HV]``.
scale (Optional[float]):
Scale factor for the KDA attention scores.
If not provided, it will default to ``1 / sqrt(K)``. Default: ``None``.
initial_state (Optional[torch.Tensor]):
Initial state of shape ``[N, HV, K, V]`` for ``N`` input sequences.
For equal-length input sequences, ``N`` equals the batch size ``B``.
Default: ``None``.
output_final_state (Optional[bool]):
Whether to output the final state of shape ``[N, HV, K, V]``. Default: ``False``.
use_qk_l2norm_in_kernel (bool):
Whether to apply L2norm to the q,k tensor internally. Default: ``False``.
use_gate_in_kernel (bool):
Whether to compute the log-space KDA decay internally.
- If ``True``:
The passed ``g`` acts as the raw input for ``-exp(A_log) * softplus(g + dt_bias.view(HV, K))``.
Note that as part of the input arguments,
``A_log`` (shape ``[HV]``) and the optional ``dt_bias`` (shape ``[HV * K]``) should be provided.
When ``lower_bound`` is set, ``A_log`` may be ``None``,
in which case the gate is ``lower_bound * sigmoid(g + dt_bias)``.
- If ``False``, ``g`` is expected to be the pre-computed decay value.
Default: ``False``.
use_beta_sigmoid_in_kernel (bool):
Whether to apply ``torch.sigmoid(beta)`` before launching the chunk kernel.
- If ``True``, the passed ``beta`` acts as the raw beta logits.
- If ``False``, ``beta`` is expected to already be in post-sigmoid space.
Default: ``False``.
allow_neg_eigval (bool):
Whether to allow negative eigenvalues by scaling ``beta`` to ``[0, 2)``.
Only takes effect together with ``use_beta_sigmoid_in_kernel=True``, in which case
the kernel computes ``2 * sigmoid(beta)`` instead of ``sigmoid(beta)``.
Default: ``False``.
safe_gate (bool):
Whether to clamp the gate to ``[lower_bound, 0)`` and enable M=16 TensorCore
acceleration for higher throughput. Requires ``lower_bound`` to be set.
Default: ``False``.
lower_bound (Optional[float]):
Lower bound for the forget gate (in log space). When set together with
``safe_gate=True``, changes the gate activation from
``-exp(A_log) * softplus(g + dt_bias)`` to
``lower_bound * sigmoid(exp(A_log) * (g + dt_bias))``,
which naturally clamps the output to ``[lower_bound, 0)``.
Recommended value: ``-5`` (i.e., per-step decay ``exp(-5) ≈ 0.0067``).
Default: ``None``.
disable_recompute (bool):
Whether to disable gradient recomputation in the kernel. When ``True``, the kernel
will save all intermediate activations for backward pass, which is beneficial
for training small models at the cost of increased memory usage. Default: ``False``.
return_intermediate_states (bool):
If True, returns intermediate state ``h`` for inference scenarios (e.g., vLLM).
Must be used within ``torch.inference_mode()`` and will return a 3-tuple instead of 2-tuple.
This is not intended for training as it bypasses autograd. Default: ``False``.
state_v_first (Optional[bool]):
Store the recurrent state in V-first ``[V, K]`` layout instead of the default ``[K, V]``. Default: ``False``.
cu_seqlens (torch.LongTensor):
Cumulative sequence lengths of shape ``[N+1]`` used for variable-length training,
consistent with the FlashAttention API.
cu_seqlens_cpu (torch.LongTensor):
Cumulative sequence lengths of shape ``[N+1]`` used for variable-length training,
consistent with the FlashAttention API.
cp_context (Optional[FLACPContext]):
Context parallel context for distributed training across multiple devices.
When provided, ``initial_state`` and ``output_final_state`` are not supported,
and ``cu_seqlens`` will be overridden by the context. Default: ``None``.
Returns:
- Normal mode (return_intermediate_states=False): A tuple (o, final_state)
o (torch.Tensor):
Outputs of shape ``[B, T, HV, V]``.
final_state (torch.Tensor):
Final state of shape ``[N, HV, K, V]`` if ``output_final_state=True`` else ``None``.
- Inference mode (return_intermediate_states=True): A tuple (o, final_state, h)
o (torch.Tensor):
Outputs of shape ``[B, T, HV, V]``.
final_state (torch.Tensor):
Final state of shape ``[N, HV, K, V]`` if ``output_final_state=True`` else ``None``.
h (torch.Tensor):
Intermediate states of shape ``[B, NT, HV, K, V]`` and dtype ``bfloat16``.
- For equal-length sequences: ``NT = ceil(T / chunk_size)``
- For variable-length sequences (cu_seqlens): B is always 1 (flattened),
NT is the total number of chunks across all sequences.
Examples::
>>> import torch
>>> import torch.nn.functional as F
>>> from einops import rearrange
>>> from fla.ops.kda import chunk_kda
# inputs with equal lengths (no GVA, HV == H)
>>> B, T, H, K, V = 4, 2048, 4, 512, 512
>>> q = torch.randn(B, T, H, K, dtype=torch.bfloat16, device='cuda')
>>> k = torch.randn(B, T, H, K, dtype=torch.bfloat16, device='cuda')
>>> v = torch.randn(B, T, H, V, dtype=torch.bfloat16, device='cuda')
>>> beta = torch.rand(B, T, H, dtype=torch.bfloat16, device='cuda')
>>> g = torch.rand(B, T, H, K, dtype=torch.bfloat16, device='cuda')
>>> h0 = torch.randn(B, H, K, V, dtype=torch.bfloat16, device='cuda')
>>> A_log = torch.randn(H, dtype=torch.float32, device='cuda')
>>> dt_bias = torch.randn(H * K, dtype=torch.float32, device='cuda')
>>> o, ht = chunk_kda(
q, k, v, g, beta,
A_log=A_log,
dt_bias=dt_bias,
use_qk_l2norm_in_kernel=True,
use_gate_in_kernel=True,
initial_state=h0,
output_final_state=True
)
# GVA mode (HV > H)
>>> HV = 8 # 2x more value heads than qk heads
>>> v = torch.randn(B, T, HV, V, dtype=torch.bfloat16, device='cuda')
>>> g = torch.rand(B, T, HV, K, dtype=torch.bfloat16, device='cuda')
>>> beta = torch.rand(B, T, HV, dtype=torch.bfloat16, device='cuda')
>>> h0 = torch.randn(B, HV, K, V, dtype=torch.bfloat16, device='cuda')
>>> A_log = torch.randn(HV, dtype=torch.float32, device='cuda')
>>> dt_bias = torch.randn(HV * K, dtype=torch.float32, device='cuda')
>>> o, ht = chunk_kda(
q, k, v, g, beta,
A_log=A_log,
dt_bias=dt_bias,
use_qk_l2norm_in_kernel=True,
use_gate_in_kernel=True,
initial_state=h0,
output_final_state=True
)
# for variable-length inputs, the batch size `B` is expected to be 1 and `cu_seqlens` is required
>>> q, k, v, beta, g = map(lambda x: rearrange(x, 'b t ... -> 1 (b t) ...'), (q, k, v, beta, g))
# for a batch with 4 sequences, `cu_seqlens` with 5 start/end positions are expected
>>> cu_seqlens = q.new_tensor([0, 2048, 4096, 6144, 8192], dtype=torch.long)
>>> o, ht = chunk_kda(
q, k, v, g, beta,
A_log=A_log,
dt_bias=dt_bias,
use_qk_l2norm_in_kernel=True,
use_gate_in_kernel=True,
initial_state=h0,
output_final_state=True,
cu_seqlens=cu_seqlens
)
"""
if 'transpose_state_layout' in kwargs:
if state_v_first:
raise ValueError("Cannot pass both `state_v_first` and the deprecated `transpose_state_layout`.")
warnings.warn(
"`transpose_state_layout` is deprecated and renamed to `state_v_first`.",
DeprecationWarning,
stacklevel=2,
)
state_v_first = kwargs.pop('transpose_state_layout')
if cp_context is not None:
assert initial_state is None, "Initial state is not supported for CP"
assert output_final_state is False, "Output final state is not supported for CP"
assert cp_context.cu_seqlens is not None, "cu_seqlens is required for CP"
# Override cu_seqlens and cu_seqlens_cpu with the ones from the context
cu_seqlens = cp_context.cu_seqlens
if cp_context.cu_seqlens_cpu is not None:
cu_seqlens_cpu = cp_context.cu_seqlens_cpu
if cu_seqlens is not None:
if q.shape[0] != 1:
raise ValueError(
f"The batch size is expected to be 1 rather than {q.shape[0]} when using `cu_seqlens`."
f"Please flatten variable-length inputs before processing.",
)
if initial_state is not None and initial_state.shape[0] != len(cu_seqlens) - 1:
raise ValueError(
f"The number of initial states is expected to be equal to the number of input sequences, "
f"i.e., {len(cu_seqlens) - 1} rather than {initial_state.shape[0]}.",
)
if initial_state is not None:
assert initial_state.dtype == torch.float32, "initial_state must be in float32."
A_log, dt_bias = None, None
if use_gate_in_kernel:
A_log, dt_bias = kwargs.get("A_log"), kwargs.get("dt_bias")
if A_log is None and lower_bound is None:
raise ValueError("`A_log` must be provided when `use_gate_in_kernel=True` and `lower_bound` is not set.")
chunk_size = kwargs.pop("chunk_size", 64)
if chunk_size not in (32, 64):
raise ValueError(f"`chunk_size` must be either 32 or 64 for KDA, got {chunk_size}.")
if safe_gate and use_gate_in_kernel:
if lower_bound is None:
raise ValueError("`lower_bound` must be specified when `safe_gate=True` and `use_gate_in_kernel=True`.")
if not (-5 <= lower_bound < 0):
raise ValueError(f"`lower_bound` must be in the safe range [-5, 0), got {lower_bound}.")
if allow_neg_eigval and not use_beta_sigmoid_in_kernel:
raise ValueError("`allow_neg_eigval=True` requires `use_beta_sigmoid_in_kernel=True`.")
# Validate head dimensions for GVA
B, T, H, K, HV = *q.shape, v.shape[2]
assert q.shape == k.shape, f"q and k must have the same shape, got q={q.shape} vs k={k.shape}"
assert K <= 256, f"Currently we only support key headdim <=256 for KDA, got {K}."
assert HV % H == 0, (
f"For GVA, num_v_heads (HV={HV}) must be evenly divisible by num_qk_heads (H={H}), "
f"but got HV % H = {HV % H}"
)
assert g.shape == (B, T, HV, K), f"g must have shape [B, T, HV, K]={[B, T, HV, K]}, got {list(g.shape)}"
assert beta.shape == (B, T, HV), f"beta must have shape [B, T, HV]={[B, T, HV]}, got {list(beta.shape)}"
if scale is None:
scale = K ** -0.5
return ChunkKDAFunction.apply(
q,
k,
v,
g,
beta,
A_log,
dt_bias,
scale,
initial_state,
output_final_state,
use_qk_l2norm_in_kernel,
use_gate_in_kernel,
use_beta_sigmoid_in_kernel,
allow_neg_eigval,
state_v_first,
cu_seqlens,
cu_seqlens_cpu,
safe_gate,
lower_bound,
chunk_size,
disable_recompute,
return_intermediate_states,
cp_context,
)
+651
View File
@@ -0,0 +1,651 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import torch
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.ops.common.chunk_delta_h import (
chunk_gated_delta_rule_bwd_dhu,
chunk_gated_delta_rule_fwd_h,
)
from kda._fla.ops.cp import FLACPContext
from kda._fla.ops.cp.chunk_delta_h import (
chunk_gated_delta_rule_bwd_dhu_pre_process,
expand_h0,
)
from kda._fla.ops.kda.chunk_intra import chunk_kda_bwd_intra
from kda._fla.ops.kda.gate import kda_gate_bwd, kda_gate_chunk_cumsum
from kda._fla.ops.kda.wy_fast import recompute_w_u_fwd
from kda._fla.ops.utils import chunk_local_cumsum, prepare_chunk_indices
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.ops.utils.constant import RCP_LN2
from kda._fla.ops.utils.op import exp2
from kda._fla.utils import (
IS_NVIDIA_HOPPER,
IS_NVIDIA_SM100,
autotune_cache_kwargs,
check_shared_mem,
)
BK_LIST = [32, 64] if check_shared_mem() else [16, 32]
BV_LIST = [64, 128] if check_shared_mem("ampere") else [16, 32]
NUM_WARPS = [2, 4] if IS_NVIDIA_HOPPER else [2, 4, 8]
@triton.heuristics(
{
"IS_VARLEN": lambda args: args["cu_seqlens"] is not None,
}
)
@fla_cache_autotune(
configs=[
triton.Config({}, num_warps=num_warps, num_stages=num_stages)
for num_warps in NUM_WARPS
for num_stages in [2, 3, 4]
],
key=["H", "HV", "K", "V", "BT", "BK", "BV"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=["T"])
def chunk_kda_bwd_kernel_dAv(
q,
k,
v,
A,
do,
dv,
dA,
cu_seqlens,
chunk_indices,
scale,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BK: tl.constexpr,
BV: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_t, i_bh = tl.program_id(0).to(tl.int64), tl.program_id(1).to(tl.int64)
i_b, i_hv = i_bh // HV, i_bh % HV
i_h = i_hv // (HV // H)
if IS_VARLEN:
i_n, i_t = (
tl.load(chunk_indices + i_t * 2).to(tl.int32),
tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64),
)
bos, eos = (
tl.load(cu_seqlens + i_n).to(tl.int64),
tl.load(cu_seqlens + i_n + 1).to(tl.int64),
)
T = eos - bos
else:
bos, eos = i_b * T, i_b * T + T
# offset calculation
q += (bos * H + i_h) * K
k += (bos * H + i_h) * K
v += (bos * HV + i_hv) * V
do += (bos * HV + i_hv) * V
dv += (bos * HV + i_hv) * V
dA += (bos * HV + i_hv) * BT
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
o_A = tl.arange(0, BT)
m_AT = (o_A[:, None] < BT) & m_t[None, :]
p_A = A + (bos * HV + i_hv) * BT + o_A[:, None] + o_t[None, :] * (HV * BT)
b_A = tl.load(p_A, mask=m_AT, other=0.0)
m_A = (o_t[:, None] <= o_t[None, :]) & (m_t[:, None] & m_t)
b_A = tl.where(m_A, b_A, 0).to(do.dtype.element_ty)
b_dA = tl.zeros([BT, BT], dtype=tl.float32)
for i_v in range(tl.cdiv(V, BV)):
o_v = i_v * BV + tl.arange(0, BV)
m_v = o_v < V
m_vT = m_v[:, None] & m_t[None, :]
m_tv = m_t[:, None] & m_v[None, :]
p_v = v + o_v[:, None] + o_t[None, :] * (HV * V)
p_do = do + o_t[:, None] * (HV * V) + o_v[None, :]
p_dv = dv + o_t[:, None] * (HV * V) + o_v[None, :]
# [BV, BT]
b_v = tl.load(p_v, mask=m_vT, other=0.0)
# [BT, BV]
b_do = tl.load(p_do, mask=m_tv, other=0.0)
# [BT, BT]
b_dA += tl.dot(b_do, b_v)
# [BT, BV]
b_dv = tl.dot(b_A.to(b_do.dtype), b_do)
tl.store(p_dv, b_dv.to(p_dv.dtype.element_ty), mask=m_tv)
m_dA = m_t[:, None] & (o_A[None, :] < BT)
p_dA = dA + o_t[:, None] * (HV * BT) + o_A[None, :]
b_dA = tl.where(o_t[:, None] >= o_t, b_dA * scale, 0.0)
tl.store(p_dA, b_dA.to(p_dA.dtype.element_ty), mask=m_dA)
@triton.heuristics(
{
"IS_VARLEN": lambda args: args["cu_seqlens"] is not None,
}
)
@fla_cache_autotune(
configs=[
triton.Config({"BK": BK, "BV": BV}, num_warps=num_warps, num_stages=num_stages)
for BK in BK_LIST
for BV in BV_LIST
for num_warps in NUM_WARPS
for num_stages in [2, 3, 4]
if not (IS_NVIDIA_HOPPER and BK == 32 and num_warps == 4)
if not (IS_NVIDIA_SM100 and BK == 32 and num_warps != 2)
],
key=["BT", "HV", "STATE_V_FIRST"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=["T"])
def chunk_kda_bwd_kernel_wy_dqkg_fused(
q,
k,
v,
v_new,
g,
beta,
A,
h,
do,
dh,
dq,
dk,
dv,
dv2,
dg,
db,
dA,
cu_seqlens,
chunk_indices,
scale,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BK: tl.constexpr,
BV: tl.constexpr,
STATE_V_FIRST: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_t, i_bh = tl.program_id(0).to(tl.int64), tl.program_id(1)
i_b, i_hv = i_bh // HV, i_bh % HV
i_h = i_hv // (HV // H)
if IS_VARLEN:
i_tg = i_t.to(tl.int64)
i_n, i_t = (
tl.load(chunk_indices + i_t * 2).to(tl.int32),
tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64),
)
bos, eos = (
tl.load(cu_seqlens + i_n).to(tl.int64),
tl.load(cu_seqlens + i_n + 1).to(tl.int64),
)
T = (eos - bos).to(tl.int32)
NT = tl.cdiv(T, BT)
else:
NT = tl.cdiv(T, BT)
i_tg = (i_b * NT + i_t).to(tl.int64)
bos, eos = (i_b * T).to(tl.int64), (i_b * T + T).to(tl.int64)
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
m_last = o_t == min(T, i_t * BT + BT) - 1
q += (bos * H + i_h) * K
k += (bos * H + i_h) * K
v += (bos * HV + i_hv) * V
v_new += (bos * HV + i_hv) * V
g += (bos * HV + i_hv) * K
beta += bos * HV + i_hv
A += (bos * HV + i_hv) * BT
h += (i_tg * HV + i_hv) * K * V
do += (bos * HV + i_hv) * V
dh += (i_tg * HV + i_hv) * K * V
dq += (bos * HV + i_hv) * K
dk += (bos * HV + i_hv) * K
dv += (bos * HV + i_hv) * V
dv2 += (bos * HV + i_hv) * V
dg += (bos * HV + i_hv) * K
db += bos * HV + i_hv
dA += (bos * HV + i_hv) * BT
p_beta = beta + o_t * HV
b_beta = tl.load(p_beta, mask=m_t, other=0.0)
o_A = tl.arange(0, BT)
m_AT = (o_A[:, None] < BT) & m_t[None, :]
p_A = A + o_A[:, None] + o_t[None, :] * (HV * BT)
b_A = tl.load(p_A, mask=m_AT, other=0.0)
b_dA = tl.zeros([BT, BT], dtype=tl.float32)
b_db = tl.zeros([BT], dtype=tl.float32)
for i_k in range(tl.cdiv(K, BK)):
o_k = i_k * BK + tl.arange(0, BK)
m_k = o_k < K
m_tk = m_t[:, None] & m_k[None, :]
p_k = k + o_t[:, None] * (H * K) + o_k[None, :]
p_g = g + o_t[:, None] * (HV * K) + o_k[None, :]
b_k = tl.load(p_k, mask=m_tk, other=0.0)
b_g = tl.load(p_g, mask=m_tk, other=0.0).to(tl.float32)
p_gn = g + (min(T, i_t * BT + BT) - 1).to(tl.int64) * HV * K + o_k
b_gn = tl.load(p_gn, mask=m_k, other=0).to(tl.float32)
b_dq = tl.zeros([BT, BK], dtype=tl.float32)
b_dk = tl.zeros([BT, BK], dtype=tl.float32)
b_dw = tl.zeros([BT, BK], dtype=tl.float32)
b_dgk = tl.zeros([BK], dtype=tl.float32)
for i_v in range(tl.cdiv(V, BV)):
o_v = i_v * BV + tl.arange(0, BV)
m_tv = m_t[:, None] & (o_v[None, :] < V)
m_h = (o_v[:, None] < V) & m_k[None, :]
p_v_new = v_new + o_t[:, None] * (HV * V) + o_v[None, :]
p_do = do + o_t[:, None] * (HV * V) + o_v[None, :]
if STATE_V_FIRST:
p_h = h + o_v[:, None] * K + o_k[None, :]
p_dh = dh + o_v[:, None] * K + o_k[None, :]
else:
p_h = h + o_v[:, None] + o_k[None, :] * V
p_dh = dh + o_v[:, None] + o_k[None, :] * V
p_dv = dv + o_t[:, None] * (HV * V) + o_v[None, :]
# [BT, BV]
b_v_new = tl.load(p_v_new, mask=m_tv, other=0.0)
b_do = tl.load(p_do, mask=m_tv, other=0.0)
# [BV, BK]
b_h = tl.load(p_h, mask=m_h, other=0.0)
b_dh = tl.load(p_dh, mask=m_h, other=0.0)
# [BT, BV]
b_dv = tl.load(p_dv, mask=m_tv, other=0.0)
b_dgk += tl.sum(b_h * b_dh, axis=0)
b_dq += tl.dot(b_do, b_h.to(b_do.dtype))
b_dk += tl.dot(b_v_new, b_dh.to(b_v_new.dtype))
b_dw += tl.dot(b_dv.to(b_v_new.dtype), b_h.to(b_v_new.dtype))
tl.debug_barrier() # DO NOT REMOVE THIS LINE!
if i_k == 0:
p_v = v + o_t[:, None] * (HV * V) + o_v[None, :]
p_dv2 = dv2 + o_t[:, None] * (HV * V) + o_v[None, :]
b_v = tl.load(p_v, mask=m_tv, other=0.0)
b_dA += tl.dot(b_dv, tl.trans(b_v))
b_dvb = tl.dot(b_A, b_dv)
b_dv2 = b_dvb * b_beta[:, None]
b_db += tl.sum(b_dvb * b_v, 1)
tl.store(p_dv2, b_dv2.to(p_dv2.dtype.element_ty), mask=m_tv)
b_gk_exp = exp2(b_g)
b_gb = b_gk_exp * b_beta[:, None]
b_dgk *= exp2(b_gn)
b_dq = b_dq * b_gk_exp * scale
b_dk = b_dk * tl.where(m_t[:, None], exp2(b_gn[None, :] - b_g), 0)
b_kg = b_k * b_gk_exp
b_dw = -b_dw.to(b_A.dtype)
b_dA += tl.dot(b_dw, tl.trans(b_kg.to(b_A.dtype)))
b_dkgb = tl.dot(b_A, b_dw)
b_db += tl.sum(b_dkgb * b_kg, 1)
p_q = q + o_t[:, None] * (H * K) + o_k[None, :]
b_q = tl.load(p_q, mask=m_tk, other=0.0)
b_kdk = b_k * b_dk
b_dgk += tl.sum(b_kdk, axis=0)
b_dg = (
b_q * b_dq
- b_kdk
+ m_last[:, None] * b_dgk
+ b_kg * b_dkgb * b_beta[:, None]
)
b_dk = b_dk + b_dkgb * b_gb
p_dq = dq + o_t[:, None] * (HV * K) + o_k[None, :]
p_dk = dk + o_t[:, None] * (HV * K) + o_k[None, :]
p_dg = dg + o_t[:, None] * (HV * K) + o_k[None, :]
tl.store(p_dq, b_dq.to(p_dq.dtype.element_ty), mask=m_tk)
tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), mask=m_tk)
tl.store(p_dg, b_dg.to(p_dg.dtype.element_ty), mask=m_tk)
m_A = (o_t[:, None] > o_t[None, :]) & (m_t[:, None] & m_t)
b_dA = tl.where(m_A, b_dA * b_beta[None, :], 0)
b_dA = tl.dot(b_dA.to(b_A.dtype), b_A)
b_dA = tl.dot(b_A, b_dA.to(b_A.dtype))
b_dA = tl.where(m_A, -b_dA, 0)
m_dA = m_t[:, None] & (o_A[None, :] < BT)
p_dA = dA + o_t[:, None] * (HV * BT) + o_A[None, :]
p_db = db + o_t * HV
tl.store(p_dA, b_dA.to(p_dA.dtype.element_ty), mask=m_dA)
tl.store(p_db, b_db.to(p_db.dtype.element_ty), mask=m_t)
@dispatch("kda")
def chunk_kda_bwd_dAv(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
do: torch.Tensor,
A: torch.Tensor | None = None,
scale: float = None,
cu_seqlens: torch.LongTensor | None = None,
chunk_size: int = 64,
chunk_indices: torch.LongTensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
B, T, H, K, HV, V = *k.shape, do.shape[2], do.shape[-1]
BT = chunk_size
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
# H100 can have larger block size
if check_shared_mem("hopper", k.device.index):
CONST_TILING = 128
elif check_shared_mem:
CONST_TILING = 64
else:
CONST_TILING = 32
BK = min(max(triton.next_power_of_2(K), 16), CONST_TILING)
BV = min(max(triton.next_power_of_2(V), 16), CONST_TILING)
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
dA = v.new_empty(B, T, HV, BT, dtype=torch.float)
dv = torch.empty_like(do)
grid = (NT, B * HV)
chunk_kda_bwd_kernel_dAv[grid](
q=q,
k=k,
v=v,
A=A,
do=do,
dv=dv,
dA=dA,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
scale=scale,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BT=BT,
BK=BK,
BV=BV,
)
return dA, dv
@dispatch("kda")
def chunk_kda_bwd_wy_dqkg_fused(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
v_new: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
A: torch.Tensor,
h: torch.Tensor,
do: torch.Tensor,
dh: torch.Tensor,
dv: torch.Tensor,
scale: float | None = None,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
chunk_size: int = 64,
chunk_indices: torch.LongTensor | None = None,
):
B, T, H, K, HV, V = *k.shape, v.shape[2], v.shape[-1]
BT = chunk_size
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
# dq, dk are allocated at HV dimension; caller reduces to H if GVA
dq = g.new_empty(B, T, HV, K, dtype=torch.float)
dk = g.new_empty(B, T, HV, K, dtype=torch.float)
dv2 = torch.empty_like(v)
dg = torch.empty_like(g, dtype=torch.float)
db = torch.empty_like(beta, dtype=torch.float)
dA = torch.empty_like(A, dtype=torch.float)
grid = (NT, B * HV)
chunk_kda_bwd_kernel_wy_dqkg_fused[grid](
q=q,
k=k,
v=v,
v_new=v_new,
g=g,
beta=beta,
A=A,
h=h,
do=do,
dh=dh,
dq=dq,
dk=dk,
dv=dv,
dv2=dv2,
dg=dg,
db=db,
dA=dA,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
scale=scale,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BT=BT,
STATE_V_FIRST=state_v_first,
)
dv = dv2
return dq, dk, dv, db, dg, dA
def chunk_kda_bwd(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
beta: torch.Tensor,
Aqk: torch.Tensor,
Akk: torch.Tensor,
scale: float,
initial_state: torch.Tensor,
do: torch.Tensor,
dht: torch.Tensor,
g: torch.Tensor | None = None,
g_org: torch.Tensor | None = None,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
chunk_indices: torch.LongTensor | None = None,
chunk_size: int = 64,
safe_gate: bool = False,
lower_bound: float | None = None,
use_gate_in_kernel: bool = False,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
disable_recompute: bool = False,
cp_context: FLACPContext | None = None,
**kwargs,
):
H, HV = q.shape[2], v.shape[2]
G = HV // H
if disable_recompute is False:
if use_gate_in_kernel:
g = kda_gate_chunk_cumsum(
g=g_org,
A_log=A_log,
dt_bias=dt_bias,
scale=RCP_LN2,
chunk_size=chunk_size,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
lower_bound=lower_bound,
)
w, u, qg, kg = recompute_w_u_fwd(
q=q,
k=k,
v=v,
beta=beta,
A=Akk,
gk=g,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
)
if cp_context is not None:
# Restore the full initial_state tensor from the compressed version.
# Only the first sequence's state is non-zero as it's the only one that could be cross-rank.
initial_state = expand_h0(initial_state, context=cp_context)
h, v_new, _ = chunk_gated_delta_rule_fwd_h(
k=kg,
w=w,
u=u,
gk=g,
initial_state=initial_state,
output_final_state=False,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
chunk_size=chunk_size,
state_v_first=state_v_first,
)
else:
w, u, qg, kg, v_new, h = (
kwargs["w"],
kwargs["u"],
kwargs["qg"],
kwargs["kg"],
kwargs["v_new"],
kwargs["h"],
)
if cp_context is not None:
# Restore the full initial_state tensor from the compressed version.
# Only the first sequence's state is non-zero as it's the only one that could be cross-rank.
initial_state = expand_h0(initial_state, context=cp_context)
# dAqk = do @ v.T
# dv = A @ do
dAqk, dv = chunk_kda_bwd_dAv(
q=q,
k=k,
v=v_new,
do=do,
A=Aqk,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_size=chunk_size,
chunk_indices=chunk_indices,
)
if cp_context is not None:
# initial_state is None in the CP mode
# We only need to compute dht of current rank and pass it to the backward kernel
dht, initial_state = chunk_gated_delta_rule_bwd_dhu_pre_process(
q=qg,
k=kg,
w=w,
do=do,
dv=dv,
gk=g,
scale=scale,
cu_seqlens=cu_seqlens,
dht=dht,
initial_state=initial_state,
context=cp_context,
chunk_size=chunk_size,
state_v_first=state_v_first,
)
dh, dh0, dv = chunk_gated_delta_rule_bwd_dhu(
q=qg,
k=kg,
w=w,
gk=g,
h0=initial_state,
dht=dht,
do=do,
dv=dv,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_size=chunk_size,
chunk_indices=chunk_indices,
state_v_first=state_v_first,
)
dq, dk, dv, db, dg, dAkk = chunk_kda_bwd_wy_dqkg_fused(
q=q,
k=k,
v=v,
v_new=v_new,
g=g,
beta=beta,
A=Akk,
h=h,
do=do,
dh=dh,
dv=dv,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_size=chunk_size,
chunk_indices=chunk_indices,
state_v_first=state_v_first,
)
dq, dk, db, dg = chunk_kda_bwd_intra(
q=q,
k=k,
g=g,
beta=beta,
dAqk=dAqk,
dAkk=dAkk,
dq=dq,
dk=dk,
db=db,
dg=dg,
cu_seqlens=cu_seqlens,
chunk_size=chunk_size,
chunk_indices=chunk_indices,
safe_gate=safe_gate,
)
# For GVA, reduce dq and dk from [B, T, HV, K] back to [B, T, H, K]
if HV > H:
dq = dq.view(*dq.shape[:2], H, G, dq.shape[-1]).sum(dim=3)
dk = dk.view(*dk.shape[:2], H, G, dk.shape[-1]).sum(dim=3)
dA, dbias = None, None
dg = chunk_local_cumsum(
dg,
chunk_size=chunk_size,
reverse=True,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
)
if use_gate_in_kernel:
dg, dA, dbias = kda_gate_bwd(
g=g_org, A_log=A_log, dt_bias=dt_bias, dyg=dg, lower_bound=lower_bound
)
return dq, dk, dv, db, dg, dh0, dA, dbias
+134
View File
@@ -0,0 +1,134 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import torch
from kda._fla.ops.common.chunk_delta_h import chunk_gated_delta_rule_fwd_h
from kda._fla.ops.cp import FLACPContext
from kda._fla.ops.cp.chunk_delta_h import chunk_gated_delta_rule_fwd_h_pre_process, compress_h0
from kda._fla.ops.gla.chunk import chunk_gla_fwd_o_gk
from kda._fla.ops.kda.chunk_intra import chunk_kda_fwd_intra
from kda._fla.ops.kda.gate import kda_gate_chunk_cumsum
from kda._fla.ops.utils import chunk_local_cumsum
from kda._fla.ops.utils.constant import RCP_LN2
def chunk_kda_fwd(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
scale: float,
initial_state: torch.Tensor,
output_final_state: bool,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
cu_seqlens_cpu: torch.LongTensor | None = None,
chunk_indices: torch.LongTensor | None = None,
chunk_size: int = 64,
safe_gate: bool = False,
lower_bound: float | None = None,
use_gate_in_kernel: bool = False,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
disable_recompute: bool = False,
return_intermediate_states: bool = False,
cp_context: FLACPContext | None = None,
):
# Apply gate activation
g_org = None
if use_gate_in_kernel:
g_org = g
g = kda_gate_chunk_cumsum(
g=g_org,
A_log=A_log,
dt_bias=dt_bias,
scale=RCP_LN2,
chunk_size=chunk_size,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
lower_bound=lower_bound,
)
else:
g = chunk_local_cumsum(
g=g,
scale=RCP_LN2,
chunk_size=chunk_size,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices
)
# qg = None if disable_recompute is False
w, u, qg, kg, Aqk, Akk = chunk_kda_fwd_intra(
q=q,
k=k,
v=v,
gk=g,
beta=beta,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_size=chunk_size,
chunk_indices=chunk_indices,
safe_gate=safe_gate,
disable_recompute=disable_recompute
)
if cp_context is not None:
initial_state = chunk_gated_delta_rule_fwd_h_pre_process(
k=kg,
w=w,
u=u,
gk=g,
cu_seqlens=cu_seqlens,
initial_state=initial_state,
context=cp_context,
chunk_size=chunk_size,
state_v_first=state_v_first,
)
h, v_new, final_state = chunk_gated_delta_rule_fwd_h(
k=kg,
w=w,
u=u,
gk=g,
initial_state=initial_state,
output_final_state=output_final_state,
cu_seqlens=cu_seqlens,
cu_seqlens_cpu=cu_seqlens_cpu,
chunk_indices=chunk_indices,
chunk_size=chunk_size,
state_v_first=state_v_first,
)
if cp_context is not None:
# In Context Parallel (CP) mode, global initial states are not supported at the entry point.
# The `initial_state` here is computed internally via inter-rank communication.
# Since only the first sequence in the local batch can be a continuation of a cross-rank sequence,
# only the first state in the tensor is relevant. We compress it to optimize memory for `save_for_backward`.
initial_state = compress_h0(initial_state, context=cp_context)
o = chunk_gla_fwd_o_gk(
q=q,
v=v_new,
g=g,
A=Aqk,
h=h,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_size=chunk_size,
chunk_indices=chunk_indices,
state_v_first=state_v_first,
)
if disable_recompute is False:
# Delete to save memory
w, u, qg, kg, v_new = None, None, None, None, None
if not return_intermediate_states:
h = None
if use_gate_in_kernel:
g = None
return o, final_state, g, Aqk, Akk, w, u, qg, kg, v_new, h, initial_state
+962
View File
@@ -0,0 +1,962 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import torch
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.ops.kda.chunk_intra_token_parallel import chunk_kda_fwd_intra_token_parallel
from kda._fla.ops.kda.wy_fast import recompute_w_u_fwd
from kda._fla.ops.utils import prepare_chunk_indices
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.ops.utils.op import exp2, gather
from kda._fla.utils import IS_GATHER_SUPPORTED, IS_TF32_SUPPORTED, autotune_cache_kwargs
if IS_TF32_SUPPORTED:
SOLVE_TRIL_DOT_PRECISION = tl.constexpr('tf32')
else:
SOLVE_TRIL_DOT_PRECISION = tl.constexpr('ieee')
################################################################################
# Fused inter + solve_tril kernel: compute off-diagonal Akk and solve in one pass
################################################################################
@triton.heuristics({
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({'BK': BK}, num_warps=num_warps)
for BK in [32, 64]
for num_warps in [1, 2, 4]
],
key=["H", "HV", "K", "BT", "BC", "NC"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_kda_fwd_kernel_inter_solve_fused(
q,
k,
g,
beta,
Aqk,
Akkd,
Akk,
scale,
cu_seqlens,
chunk_indices,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
BT: tl.constexpr,
BC: tl.constexpr,
NC: tl.constexpr,
BK: tl.constexpr,
IS_VARLEN: tl.constexpr,
USE_SAFE_GATE: tl.constexpr,
):
"""
Fused kernel: compute inter-subchunk Akk + solve_tril in one pass.
Prerequisite: token_parallel has already computed diagonal Akk blocks in Akkd.
This kernel:
1. Computes off-diagonal Aqk blocks -> writes to global
2. Computes off-diagonal Akk blocks -> keeps in registers
3. Loads diagonal Akk blocks from Akkd (fp32)
4. Does forward substitution on diagonals
5. Computes merged Akk_inv
6. Writes Akk_inv to Akk
"""
i_t, i_bh = tl.program_id(0).to(tl.int64), tl.program_id(1).to(tl.int64)
i_b, i_hv = i_bh // HV, i_bh % HV
i_h = i_hv // (HV // H)
if IS_VARLEN:
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
else:
bos, eos = i_b * T, i_b * T + T
if i_t * BT >= T:
return
i_tc0 = i_t * BT
i_tc1 = i_t * BT + BC
i_tc2 = i_t * BT + 2 * BC
i_tc3 = i_t * BT + 3 * BC
q += (bos * H + i_h) * K
k += (bos * H + i_h) * K
g += (bos * HV + i_hv) * K
Aqk += (bos * HV + i_hv) * BT
Akk += (bos * HV + i_hv) * BT
Akkd += (bos * HV + i_hv) * BC
o_i = tl.arange(0, BC)
m_tc1 = (i_tc1 + o_i) < T
m_tc2 = (i_tc2 + o_i) < T
m_tc3 = (i_tc3 + o_i) < T
o_c0 = i_tc0 + o_i
o_c1 = i_tc1 + o_i
o_c2 = i_tc2 + o_i
o_c3 = i_tc3 + o_i
m_tc0 = o_c0 < T
m_A0 = m_tc0[:, None] & (o_i[None, :] < BT)
m_A1 = m_tc1[:, None] & (o_i[None, :] < BT)
m_A2 = m_tc2[:, None] & (o_i[None, :] < BT)
m_A3 = m_tc3[:, None] & (o_i[None, :] < BT)
b_Aqk10 = tl.zeros([BC, BC], dtype=tl.float32)
b_Akk10 = tl.zeros([BC, BC], dtype=tl.float32)
b_Aqk20 = tl.zeros([BC, BC], dtype=tl.float32)
b_Akk20 = tl.zeros([BC, BC], dtype=tl.float32)
b_Aqk21 = tl.zeros([BC, BC], dtype=tl.float32)
b_Akk21 = tl.zeros([BC, BC], dtype=tl.float32)
b_Aqk30 = tl.zeros([BC, BC], dtype=tl.float32)
b_Akk30 = tl.zeros([BC, BC], dtype=tl.float32)
b_Aqk31 = tl.zeros([BC, BC], dtype=tl.float32)
b_Akk31 = tl.zeros([BC, BC], dtype=tl.float32)
b_Aqk32 = tl.zeros([BC, BC], dtype=tl.float32)
b_Akk32 = tl.zeros([BC, BC], dtype=tl.float32)
################################################################################
# off-diagonal blocks
################################################################################
for i_k in range(tl.cdiv(K, BK)):
o_k = i_k * BK + tl.arange(0, BK)
m_k = o_k < K
m_ck0 = m_tc0[:, None] & m_k[None, :]
p_k0 = k + o_c0[:, None] * (H*K) + o_k[None, :]
p_g0 = g + o_c0[:, None] * (HV*K) + o_k[None, :]
b_k0 = tl.load(p_k0, mask=m_ck0, other=0.0).to(tl.float32)
b_g0 = tl.load(p_g0, mask=m_ck0, other=0.0).to(tl.float32)
if i_tc1 < T:
m_ck1 = m_tc1[:, None] & m_k[None, :]
p_q1 = q + o_c1[:, None] * (H*K) + o_k[None, :]
p_k1 = k + o_c1[:, None] * (H*K) + o_k[None, :]
p_g1 = g + o_c1[:, None] * (HV*K) + o_k[None, :]
# [BC, BK]
b_q1 = tl.load(p_q1, mask=m_ck1, other=0.0).to(tl.float32)
b_k1 = tl.load(p_k1, mask=m_ck1, other=0.0).to(tl.float32)
b_g1 = tl.load(p_g1, mask=m_ck1, other=0.0).to(tl.float32)
# [BK]
b_gn1 = tl.load(g + i_tc1 * HV*K + o_k, mask=m_k, other=0).to(tl.float32)
# [BC, BK]
b_gqn = tl.where(m_tc1[:, None], exp2(b_g1 - b_gn1[None, :]), 0)
# [BK, BC]
b_kgt = tl.trans(b_k0 * exp2(b_gn1[None, :] - b_g0))
# [BC, BC]
b_Aqk10 += tl.dot(b_q1 * b_gqn, b_kgt)
b_Akk10 += tl.dot(b_k1 * b_gqn, b_kgt)
if NC >= 3 and i_tc2 < T:
m_ck2 = m_tc2[:, None] & m_k[None, :]
p_q2 = q + o_c2[:, None] * (H*K) + o_k[None, :]
p_k2 = k + o_c2[:, None] * (H*K) + o_k[None, :]
p_g2 = g + o_c2[:, None] * (HV*K) + o_k[None, :]
# [BC, BK]
b_q2 = tl.load(p_q2, mask=m_ck2, other=0.0).to(tl.float32)
b_k2 = tl.load(p_k2, mask=m_ck2, other=0.0).to(tl.float32)
b_g2 = tl.load(p_g2, mask=m_ck2, other=0.0).to(tl.float32)
# [BK]
b_gn2 = tl.load(g + i_tc2 * HV*K + o_k, mask=m_k, other=0).to(tl.float32)
# [BC, BK]
b_gqn2 = tl.where(m_tc2[:, None], exp2(b_g2 - b_gn2[None, :]), 0)
b_qg2 = b_q2 * b_gqn2
b_kg2 = b_k2 * b_gqn2
# [BK, BC]
b_kgt = tl.trans(b_k0 * exp2(b_gn2[None, :] - b_g0))
b_Aqk20 += tl.dot(b_qg2, b_kgt)
b_Akk20 += tl.dot(b_kg2, b_kgt)
# [BC, BC]
b_kgt = tl.trans(b_k1 * exp2(b_gn2[None, :] - b_g1))
# [BC, BC]
b_Aqk21 += tl.dot(b_qg2, b_kgt)
b_Akk21 += tl.dot(b_kg2, b_kgt)
if NC >= 4 and i_tc3 < T:
m_ck3 = m_tc3[:, None] & m_k[None, :]
p_q3 = q + o_c3[:, None] * (H*K) + o_k[None, :]
p_k3 = k + o_c3[:, None] * (H*K) + o_k[None, :]
p_g3 = g + o_c3[:, None] * (HV*K) + o_k[None, :]
# [BC, BK]
b_q3 = tl.load(p_q3, mask=m_ck3, other=0.0).to(tl.float32)
b_k3 = tl.load(p_k3, mask=m_ck3, other=0.0).to(tl.float32)
b_g3 = tl.load(p_g3, mask=m_ck3, other=0.0).to(tl.float32)
# [BK]
b_gn3 = tl.load(g + i_tc3 * HV*K + o_k, mask=m_k, other=0).to(tl.float32)
# [BC, BK]
b_gqn3 = tl.where(m_tc3[:, None], exp2(b_g3 - b_gn3[None, :]), 0)
b_qg3 = b_q3 * b_gqn3
b_kg3 = b_k3 * b_gqn3
# [BK, BC]
b_kgt = tl.trans(b_k0 * exp2(b_gn3[None, :] - b_g0))
# [BC, BC]
b_Aqk30 += tl.dot(b_qg3, b_kgt)
b_Akk30 += tl.dot(b_kg3, b_kgt)
# [BK, BC]
b_kgt = tl.trans(b_k1 * exp2(b_gn3[None, :] - b_g1))
# [BC, BC]
b_Aqk31 += tl.dot(b_qg3, b_kgt)
b_Akk31 += tl.dot(b_kg3, b_kgt)
# [BK, BC]
b_kgt = tl.trans(b_k2 * exp2(b_gn3[None, :] - b_g2))
# [BC, BC]
b_Aqk32 += tl.dot(b_qg3, b_kgt)
b_Akk32 += tl.dot(b_kg3, b_kgt)
################################################################################
# save off-diagonal Aqk blocks and prepare Akk
################################################################################
if i_tc1 < T:
p_Aqk10 = Aqk + o_c1[:, None] * (HV*BT) + o_i[None, :]
tl.store(p_Aqk10, (b_Aqk10 * scale).to(Aqk.dtype.element_ty), mask=m_A1)
p_b1 = beta + bos * HV + i_hv + o_c1 * HV
b_b1 = tl.load(p_b1, mask=m_tc1, other=0.0).to(tl.float32)
b_Akk10 = b_Akk10 * b_b1[:, None]
if NC >= 3 and i_tc2 < T:
p_Aqk20 = Aqk + o_c2[:, None] * (HV*BT) + o_i[None, :]
p_Aqk21 = Aqk + o_c2[:, None] * (HV*BT) + (o_i + BC)[None, :]
tl.store(p_Aqk20, (b_Aqk20 * scale).to(Aqk.dtype.element_ty), mask=m_A2)
tl.store(p_Aqk21, (b_Aqk21 * scale).to(Aqk.dtype.element_ty), mask=m_A2)
p_b2 = beta + bos * HV + i_hv + o_c2 * HV
b_b2 = tl.load(p_b2, mask=m_tc2, other=0.0).to(tl.float32)
b_Akk20 = b_Akk20 * b_b2[:, None]
b_Akk21 = b_Akk21 * b_b2[:, None]
if NC >= 4 and i_tc3 < T:
p_Aqk30 = Aqk + o_c3[:, None] * (HV*BT) + o_i[None, :]
p_Aqk31 = Aqk + o_c3[:, None] * (HV*BT) + (o_i + BC)[None, :]
p_Aqk32 = Aqk + o_c3[:, None] * (HV*BT) + (o_i + 2*BC)[None, :]
tl.store(p_Aqk30, (b_Aqk30 * scale).to(Aqk.dtype.element_ty), mask=m_A3)
tl.store(p_Aqk31, (b_Aqk31 * scale).to(Aqk.dtype.element_ty), mask=m_A3)
tl.store(p_Aqk32, (b_Aqk32 * scale).to(Aqk.dtype.element_ty), mask=m_A3)
p_b3 = beta + bos * HV + i_hv + o_c3 * HV
b_b3 = tl.load(p_b3, mask=m_tc3, other=0.0).to(tl.float32)
b_Akk30 = b_Akk30 * b_b3[:, None]
b_Akk31 = b_Akk31 * b_b3[:, None]
b_Akk32 = b_Akk32 * b_b3[:, None]
p_Akk00 = Akkd + o_c0[:, None] * (HV*BC) + o_i[None, :]
p_Akk11 = Akkd + o_c1[:, None] * (HV*BC) + o_i[None, :]
b_Ai00 = tl.load(p_Akk00, mask=m_A0, other=0.0).to(tl.float32)
b_Ai11 = tl.load(p_Akk11, mask=m_A1, other=0.0).to(tl.float32)
if NC >= 3:
p_Akk22 = Akkd + o_c2[:, None] * (HV*BC) + o_i[None, :]
b_Ai22 = tl.load(p_Akk22, mask=m_A2, other=0.0).to(tl.float32)
if NC >= 4:
p_Akk33 = Akkd + o_c3[:, None] * (HV*BC) + o_i[None, :]
b_Ai33 = tl.load(p_Akk33, mask=m_A3, other=0.0).to(tl.float32)
################################################################################
# forward substitution on diagonals
################################################################################
if not USE_SAFE_GATE:
m_A = o_i[:, None] > o_i[None, :]
m_I = o_i[:, None] == o_i[None, :]
b_Ai00 = -tl.where(m_A, b_Ai00, 0)
b_Ai11 = -tl.where(m_A, b_Ai11, 0)
if NC >= 3:
b_Ai22 = -tl.where(m_A, b_Ai22, 0)
if NC >= 4:
b_Ai33 = -tl.where(m_A, b_Ai33, 0)
for i in range(2, min(BC, T - i_tc0)):
b_a00 = -tl.load(Akkd + (i_tc0 + i) * HV*BC + o_i)
b_a00 = tl.where(o_i < i, b_a00, 0.)
b_a00 += tl.sum(b_a00[:, None] * b_Ai00, 0)
b_Ai00 = tl.where((o_i == i)[:, None], b_a00, b_Ai00)
for i in range(BC + 2, min(2*BC, T - i_tc0)):
b_a11 = -tl.load(Akkd + (i_tc0 + i) * HV*BC + o_i)
b_a11 = tl.where(o_i < i - BC, b_a11, 0.)
b_a11 += tl.sum(b_a11[:, None] * b_Ai11, 0)
b_Ai11 = tl.where((o_i == i - BC)[:, None], b_a11, b_Ai11)
if NC >= 3:
for i in range(2*BC + 2, min(3*BC, T - i_tc0)):
b_a22 = -tl.load(Akkd + (i_tc0 + i) * HV*BC + o_i)
b_a22 = tl.where(o_i < i - 2*BC, b_a22, 0.)
b_a22 += tl.sum(b_a22[:, None] * b_Ai22, 0)
b_Ai22 = tl.where((o_i == i - 2*BC)[:, None], b_a22, b_Ai22)
if NC >= 4:
for i in range(3*BC + 2, min(4*BC, T - i_tc0)):
b_a33 = -tl.load(Akkd + (i_tc0 + i) * HV*BC + o_i)
b_a33 = tl.where(o_i < i - 3*BC, b_a33, 0.)
b_a33 += tl.sum(b_a33[:, None] * b_Ai33, 0)
b_Ai33 = tl.where((o_i == i - 3*BC)[:, None], b_a33, b_Ai33)
b_Ai00 += m_I
b_Ai11 += m_I
if NC >= 3:
b_Ai22 += m_I
if NC >= 4:
b_Ai33 += m_I
################################################################################
# compute merged inverse using off-diagonals
################################################################################
# we used tf32 to maintain matrix inverse's precision whenever possible.
b_Ai10 = -tl.dot(
tl.dot(b_Ai11, b_Akk10, input_precision=SOLVE_TRIL_DOT_PRECISION),
b_Ai00,
input_precision=SOLVE_TRIL_DOT_PRECISION
)
if NC >= 3:
b_Ai21 = -tl.dot(
tl.dot(b_Ai22, b_Akk21, input_precision=SOLVE_TRIL_DOT_PRECISION),
b_Ai11,
input_precision=SOLVE_TRIL_DOT_PRECISION
)
b_Ai20 = -tl.dot(
b_Ai22,
tl.dot(b_Akk20, b_Ai00, input_precision=SOLVE_TRIL_DOT_PRECISION) +
tl.dot(b_Akk21, b_Ai10, input_precision=SOLVE_TRIL_DOT_PRECISION),
input_precision=SOLVE_TRIL_DOT_PRECISION
)
if NC >= 4:
b_Ai32 = -tl.dot(
tl.dot(b_Ai33, b_Akk32, input_precision=SOLVE_TRIL_DOT_PRECISION),
b_Ai22,
input_precision=SOLVE_TRIL_DOT_PRECISION
)
b_Ai31 = -tl.dot(
b_Ai33,
tl.dot(b_Akk31, b_Ai11, input_precision=SOLVE_TRIL_DOT_PRECISION) +
tl.dot(b_Akk32, b_Ai21, input_precision=SOLVE_TRIL_DOT_PRECISION),
input_precision=SOLVE_TRIL_DOT_PRECISION
)
b_Ai30 = -tl.dot(
b_Ai33,
tl.dot(b_Akk30, b_Ai00, input_precision=SOLVE_TRIL_DOT_PRECISION) +
tl.dot(b_Akk31, b_Ai10, input_precision=SOLVE_TRIL_DOT_PRECISION) +
tl.dot(b_Akk32, b_Ai20, input_precision=SOLVE_TRIL_DOT_PRECISION),
input_precision=SOLVE_TRIL_DOT_PRECISION
)
################################################################################
# store full Akk_inv to Akk
################################################################################
p_Akk00 = Akk + o_c0[:, None] * (HV*BT) + o_i[None, :]
p_Akk10 = Akk + o_c1[:, None] * (HV*BT) + o_i[None, :]
p_Akk11 = Akk + o_c1[:, None] * (HV*BT) + (o_i + BC)[None, :]
tl.store(p_Akk00, b_Ai00.to(Akk.dtype.element_ty), mask=m_A0)
tl.store(p_Akk10, b_Ai10.to(Akk.dtype.element_ty), mask=m_A1)
tl.store(p_Akk11, b_Ai11.to(Akk.dtype.element_ty), mask=m_A1)
if NC >= 3:
p_Akk20 = Akk + o_c2[:, None] * (HV*BT) + o_i[None, :]
p_Akk21 = Akk + o_c2[:, None] * (HV*BT) + (o_i + BC)[None, :]
p_Akk22 = Akk + o_c2[:, None] * (HV*BT) + (o_i + 2*BC)[None, :]
tl.store(p_Akk20, b_Ai20.to(Akk.dtype.element_ty), mask=m_A2)
tl.store(p_Akk21, b_Ai21.to(Akk.dtype.element_ty), mask=m_A2)
tl.store(p_Akk22, b_Ai22.to(Akk.dtype.element_ty), mask=m_A2)
if NC >= 4:
p_Akk30 = Akk + o_c3[:, None] * (HV*BT) + o_i[None, :]
p_Akk31 = Akk + o_c3[:, None] * (HV*BT) + (o_i + BC)[None, :]
p_Akk32 = Akk + o_c3[:, None] * (HV*BT) + (o_i + 2*BC)[None, :]
p_Akk33 = Akk + o_c3[:, None] * (HV*BT) + (o_i + 3*BC)[None, :]
tl.store(p_Akk30, b_Ai30.to(Akk.dtype.element_ty), mask=m_A3)
tl.store(p_Akk31, b_Ai31.to(Akk.dtype.element_ty), mask=m_A3)
tl.store(p_Akk32, b_Ai32.to(Akk.dtype.element_ty), mask=m_A3)
tl.store(p_Akk33, b_Ai33.to(Akk.dtype.element_ty), mask=m_A3)
@triton.heuristics({
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({}, num_warps=num_warps, num_stages=num_stages)
for num_warps in [1, 2, 4, 8]
for num_stages in [2, 3, 4]
],
key=['BK', 'NC', 'BT', 'HV'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['B', 'T'])
def chunk_kda_bwd_kernel_intra(
q,
k,
g,
beta,
dAqk,
dAkk,
dq,
dq2,
dk,
dk2,
dg,
dg2,
db,
cu_seqlens,
chunk_indices,
B,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
BT: tl.constexpr,
BC: tl.constexpr,
BK: tl.constexpr,
NC: tl.constexpr,
IS_VARLEN: tl.constexpr,
SAFE_GATE: tl.constexpr,
USE_GATHER: tl.constexpr,
):
i_kc, i_t, i_bh = tl.program_id(0), tl.program_id(1).to(tl.int64), tl.program_id(2).to(tl.int64)
i_b, i_hv = i_bh // HV, i_bh % HV
i_h = i_hv // (HV // H)
i_k, i_i = i_kc // NC, i_kc % NC
all = B * T
if IS_VARLEN:
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
else:
bos, eos = i_b * T, i_b * T + T
T = eos - bos
i_ti = i_t * BT + i_i * BC
if i_ti >= T:
return
o_k = i_k * BK + tl.arange(0, BK)
m_k = o_k < K
q += (bos * H + i_h) * K
k += (bos * H + i_h) * K
g += (bos * HV + i_hv) * K
beta += bos * HV + i_hv
dAqk += (bos * HV + i_hv) * BT
dAkk += (bos * HV + i_hv) * BT
dq += (bos * HV + i_hv) * K
dq2 += (bos * HV + i_hv) * K
dk += (bos * HV + i_hv) * K
dk2 += (bos * HV + i_hv) * K
dg += (bos * HV + i_hv) * K
dg2 += (bos * HV + i_hv) * K
db += (i_k * all + bos) * HV + i_hv
o_i = tl.arange(0, BC)
o_c = i_ti + o_i
m_c = o_c < T
m_ck = m_c[:, None] & m_k[None, :]
m_dAf = m_c[:, None] & (o_i[None, :] < BT)
m_dAt = (o_i[:, None] < BT) & m_c[None, :]
p_g = g + o_c[:, None] * (HV*K) + o_k[None, :]
b_g = tl.load(p_g, mask=m_ck, other=0.0).to(tl.float32)
p_b = beta + o_c * HV
b_b = tl.load(p_b, mask=m_c, other=0.0)
b_dq2 = tl.zeros([BC, BK], dtype=tl.float32)
b_dk2 = tl.zeros([BC, BK], dtype=tl.float32)
if i_i > 0:
p_gn = g + i_ti * HV*K + o_k
# [BK,]
b_gn = tl.load(p_gn, mask=m_k, other=0).to(tl.float32)[None, :]
for i_j in range(0, i_i):
o_j = i_t * BT + i_j * BC + o_i
m_jk = (o_j < T)[:, None] & m_k[None, :]
p_k = k + o_j[:, None] * (H*K) + o_k[None, :]
p_gk = g + o_j[:, None] * (HV*K) + o_k[None, :]
p_dAqk = dAqk + o_c[:, None] * (HV*BT) + (i_j * BC + o_i)[None, :]
p_dAkk = dAkk + o_c[:, None] * (HV*BT) + (i_j * BC + o_i)[None, :]
# [BC, BK]
b_k = tl.load(p_k, mask=m_jk, other=0.0)
b_gk = tl.load(p_gk, mask=m_jk, other=0.0)
b_kg = b_k * exp2(b_gn - b_gk)
# [BC, BC]
b_dAqk = tl.load(p_dAqk, mask=m_dAf, other=0.0)
b_dAkk = tl.load(p_dAkk, mask=m_dAf, other=0.0)
# [BC, BK]
b_dq2 += tl.dot(b_dAqk, b_kg)
b_dk2 += tl.dot(b_dAkk, b_kg)
b_gqn = exp2(b_g - b_gn)
b_dq2 *= b_gqn
b_dk2 *= b_gqn
o_i = tl.arange(0, BC)
m_dA = (i_ti + o_i) < T
o_dA = (i_ti + o_i) * HV*BT + i_i * BC
p_kj = k + i_ti * H*K + o_k
p_gkj = g + i_ti * HV*K + o_k
p_q = q + o_c[:, None] * (H*K) + o_k[None, :]
p_k = k + o_c[:, None] * (H*K) + o_k[None, :]
b_q = tl.load(p_q, mask=m_ck, other=0.0)
b_k = tl.load(p_k, mask=m_ck, other=0.0)
if SAFE_GATE:
if USE_GATHER:
b_gn = gather(b_g, tl.full([1, BK], min(BC//2, T - i_ti - 1), dtype=tl.int16), axis=0)
else:
p_gn = g + (i_ti + min(BC // 2, T - i_ti - 1)) * HV*K + o_k
b_gn = tl.load(p_gn, mask=m_k, other=0)[None, :]
p_dAqk = dAqk + o_c[:, None] * (HV*BT) + (i_i * BC + o_i)[None, :]
p_dAkk = dAkk + o_c[:, None] * (HV*BT) + (i_i * BC + o_i)[None, :]
b_dAqk_diag_qk = tl.load(p_dAqk, mask=m_dAf, other=0.0).to(tl.float32)
b_dAkk_diag_qk = tl.load(p_dAkk, mask=m_dAf, other=0.0).to(tl.float32)
m_i_diag_qk = (o_i[:, None] >= o_i[None, :]) & ((i_ti + o_i[:, None]) < T) & ((i_ti + o_i[None, :]) < T)
m_j_diag_qk = (i_ti + o_i[:, None]) < T
b_dAqk_diag_qk = tl.where(m_i_diag_qk, b_dAqk_diag_qk, 0.)
b_dAkk_diag_qk = tl.where(m_i_diag_qk, b_dAkk_diag_qk, 0.)
b_g_diag_qk = tl.where(m_j_diag_qk, b_g - b_gn, 0.)
exp_b_g_diag_qk = tl.where(m_j_diag_qk, exp2(b_g_diag_qk), 0.)
exp_neg_b_g_diag_qk = tl.where(m_j_diag_qk, exp2(-b_g_diag_qk), 0.)
b_k_exp_diag_qk = b_k * exp_neg_b_g_diag_qk
b_dq2 += tl.dot(b_dAqk_diag_qk, b_k_exp_diag_qk) * exp_b_g_diag_qk
b_dk2 += tl.dot(b_dAkk_diag_qk, b_k_exp_diag_qk) * exp_b_g_diag_qk
else:
for j in range(0, min(BC, T - i_t * BT - i_i * BC)):
# [BC]
b_dAqk = tl.load(dAqk + o_dA + j, mask=m_dA, other=0)
b_dAkk = tl.load(dAkk + o_dA + j, mask=m_dA, other=0)
# [BK]
b_kj = tl.load(p_kj, mask=m_k, other=0).to(tl.float32)
b_gkj = tl.load(p_gkj, mask=m_k, other=0).to(tl.float32)
# [BC, BK]
m_i = o_i[:, None] >= j
# [BC, BK]
b_gqk = exp2(b_g - b_gkj[None, :])
b_dq2 += tl.where(m_i, b_dAqk[:, None] * b_kj[None, :] * b_gqk, 0.)
b_dk2 += tl.where(m_i, b_dAkk[:, None] * b_kj[None, :] * b_gqk, 0.)
p_kj += H*K
p_gkj += HV*K
b_db = tl.sum(b_dk2 * b_k, 1)
b_dk2 *= b_b[:, None]
p_dq = dq + o_c[:, None] * (HV*K) + o_k[None, :]
p_dq2 = dq2 + o_c[:, None] * (HV*K) + o_k[None, :]
p_db = db + o_c * HV
b_dg2 = b_q * b_dq2
b_dq2 = b_dq2 + tl.load(p_dq, mask=m_ck, other=0.0)
tl.store(p_dq2, b_dq2.to(p_dq2.dtype.element_ty), mask=m_ck)
tl.store(p_db, b_db.to(p_db.dtype.element_ty), mask=m_c)
tl.debug_barrier()
b_dkt = tl.zeros([BC, BK], dtype=tl.float32)
NC = min(NC, tl.cdiv(T - i_t * BT, BC))
if i_i < NC - 1:
p_gn = g + (min(i_ti + BC, T) - 1) * HV*K + o_k
# [BK,]
b_gn = tl.load(p_gn, mask=m_k, other=0).to(tl.float32)[None, :]
for i_j in range(i_i + 1, NC):
o_j = i_t * BT + i_j * BC + o_i
m_j = o_j < T
m_jk = m_j[:, None] & m_k[None, :]
m_dAj = (o_i[:, None] < BT) & m_j[None, :]
p_q = q + o_j[:, None] * (H*K) + o_k[None, :]
p_k = k + o_j[:, None] * (H*K) + o_k[None, :]
p_gk = g + o_j[:, None] * (HV*K) + o_k[None, :]
p_b = beta + o_j * HV
p_dAqk = dAqk + (i_i * BC + o_i)[:, None] + o_j[None, :] * (HV*BT)
p_dAkk = dAkk + (i_i * BC + o_i)[:, None] + o_j[None, :] * (HV*BT)
# [BC]
b_b = tl.load(p_b, mask=m_j, other=0.0)
# [BC, BK]
b_q = tl.load(p_q, mask=m_jk, other=0.0)
b_kb = tl.load(p_k, mask=m_jk, other=0.0) * b_b[:, None]
b_gk = tl.load(p_gk, mask=m_jk, other=0.0).to(tl.float32)
# [BC, BC]
b_dAqk = tl.load(p_dAqk, mask=m_dAj, other=0.0)
b_dAkk = tl.load(p_dAkk, mask=m_dAj, other=0.0)
# [BC, BK]
b_gkn = exp2(b_gk - b_gn)
b_qg = b_q * tl.where(m_j[:, None], b_gkn, 0)
b_kbg = b_kb * tl.where(m_j[:, None], b_gkn, 0)
# [BC, BK]
# (SY 09/17) important to not use bf16 here to have a good precision.
b_dkt += tl.dot(b_dAqk, b_qg)
b_dkt += tl.dot(b_dAkk, b_kbg)
b_dkt *= exp2(b_gn - b_g)
o_dA = i_ti * HV*BT + i_i * BC + o_i
p_qj = q + i_ti * H*K + o_k
p_kj = k + i_ti * H*K + o_k
p_gkj = g + i_ti * HV*K + o_k
p_bj = beta + i_ti * HV
if SAFE_GATE:
if USE_GATHER:
b_gn = gather(b_g, tl.full([1, BK], min(BC//2, T - i_ti - 1), dtype=tl.int16), axis=0)
else:
p_gn = g + (i_ti + min(BC // 2, T - i_ti - 1)) * HV*K + o_k
b_gn = tl.load(p_gn, mask=m_k, other=0).to(tl.float32)[None, :]
p_q = q + o_c[:, None] * (H*K) + o_k[None, :]
b_q = tl.load(p_q, mask=m_ck, other=0.0)
p_b = beta + o_c * HV
b_b = tl.load(p_b, mask=m_c, other=0.0)
p_dAqk = dAqk + (i_i * BC + o_i)[:, None] + o_c[None, :] * (HV*BT)
p_dAkk = dAkk + (i_i * BC + o_i)[:, None] + o_c[None, :] * (HV*BT)
b_dAqk_diag_kk = tl.load(p_dAqk, mask=m_dAt, other=0.0).to(tl.float32)
b_dAkk_diag_kk = tl.load(p_dAkk, mask=m_dAt, other=0.0).to(tl.float32)
m_i_diag_kk = (o_i[:, None] <= o_i[None, :]) & ((i_ti + o_i[:, None]) < T) & ((i_ti + o_i[None, :]) < T)
m_j_diag_kk = (i_ti + o_i[:, None]) < T
b_dAqk_diag_kk = tl.where(m_i_diag_kk, b_dAqk_diag_kk, 0.)
b_dAkk_diag_kk = tl.where(m_i_diag_kk, b_dAkk_diag_kk, 0.)
# ensure numerical stability
b_g_diag_kk = tl.where(m_j_diag_kk, b_g - b_gn, 0.)
exp_b_g_diag_kk = tl.where(m_j_diag_kk, exp2(b_g_diag_kk), 0.)
exp_neg_b_g_diag_kk = tl.where(m_j_diag_kk, exp2(-b_g_diag_kk), 0.)
b_q_exp = b_q * exp_b_g_diag_kk
b_kb_exp = b_k * b_b[:, None] * exp_b_g_diag_kk
b_dkt += tl.dot(b_dAqk_diag_kk, b_q_exp) * exp_neg_b_g_diag_kk
b_dkt += tl.dot(b_dAkk_diag_kk, b_kb_exp) * exp_neg_b_g_diag_kk
else:
for j in range(0, min(BC, T - i_t * BT - i_i * BC)):
# [BC,]
b_dAqk = tl.load(dAqk + o_dA + j * HV*BT)
b_dAkk = tl.load(dAkk + o_dA + j * HV*BT)
# [BK,]
b_qj = tl.load(p_qj, mask=m_k, other=0).to(tl.float32)
b_kbj = tl.load(p_kj, mask=m_k, other=0).to(tl.float32) * tl.load(p_bj)
b_gkj = tl.load(p_gkj, mask=m_k, other=0).to(tl.float32)
# [BC, BK]
m_i = o_i[:, None] <= j
b_gkq = exp2(b_gkj[None, :] - b_g)
b_dkt += tl.where(m_i, b_dAqk[:, None] * b_qj[None, :] * b_gkq, 0.)
b_dkt += tl.where(m_i, b_dAkk[:, None] * b_kbj[None, :] * b_gkq, 0.)
p_qj += H*K
p_kj += H*K
p_gkj += HV*K
p_bj += HV
p_dk = dk + o_c[:, None] * (HV*K) + o_k[None, :]
p_dk2 = dk2 + o_c[:, None] * (HV*K) + o_k[None, :]
p_dg = dg + o_c[:, None] * (HV*K) + o_k[None, :]
p_dg2 = dg2 + o_c[:, None] * (HV*K) + o_k[None, :]
b_dg2 += (b_dk2 - b_dkt) * b_k + tl.load(p_dg, mask=m_ck, other=0.0)
b_dk2 += tl.load(p_dk, mask=m_ck, other=0.0)
b_dk2 += b_dkt
tl.store(p_dk2, b_dk2.to(p_dk2.dtype.element_ty), mask=m_ck)
tl.store(p_dg2, b_dg2.to(p_dg2.dtype.element_ty), mask=m_ck)
@triton.heuristics({
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({}, num_warps=num_warps, num_stages=num_stages)
for num_warps in [1, 2, 4, 8]
for num_stages in [2, 3, 4]
],
key=["BT", "BC", "HV"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_kda_fwd_kernel_intra_sub_chunk(
q,
k,
g,
beta,
Aqk,
Akk,
scale,
cu_seqlens,
chunk_indices,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
BT: tl.constexpr,
BC: tl.constexpr,
BK: tl.constexpr,
IS_VARLEN: tl.constexpr,
USE_GATHER: tl.constexpr,
):
i_t, i_i, i_bh = tl.program_id(0).to(tl.int64), tl.program_id(1), tl.program_id(2).to(tl.int64)
i_b, i_hv = i_bh // HV, i_bh % HV
i_h = i_hv // (HV // H)
if IS_VARLEN:
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
else:
bos, eos = i_b * T, i_b * T + T
i_ti = i_t * BT + i_i * BC
if i_ti >= T:
return
o_c = i_ti + tl.arange(0, BC)
m_c = o_c < T
q = q + (bos * H + i_h) * K
k = k + (bos * H + i_h) * K
g = g + (bos * HV + i_hv) * K
beta = beta + bos * HV + i_hv
Aqk = Aqk + (bos * HV + i_hv) * BT
Akk = Akk + (bos * HV + i_hv) * BC
o_k = tl.arange(0, BK)
m_k = o_k < K
m_ck = m_c[:, None] & m_k[None, :]
p_q = q + o_c[:, None] * (H*K) + o_k[None, :]
p_k = k + o_c[:, None] * (H*K) + o_k[None, :]
p_g = g + o_c[:, None] * (HV*K) + o_k[None, :]
p_beta = beta + o_c * HV
b_q = tl.load(p_q, mask=m_ck, other=0.0)
b_k = tl.load(p_k, mask=m_ck, other=0.0)
b_g = tl.load(p_g, mask=m_ck, other=0.0)
b_beta = tl.load(p_beta, mask=m_c, other=0.0)
if USE_GATHER:
b_gn = gather(b_g, tl.full([1, BK], min(BC//2, T - i_ti - 1), dtype=tl.int16), axis=0)
else:
# caculate offset
p_gn = g + (i_ti + min(BC // 2, T - i_ti - 1)) * HV*K + tl.arange(0, BK)
b_gn = tl.load(p_gn, mask=tl.arange(0, BK) < K, other=0.0)
b_gn = b_gn[None, :]
# current block, keep numerical stability by subtracting the left boundary
# less than 85 to avoid overflow in exp2
b_gm = (b_g - b_gn).to(tl.float32)
b_gq = tl.where(m_c[:, None], exp2(b_gm), 0.)
b_gk = tl.where(m_c[:, None], exp2(-b_gm), 0.)
b_kgt = tl.trans(b_k * b_gk)
b_Aqk = tl.dot(b_q * b_gq, b_kgt) * scale
b_Akk = tl.dot(b_k * b_gq, b_kgt) * b_beta[:, None]
o_i = tl.arange(0, BC)
m_Aqk = o_i[:, None] >= o_i[None, :]
m_Akk = o_i[:, None] > o_i[None, :]
m_I = o_i[:, None] == o_i[None, :]
b_Aqk = tl.where(m_Aqk, b_Aqk, 0.0)
b_Akk = tl.where(m_Akk, b_Akk, 0.0)
m_Aqk_st = m_c[:, None] & (o_i[None, :] < BT)
m_Akk_st = m_c[:, None] & (o_i[None, :] < BC)
p_Aqk = Aqk + o_c[:, None] * (HV*BT) + (i_i * BC + o_i)[None, :]
p_Akk = Akk + o_c[:, None] * (HV*BC) + o_i[None, :]
tl.store(p_Aqk, b_Aqk.to(Aqk.dtype.element_ty), mask=m_Aqk_st)
tl.store(p_Akk, b_Akk.to(Akk.dtype.element_ty), mask=m_Akk_st)
tl.debug_barrier()
################################################################################
# forward substitution
################################################################################
b_Ai = -b_Akk
for i in range(2, min(BC, T - i_ti)):
b_a = -tl.load(Akk + (i_ti + i) * HV*BC + o_i)
b_a = tl.where(o_i < i, b_a, 0.)
b_a += tl.sum(b_a[:, None] * b_Ai, 0)
b_Ai = tl.where((o_i == i)[:, None], b_a, b_Ai)
b_Ai += m_I
tl.store(p_Akk, b_Ai.to(Akk.dtype.element_ty), mask=m_Akk_st)
@dispatch('kda')
def chunk_kda_fwd_intra(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
gk: torch.Tensor | None = None,
beta: torch.Tensor | None = None,
scale: float | None = None,
cu_seqlens: torch.LongTensor | None = None,
chunk_size: int = 64,
chunk_indices: torch.LongTensor | None = None,
safe_gate: bool = False,
disable_recompute: bool = False,
):
B, T, H, K, HV = *k.shape, gk.shape[2]
BT = chunk_size
if BT not in (32, 64):
raise ValueError(f"KDA intra chunk kernel only supports chunk_size 32 or 64, got {BT}.")
BC = 16
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
NC = triton.cdiv(BT, BC)
Aqk = torch.empty(B, T, HV, BT, device=k.device, dtype=k.dtype)
# Akk must be zero-initialized - kernel only writes lower triangular
Akk = torch.zeros(B, T, HV, BT, device=k.device, dtype=k.dtype)
# Separate fp32 buffer for diagonal 16x16 blocks (for precision in solve_tril)
Akkd = torch.empty(B, T, HV, BC, device=k.device, dtype=torch.float32)
# Step 1: Run token_parallel first to compute diagonal blocks into Akkd (fp32)
# Step 1: compute diagonal blocks into Akk_diag (fp32)
if safe_gate:
grid = (NT, NC, B * HV)
BK = triton.next_power_of_2(K)
chunk_kda_fwd_kernel_intra_sub_chunk[grid](
q=q,
k=k,
g=gk,
beta=beta,
Aqk=Aqk,
Akk=Akkd,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
T=T,
H=H,
HV=HV,
K=K,
BT=BT,
BC=BC,
BK=BK,
USE_GATHER=IS_GATHER_SUPPORTED,
)
else:
Aqk, Akkd = chunk_kda_fwd_intra_token_parallel(
q=q,
k=k,
gk=gk,
beta=beta,
Aqk=Aqk,
Akk=Akkd,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_size=BT,
sub_chunk_size=BC,
)
# Step 2: Fused inter + solve_tril (works for both fixed-len and varlen)
grid = (NT, B * HV)
chunk_kda_fwd_kernel_inter_solve_fused[grid](
q=q,
k=k,
g=gk,
beta=beta,
Aqk=Aqk,
Akkd=Akkd,
Akk=Akk,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
T=T,
H=H,
HV=HV,
K=K,
BT=BT,
BC=BC,
NC=NC,
USE_SAFE_GATE=safe_gate,
)
w, u, qg, kg = recompute_w_u_fwd(
k=k,
v=v,
beta=beta,
A=Akk,
q=q if disable_recompute else None,
gk=gk,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
)
return w, u, qg, kg, Aqk, Akk
@dispatch('kda')
def chunk_kda_bwd_intra(
q: torch.Tensor,
k: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
dAqk: torch.Tensor,
dAkk: torch.Tensor,
dq: torch.Tensor,
dk: torch.Tensor,
db: torch.Tensor,
dg: torch.Tensor,
cu_seqlens: torch.LongTensor | None = None,
chunk_indices: torch.LongTensor | None = None,
chunk_size: int = 64,
safe_gate: bool = False,
):
B, T, H, K, HV = *k.shape, g.shape[2]
BT = chunk_size
BC = min(16, BT)
BK = min(32, triton.next_power_of_2(K))
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
NC = triton.cdiv(BT, BC)
NK = triton.cdiv(K, BK)
dq2 = torch.empty_like(dq)
dk2 = torch.empty_like(dk)
db2 = beta.new_empty(NK, *beta.shape, dtype=torch.float)
dg2 = torch.empty_like(dg, dtype=torch.float)
grid = (NK * NC, NT, B * HV)
chunk_kda_bwd_kernel_intra[grid](
q=q,
k=k,
g=g,
beta=beta,
dAqk=dAqk,
dAkk=dAkk,
dq=dq,
dq2=dq2,
dk=dk,
dk2=dk2,
dg=dg,
dg2=dg2,
db=db2,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
B=B,
T=T,
H=H,
HV=HV,
K=K,
BT=BT,
BC=BC,
BK=BK,
NC=NC,
SAFE_GATE=safe_gate,
USE_GATHER=IS_GATHER_SUPPORTED,
)
dq = dq2
dk = dk2
db = db2.sum(0).add_(db)
dg = dg2
return dq, dk, db, dg
@@ -0,0 +1,182 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
# Token-parallel implementation of KDA intra chunk kernel
import torch
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.ops.utils.op import exp2
from kda._fla.utils import autotune_cache_kwargs
@triton.heuristics({
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({'BH': BH}, num_warps=num_warps)
for BH in [1, 2, 4, 8]
for num_warps in [1, 2, 4, 8]
],
key=["K", "H", "HV"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T', 'N'])
def chunk_kda_fwd_kernel_intra_token_parallel(
q,
k,
g,
beta,
Aqk,
Akk,
scale,
cu_seqlens,
N,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
BK: tl.constexpr,
BT: tl.constexpr,
BC: tl.constexpr,
BH: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_tg, i_hg = tl.program_id(0).to(tl.int64), tl.program_id(1)
if IS_VARLEN:
i_n = 0
left, right = 0, N
# Unrolled binary search (max B=2^32)
# We can limit iterations based on expected max batch size if needed
# 20 iterations covers B=1M, usually enough
for _ in range(20):
if left < right:
mid = (left + right) // 2
if i_tg < tl.load(cu_seqlens + mid + 1).to(tl.int32):
right = mid
else:
left = mid + 1
i_n = left
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
i_t = i_tg - bos
else:
bos = (i_tg // T) * T
i_t = i_tg % T
if i_t >= T:
return
i_c = i_t // BT
i_s = (i_t % BT) // BC
i_tc = i_c * BT
i_ts = i_tc + i_s * BC
G: tl.constexpr = HV // H
q += bos * H*K
k += bos * H*K
g += bos * HV*K
Aqk += bos * HV*BT
Akk += bos * HV*BC
beta += bos * HV
o_hv = i_hg * BH + tl.arange(0, BH)
o_h = o_hv // G
o_k = tl.arange(0, BK)
m_hv = o_hv < HV
m_k = o_k < K
m_hk = m_hv[:, None] & m_k[None, :]
# q/k: [B, T, H, K], manual load via mapped qk head index
p_qk = o_h[:, None] * K + o_k[None, :]
b_q = tl.load(q + i_t * H * K + p_qk, mask=m_hk, other=0).to(tl.float32)
b_k = tl.load(k + i_t * H * K + p_qk, mask=m_hk, other=0).to(tl.float32)
# g: [B, T, HV, K], beta: [B, T, HV]
p_g = g + i_t * HV * K + o_hv[:, None] * K + o_k[None, :]
p_beta = beta + i_t * HV + o_hv
b_g = tl.load(p_g, mask=m_hk, other=0.0).to(tl.float32)
b_k = b_k * tl.load(p_beta, mask=m_hv, other=0.0).to(tl.float32)[:, None]
for j in range(i_ts, min(i_t + 1, min(T, i_ts + BC))):
b_kj = tl.load(k + j * H * K + p_qk, mask=m_hk, other=0).to(tl.float32)
p_gj = g + j * HV * K + o_hv[:, None] * K + o_k[None, :]
b_gj = tl.load(p_gj, mask=m_hk, other=0.0).to(tl.float32)
b_kgj = tl.where(m_k[None, :], b_kj * exp2(b_g - b_gj), 0.0)
b_Aqk = tl.sum(b_q * b_kgj, axis=1) * scale
b_Akk = tl.sum(b_k * b_kgj, axis=1) * tl.where(j < i_t, 1.0, 0.0)
tl.store(Aqk + i_t * HV * BT + o_hv * BT + j % BT, b_Aqk.to(Aqk.dtype.element_ty), mask=m_hv)
tl.store(Akk + i_t * HV * BC + o_hv * BC + j - i_ts, b_Akk.to(Akk.dtype.element_ty), mask=m_hv)
@dispatch('kda')
def chunk_kda_fwd_intra_token_parallel(
q: torch.Tensor,
k: torch.Tensor,
gk: torch.Tensor,
beta: torch.Tensor,
Aqk: torch.Tensor,
Akk: torch.Tensor,
scale: float,
cu_seqlens: torch.LongTensor | None = None,
chunk_size: int = 64,
sub_chunk_size: int = 16,
) -> None:
"""
Token-parallel implementation: each token gets its own thread block.
Supports both fixed-length and variable-length sequences.
Reduces wasted computation on padding.
Writes directly to Aqk and Akk tensors (in-place).
Args:
q: [B, T, H, K]
k: [B, T, H, K]
gk: [B, T, HV, K] cumsum of gates (HV >= H for GVA)
beta: [B, T, HV]
Aqk: [B, T, HV, BT] output tensor to write to
Akk: [B, T, HV, BC] output tensor for diagonal blocks (fp32)
scale: attention scale
chunk_size: BT (default 64)
sub_chunk_size: BC (default 16)
"""
B, T, H, K, HV = *q.shape, gk.shape[2]
N = len(cu_seqlens) - 1 if cu_seqlens is not None else B
BT = chunk_size
BC = sub_chunk_size
BK = triton.next_power_of_2(K)
def grid(meta): return (B * T, triton.cdiv(HV, meta['BH']))
chunk_kda_fwd_kernel_intra_token_parallel[grid](
q=q,
k=k,
g=gk,
beta=beta,
Aqk=Aqk,
Akk=Akk,
scale=scale,
cu_seqlens=cu_seqlens,
N=N,
T=T,
H=H,
HV=HV,
K=K,
BK=BK,
BT=BT,
BC=BC,
)
return Aqk, Akk
+491
View File
@@ -0,0 +1,491 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
# This kernel is modified from the Decode kernel of the vllm gdn/kda model.
import warnings
import torch
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.ops.utils.op import exp
from kda._fla.ops.utils.softplus import softplus
from kda._fla.utils import input_guard
@triton.heuristics(
{
"USE_INITIAL_STATE": lambda args: args["h0"] is not None,
"STORE_FINAL_STATE": lambda args: args["ht"] is not None,
"IS_VARLEN": lambda args: args["cu_seqlens"] is not None,
"IS_CONTINUOUS_BATCHING": lambda args: args["ssm_state_indices"] is not None,
"IS_SPEC_DECODING": lambda args: args["num_accepted_tokens"] is not None,
"HAS_A": lambda args: args["A_log"] is not None,
"HAS_BIAS": lambda args: args["dt_bias"] is not None,
"USE_LOWER_BOUND": lambda args: args["lower_bound"] is not None,
}
)
@triton.jit(do_not_specialize=["N", "T"])
def fused_recurrent_kda_fwd_kernel(
q,
k,
v,
g,
beta,
A_log,
dt_bias,
o,
h0,
ht,
cu_seqlens,
ssm_state_indices,
num_accepted_tokens,
lower_bound,
scale: tl.constexpr,
N: tl.int64, # num of sequences
T: tl.int64, # num of tokens
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BK: tl.constexpr,
BV: tl.constexpr,
stride_init_state_token: tl.constexpr,
stride_final_state_token: tl.constexpr,
stride_indices_seq: tl.constexpr,
stride_indices_tok: tl.constexpr,
USE_INITIAL_STATE: tl.constexpr, # whether to use initial state
INPLACE_FINAL_STATE: tl.constexpr, # whether to store final state inplace
IS_BETA_HEADWISE: tl.constexpr, # whether beta is headwise vector or scalar,
USE_QK_L2NORM_IN_KERNEL: tl.constexpr,
IS_VARLEN: tl.constexpr,
IS_CONTINUOUS_BATCHING: tl.constexpr,
IS_SPEC_DECODING: tl.constexpr,
STORE_FINAL_STATE: tl.constexpr,
HAS_A: tl.constexpr,
HAS_BIAS: tl.constexpr,
USE_GATE_IN_KERNEL: tl.constexpr,
USE_LOWER_BOUND: tl.constexpr,
APPLY_BETA_SIGMOID: tl.constexpr,
ALLOW_NEG_EIGVAL: tl.constexpr,
STATE_V_FIRST: tl.constexpr,
num_stages: tl.constexpr,
):
pid = tl.program_id(0).to(tl.int64)
NV = tl.cdiv(V, BV)
NK = tl.cdiv(K, BK)
i_k = pid % NK
pid_rest = pid // NK
i_v = pid_rest % NV
i_nh = pid_rest // NV
i_n, i_hv = i_nh // HV, i_nh % HV
i_h = i_hv // (HV // H)
if IS_VARLEN:
bos, eos = (
tl.load(cu_seqlens + i_n).to(tl.int64),
tl.load(cu_seqlens + i_n + 1).to(tl.int64),
)
T = eos - bos
else:
bos, eos = i_n * T, i_n * T + T
if T == 0:
# no tokens to process for this sequence
return
o_k = i_k * BK + tl.arange(0, BK)
o_v = i_v * BV + tl.arange(0, BV)
p_q = q + (bos * H + i_h) * K + o_k
p_k = k + (bos * H + i_h) * K + o_k
p_v = v + (bos * HV + i_hv) * V + o_v
if IS_BETA_HEADWISE:
p_beta = beta + (bos * HV + i_hv) * V + o_v
else:
p_beta = beta + bos * HV + i_hv
p_g = g + (bos * HV + i_hv) * K + o_k
p_o = o + (bos * HV + i_hv) * V + o_v
mask_k = o_k < K
mask_v = o_v < V
if STATE_V_FIRST:
mask_h = mask_v[:, None] & mask_k[None, :]
else:
mask_h = mask_k[:, None] & mask_v[None, :]
if STATE_V_FIRST:
b_h = tl.zeros([BV, BK], dtype=tl.float32)
else:
b_h = tl.zeros([BK, BV], dtype=tl.float32)
if USE_INITIAL_STATE:
if IS_CONTINUOUS_BATCHING:
if IS_SPEC_DECODING:
i_t = tl.load(num_accepted_tokens + i_n).to(tl.int64) - 1
else:
i_t = 0
p_h0 = (
h0
+ tl.load(ssm_state_indices + i_n * stride_indices_seq + i_t).to(
tl.int64
)
* stride_init_state_token
)
if STATE_V_FIRST:
p_h0 = p_h0 + i_hv * K * V + o_v[:, None] * K + o_k[None, :]
else:
p_h0 = p_h0 + i_hv * K * V + o_k[:, None] * V + o_v[None, :]
else:
if STATE_V_FIRST:
p_h0 = h0 + (i_n * HV + i_hv) * K * V + o_v[:, None] * K + o_k[None, :]
else:
p_h0 = h0 + (i_n * HV + i_hv) * K * V + o_k[:, None] * V + o_v[None, :]
b_h += tl.load(p_h0, mask=mask_h, other=0).to(tl.float32)
for i_t in tl.range(0, T, num_stages=num_stages):
b_q = tl.load(p_q, mask=mask_k, other=0, eviction_policy='evict_last').to(tl.float32)
b_k = tl.load(p_k, mask=mask_k, other=0, eviction_policy='evict_last').to(tl.float32)
b_v = tl.load(p_v, mask=mask_v, other=0, eviction_policy='evict_first').to(tl.float32)
if USE_QK_L2NORM_IN_KERNEL:
b_q = b_q / tl.sqrt(tl.sum(b_q * b_q) + 1e-6)
b_k = b_k / tl.sqrt(tl.sum(b_k * b_k) + 1e-6)
b_q = b_q * scale
b_g = tl.load(p_g, mask=mask_k, other=0, eviction_policy='evict_last').to(tl.float32)
if USE_GATE_IN_KERNEL:
b_A = tl.load(A_log + i_hv).to(tl.float32) if HAS_A else 1.0
if HAS_BIAS:
b_bias = tl.load(dt_bias + i_hv * K + o_k, mask=mask_k, other=0).to(tl.float32)
b_g = b_g + b_bias
if USE_LOWER_BOUND:
b_gk = lower_bound * tl.sigmoid((exp(b_A) if HAS_A else b_A) * b_g)
else:
b_gk = -exp(b_A) * softplus(b_g)
else:
b_gk = b_g
if STATE_V_FIRST:
b_h *= exp(b_gk[None, :])
else:
b_h *= exp(b_gk[:, None])
if STATE_V_FIRST:
b_v -= tl.sum(b_h * b_k[None, :], 1)
else:
b_v -= tl.sum(b_h * b_k[:, None], 0)
if IS_BETA_HEADWISE:
b_beta = tl.load(p_beta, mask=mask_v, other=0, eviction_policy='evict_first').to(tl.float32)
else:
b_beta = tl.load(p_beta, eviction_policy='evict_last').to(tl.float32)
if APPLY_BETA_SIGMOID:
b_beta = tl.sigmoid(b_beta)
if ALLOW_NEG_EIGVAL:
b_beta = b_beta * 2
b_v *= b_beta
if STATE_V_FIRST:
b_h += b_v[:, None] * b_k[None, :]
b_o = tl.sum(b_h * b_q[None, :], 1)
else:
b_h += b_k[:, None] * b_v[None, :]
b_o = tl.sum(b_h * b_q[:, None], 0)
tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=mask_v, eviction_policy='evict_first')
if IS_CONTINUOUS_BATCHING:
if INPLACE_FINAL_STATE:
p_ht = (
ht
+ tl.load(ssm_state_indices + i_n * stride_indices_seq + i_t).to(
tl.int64
)
* stride_final_state_token
)
else:
p_ht = ht + (bos + i_t) * stride_final_state_token
if STATE_V_FIRST:
p_ht = p_ht + i_hv * K * V + o_v[:, None] * K + o_k[None, :]
else:
p_ht = p_ht + i_hv * K * V + o_k[:, None] * V + o_v[None, :]
tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), mask=mask_h)
p_q += H * K
p_k += H * K
p_o += HV * V
p_v += HV * V
p_g += HV * K
p_beta += HV * (V if IS_BETA_HEADWISE else 1)
if not IS_CONTINUOUS_BATCHING:
if STORE_FINAL_STATE:
if STATE_V_FIRST:
p_ht = ht + (i_n * HV + i_hv) * K * V + o_v[:, None] * K + o_k[None, :]
else:
p_ht = ht + (i_n * HV + i_hv) * K * V + o_k[:, None] * V + o_v[None, :]
tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), mask=mask_h)
@dispatch("kda")
def fused_recurrent_kda_fwd(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
initial_state: torch.Tensor | None = None,
scale: float | None = None,
output_final_state: bool = False,
inplace_final_state: bool = True,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
ssm_state_indices: torch.Tensor | None = None,
num_accepted_tokens: torch.Tensor | None = None,
use_qk_l2norm_in_kernel: bool = False,
use_gate_in_kernel: bool = False,
use_beta_sigmoid_in_kernel: bool = False,
allow_neg_eigval: bool = False,
lower_bound: float | None = None,
out: torch.Tensor | None = None,
**kwargs,
) -> tuple[torch.Tensor, torch.Tensor]:
if scale is None:
scale = k.shape[-1] ** -0.5
B, T, H, K, V = *k.shape, v.shape[-1]
HV = v.shape[2]
N = B if cu_seqlens is None else len(cu_seqlens) - 1
BK = triton.next_power_of_2(K)
BV = 32
if out is None:
out = torch.zeros_like(v)
else:
assert out.shape == v.shape
if inplace_final_state:
assert initial_state is not None
final_state = initial_state
elif output_final_state:
if state_v_first:
final_state = q.new_empty(N, HV, V, K, dtype=torch.float32)
else:
final_state = q.new_empty(N, HV, K, V, dtype=torch.float32)
else:
final_state = None
stride_init_state_token = initial_state.stride(0) if initial_state is not None else 1
stride_final_state_token = final_state.stride(0) if final_state is not None else 1
if ssm_state_indices is None:
stride_indices_seq, stride_indices_tok = 1, 1
elif ssm_state_indices.ndim == 1:
stride_indices_seq, stride_indices_tok = ssm_state_indices.stride(0), 1
else:
stride_indices_seq, stride_indices_tok = ssm_state_indices.stride()
grid = (triton.cdiv(V, BV) * N * HV, )
fused_recurrent_kda_fwd_kernel[grid](
q=q,
k=k,
v=v,
g=g,
beta=beta,
A_log=A_log,
dt_bias=dt_bias,
o=out,
h0=initial_state,
ht=final_state,
cu_seqlens=cu_seqlens,
ssm_state_indices=ssm_state_indices,
num_accepted_tokens=num_accepted_tokens,
lower_bound=lower_bound,
scale=scale,
N=N,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BK=BK,
BV=BV,
stride_init_state_token=stride_init_state_token,
stride_final_state_token=stride_final_state_token,
stride_indices_seq=stride_indices_seq,
stride_indices_tok=stride_indices_tok,
IS_BETA_HEADWISE=beta.ndim == v.ndim,
USE_QK_L2NORM_IN_KERNEL=use_qk_l2norm_in_kernel,
INPLACE_FINAL_STATE=inplace_final_state,
USE_GATE_IN_KERNEL=use_gate_in_kernel,
APPLY_BETA_SIGMOID=use_beta_sigmoid_in_kernel,
ALLOW_NEG_EIGVAL=allow_neg_eigval,
STATE_V_FIRST=state_v_first,
num_warps=4,
num_stages=2,
)
return out, final_state
@input_guard
def fused_recurrent_kda(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
scale: float | None = None,
initial_state: torch.Tensor = None,
output_final_state: bool = False,
use_qk_l2norm_in_kernel: bool = False,
use_gate_in_kernel: bool = False,
use_beta_sigmoid_in_kernel: bool = False,
allow_neg_eigval: bool = False,
lower_bound: float | None = None,
state_v_first: bool = False,
cu_seqlens: torch.LongTensor | None = None,
**kwargs,
) -> tuple[torch.Tensor, torch.Tensor]:
r"""
Args:
q (torch.Tensor):
queries of shape `[B, T, H, K]`.
k (torch.Tensor):
keys of shape `[B, T, H, K]`.
v (torch.Tensor):
values of shape `[B, T, HV, V]`.
GVA is applied if `HV > H`.
g (torch.Tensor):
g (decays) of shape `[B, T, HV, K]`.
beta (torch.Tensor):
betas of shape `[B, T, HV]`.
A_log (Optional[torch.Tensor]):
Decay parameter of shape `[HV]`.
When `use_gate_in_kernel=True` together with `lower_bound`,
may be `None` to use `lower_bound * sigmoid(g + dt_bias)`.
dt_bias (Optional[torch.Tensor]):
Bias added to `g` before activation, of shape `[HV]`. Only used when `use_gate_in_kernel=True`.
scale (Optional[float]):
Scale factor for the RetNet attention scores.
If not provided, it will default to `1 / sqrt(K)`. Default: `None`.
initial_state (Optional[torch.Tensor]):
Initial state of shape `[N, HV, K, V]` for `N` input sequences.
For equal-length input sequences, `N` equals the batch size `B`.
Default: `None`.
output_final_state (Optional[bool]):
Whether to output the final state of shape `[N, HV, K, V]`. Default: `False`.
use_qk_l2norm_in_kernel (Optional[bool]):
Whether to use L2 normalization in the kernel. Default: `False`.
use_gate_in_kernel (Optional[bool]):
Whether to compute the log-space KDA decay internally.
When `True`, `g` is the raw input and the kernel fuses gate activation into the recurrence.
Default: `False`.
use_beta_sigmoid_in_kernel (Optional[bool]):
Whether to apply `torch.sigmoid(beta)` inside the kernel.
- If `True`, the passed `beta` acts as the raw beta logits.
- If `False`, `beta` is expected to already be in post-sigmoid space.
Default: `False`.
allow_neg_eigval (Optional[bool]):
Whether to allow negative eigenvalues by scaling `beta` to `[0, 2)`.
Only takes effect together with `use_beta_sigmoid_in_kernel=True`, in which case
the kernel computes `2 * sigmoid(beta)` instead of `sigmoid(beta)`. Default: `False`.
lower_bound (Optional[float]):
Lower bound for the forget gate (in log space). Only used when `use_gate_in_kernel=True`. Default: `None`.
state_v_first (Optional[bool]):
Store the recurrent state in V-first ``[V, K]`` layout instead of the default ``[K, V]``. Default: ``False``.
cu_seqlens (torch.LongTensor):
Cumulative sequence lengths of shape `[N+1]` used for variable-length training,
consistent with the FlashAttention API.
Returns:
o (torch.Tensor):
Outputs of shape `[B, T, HV, V]`.
final_state (torch.Tensor):
Final state of shape `[N, HV, K, V]` if `output_final_state=True` else `None`.
Examples::
>>> import torch
>>> import torch.nn.functional as F
>>> from einops import rearrange
>>> from fla.ops.kda import fused_recurrent_kda
# inputs with equal lengths
>>> B, T, H, HV, K, V = 4, 2048, 4, 8, 512, 512
>>> q = torch.randn(B, T, H, K, device='cuda')
>>> k = F.normalize(torch.randn(B, T, H, K, device='cuda'), p=2, dim=-1)
>>> v = torch.randn(B, T, HV, V, device='cuda')
>>> g = F.logsigmoid(torch.rand(B, T, HV, K, device='cuda'))
>>> beta = torch.rand(B, T, HV, device='cuda').sigmoid()
>>> h0 = torch.randn(B, HV, K, V, device='cuda')
>>> o, ht = fused_recurrent_kda(
q, k, v, g, beta,
initial_state=h0,
output_final_state=True
)
# for variable-length inputs, the batch size `B` is expected to be 1 and `cu_seqlens` is required
>>> q, k, v, g, beta = map(lambda x: rearrange(x, 'b t ... -> 1 (b t) ...'), (q, k, v, g, beta))
# for a batch with 4 sequences, `cu_seqlens` with 5 start/end positions are expected
>>> cu_seqlens = q.new_tensor([0, 2048, 4096, 6144, 8192], dtype=torch.long)
>>> o_var, ht_var = fused_recurrent_kda(
q, k, v, g, beta,
initial_state=h0,
output_final_state=True,
cu_seqlens=cu_seqlens
)
"""
if 'transpose_state_layout' in kwargs:
if state_v_first:
raise ValueError("Cannot pass both `state_v_first` and the deprecated `transpose_state_layout`.")
warnings.warn(
"`transpose_state_layout` is deprecated and renamed to `state_v_first`.",
DeprecationWarning,
stacklevel=2,
)
state_v_first = kwargs.pop('transpose_state_layout')
if cu_seqlens is not None:
if q.shape[0] != 1:
raise ValueError(
f"The batch size is expected to be 1 rather than {q.shape[0]} when using `cu_seqlens`."
f"Please flatten variable-length inputs before processing.",
)
if initial_state is not None and initial_state.shape[0] != len(cu_seqlens) - 1:
raise ValueError(
f"The number of initial states is expected to be equal to the number of input sequences, "
f"i.e., {len(cu_seqlens) - 1} rather than {initial_state.shape[0]}.",
)
if scale is None:
scale = k.shape[-1] ** -0.5
if allow_neg_eigval and not use_beta_sigmoid_in_kernel:
raise ValueError("`allow_neg_eigval=True` requires `use_beta_sigmoid_in_kernel=True`.")
o, final_state = fused_recurrent_kda_fwd(
q=q,
k=k,
v=v,
g=g,
beta=beta,
A_log=A_log,
dt_bias=dt_bias,
scale=scale,
initial_state=initial_state,
inplace_final_state=False,
output_final_state=output_final_state,
use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel,
use_gate_in_kernel=use_gate_in_kernel,
use_beta_sigmoid_in_kernel=use_beta_sigmoid_in_kernel,
allow_neg_eigval=allow_neg_eigval,
lower_bound=lower_bound,
cu_seqlens=cu_seqlens,
state_v_first=state_v_first,
)
return o, final_state
+514
View File
@@ -0,0 +1,514 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
# This file is modified and supported by the Moonshot AI Team
import torch
import torch.nn.functional as F
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.ops.utils.index import prepare_chunk_indices
from kda._fla.ops.utils.op import exp
from kda._fla.ops.utils.softplus import softplus
from kda._fla.utils import IS_AMD, autocast_custom_bwd, autocast_custom_fwd, autotune_cache_kwargs, check_shared_mem, input_guard
BS_LIST = [32, 64] if check_shared_mem() else [16, 32]
BT_LIST_AUTOTUNE = [32, 64, 128]
NUM_WARPS_AUTOTUNE = [2, 4, 8, 16] if IS_AMD else [4, 8, 16, 32]
def naive_kda_gate(
g: torch.Tensor,
A_log: torch.Tensor,
dt_bias: torch.Tensor | None = None,
output_dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
"""
Torch reference implementation for KDA gate computation.
Computes: g = -A_log.exp().unsqueeze(-1) * softplus(g + dt_bias.view(g.shape[-2:]))
Args:
g (torch.Tensor):
Input tensor of shape `[..., H, K]`.
A_log (torch.Tensor):
Parameter tensor with `H` elements.
dt_bias (torch.Tensor | None):
Optional bias tensor added to `g` before activation, shape `[H * K]`.
Returns:
Output tensor of shape `[..., H, K]` .
"""
H, _ = g.shape[-2:]
g = g.float()
if dt_bias is not None:
g = g + dt_bias.view(H, -1)
g = (-A_log.view(H, 1).float().exp() * F.softplus(g.float())).to(output_dtype)
return g
def naive_kda_lowerbound_gate(
g: torch.Tensor,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
lower_bound: float = -5.0,
output_dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
"""
Torch reference implementation for KDA lowerbound gate computation.
Computes: ``g = lower_bound * sigmoid(exp(A_log) * (g + dt_bias))``.
When ``A_log`` is ``None``: ``g = lower_bound * sigmoid(g + dt_bias)``.
Args:
g (torch.Tensor):
Input tensor of shape `[..., H, K]`.
A_log (torch.Tensor | None):
Optional parameter tensor with `H` elements.
dt_bias (torch.Tensor | None):
Optional bias tensor added to `g` before activation, shape `[H * K]`.
lower_bound (float):
Lower bound for the gate output. Default: `-5.0`.
output_dtype (torch.dtype):
The dtype of the output tensor. Default: `torch.float32`.
Returns:
Output tensor of shape `[..., H, K]`.
"""
H, _ = g.shape[-2:]
g = g.float()
if dt_bias is not None:
g = g + dt_bias.view(H, -1)
if A_log is not None:
g = A_log.view(H, 1).float().exp() * g
g = lower_bound * F.sigmoid(g)
return g.to(output_dtype)
@triton.heuristics({
"HAS_A": lambda args: args["A_log"] is not None,
"HAS_BIAS": lambda args: args["dt_bias"] is not None,
"HAS_BETA": lambda args: args["beta"] is not None,
'USE_LOWER_BOUND': lambda args: args['lower_bound'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({"BT": BT}, num_warps=num_warps, num_stages=num_stages)
for BT in BT_LIST_AUTOTUNE
for num_warps in NUM_WARPS_AUTOTUNE
for num_stages in [2, 3]
],
key=["H", "D"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def kda_gate_fwd_kernel(
g,
A_log,
dt_bias,
beta,
yg,
yb,
lower_bound,
T,
H: tl.constexpr,
D: tl.constexpr,
BT: tl.constexpr,
BD: tl.constexpr,
HAS_A: tl.constexpr,
HAS_BIAS: tl.constexpr,
HAS_BETA: tl.constexpr,
USE_LOWER_BOUND: tl.constexpr,
):
i_t, i_h = tl.program_id(0).to(tl.int64), tl.program_id(1)
b_A = tl.load(A_log + i_h).to(tl.float32) if HAS_A else 1.0
o_t = i_t * BT + tl.arange(0, BT)
o_d = tl.arange(0, BD)
m_t = o_t < T
m_g = m_t[:, None] & (o_d[None, :] < D)
p_g = g + i_h * D + o_t[:, None] * (H * D) + o_d[None, :]
p_yg = yg + i_h * D + o_t[:, None] * (H * D) + o_d[None, :]
# [BT, BD]
b_g = tl.load(p_g, mask=m_g, other=0.0).to(tl.float32)
if HAS_BIAS:
o_b = i_h * D + tl.arange(0, BD)
b_g = b_g + tl.load(dt_bias + o_b, mask=o_b < H * D, other=0.0).to(tl.float32)
if not USE_LOWER_BOUND:
b_yg = -exp(b_A) * softplus(b_g)
else:
b_yg = lower_bound * tl.sigmoid((exp(b_A) if HAS_A else b_A) * b_g)
tl.store(p_yg, b_yg.to(p_yg.dtype.element_ty), mask=m_g)
if HAS_BETA:
p_b = beta + i_h + o_t * H
p_yb = yb + i_h + o_t * H
b_yb = tl.sigmoid(tl.load(p_b, mask=m_t, other=0.0).to(tl.float32))
tl.store(p_yb, b_yb.to(p_yb.dtype.element_ty), mask=m_t)
@triton.heuristics({
"HAS_A": lambda args: args["A_log"] is not None,
"HAS_BIAS": lambda args: args["dt_bias"] is not None,
"HAS_BETA": lambda args: args["beta"] is not None,
'USE_LOWER_BOUND': lambda args: args['lower_bound'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({}, num_warps=num_warps, num_stages=num_stages)
for num_warps in NUM_WARPS_AUTOTUNE
for num_stages in [2, 3]
],
key=["H", "D"],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def kda_gate_bwd_kernel(
g,
A_log,
dt_bias,
beta,
dyg,
dyb,
dg,
dA,
dbeta,
lower_bound,
T,
H: tl.constexpr,
D: tl.constexpr,
BT: tl.constexpr,
BD: tl.constexpr,
HAS_A: tl.constexpr,
HAS_BIAS: tl.constexpr,
HAS_BETA: tl.constexpr,
USE_LOWER_BOUND: tl.constexpr,
):
i_t, i_h = tl.program_id(0).to(tl.int64), tl.program_id(1)
b_A = tl.load(A_log + i_h).to(tl.float32) if HAS_A else 1.0
o_t = i_t * BT + tl.arange(0, BT)
o_d = tl.arange(0, BD)
m_t = o_t < T
m_g = m_t[:, None] & (o_d[None, :] < D)
p_g = g + i_h * D + o_t[:, None] * (H * D) + o_d[None, :]
p_dg = dg + i_h * D + o_t[:, None] * (H * D) + o_d[None, :]
p_dyg = dyg + i_h * D + o_t[:, None] * (H * D) + o_d[None, :]
# [BT, BD]
b_g = tl.load(p_g, mask=m_g, other=0.0).to(tl.float32)
b_dyg = tl.load(p_dyg, mask=m_g, other=0.0).to(tl.float32)
if HAS_BIAS:
o_b = i_h * D + tl.arange(0, BD)
b_g = b_g + tl.load(dt_bias + o_b, mask=o_b < H * D, other=0.0).to(tl.float32)
# [BT, BD]
if not USE_LOWER_BOUND:
b_A = -exp(b_A)
b_yg = b_A * softplus(b_g)
b_dg = b_A * (b_dyg * tl.sigmoid(b_g))
b_dA = tl.sum(tl.sum(b_dyg * b_yg, 1), 0)
else:
b_A = exp(b_A) if HAS_A else b_A
b_inner = b_A * b_g
b_sig = tl.sigmoid(b_inner)
b_dsig = b_sig * (1.0 - b_sig)
# Common term: dy * (LB * dsig)
b_d_inner_term = b_dyg * (lower_bound * b_dsig)
# dg = d_inner_term * A
b_dg = b_d_inner_term * b_A
b_dA = tl.sum(tl.sum(b_dg * b_g, 1), 0) if HAS_A else 0.0
tl.store(p_dg, b_dg.to(p_dg.dtype.element_ty), mask=m_g)
if HAS_A:
tl.store(dA + i_t * H + i_h, b_dA)
if HAS_BETA:
p_b = beta + i_h + o_t * H
p_db = dbeta + i_h + o_t * H
p_dyb = dyb + i_h + o_t * H
b_b = tl.load(p_b, mask=m_t, other=0.0).to(tl.float32)
b_db = tl.load(p_dyb, mask=m_t, other=0.0).to(tl.float32) * b_b * (1.0 - b_b)
tl.store(p_db, b_db.to(p_db.dtype.element_ty), mask=m_t)
@dispatch('kda')
def kda_gate_fwd(
g: torch.Tensor,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
lower_bound: float | None = None,
output_dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
H, K = g.shape[-2:]
T = g.numel() // (H * K)
yg = torch.empty_like(g, dtype=output_dtype)
def grid(meta):
return (triton.cdiv(T, meta["BT"]), H)
kda_gate_fwd_kernel[grid](
g=g,
A_log=A_log,
dt_bias=dt_bias,
beta=None,
yg=yg,
yb=None,
T=T,
H=H,
D=K,
BD=triton.next_power_of_2(K),
lower_bound=lower_bound,
)
return yg
@dispatch('kda')
def kda_gate_bwd(
g: torch.Tensor,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
dyg: torch.Tensor | None = None,
lower_bound: float | None = None,
) -> tuple[torch.Tensor, torch.Tensor | None, torch.Tensor | None]:
H, K = g.shape[-2:]
T = g.numel() // (H * K)
BT = 32
NT = triton.cdiv(T, BT)
dg = torch.empty_like(g, dtype=torch.float32)
dA = g.new_empty(NT, H, dtype=torch.float32) if A_log is not None else None
grid = (triton.cdiv(T, BT), H)
kda_gate_bwd_kernel[grid](
g=g,
A_log=A_log,
dt_bias=dt_bias,
beta=None,
dyg=dyg,
dyb=None,
dg=dg,
dA=dA,
dbeta=None,
T=T,
H=H,
D=K,
BT=BT,
BD=triton.next_power_of_2(K),
lower_bound=lower_bound,
)
dg = dg.view_as(g).type_as(g)
dA = dA.sum(0).view_as(A_log).type_as(A_log) if A_log is not None else None
# dt_bias is [HV, K] in KDAAttention and [HV*K] in some FLA call sites.
dbias = (
dg.view(-1, H * K).sum(0).reshape_as(dt_bias).type_as(dt_bias)
if dt_bias is not None
else None
)
return dg, dA, dbias
class KDAGateFunction(torch.autograd.Function):
@staticmethod
@input_guard
@autocast_custom_fwd
def forward(
ctx,
g: torch.Tensor,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
lower_bound: float | None = None,
output_dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
yg = kda_gate_fwd(
g=g,
A_log=A_log,
dt_bias=dt_bias,
lower_bound=lower_bound,
output_dtype=output_dtype
)
ctx.save_for_backward(g, A_log, dt_bias)
ctx.lower_bound = lower_bound
return yg
@staticmethod
@input_guard
@autocast_custom_bwd
def backward(ctx, dyg: torch.Tensor):
g, A_log, dt_bias = ctx.saved_tensors
dg, dA, dbias = kda_gate_bwd(
g=g,
A_log=A_log,
dt_bias=dt_bias,
dyg=dyg,
lower_bound=ctx.lower_bound
)
return dg, dA, dbias, None, None
@dispatch('kda')
@torch.compiler.disable
def fused_kda_gate(
g: torch.Tensor,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
lower_bound: float | None = None,
output_dtype: torch.dtype = torch.float32,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
"""
Fused KDA gate computation with autograd support.
Computes: g = -A_log.exp().unsqueeze(-1) * softplus(g + dt_bias.view(g.shape[-2:]))
When ``lower_bound`` is set: g = lower_bound * sigmoid(exp(A_log) * (g + dt_bias)).
When ``A_log`` is ``None`` (requires ``lower_bound``): g = lower_bound * sigmoid(g + dt_bias).
Args:
g (torch.Tensor):
Input tensor of shape `[..., H, K]`.
A_log (torch.Tensor | None):
Optional parameter tensor with `H` elements.
When ``None``, the gate reduces to ``lower_bound * sigmoid(g + dt_bias)`` (requires ``lower_bound``).
dt_bias (torch.Tensor | None):
Optional bias tensor added to `g` before activation, shape `[H * K]`.
Returns:
Output tensor of shape `[..., H, K]`.
"""
return KDAGateFunction.apply(g, A_log, dt_bias, lower_bound, output_dtype)
@triton.heuristics({
"HAS_A": lambda args: args["A_log"] is not None,
"HAS_BIAS": lambda args: args["dt_bias"] is not None,
'HAS_SCALE': lambda args: args['scale'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
'USE_LOWER_BOUND': lambda args: args['lower_bound'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({'BS': BS}, num_warps=num_warps)
for BS in BS_LIST
for num_warps in [2, 4, 8]
],
key=['H', 'S', 'BT', 'IS_VARLEN', 'REVERSE'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def kda_gate_chunk_cumsum_vector_kernel(
s,
A_log,
dt_bias,
o,
scale,
cu_seqlens,
chunk_indices,
lower_bound,
T,
H: tl.constexpr,
S: tl.constexpr,
BT: tl.constexpr,
BS: tl.constexpr,
REVERSE: tl.constexpr,
HAS_A: tl.constexpr,
HAS_BIAS: tl.constexpr,
HAS_SCALE: tl.constexpr,
IS_VARLEN: tl.constexpr,
USE_LOWER_BOUND: tl.constexpr,
):
i_s, i_t, i_bh = tl.program_id(0), tl.program_id(1).to(tl.int64), tl.program_id(2).to(tl.int64)
i_b, i_h = i_bh // H, i_bh % H
if IS_VARLEN:
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
else:
bos, eos = i_b * T, i_b * T + T
o_t = i_t * BT + tl.arange(0, BT)
o_s = i_s * BS + tl.arange(0, BS)
m_s = (o_t[:, None] < T) & (o_s[None, :] < S)
p_s = s + (bos * H + i_h) * S + o_t[:, None] * (H*S) + o_s[None, :]
p_o = o + (bos * H + i_h) * S + o_t[:, None] * (H*S) + o_s[None, :]
# [BT, BS]
b_s = tl.load(p_s, mask=m_s, other=0.0).to(tl.float32)
# Apply dt_bias if exists
if HAS_BIAS:
b_bias = tl.load(dt_bias + i_h * S + o_s, mask=o_s < S, other=0.0).to(tl.float32)
b_s = b_s + b_bias[None, :]
b_A = tl.load(A_log + i_h).to(tl.float32) if HAS_A else 1.0
if not USE_LOWER_BOUND:
# Apply gate: -exp(A_log) * softplus(g + bias)
b_gate = -exp(b_A) * softplus(b_s)
else:
b_gate = lower_bound * tl.sigmoid((exp(b_A) if HAS_A else b_A) * b_s)
# Apply chunk local cumsum
if REVERSE:
b_o = tl.cumsum(b_gate, axis=0, reverse=True)
else:
b_o = tl.cumsum(b_gate, axis=0)
if HAS_SCALE:
b_o *= scale
tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=m_s)
@input_guard
@dispatch('kda')
def kda_gate_chunk_cumsum(
g: torch.Tensor,
A_log: torch.Tensor | None,
chunk_size: int,
scale: float = None,
dt_bias: torch.Tensor | None = None,
cu_seqlens: torch.Tensor | None = None,
output_dtype: torch.dtype | None = torch.float,
chunk_indices: torch.LongTensor | None = None,
lower_bound: float | None = None,
**kwargs,
) -> torch.Tensor:
if cu_seqlens is not None:
assert g.shape[0] == 1, "Only batch size 1 is supported when cu_seqlens are provided"
assert len(g.shape) == 4
B, T, H, S = g.shape
BT = chunk_size
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
assert chunk_size == 2**(chunk_size.bit_length()-1), "chunk_size must be a power of 2"
g_org, g = g, torch.empty_like(g, dtype=output_dtype or g.dtype)
def grid(meta): return (triton.cdiv(meta['S'], meta['BS']), NT, B * H)
kda_gate_chunk_cumsum_vector_kernel[grid](
s=g_org,
A_log=A_log,
dt_bias=dt_bias,
o=g,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
lower_bound=lower_bound,
T=T,
H=H,
S=S,
BT=BT,
REVERSE=False,
)
return g
+369
View File
@@ -0,0 +1,369 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import torch
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.ops.utils import prepare_chunk_indices
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.ops.utils.op import exp2
from kda._fla.utils import autotune_cache_kwargs, check_shared_mem
@triton.heuristics({
'STORE_QG': lambda args: args['qg'] is not None,
'STORE_KG': lambda args: args['kg'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({}, num_warps=num_warps, num_stages=num_stages)
for num_warps in [2, 4, 8]
for num_stages in [2, 3, 4]
],
key=['H', 'HV', 'K', 'V', 'BT', 'BK', 'BV', 'IS_VARLEN'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def recompute_w_u_fwd_kda_kernel(
q,
k,
qg,
kg,
v,
beta,
w,
u,
A,
gk,
cu_seqlens,
chunk_indices,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BK: tl.constexpr,
BV: tl.constexpr,
STORE_QG: tl.constexpr,
STORE_KG: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_t, i_bh = tl.program_id(0).to(tl.int64), tl.program_id(1).to(tl.int64)
i_b, i_hv = i_bh // HV, i_bh % HV
i_h = i_hv // (HV // H)
if IS_VARLEN:
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
else:
bos, eos = i_b * T, i_b * T + T
k += (bos * H + i_h) * K
v += (bos * HV + i_hv) * V
u += (bos * HV + i_hv) * V
w += (bos * HV + i_hv) * K
gk += (bos * HV + i_hv) * K
beta += bos * HV + i_hv
A += (bos * HV + i_hv) * BT
if STORE_QG:
q += (bos * H + i_h) * K
qg += (bos * HV + i_hv) * K
if STORE_KG:
kg += (bos * HV + i_hv) * K
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
p_b = beta + o_t * HV
b_b = tl.load(p_b, mask=m_t, other=0.0)
o_A = tl.arange(0, BT)
m_A = m_t[:, None] & (o_A[None, :] < BT)
p_A = A + o_t[:, None] * (HV*BT) + o_A[None, :]
b_A = tl.load(p_A, mask=m_A, other=0.0)
for i_v in range(tl.cdiv(V, BV)):
o_v = i_v * BV + tl.arange(0, BV)
m_v = m_t[:, None] & (o_v[None, :] < V)
p_v = v + o_t[:, None] * (HV*V) + o_v[None, :]
p_u = u + o_t[:, None] * (HV*V) + o_v[None, :]
b_v = tl.load(p_v, mask=m_v, other=0.0)
b_vb = (b_v * b_b[:, None]).to(b_v.dtype)
b_u = tl.dot(b_A, b_vb)
tl.store(p_u, b_u.to(p_u.dtype.element_ty), mask=m_v)
for i_k in range(tl.cdiv(K, BK)):
o_k = i_k * BK + tl.arange(0, BK)
m_k = o_k < K
m_tk = m_t[:, None] & m_k[None, :]
p_w = w + o_t[:, None] * (HV*K) + o_k[None, :]
p_k = k + o_t[:, None] * (H*K) + o_k[None, :]
b_k = tl.load(p_k, mask=m_tk, other=0.0)
b_kb = b_k * b_b[:, None]
p_gk = gk + o_t[:, None] * (HV*K) + o_k[None, :]
b_gk = tl.load(p_gk, mask=m_tk, other=0.0).to(tl.float32)
b_kb *= exp2(b_gk)
if STORE_QG:
p_q = q + o_t[:, None] * (H*K) + o_k[None, :]
p_qg = qg + o_t[:, None] * (HV*K) + o_k[None, :]
b_q = tl.load(p_q, mask=m_tk, other=0.0)
b_qg = b_q * exp2(b_gk)
tl.store(p_qg, b_qg.to(p_qg.dtype.element_ty), mask=m_tk)
if STORE_KG:
last_idx = min(i_t * BT + BT, T) - 1
b_gn = tl.load(gk + last_idx * HV*K + o_k, mask=m_k, other=0.).to(tl.float32)
b_kg = b_k * tl.where((i_t * BT + tl.arange(0, BT) < T)[:, None], exp2(b_gn[None, :] - b_gk), 0)
p_kg = kg + o_t[:, None] * (HV*K) + o_k[None, :]
tl.store(p_kg, b_kg.to(p_kg.dtype.element_ty), mask=m_tk)
b_w = tl.dot(b_A, b_kb.to(b_k.dtype))
tl.store(p_w, b_w.to(p_w.dtype.element_ty), mask=m_tk)
@triton.heuristics({
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({}, num_warps=num_warps, num_stages=num_stages)
for num_warps in [2, 4]
for num_stages in [2, 3, 4]
],
key=['H', 'HV', 'K', 'V', 'BT', 'BK', 'BV', 'IS_VARLEN'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def prepare_wy_repr_bwd_kda_kernel(
k,
v,
beta,
gk,
A,
dA,
dw,
du,
dk,
dk2,
dv,
db,
dg,
dg2,
cu_seqlens,
chunk_indices,
T,
H: tl.constexpr,
HV: tl.constexpr,
K: tl.constexpr,
V: tl.constexpr,
BT: tl.constexpr,
BK: tl.constexpr,
BV: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_t, i_bh = tl.program_id(0).to(tl.int64), tl.program_id(1).to(tl.int64)
i_b, i_hv = i_bh // HV, i_bh % HV
i_h = i_hv // (HV // H)
if IS_VARLEN:
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
else:
bos, eos = i_b * T, i_b * T + T
k += (bos * H + i_h) * K
v += (bos * HV + i_hv) * V
beta += bos * HV + i_hv
gk += (bos * HV + i_hv) * K
A += (bos * HV + i_hv) * BT
dA += (bos * HV + i_hv) * BT
dk += (bos * HV + i_hv) * K
dk2 += (bos * HV + i_hv) * K
dw += (bos * HV + i_hv) * K
du += (bos * HV + i_hv) * V
dv += (bos * HV + i_hv) * V
db += bos * HV + i_hv
dg += (bos * HV + i_hv) * K
dg2 += (bos * HV + i_hv) * K
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
p_b = beta + o_t * HV
p_db = db + o_t * HV
o_A = tl.arange(0, BT)
m_AT = (o_A[:, None] < BT) & m_t[None, :]
p_A = A + o_A[:, None] + o_t[None, :] * (HV*BT)
b_b = tl.load(p_b, mask=m_t, other=0.0)
b_db = tl.zeros([BT], dtype=tl.float32)
b_A = tl.load(p_A, mask=m_AT, other=0.0)
b_dA = tl.zeros([BT, BT], dtype=tl.float32)
for i_k in range(tl.cdiv(K, BK)):
o_k = i_k * BK + tl.arange(0, BK)
m_k = m_t[:, None] & (o_k[None, :] < K)
p_k = k + o_t[:, None] * (H*K) + o_k[None, :]
p_dk = dk + o_t[:, None] * (HV*K) + o_k[None, :]
p_dk2 = dk2 + o_t[:, None] * (HV*K) + o_k[None, :]
p_dw = dw + o_t[:, None] * (HV*K) + o_k[None, :]
p_dg = dg + o_t[:, None] * (HV*K) + o_k[None, :]
p_dg2 = dg2 + o_t[:, None] * (HV*K) + o_k[None, :]
# [BT, BK]
b_k = tl.load(p_k, mask=m_k, other=0.0)
p_gk = gk + o_t[:, None] * (HV*K) + o_k[None, :]
b_gk_exp = exp2(tl.load(p_gk, mask=m_k, other=0.0))
b_kbg = b_k * b_b[:, None] * b_gk_exp
b_dw = tl.load(p_dw, mask=m_k, other=0.0)
b_dA += tl.dot(b_dw, tl.trans(b_kbg).to(b_dw.dtype))
b_dkbg = tl.dot(b_A, b_dw)
b_dk = b_dkbg * b_gk_exp * b_b[:, None] + tl.load(p_dk, mask=m_k, other=0.0)
b_db += tl.sum(b_dkbg * b_k * b_gk_exp, 1)
b_dg = b_kbg * b_dkbg + tl.load(p_dg, mask=m_k, other=0.0)
tl.store(p_dk2, b_dk.to(p_dk2.dtype.element_ty), mask=m_k)
tl.store(p_dg2, b_dg.to(p_dg2.dtype.element_ty), mask=m_k)
for i_v in range(tl.cdiv(V, BV)):
o_v = i_v * BV + tl.arange(0, BV)
m_v = m_t[:, None] & (o_v[None, :] < V)
p_v = v + o_t[:, None] * (HV*V) + o_v[None, :]
p_dv = dv + o_t[:, None] * (HV*V) + o_v[None, :]
p_du = du + o_t[:, None] * (HV*V) + o_v[None, :]
b_v = tl.load(p_v, mask=m_v, other=0.0)
b_vb = (b_v * b_b[:, None]).to(b_v.dtype)
b_du = tl.load(p_du, mask=m_v, other=0.0)
b_dA += tl.dot(b_du, tl.trans(b_vb))
b_dvb = tl.dot(b_A, b_du)
b_dv = b_dvb * b_b[:, None]
b_db += tl.sum(b_dvb * b_v, 1)
tl.store(p_dv, b_dv.to(p_dv.dtype.element_ty), mask=m_v)
m_A = (o_t[:, None] > o_t[None, :]) & (m_t[:, None] & m_t)
b_dA = tl.where(m_A, b_dA, 0)
b_dA = tl.dot(b_dA.to(b_A.dtype), b_A)
b_dA = tl.dot(b_A, b_dA.to(b_A.dtype))
b_dA = tl.where(m_A, -b_dA, 0)
m_dA = m_t[:, None] & (o_A[None, :] < BT)
p_dA = dA + o_t[:, None] * (HV*BT) + o_A[None, :]
tl.store(p_dA, b_dA.to(p_dA.dtype.element_ty), mask=m_dA)
tl.store(p_db, b_db.to(p_db.dtype.element_ty), mask=m_t)
@dispatch('kda')
def recompute_w_u_fwd(
k: torch.Tensor,
v: torch.Tensor,
beta: torch.Tensor,
A: torch.Tensor,
gk: torch.Tensor,
q: torch.Tensor | None = None,
cu_seqlens: torch.LongTensor | None = None,
chunk_indices: torch.LongTensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None, torch.Tensor | None]:
B, T, H, K, V = *k.shape, v.shape[-1]
HV = v.shape[2]
BT = A.shape[-1]
BK = 64
BV = 64
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
w = torch.empty(B, T, HV, K, device=k.device, dtype=k.dtype)
u = torch.empty_like(v)
qg = torch.empty(B, T, HV, K, device=k.device, dtype=k.dtype) if q is not None else None
kg = torch.empty(B, T, HV, K, device=k.device, dtype=k.dtype)
recompute_w_u_fwd_kda_kernel[(NT, B*HV)](
q=q,
k=k,
qg=qg,
kg=kg,
v=v,
beta=beta,
w=w,
u=u,
A=A,
gk=gk,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BT=BT,
BK=BK,
BV=BV,
)
return w, u, qg, kg
def prepare_wy_repr_bwd(
k: torch.Tensor,
v: torch.Tensor,
beta: torch.Tensor,
gk: torch.Tensor,
A: torch.Tensor,
dk: torch.Tensor,
dw: torch.Tensor,
du: torch.Tensor,
dg: torch.Tensor,
cu_seqlens: torch.LongTensor | None = None,
chunk_indices: torch.LongTensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
B, T, H, K, V = *k.shape, v.shape[-1]
HV = v.shape[2]
BT = A.shape[-1]
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
CONST_TILING = 64 if check_shared_mem() else 32
BK = min(max(triton.next_power_of_2(K), 16), CONST_TILING)
BV = min(max(triton.next_power_of_2(V), 16), CONST_TILING)
dk2 = torch.empty_like(dk, dtype=torch.float)
dv = torch.empty_like(v)
dg2 = torch.empty_like(gk, dtype=torch.float)
dA = torch.empty_like(A, dtype=torch.float)
db = torch.empty_like(beta, dtype=torch.float)
prepare_wy_repr_bwd_kda_kernel[(NT, B * HV)](
k=k,
v=v,
beta=beta,
gk=gk,
A=A,
dA=dA,
dw=dw,
du=du,
dk=dk,
dk2=dk2,
dv=dv,
db=db,
dg=dg,
dg2=dg2,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
T=T,
H=H,
HV=HV,
K=K,
V=V,
BT=BT,
BK=BK,
BV=BV,
)
dk = dk2
dg = dg2
return dk, dv, db, dg, dA
+14
View File
@@ -0,0 +1,14 @@
from .cumsum import (
chunk_local_cumsum,
chunk_local_cumsum_scalar,
chunk_local_cumsum_vector,
)
from .index import prepare_chunk_indices, prepare_chunk_offsets
__all__ = [
"chunk_local_cumsum",
"chunk_local_cumsum_scalar",
"chunk_local_cumsum_vector",
"prepare_chunk_indices",
"prepare_chunk_offsets",
]
+449
View File
@@ -0,0 +1,449 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import dataclasses
import enum
import json
import logging
import os
import re
from functools import cache, lru_cache
from pathlib import Path
from typing import Any
import torch
import triton
from packaging import version
from triton.runtime.autotuner import Autotuner
TRITON_ABOVE_3_5_1 = version.parse(triton.__version__) >= version.parse("3.5.1")
TRITON_ABOVE_3_4_0 = version.parse(triton.__version__) >= version.parse("3.4.0")
class FlaCacheMode(enum.Enum):
"""Controls how FLA loads kernel configs from its config cache (FLA_CACHE_MODE env var).
DISABLED — skip all cache lookups, always fall back to Triton autotune (default when FLA_CACHE_MODE is unset)
STRICT — exact key match only; falls back to Triton autotune if no match
FUZZY — exact key match → fuzzy key match; falls back to Triton autotune if no match
FULL — exact key match → fuzzy key match → default_config fallback
DEFAULT — use only the top-level default_config field, skip key-based lookup
ALWAYS — like DEFAULT, but re-reads config files on every kernel call;
useful for debugging: edit default_config in a JSON file and the next
kernel call picks it up without restarting the process
"""
DISABLED = "disabled"
STRICT = "strict"
FUZZY = "fuzzy"
FULL = "full"
DEFAULT = "default"
ALWAYS = "always"
def uses_default_config(self) -> bool:
"""Return True for modes that may fall back to default_config (FULL, DEFAULT, ALWAYS)."""
return self in (FlaCacheMode.FULL, FlaCacheMode.DEFAULT, FlaCacheMode.ALWAYS)
@classmethod
def from_env(cls) -> "FlaCacheMode":
mode_str = os.environ.get("FLA_CACHE_MODE", cls.DISABLED.value)
try:
return cls(mode_str)
except ValueError:
valid = [m.value for m in cls]
raise ValueError(
f"Invalid FLA_CACHE_MODE={mode_str!r}. Valid values: {valid}"
) from None
FLA_CACHE_MODE: FlaCacheMode = FlaCacheMode.from_env()
logger = logging.getLogger(__name__)
def sanitize_gpu_name(gpu_name: str) -> str:
sanitized = re.sub(r"[^0-9A-Za-z]+", "_", gpu_name)
sanitized = sanitized.strip("_")
return sanitized or "unknown_gpu"
@lru_cache(maxsize=1)
def get_gpu_info():
"""Get GPU model information.
This function detects the GPU model and returns a sanitized string identifier.
It prioritizes FLA_GPU_NAME environment variable if set, then detects from
available hardware (CUDA, ROCm, Intel GPU, or CPU).
"""
# Check if GPU name is overridden via environment variable
gpu_name = None
# Check if GPU name is overridden via environment variable
if "FLA_GPU_NAME" in os.environ:
gpu_name = os.environ["FLA_GPU_NAME"]
# Try to get device name based on availability
elif torch.cuda.is_available():
# Works for both NVIDIA and AMD GPUs (ROCm)
gpu_name = torch.cuda.get_device_name(0)
elif hasattr(torch, 'xpu') and torch.xpu.is_available():
gpu_name = torch.xpu.get_device_name(0)
if gpu_name:
return sanitize_gpu_name(gpu_name)
# Default to CPU if no GPU available
return "cpu"
def get_fla_config_dir() -> Path:
"""Get FLA's configs directory.
The directory can be overridden by setting the FLA_CONFIG_DIR environment variable.
If set, configs will be loaded directly from $FLA_CONFIG_DIR/. Otherwise FLA
falls back to the default fla/configs/{GPU}/ directory in the project.
"""
# Check if custom config dir is set via environment variable
if "FLA_CONFIG_DIR" in os.environ:
return Path(os.environ["FLA_CONFIG_DIR"])
# Default: project_dir/fla/configs/{GPU}/
project_dir = Path(__file__).parent.parent.parent
return project_dir / "configs" / get_gpu_info()
@dataclasses.dataclass(frozen=True)
class AutotuneKey:
"""Autotune key with exact/fuzzy matching, serialization, and construction helpers."""
autotune_key: tuple[Any, ...]
@staticmethod
def normalize_autotune_key(value: Any) -> Any:
if isinstance(value, (list, tuple)):
return [AutotuneKey.normalize_autotune_key(v) for v in value]
if isinstance(value, dict):
return {k: AutotuneKey.normalize_autotune_key(v) for k, v in value.items()}
return value
@staticmethod
def serialize(key: Any) -> str:
return json.dumps(AutotuneKey.normalize_autotune_key(key), separators=(",", ":"), sort_keys=True)
@staticmethod
def key_hash(key: Any) -> str:
import hashlib
return hashlib.md5(AutotuneKey.serialize(key).encode()).hexdigest()
@staticmethod
def is_numeric(value: Any) -> bool:
return isinstance(value, (int, float)) and not isinstance(value, bool)
@staticmethod
def keys_fuzzy_match(cached_key: Any, requested_key: Any) -> bool:
# Fuzzy match: numeric leaves are compatible regardless of their actual numeric values
# (e.g. a config tuned for seq_len=1024 can apply to seq_len=2048).
# Structure (type, length, dict keys) must still match exactly.
if AutotuneKey.is_numeric(cached_key) and AutotuneKey.is_numeric(requested_key):
return True
if isinstance(cached_key, (list, tuple)) and isinstance(requested_key, (list, tuple)):
return len(cached_key) == len(requested_key) and all(
AutotuneKey.keys_fuzzy_match(c, r) for c, r in zip(cached_key, requested_key)
)
if isinstance(cached_key, dict) and isinstance(requested_key, dict):
return cached_key.keys() == requested_key.keys() and all(
AutotuneKey.keys_fuzzy_match(cached_key[k], requested_key[k]) for k in cached_key
)
return cached_key == requested_key
@classmethod
def build(
cls,
arg_names: list[str],
key_names: list[str],
positional_args: tuple[Any, ...],
runtime_kwargs: dict[str, Any],
) -> "AutotuneKey":
named_args = dict(zip(arg_names, positional_args))
all_args = {**named_args, **runtime_kwargs}
tracked_args = {k: v for (k, v) in all_args.items() if k in arg_names}
tuning_key = [tracked_args[name] for name in key_names if name in tracked_args]
for arg in tracked_args.values():
if hasattr(arg, "dtype"):
tuning_key.append(str(arg.dtype))
return cls(autotune_key=tuple(tuning_key))
def exact_matches(self, entry_key: Any) -> bool:
return self.serialize(self.autotune_key) == self.serialize(entry_key)
def fuzzy_matches(self, entry_key: Any) -> bool:
self_normalized = self.normalize_autotune_key(self.autotune_key)
entry_normalized = self.normalize_autotune_key(entry_key)
return (
isinstance(self_normalized, list)
and isinstance(entry_normalized, list)
and len(self_normalized) == len(entry_normalized)
and AutotuneKey.keys_fuzzy_match(self_normalized, entry_normalized)
)
@dataclasses.dataclass(frozen=True)
class KernelConfigFile:
"""Validated in-memory representation of a {kernel_name}.json config file."""
kernel_name: str | None
triton_version: str | None
autotune_entries: dict[str, dict[str, Any]] | None
default_config: dict[str, Any] | None
@classmethod
def from_dict(cls, config_file: Path, data: Any) -> "KernelConfigFile | None":
"""Parse and validate a raw JSON dict. Returns None (with a warning) if malformed."""
def fail(msg, *args):
logger.warning(msg, *args)
raise ValueError
try:
if not isinstance(data, dict):
fail("Malformed config %s: root is %s, expected dict", config_file, type(data).__name__)
raw_entries = data.get("autotune_entries")
entries: dict[str, dict[str, Any]] | None = None
if raw_entries is not None:
if not isinstance(raw_entries, dict):
fail("Malformed config %s: 'autotune_entries' is %s, expected dict",
config_file, type(raw_entries).__name__)
for h, entry in raw_entries.items():
if not isinstance(entry, dict):
fail("Malformed config %s: autotune_entries[%r] is %s, expected dict",
config_file, h, type(entry).__name__)
if not isinstance(entry.get("config"), dict):
fail("Malformed config %s: autotune_entries[%r] missing valid 'config' field", config_file, h)
entries = raw_entries
default_config = data.get("default_config")
if default_config is not None and not isinstance(default_config, dict):
fail("Malformed config %s: 'default_config' is %s, expected dict", config_file, type(default_config).__name__)
return cls(
kernel_name=data.get("kernel_name"),
triton_version=data.get("triton_version"),
autotune_entries=entries,
default_config=default_config,
)
except ValueError:
return None
@classmethod
def from_file(cls, config_file: Path) -> "KernelConfigFile | None":
"""Read and validate a config file. Returns None if the file is missing or malformed."""
config_data = read_config_file(config_file)
if config_data is None:
return None
return cls.from_dict(config_file, config_data)
def lookup_exact(self, key: AutotuneKey) -> dict[str, Any] | None:
if self.autotune_entries is None:
return None
return self.autotune_entries.get(AutotuneKey.key_hash(key.autotune_key))
def lookup_fuzzy(self, key: AutotuneKey) -> dict[str, Any] | None:
if self.autotune_entries is None:
return None
for entry in self.autotune_entries.values():
if key.fuzzy_matches(entry.get("autotune_key")):
return entry
return None
@cache
def load_config_file(config_file: Path) -> dict[str, Any] | None:
try:
with open(config_file) as f:
return json.load(f)
except Exception as e:
logger.warning("Error reading config file %s: %s", config_file, e)
return None
def read_config_file(config_file: Path) -> dict[str, Any] | None:
"""Read a config file, bypassing the in-process cache in ALWAYS mode."""
if FLA_CACHE_MODE is FlaCacheMode.ALWAYS:
return load_config_file.__wrapped__(config_file)
return load_config_file(config_file)
def load_cached_config(kernel_name: str, autotune_key: AutotuneKey | None = None) -> dict[str, Any] | None:
"""
Load cached best config for a kernel from FLA configs directory.
This function loads the cached best configuration for a given kernel name
from get_fla_config_dir()/{kernel_name}.json.
Cache files may contain multiple autotune entries keyed by Triton's
runtime tuning key plus a top-level default config.
If the config file is not found or cannot be loaded, a warning is printed
and None is returned, allowing fallback to Triton's autotune.
The lookup mode is controlled by the FLA_CACHE_MODE environment variable (see FlaCacheMode).
Args:
kernel_name: Name of the kernel (e.g., "causal_conv1d_fwd_kernel")
autotune_key: Triton autotune key for the current invocation
Returns:
Best config dictionary or None if not found or disabled
"""
if FLA_CACHE_MODE is FlaCacheMode.DISABLED:
return None
config_dir = get_fla_config_dir()
config_file = config_dir / f"{kernel_name}.json"
if not config_file.exists():
return None
config_data = read_config_file(config_file)
if config_data is None:
return None
config = KernelConfigFile.from_dict(config_file, config_data)
if config is None:
return None
if FLA_CACHE_MODE is FlaCacheMode.DEFAULT or FLA_CACHE_MODE is FlaCacheMode.ALWAYS:
return config.default_config
# STRICT mode: exact match only, no fuzzy fallback
if FLA_CACHE_MODE is FlaCacheMode.STRICT:
if autotune_key is not None:
entry = config.lookup_exact(autotune_key)
if entry is not None:
return entry["config"]
return None
# FULL and FUZZY modes: try exact key match first, then fuzzy match
if autotune_key is not None:
entry = config.lookup_exact(autotune_key) or config.lookup_fuzzy(autotune_key)
if entry is not None:
return entry["config"]
if FLA_CACHE_MODE is FlaCacheMode.FUZZY:
return None
# FULL mode: fall back to default_config, then legacy raw config (no autotune_entries)
if config.default_config is not None:
return config.default_config
if config.autotune_entries is not None:
return None
return config_data
class CachedAutotuner(Autotuner):
"""
A modified autotuner that loads best config from FLA's config directory.
This class extends Triton's Autotuner but overrides the run method to
try loading cached configuration first before falling back to autotune.
"""
def __init__(self, fn, arg_names, configs, key, reset_to_zero, restore_value, **kwargs):
super().__init__(fn, arg_names, configs, key, reset_to_zero, restore_value, **kwargs)
self.kernel_name = fn.fn.__name__ if hasattr(fn, 'fn') else fn.__name__
# None-safe pre/post hooks: Triton's defaults crash when a restore_value / reset_to_zero arg
# is None (idiomatic for optional pointers gated by a tl.constexpr flag).
# Fixed upstream in triton-lang/triton#10295 — remove this override once FLA's minimum Triton version has it.
if not self.user_defined_pre_hook and (self.reset_to_zero or self.restore_value):
def _pre_hook(kw, reset_only=False):
for n in self.reset_to_zero:
if kw[n] is not None:
kw[n].zero_()
if not reset_only:
self.restore_copies = {n: kw[n].clone() for n in self.restore_value if kw[n] is not None}
self.pre_hook = _pre_hook
if not self.user_defined_post_hook and self.restore_value:
def _post_hook(kw, exception):
for n, copy in self.restore_copies.items():
kw[n].copy_(copy)
self.restore_copies = {}
self.post_hook = _post_hook
def should_check_fla_cache(self, key: AutotuneKey) -> bool:
if FLA_CACHE_MODE is FlaCacheMode.DISABLED:
return False
if FLA_CACHE_MODE is FlaCacheMode.ALWAYS:
return True
return key.autotune_key not in self.cache
def run(self, *args, **kwargs):
key = AutotuneKey.build(self.arg_names, self.keys, args, kwargs)
if self.should_check_fla_cache(key):
self.maybe_load_cached_config(key)
return super().run(*args, **kwargs)
def maybe_load_cached_config(self, key: AutotuneKey):
best_config = load_cached_config(self.kernel_name, key)
if best_config is not None:
kw = best_config["kwargs"]
num_warps = best_config["num_warps"]
num_stages = best_config["num_stages"]
extra = {
"num_ctas": best_config["num_ctas"],
"maxnreg": best_config.get("maxnreg"),
"pre_hook": None,
"ir_override": best_config.get("ir_override"),
} if TRITON_ABOVE_3_5_1 else {}
cfg = triton.Config(kw, num_warps=num_warps, num_stages=num_stages, **extra)
self.cache[key.autotune_key] = cfg
else:
logger.debug(
"No cached config found for kernel %s and key %s; falling back to Triton autotune",
self.kernel_name,
list(key.autotune_key),
)
def fla_cache_autotune(configs, key=None, prune_configs_by=None, reset_to_zero=None, restore_value=None,
pre_hook=None, post_hook=None, warmup=None, rep=None, use_cuda_graph=False,
do_bench=None, cache_results=False):
"""
Decorator for auto-tuning a :code:`triton.jit`'d function with FLA config support.
Extends Triton's autotune to load best configurations from FLA's config directory
(default: fla/configs/{GPU}/, or FLA_CONFIG_DIR/ when overridden), keyed by kernel
name from {kernel_name}.json. Lookup behaviour is controlled by FLA_CACHE_MODE.
Falls back to normal Triton autotuning when no cached config is found.
"""
# key can be None when we want to use cache only (no fallback autotune)
if key is None:
key = []
def decorator(fn):
kwargs = {}
if TRITON_ABOVE_3_4_0:
kwargs = {"cache_results": cache_results}
return CachedAutotuner(fn, fn.arg_names, configs, key, reset_to_zero, restore_value,
pre_hook=pre_hook, post_hook=post_hook,
prune_configs_by=prune_configs_by, warmup=warmup, rep=rep,
use_cuda_graph=use_cuda_graph, do_bench=do_bench,
**kwargs,
)
return decorator
def configure_fla_cache_autotune():
triton.autotune = fla_cache_autotune
logger.info(
"configure_fla_cache_autotune() is enabling FLA fla_cache_autotune; "
"triton.autotune will be replaced with fla_cache_autotune."
)
def restore_autotune_backend():
from triton.runtime.autotuner import autotune as original_autotune
triton.autotune = original_autotune
logger.info(
"restore_autotune_backend() is restoring Triton's original autotune; "
"triton.autotune will be replaced with triton.runtime.autotuner.autotune."
)
+10
View File
@@ -0,0 +1,10 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
# Approximate value of 1/ln(2), used for log/exp base conversion
# Best FP32 approximation: 1.4426950216 (hex 0x3FB8AA3B)
RCP_LN2 = 1.4426950216
+468
View File
@@ -0,0 +1,468 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import torch
import triton
import triton.language as tl
from kda._fla.ops.backends import dispatch
from kda._fla.ops.utils.cache import fla_cache_autotune
from kda._fla.ops.utils.index import prepare_chunk_indices
from kda._fla.utils import autotune_cache_kwargs, check_shared_mem, input_guard
BS_LIST = [32, 64] if check_shared_mem() else [16, 32]
@triton.heuristics({
'HAS_SCALE': lambda args: args['scale'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({}, num_warps=num_warps)
for num_warps in [1, 2, 4, 8]
],
key=['B', 'H', 'BT', 'IS_VARLEN', 'REVERSE'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_local_cumsum_scalar_kernel(
s,
o,
scale,
cu_seqlens,
chunk_indices,
T,
B: tl.constexpr,
H: tl.constexpr,
BT: tl.constexpr,
REVERSE: tl.constexpr,
HAS_SCALE: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_t, i_bh = tl.program_id(0).to(tl.int64), tl.program_id(1).to(tl.int64)
i_b, i_h = i_bh // H, i_bh % H
if IS_VARLEN:
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
else:
bos, eos = i_b * T, i_b * T + T
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
p_s = s + bos*H + i_h + o_t * H
p_o = o + bos*H + i_h + o_t * H
# [BT]
b_s = tl.load(p_s, mask=m_t, other=0.0).to(tl.float32)
if REVERSE:
b_o = tl.cumsum(b_s, axis=0, reverse=True)
else:
b_o = tl.cumsum(b_s, axis=0)
if HAS_SCALE:
b_o *= scale
tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=m_t)
@triton.heuristics({
'HAS_SCALE': lambda args: args['scale'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@fla_cache_autotune(
configs=[
triton.Config({'BS': BS}, num_warps=num_warps)
for BS in BS_LIST
for num_warps in [2, 4, 8]
],
key=['B', 'H', 'S', 'BT', 'IS_VARLEN', 'REVERSE'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_local_cumsum_vector_kernel(
s,
o,
scale,
cu_seqlens,
chunk_indices,
T,
B: tl.constexpr,
H: tl.constexpr,
S: tl.constexpr,
BT: tl.constexpr,
BS: tl.constexpr,
REVERSE: tl.constexpr,
HAS_SCALE: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_s, i_t, i_bh = tl.program_id(0), tl.program_id(1).to(tl.int64), tl.program_id(2).to(tl.int64)
i_b, i_h = i_bh // H, i_bh % H
if IS_VARLEN:
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
T = eos - bos
else:
bos, eos = i_b * T, i_b * T + T
o_t = i_t * BT + tl.arange(0, BT)
o_s = i_s * BS + tl.arange(0, BS)
m_s = (o_t[:, None] < T) & (o_s[None, :] < S)
p_s = s + (bos * H + i_h) * S + o_t[:, None] * (H*S) + o_s[None, :]
p_o = o + (bos * H + i_h) * S + o_t[:, None] * (H*S) + o_s[None, :]
# [BT, BS]
b_s = tl.load(p_s, mask=m_s, other=0.0).to(tl.float32)
if REVERSE:
b_o = tl.cumsum(b_s, axis=0, reverse=True)
else:
b_o = tl.cumsum(b_s, axis=0)
if HAS_SCALE:
b_o *= scale
tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=m_s)
@triton.heuristics({
'HAS_SCALE': lambda args: args['scale'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@triton.autotune(
configs=[
triton.Config({'BT': BT}, num_warps=num_warps, num_stages=num_stages)
for BT in [32, 64, 128, 256]
for num_warps in [2, 4, 8]
for num_stages in [1, 2, 3, 4]
],
key=['B', 'H', 'IS_VARLEN', 'REVERSE'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_global_cumsum_scalar_kernel(
s,
o,
scale,
cu_seqlens,
T,
B: tl.constexpr,
H: tl.constexpr,
BT: tl.constexpr,
REVERSE: tl.constexpr,
HAS_SCALE: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_nh = tl.program_id(0).to(tl.int64)
i_n, i_h = i_nh // H, i_nh % H
if IS_VARLEN:
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
else:
bos, eos = i_n * T, i_n * T + T
T = eos - bos
b_z = tl.zeros([], dtype=tl.float32)
NT = tl.cdiv(T, BT)
for i_c in range(NT):
i_t = NT - 1 - i_c if REVERSE else i_c
o_t = i_t * BT + tl.arange(0, BT)
m_t = o_t < T
p_s = s + bos*H + i_h + o_t * H
p_o = o + bos*H + i_h + o_t * H
b_s = tl.load(p_s, mask=m_t, other=0.0).to(tl.float32)
if REVERSE:
b_o = tl.cumsum(b_s, axis=0, reverse=True)
else:
b_o = tl.cumsum(b_s, axis=0)
b_ss = tl.sum(b_s, 0)
b_o += b_z
if i_c >= 0:
b_z += b_ss
if HAS_SCALE:
b_o *= scale
tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=m_t)
@triton.heuristics({
'HAS_SCALE': lambda args: args['scale'] is not None,
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@triton.autotune(
configs=[
triton.Config({'BT': BT}, num_warps=num_warps, num_stages=num_stages)
for BT in [16, 32, 64, 128]
for num_warps in [2, 4, 8]
for num_stages in [1, 2, 3, 4]
],
key=['B', 'H', 'S', 'IS_VARLEN', 'REVERSE'],
**autotune_cache_kwargs,
)
@triton.jit(do_not_specialize=['T'])
def chunk_global_cumsum_vector_kernel(
s,
o,
scale,
cu_seqlens,
T,
B: tl.constexpr,
H: tl.constexpr,
S: tl.constexpr,
BT: tl.constexpr,
BS: tl.constexpr,
REVERSE: tl.constexpr,
HAS_SCALE: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_s, i_nh = tl.program_id(0), tl.program_id(1).to(tl.int64)
i_n, i_h = i_nh // H, i_nh % H
if IS_VARLEN:
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
else:
bos, eos = i_n * T, i_n * T + T
T = eos - bos
b_z = tl.zeros([BS], dtype=tl.float32)
NT = tl.cdiv(T, BT)
for i_c in range(NT):
i_t = NT - 1 - i_c if REVERSE else i_c
o_t = i_t * BT + tl.arange(0, BT)
o_s = i_s * BS + tl.arange(0, BS)
m_s = (o_t[:, None] < T) & (o_s[None, :] < S)
p_s = s + (bos * H + i_h) * S + o_t[:, None] * (H*S) + o_s[None, :]
p_o = o + (bos * H + i_h) * S + o_t[:, None] * (H*S) + o_s[None, :]
# [BT, BS]
b_s = tl.load(p_s, mask=m_s, other=0.0).to(tl.float32)
if REVERSE:
b_c = b_z[None, :] + tl.cumsum(b_s, axis=0, reverse=True)
else:
b_c = b_z[None, :] + tl.cumsum(b_s, axis=0)
if HAS_SCALE:
b_c *= scale
tl.store(p_o, b_c.to(p_o.dtype.element_ty), mask=m_s)
b_z += tl.sum(b_s, 0)
def chunk_local_cumsum_scalar(
g: torch.Tensor,
chunk_size: int,
reverse: bool = False,
scale: float = None,
cu_seqlens: torch.Tensor | None = None,
output_dtype: torch.dtype | None = torch.float,
chunk_indices: torch.LongTensor | None = None,
**kwargs,
) -> torch.Tensor:
if 'head_first' in kwargs:
raise DeprecationWarning(
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
)
B, T, H = g.shape
assert chunk_size == 2**(chunk_size.bit_length()-1), "chunk_size must be a power of 2"
BT = chunk_size
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
g_org, g = g, torch.empty_like(g, dtype=output_dtype or g.dtype)
grid = (NT, B * H)
chunk_local_cumsum_scalar_kernel[grid](
s=g_org,
o=g,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
T=T,
B=B,
H=H,
BT=BT,
REVERSE=reverse,
)
return g
def chunk_local_cumsum_vector(
g: torch.Tensor,
chunk_size: int,
reverse: bool = False,
scale: float = None,
cu_seqlens: torch.Tensor | None = None,
output_dtype: torch.dtype | None = torch.float,
chunk_indices: torch.LongTensor | None = None,
**kwargs,
) -> torch.Tensor:
if 'head_first' in kwargs:
raise DeprecationWarning(
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
)
B, T, H, S = g.shape
BT = chunk_size
if chunk_indices is None and cu_seqlens is not None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
assert chunk_size == 2**(chunk_size.bit_length()-1), "chunk_size must be a power of 2"
g_org, g = g, torch.empty_like(g, dtype=output_dtype or g.dtype)
def grid(meta): return (triton.cdiv(meta['S'], meta['BS']), NT, B * H)
# keep cummulative normalizer in fp32
# this kernel is equivalent to
# g = g.view(B, H, NT, BT, -1).cumsum(-2).view(B, H, T, -1)
chunk_local_cumsum_vector_kernel[grid](
s=g_org,
o=g,
scale=scale,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
T=T,
B=B,
H=H,
S=S,
BT=BT,
REVERSE=reverse,
)
return g
@input_guard
def chunk_global_cumsum_scalar(
s: torch.Tensor,
reverse: bool = False,
cu_seqlens: torch.Tensor | None = None,
scale: float = None,
output_dtype: torch.dtype | None = torch.float,
**kwargs,
) -> torch.Tensor:
if 'head_first' in kwargs:
raise DeprecationWarning(
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
)
B, T, H = s.shape
N = len(cu_seqlens) - 1 if cu_seqlens is not None else B
z = torch.empty_like(s, dtype=output_dtype or s.dtype)
grid = (N * H,)
chunk_global_cumsum_scalar_kernel[grid](
s=s,
o=z,
scale=scale,
cu_seqlens=cu_seqlens,
T=T,
B=B,
H=H,
REVERSE=reverse,
)
return z
@input_guard
def chunk_global_cumsum_vector(
s: torch.Tensor,
reverse: bool = False,
cu_seqlens: torch.Tensor | None = None,
scale: float = None,
output_dtype: torch.dtype | None = torch.float,
**kwargs,
) -> torch.Tensor:
if 'head_first' in kwargs:
raise DeprecationWarning(
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
)
B, T, H, S = s.shape
N = len(cu_seqlens) - 1 if cu_seqlens is not None else B
BS = min(32, triton.next_power_of_2(S))
z = torch.empty_like(s, dtype=output_dtype or s.dtype)
grid = (triton.cdiv(S, BS), N * H)
chunk_global_cumsum_vector_kernel[grid](
s=s,
o=z,
scale=scale,
cu_seqlens=cu_seqlens,
T=T,
B=B,
H=H,
S=S,
BS=BS,
REVERSE=reverse,
)
return z
@input_guard
@dispatch('utils')
def chunk_global_cumsum(
s: torch.Tensor,
reverse: bool = False,
cu_seqlens: torch.Tensor | None = None,
scale: float = None,
output_dtype: torch.dtype | None = torch.float,
**kwargs,
) -> torch.Tensor:
if 'head_first' in kwargs:
raise DeprecationWarning(
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
)
if cu_seqlens is not None:
assert s.shape[0] == 1, "Only batch size 1 is supported when cu_seqlens are provided"
if len(s.shape) == 3:
return chunk_global_cumsum_scalar(
s=s,
reverse=reverse,
cu_seqlens=cu_seqlens,
scale=scale,
output_dtype=output_dtype,
)
elif len(s.shape) == 4:
return chunk_global_cumsum_vector(
s=s,
reverse=reverse,
cu_seqlens=cu_seqlens,
scale=scale,
output_dtype=output_dtype,
)
else:
raise ValueError(
f"Unsupported input shape {s.shape}, "
f"which should be [B, T, H] or [B, T, H, D]",
)
@input_guard
@dispatch('utils')
def chunk_local_cumsum(
g: torch.Tensor,
chunk_size: int,
reverse: bool = False,
scale: float = None,
cu_seqlens: torch.Tensor | None = None,
output_dtype: torch.dtype | None = torch.float,
chunk_indices: torch.LongTensor | None = None,
**kwargs,
) -> torch.Tensor:
if 'head_first' in kwargs:
raise DeprecationWarning(
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
)
if cu_seqlens is not None:
assert g.shape[0] == 1, "Only batch size 1 is supported when cu_seqlens are provided"
if len(g.shape) == 3:
return chunk_local_cumsum_scalar(
g=g,
chunk_size=chunk_size,
reverse=reverse,
scale=scale,
cu_seqlens=cu_seqlens,
output_dtype=output_dtype,
chunk_indices=chunk_indices,
)
elif len(g.shape) == 4:
return chunk_local_cumsum_vector(
g=g,
chunk_size=chunk_size,
reverse=reverse,
scale=scale,
cu_seqlens=cu_seqlens,
output_dtype=output_dtype,
chunk_indices=chunk_indices,
)
else:
raise ValueError(
f"Unsupported input shape {g.shape}, "
f"which should be (B, T, H) or (B, T, H, D)",
)
+183
View File
@@ -0,0 +1,183 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import torch
import torch.nn.functional as F
import triton
import triton.language as tl
from kda._fla.utils import autotune_cache_kwargs, tensor_cache
@triton.autotune(
configs=[
triton.Config({}, num_warps=num_warps)
for num_warps in [4, 8, 16, 32]
],
key=['B'],
**autotune_cache_kwargs,
)
@triton.jit
def prepare_position_ids_kernel(
y,
cu_seqlens,
B: tl.constexpr,
):
i_n = tl.program_id(0)
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
T = eos - bos
o = tl.arange(0, B)
for i in range(0, tl.cdiv(T, B) * B, B):
o_i = o + i
tl.store(y + bos + o_i, o_i, o_i < T)
@tensor_cache
def prepare_lens(cu_seqlens: torch.LongTensor) -> torch.LongTensor:
return torch.diff(cu_seqlens)
@tensor_cache
def prepare_lens_from_mask(mask: torch.BoolTensor) -> torch.LongTensor:
return mask.sum(dim=-1, dtype=torch.int32)
@tensor_cache
def prepare_cu_seqlens_from_lens(
lens: torch.LongTensor,
dtype: torch.dtype | None = torch.int32,
) -> torch.LongTensor:
return F.pad(lens.cumsum(dim=0, dtype=dtype), (1, 0))
@tensor_cache
def prepare_cu_seqlens_from_mask(
mask: torch.BoolTensor,
dtype: torch.dtype | None = torch.int32,
) -> torch.LongTensor:
return prepare_cu_seqlens_from_lens(prepare_lens_from_mask(mask), dtype)
@tensor_cache
def prepare_split_cu_seqlens(
batch_size: int | None = None,
seq_len: int | None = None,
split_size: int | None = None,
cu_seqlens: torch.LongTensor | None = None,
dtype: torch.dtype | None = torch.int32,
device: torch.device | None = torch.device('cpu'),
) -> torch.LongTensor:
"""Sub-split a (optionally packed) batch along the token axis.
Two calling modes:
- **Rectangular batch**: pass `batch_size` and `seq_len`, leave
`cu_seqlens=None`. Internally synthesizes `[0, L, 2L, ..., B*L]`.
- **Packed varlen**: pass `cu_seqlens`. `batch_size` and `seq_len` are
ignored (kept as optional kwargs for backward-compat with callers
that used to pass dummies).
`split_size` is always required.
The legacy positional signature `(batch_size, seq_len, split_size, ...)`
continues to work — the first two args retain their position but may now
be omitted when `cu_seqlens` is supplied.
"""
if split_size is None:
raise TypeError("prepare_split_cu_seqlens() requires `split_size`")
if cu_seqlens is None:
if batch_size is None or seq_len is None:
raise TypeError(
"prepare_split_cu_seqlens() requires either `cu_seqlens`, "
"or both `batch_size` and `seq_len`"
)
total_tokens = batch_size * seq_len
cu_seqlens = list(range(0, total_tokens, seq_len)) + [total_tokens]
else:
cu_seqlens = cu_seqlens.tolist()
return torch.tensor(
[
i
for bos, eos in zip(cu_seqlens[:-1], cu_seqlens[1:], strict=False)
for i in range(bos, eos, split_size)
] + [cu_seqlens[-1]],
dtype=dtype,
device=device,
)
def _segmented_arange(counts: torch.LongTensor) -> tuple[torch.LongTensor, torch.LongTensor]:
"""Expand per-segment counts into flat per-slot index tensors.
Given segment sizes ``counts = [c0, c1, ...]``, return two 1-D tensors of
length ``counts.sum()`` that together label every slot with its segment and
its position within that segment.
Example -- ``counts = [2, 3]`` (segment 0 spans 2 slots, segment 1 spans 3)::
seg_id = [0, 0, 1, 1, 1] # which segment each slot belongs to
intra_idx = [0, 1, 0, 1, 2] # running index within that segment
With CUDA ``counts``, ``repeat_interleave`` reads ``counts.sum()`` on the
host (one device sync). Pass host-side counts to avoid it.
"""
seg_id = torch.repeat_interleave(
torch.arange(counts.numel(), device=counts.device, dtype=counts.dtype),
counts,
)
seg_start = F.pad(counts.cumsum(0), (1, 0))[:-1]
intra_idx = torch.arange(seg_id.shape[0], device=counts.device, dtype=counts.dtype) - seg_start[seg_id]
return seg_id, intra_idx
@tensor_cache
def prepare_position_ids(cu_seqlens: torch.LongTensor, cu_seqlens_cpu: torch.LongTensor | None = None) -> torch.LongTensor:
src = cu_seqlens_cpu if cu_seqlens_cpu is not None else cu_seqlens
_, position_ids = _segmented_arange(prepare_lens(src))
return position_ids.to(cu_seqlens)
@tensor_cache
def prepare_sequence_ids(cu_seqlens: torch.LongTensor, cu_seqlens_cpu: torch.LongTensor | None = None) -> torch.LongTensor:
return prepare_position_ids(cu_seqlens, cu_seqlens_cpu).eq(0).cumsum(0) - 1
@tensor_cache
def prepare_token_indices(cu_seqlens: torch.LongTensor, cu_seqlens_cpu: torch.LongTensor | None = None) -> torch.LongTensor:
position_ids = prepare_position_ids(cu_seqlens, cu_seqlens_cpu)
return torch.stack([prepare_sequence_ids(cu_seqlens, cu_seqlens_cpu), position_ids], 1).to(cu_seqlens)
@tensor_cache
def prepare_chunk_indices(
cu_seqlens: torch.LongTensor,
chunk_size: int,
cu_seqlens_cpu: torch.LongTensor | None = None,
) -> torch.LongTensor:
src = cu_seqlens_cpu if cu_seqlens_cpu is not None else cu_seqlens
chunk_counts = (prepare_lens(src) + (chunk_size - 1)).div(chunk_size, rounding_mode='floor')
seg_id, intra_chunk_idx = _segmented_arange(chunk_counts)
return torch.stack([seg_id, intra_chunk_idx], 1).to(cu_seqlens)
@tensor_cache
def prepare_chunk_offsets(
cu_seqlens: torch.LongTensor,
chunk_size: int,
) -> torch.LongTensor:
return F.pad(triton.cdiv(prepare_lens(cu_seqlens), chunk_size), (1, 0), value=0).cumsum(-1)
@tensor_cache
def get_max_num_splits(
cu_seqlens: torch.LongTensor,
chunk_size: int,
cu_seqlens_cpu: torch.LongTensor | None = None
) -> int:
if cu_seqlens_cpu is not None:
return triton.cdiv(int(max(prepare_lens(cu_seqlens_cpu))), chunk_size)
return triton.cdiv(int(max(prepare_lens(cu_seqlens))), chunk_size)
+101
View File
@@ -0,0 +1,101 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import os
import triton
import triton.language as tl
import triton.language.extra.libdevice as tldevice
from kda._fla.utils import IS_GATHER_SUPPORTED, IS_NVIDIA_BLACKWELL
if os.environ.get('FLA_USE_FAST_OPS', '0') == '1':
@triton.jit
def exp(x): return tldevice.fast_expf(x.to(tl.float32))
@triton.jit
def exp2(x): return tldevice.exp2(x.to(tl.float32))
@triton.jit
def log(x): return tldevice.fast_logf(x.to(tl.float32))
@triton.jit
def log2(x): return tldevice.fast_log2f(x.to(tl.float32))
@triton.jit
def tanh(x): return tldevice.fast_tanhf(x.to(tl.float32))
else:
@triton.jit
def exp(x): return tl.exp(x.to(tl.float32))
@triton.jit
def exp2(x): return tl.math.exp2(x.to(tl.float32))
@triton.jit
def log(x): return tl.log(x.to(tl.float32))
@triton.jit
def log2(x): return tl.log2(x.to(tl.float32))
@triton.jit
def tanh(x): return tldevice.tanh(x.to(tl.float32))
if IS_NVIDIA_BLACKWELL:
"""
Compute tl.dot with Blackwell workaround.
On SM100 datacenter and SM120 consumer Blackwell GPUs, wraps the result in
inline assembly to prevent the TritonGPUHoistTMEMAlloc pass from incorrectly
fusing add and dot operations.
See: https://github.com/fla-org/flash-linear-attention/issues/638
TODO: Remove this workaround once the Triton compiler bug is fixed.
Track upstream issue at: https://github.com/triton-lang/triton/issues/8695
"""
@triton.jit
def safe_dot(a, b, allow_tf32: tl.constexpr = None):
return tl.inline_asm_elementwise(
asm="mov.f32 $0, $1;",
constraints="=r,r",
args=[tl.dot(a, b, allow_tf32=allow_tf32)],
dtype=tl.float32,
is_pure=True,
pack=1,
)
else:
@triton.jit
def safe_dot(a, b, allow_tf32: tl.constexpr = None):
return tl.dot(a, b, allow_tf32=allow_tf32)
if not IS_GATHER_SUPPORTED:
@triton.jit
def gather(src, index, axis, _builder=None):
"""
Gather operation that works when tl.gather is not supported.
This is a fallback implementation that returns None.
Just to make triton compiler happy.
"""
return None
else:
gather = tl.gather
if hasattr(triton.language, '_experimental_make_tensor_descriptor'):
# For Triton 3.3.x
make_tensor_descriptor = triton.language._experimental_make_tensor_descriptor
elif hasattr(triton.language, 'make_tensor_descriptor'):
# For Triton 3.4.x and later
make_tensor_descriptor = triton.language.make_tensor_descriptor
else:
"""
Fallback implementation when TMA is not supported.
Returns None to indicate TMA descriptors are unavailable.
Just make triton compiler happy.
"""
@triton.jit
def make_tensor_descriptor(
base,
shape,
strides,
block_shape,
_builder=None,
):
return None
+115
View File
@@ -0,0 +1,115 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
# REVISED FROM
# https://github.com/shawntan/stickbreaking-attention/blob/main/stickbreaking_attention/sb_varlen/softplus.py
import triton
from triton import language as tl
from kda._fla.utils import IS_NVIDIA
def _generate_softplus(num_pack):
template = """
.reg .pred p;
setp.gt.f32 p, ${in_reg}, 20.;
@p mov.f32 ${out_reg}, ${in_reg};
@!p mul.f32 ${out_reg}, ${in_reg}, 1.4426950408889634;
@!p ex2.approx.ftz.f32 ${out_reg}, ${out_reg};
@!p add.f32 ${out_reg}, ${out_reg}, 1.0;
@!p lg2.approx.ftz.f32 ${out_reg}, ${out_reg};
@!p mul.f32 ${out_reg}, ${out_reg}, 0.6931471805599453;
"""
out_str = ""
for i in range(num_pack):
inner_str = template.format(out_reg=i, in_reg=i + num_pack)
out_str += "{" + inner_str + "}\n"
# flatten out because torch.compile doesn't like newlines
out_str = " ".join(out_str.split("\n"))
return out_str
def _generate_softplus2(num_pack):
template = """
.reg .pred p;
setp.gt.f32 p, ${in_reg}, 15.;
@p mov.f32 ${out_reg}, ${in_reg};
@!p ex2.approx.ftz.f32 ${out_reg}, ${in_reg};
@!p add.f32 ${out_reg}, ${out_reg}, 1.0;
@!p lg2.approx.ftz.f32 ${out_reg}, ${out_reg};
"""
out_str = ""
for i in range(num_pack):
inner_str = template.format(out_reg=i, in_reg=i + num_pack)
out_str += "{" + inner_str + "}\n"
# flatten out because torch.compile doesn't like newlines
out_str = " ".join(out_str.split("\n"))
return out_str
def _generate_constraints(num_pack):
return ",".join("=r" for i in range(num_pack)) + "," + ",".join("r" for i in range(num_pack))
_NUM_REG = 1
s_softplus: tl.constexpr = tl.constexpr(_generate_softplus(_NUM_REG))
s_softplus2: tl.constexpr = tl.constexpr(_generate_softplus2(_NUM_REG))
s_constraints: tl.constexpr = tl.constexpr(_generate_constraints(_NUM_REG))
NUM_REG: tl.constexpr = tl.constexpr(_NUM_REG)
@triton.jit
def softplus_nv(x):
# equivalent to:
# return tl.where(x < 20.0, tl.math.log(1 + tl.math.exp(x)), x)
return tl.inline_asm_elementwise(
asm=s_softplus,
constraints=s_constraints,
pack=NUM_REG,
args=[
x,
],
dtype=tl.float32,
is_pure=True,
)
@triton.jit
def softplus_triton(x):
return tl.where(x < 20.0, tl.math.log(1 + tl.math.exp(x)), x)
@triton.jit
def softplus2_nv(x):
# equivalent to:
# return tl.where(x < 15.0, tl.math.log2(1 + tl.math.exp2(x)), x)
return tl.inline_asm_elementwise(
asm=s_softplus2,
constraints=s_constraints,
pack=NUM_REG,
args=[
x,
],
dtype=tl.float32,
is_pure=True,
)
@triton.jit
def softplus2_triton(x):
return tl.where(x < 15.0, tl.math.log2(1 + tl.math.exp2(x)), x)
if IS_NVIDIA:
softplus = softplus_nv
softplus2 = softplus2_nv
else:
softplus = softplus_triton
softplus2 = softplus2_triton