use std::ffi::c_void;
use std::sync::Arc;
use cudarc::driver::{LaunchConfig, PushKernelArg};
use onnx_runtime_ep_api::{EpError, Kernel, KernelFactory, Result, TensorMut, TensorView};
use onnx_runtime_ir::{DataType, Node};
use crate::error::{driver_err, not_implemented};
use crate::runtime::{CudaRuntime, cuptr};
const BLOCK: u32 = 256;
const LA_WARP: u32 = 32;
const MAX_D_K: usize = 256;
pub(crate) const FUSE_BETA_SIGMOID_ATTR: &str = "com.microsoft.cuda_fuse_beta_sigmoid";
pub(crate) const FUSE_DECAY_SOFTPLUS_ATTR: &str = "com.microsoft.cuda_fuse_decay_softplus";
pub(crate) const FUSE_NEG_EXP_ATTR: &str = "com.microsoft.cuda_fuse_neg_exp";
fn linattn_warp_coop_disabled() -> bool {
matches!(
std::env::var("ONNX_GENAI_CUDA_DISABLE_LINATTN_WARP_COOP")
.ok()
.as_deref(),
Some("1") | Some("true") | Some("on")
)
}
const SOURCE: &str = r#"
#include <cuda_fp16.h>
#include <cuda_bf16.h>
#define MAX_D_K 256
#define LA_WARP 32
// Registers per lane for the warp-cooperative state column: each lane owns the
// d_k rows i = lane, lane+32, ... so it needs ceil(MAX_D_K / warp) slots.
#define LA_MAX_SLOTS ((MAX_D_K + LA_WARP - 1) / LA_WARP)
__device__ __forceinline__ float to_f(float x) { return x; }
__device__ __forceinline__ float to_f(__half x) { return __half2float(x); }
__device__ __forceinline__ float to_f(__nv_bfloat16 x) { return __bfloat162float(x); }
__device__ __forceinline__ float from_f_val(float x, float*) { return x; }
__device__ __forceinline__ __half from_f_val(float x, __half*) { return __float2half_rn(x); }
__device__ __forceinline__ __nv_bfloat16 from_f_val(float x, __nv_bfloat16*) {
return __float2bfloat16_rn(x);
}
// Round `x` through the storage dtype `T` and widen back to f32. Reproduces the
// per-op narrow rounding a standalone `Sigmoid`/`Softplus`/`Add`/`Mul` kernel
// applies at each fused gate's boundary, so a folded gate stays byte-identical
// for narrow (f16/bf16) I/O and is an exact no-op for f32.
template <typename T>
__device__ __forceinline__ float round_store(float x) {
return to_f(from_f_val(x, (T*)0));
}
// Byte-exact port of the standalone `Sigmoid` op device function
// (`kernels/elementwise.rs::op_sigmoid`): the `exp` is evaluated in double, so a
// folded `beta = Sigmoid(x)` gate reproduces the unfused kernel bit-for-bit.
__device__ __forceinline__ float la_sigmoid(float x) {
if (x >= 0.0f) return 1.0f / (1.0f + (float)exp((double)-x));
const float e = (float)exp((double)x);
return e / (1.0f + e);
}
// Byte-exact port of the standalone `Softplus` op device function
// (`kernels/pointwise.rs::op_softplus`).
__device__ __forceinline__ float la_softplus(float x) {
return fmaxf(x, 0.0f) + log1pf(expf(-fabsf(x)));
}
// Byte-exact port of the standalone `Exp` op device function
// (`kernels/pointwise.rs::op_exp`).
__device__ __forceinline__ float la_exp(float x) { return expf(x); }
// The per-head decay coefficient `neg_exp_A = -exp(A_log)`. When `fuse_neg_exp`
// is set the trailing decay operand holds the raw `A_log` constant instead of a
// precomputed `neg_exp_A`, so the kernel folds the exported `Neg(Exp(A_log))`
// chain inline. It is byte-identical to the two standalone ops: `Exp` rounds
// `expf(A_log)` to the storage dtype, then `Neg` negates that (an exact bf16/f16
// sign flip, so re-rounding is a no-op) — i.e. `-round_store<T>(exp(A_log))`.
// When `fuse_neg_exp` is clear, `raw` already is `neg_exp_A` and is used as-is.
template <typename T>
__device__ __forceinline__ float la_neg_exp(float raw, int fuse_neg_exp) {
return fuse_neg_exp ? -round_store<T>(la_exp(raw)) : raw;
}
// Full-warp butterfly sum: every lane returns Σ over all 32 lanes. Used by the
// warp-cooperative kernel to reduce the d_k dot products (retrieval r = Sáµ€k and
// readout o = qáµ€S) that the serial kernel walks with a 128-iteration loop. The
// tree order differs from the serial left-to-right sum, so fp32 results shift at
// the ULP level (accumulation stays in f32); greedy argmax is unaffected.
__device__ __forceinline__ float la_warp_reduce_sum(float v) {
#pragma unroll
for (int off = LA_WARP / 2; off > 0; off >>= 1) {
v += __shfl_xor_sync(0xffffffffu, v, off);
}
return v;
}
template <typename T>
__device__ void linear_attention_core(
const T* q, const T* k, const T* v,
const T* past_state, const T* decay, const T* beta,
const T* dt_bias, const T* neg_exp_A,
T* output, T* present_state,
unsigned long long batch, unsigned long long seq,
unsigned long long d_k, unsigned long long d_v,
unsigned long long q_num_heads, unsigned long long kv_num_heads,
unsigned long long n_k_heads, unsigned long long heads_per_group,
unsigned long long kv_per_k_head, unsigned long long output_hidden,
float scale, int needs_decay, int decay_per_key_dim,
int needs_delta, int beta_per_head,
int fuse_beta_sigmoid, int fuse_decay_softplus,
int fuse_neg_exp) {
const unsigned long long total = batch * kv_num_heads * d_v;
const unsigned long long stride =
(unsigned long long)gridDim.x * blockDim.x;
for (unsigned long long tid =
(unsigned long long)blockIdx.x * blockDim.x + threadIdx.x;
tid < total; tid += stride) {
const unsigned long long j = tid % d_v;
const unsigned long long hk_flat = tid / d_v; // b * kv_num_heads + h_kv
const unsigned long long b = hk_flat / kv_num_heads;
const unsigned long long h_kv = hk_flat % kv_num_heads;
const unsigned long long h_k = h_kv / kv_per_k_head;
const unsigned long long sbase = hk_flat * d_k * d_v; // + i*d_v + j
float sc[MAX_D_K];
for (unsigned long long i = 0; i < d_k; ++i) {
sc[i] = past_state ? to_f(past_state[sbase + i * d_v + j]) : 0.0f;
}
for (unsigned long long t = 0; t < seq; ++t) {
const unsigned long long row = b * seq + t;
// Step 1: decay S *= exp(g_t)
if (needs_decay) {
if (decay_per_key_dim) {
const T* g = decay + row * (kv_num_heads * d_k) + h_kv * d_k;
if (fuse_decay_softplus) {
const T* dtb = dt_bias + h_kv * d_k;
const T* na = neg_exp_A + h_kv * d_k;
for (unsigned long long i = 0; i < d_k; ++i) {
float a = round_store<T>(to_f(g[i]) + to_f(dtb[i]));
a = round_store<T>(la_softplus(a));
a = round_store<T>(la_neg_exp<T>(to_f(na[i]), fuse_neg_exp) * a);
sc[i] *= expf(a);
}
} else {
for (unsigned long long i = 0; i < d_k; ++i) sc[i] *= expf(to_f(g[i]));
}
} else {
float g_val;
if (fuse_decay_softplus) {
float a = round_store<T>(
to_f(decay[row * kv_num_heads + h_kv]) + to_f(dt_bias[h_kv]));
a = round_store<T>(la_softplus(a));
g_val = round_store<T>(la_neg_exp<T>(to_f(neg_exp_A[h_kv]), fuse_neg_exp) * a);
} else {
g_val = to_f(decay[row * kv_num_heads + h_kv]);
}
const float eg = expf(g_val);
for (unsigned long long i = 0; i < d_k; ++i) sc[i] *= eg;
}
}
const float vt = to_f(v[row * (kv_num_heads * d_v) + h_kv * d_v + j]);
const T* kt = k + row * (n_k_heads * d_k) + h_k * d_k;
if (needs_delta) {
// Step 2: retrieval r = Sáµ€ k_t (over d_k)
float r = 0.0f;
for (unsigned long long i = 0; i < d_k; ++i) r += sc[i] * to_f(kt[i]);
// Step 3: delta update S += k_t ⊗ (beta·(v_t − r))
float bt = beta_per_head ? to_f(beta[row * kv_num_heads + h_kv]) : to_f(beta[row]);
if (fuse_beta_sigmoid) bt = round_store<T>(la_sigmoid(bt));
const float d = bt * (vt - r);
for (unsigned long long i = 0; i < d_k; ++i) sc[i] += to_f(kt[i]) * d;
} else {
// linear / gated: S += k_t ⊗ v_t
for (unsigned long long i = 0; i < d_k; ++i) sc[i] += to_f(kt[i]) * vt;
}
// Step 4: readout o_t = scale · q_tᵀ S (updated S)
if (heads_per_group > 0) {
for (unsigned long long g = 0; g < heads_per_group; ++g) {
const unsigned long long h_q = h_kv * heads_per_group + g;
const T* qt = q + row * (q_num_heads * d_k) + h_q * d_k;
float o = 0.0f;
for (unsigned long long i = 0; i < d_k; ++i) o += to_f(qt[i]) * sc[i];
output[row * output_hidden + h_q * d_v + j] = from_f_val(o * scale, output);
}
} else {
// Inverse GQA: output slot is h_kv, query head h_kv·H_q/H_kv.
const unsigned long long h_q = h_kv * q_num_heads / kv_num_heads;
const T* qt = q + row * (q_num_heads * d_k) + h_q * d_k;
float o = 0.0f;
for (unsigned long long i = 0; i < d_k; ++i) o += to_f(qt[i]) * sc[i];
output[row * output_hidden + h_kv * d_v + j] = from_f_val(o * scale, output);
}
}
for (unsigned long long i = 0; i < d_k; ++i) {
present_state[sbase + i * d_v + j] = from_f_val(sc[i], present_state);
}
}
}
extern "C" __global__ void linear_attention_f32(
const float* q, const float* k, const float* v,
const float* past_state, const float* decay, const float* beta,
const float* dt_bias, const float* neg_exp_A,
float* output, float* present_state,
unsigned long long batch, unsigned long long seq,
unsigned long long d_k, unsigned long long d_v,
unsigned long long q_num_heads, unsigned long long kv_num_heads,
unsigned long long n_k_heads, unsigned long long heads_per_group,
unsigned long long kv_per_k_head, unsigned long long output_hidden,
float scale, int needs_decay, int decay_per_key_dim,
int needs_delta, int beta_per_head,
int fuse_beta_sigmoid, int fuse_decay_softplus,
int fuse_neg_exp) {
linear_attention_core<float>(
q, k, v, past_state, decay, beta, dt_bias, neg_exp_A, output,
present_state, batch, seq, d_k, d_v, q_num_heads, kv_num_heads, n_k_heads,
heads_per_group, kv_per_k_head, output_hidden, scale, needs_decay,
decay_per_key_dim, needs_delta, beta_per_head, fuse_beta_sigmoid,
fuse_decay_softplus, fuse_neg_exp);
}
extern "C" __global__ void linear_attention_f16(
const __half* q, const __half* k, const __half* v,
const __half* past_state, const __half* decay, const __half* beta,
const __half* dt_bias, const __half* neg_exp_A,
__half* output, __half* present_state,
unsigned long long batch, unsigned long long seq,
unsigned long long d_k, unsigned long long d_v,
unsigned long long q_num_heads, unsigned long long kv_num_heads,
unsigned long long n_k_heads, unsigned long long heads_per_group,
unsigned long long kv_per_k_head, unsigned long long output_hidden,
float scale, int needs_decay, int decay_per_key_dim,
int needs_delta, int beta_per_head,
int fuse_beta_sigmoid, int fuse_decay_softplus,
int fuse_neg_exp) {
linear_attention_core<__half>(
q, k, v, past_state, decay, beta, dt_bias, neg_exp_A, output,
present_state, batch, seq, d_k, d_v, q_num_heads, kv_num_heads, n_k_heads,
heads_per_group, kv_per_k_head, output_hidden, scale, needs_decay,
decay_per_key_dim, needs_delta, beta_per_head, fuse_beta_sigmoid,
fuse_decay_softplus, fuse_neg_exp);
}
extern "C" __global__ void linear_attention_bf16(
const __nv_bfloat16* q, const __nv_bfloat16* k, const __nv_bfloat16* v,
const __nv_bfloat16* past_state, const __nv_bfloat16* decay,
const __nv_bfloat16* beta,
const __nv_bfloat16* dt_bias, const __nv_bfloat16* neg_exp_A,
__nv_bfloat16* output,
__nv_bfloat16* present_state,
unsigned long long batch, unsigned long long seq,
unsigned long long d_k, unsigned long long d_v,
unsigned long long q_num_heads, unsigned long long kv_num_heads,
unsigned long long n_k_heads, unsigned long long heads_per_group,
unsigned long long kv_per_k_head, unsigned long long output_hidden,
float scale, int needs_decay, int decay_per_key_dim,
int needs_delta, int beta_per_head,
int fuse_beta_sigmoid, int fuse_decay_softplus,
int fuse_neg_exp) {
linear_attention_core<__nv_bfloat16>(
q, k, v, past_state, decay, beta, dt_bias, neg_exp_A, output,
present_state, batch, seq, d_k, d_v, q_num_heads, kv_num_heads, n_k_heads,
heads_per_group, kv_per_k_head, output_hidden, scale, needs_decay,
decay_per_key_dim, needs_delta, beta_per_head, fuse_beta_sigmoid,
fuse_decay_softplus, fuse_neg_exp);
}
// ── Warp-cooperative variant ────────────────────────────────────────────────
//
// Same math as `linear_attention_core`, but one WARP (not one thread) owns each
// state column (b, h_kv, j): lane `l` holds the d_k rows i = l, l+warp, ... in
// registers (`sc[LA_MAX_SLOTS]`, ≤8 f32) instead of a per-thread `sc[MAX_D_K]`
// array that spills to local memory. The two d_k dot products (retrieval and
// readout) become `__shfl_xor` warp reductions instead of a 128-iteration serial
// loop. This raises the launch from `batch·H_kv·d_v` threads to that many WARPS
// (32× more resident work → fills the SMs) and removes the local-memory spill
// that dominated the serial kernel's Long-Scoreboard stalls. Generic over the
// per-op dtype, head geometry (GQA / inverse GQA / dual head sizes) and any
// `d_k ≤ MAX_D_K` (the `i < d_k` guards handle d_k not a multiple of the warp).
template <typename T>
__device__ void linear_attention_core_coop(
const T* q, const T* k, const T* v,
const T* past_state, const T* decay, const T* beta,
const T* dt_bias, const T* neg_exp_A,
T* output, T* present_state,
unsigned long long batch, unsigned long long seq,
unsigned long long d_k, unsigned long long d_v,
unsigned long long q_num_heads, unsigned long long kv_num_heads,
unsigned long long n_k_heads, unsigned long long heads_per_group,
unsigned long long kv_per_k_head, unsigned long long output_hidden,
float scale, int needs_decay, int decay_per_key_dim,
int needs_delta, int beta_per_head,
int fuse_beta_sigmoid, int fuse_decay_softplus,
int fuse_neg_exp) {
const unsigned lane = threadIdx.x & (LA_WARP - 1);
const unsigned warps_per_block = blockDim.x / LA_WARP;
const unsigned long long total_warps = batch * kv_num_heads * d_v;
const unsigned long long warp_stride =
(unsigned long long)gridDim.x * warps_per_block;
// The loop bound `w` is warp-uniform (every lane shares it), so an out-of-range
// warp exits collectively — no partial-warp `__shfl_xor_sync` divergence.
for (unsigned long long w =
(unsigned long long)blockIdx.x * warps_per_block +
(threadIdx.x / LA_WARP);
w < total_warps; w += warp_stride) {
const unsigned long long j = w % d_v;
const unsigned long long hk_flat = w / d_v; // b * kv_num_heads + h_kv
const unsigned long long b = hk_flat / kv_num_heads;
const unsigned long long h_kv = hk_flat % kv_num_heads;
const unsigned long long h_k = h_kv / kv_per_k_head;
const unsigned long long sbase = hk_flat * d_k * d_v; // + i*d_v + j
float sc[LA_MAX_SLOTS];
#pragma unroll
for (int s = 0; s < LA_MAX_SLOTS; ++s) {
const unsigned long long i = lane + (unsigned long long)s * LA_WARP;
sc[s] = (i < d_k && past_state) ? to_f(past_state[sbase + i * d_v + j]) : 0.0f;
}
for (unsigned long long t = 0; t < seq; ++t) {
const unsigned long long row = b * seq + t;
// Step 1: decay S *= exp(g_t)
if (needs_decay) {
if (decay_per_key_dim) {
const T* g = decay + row * (kv_num_heads * d_k) + h_kv * d_k;
if (fuse_decay_softplus) {
const T* dtb = dt_bias + h_kv * d_k;
const T* na = neg_exp_A + h_kv * d_k;
#pragma unroll
for (int s = 0; s < LA_MAX_SLOTS; ++s) {
const unsigned long long i = lane + (unsigned long long)s * LA_WARP;
if (i < d_k) {
float a = round_store<T>(to_f(g[i]) + to_f(dtb[i]));
a = round_store<T>(la_softplus(a));
a = round_store<T>(la_neg_exp<T>(to_f(na[i]), fuse_neg_exp) * a);
sc[s] *= expf(a);
}
}
} else {
#pragma unroll
for (int s = 0; s < LA_MAX_SLOTS; ++s) {
const unsigned long long i = lane + (unsigned long long)s * LA_WARP;
if (i < d_k) sc[s] *= expf(to_f(g[i]));
}
}
} else {
float g_val;
if (fuse_decay_softplus) {
float a = round_store<T>(
to_f(decay[row * kv_num_heads + h_kv]) + to_f(dt_bias[h_kv]));
a = round_store<T>(la_softplus(a));
g_val = round_store<T>(la_neg_exp<T>(to_f(neg_exp_A[h_kv]), fuse_neg_exp) * a);
} else {
g_val = to_f(decay[row * kv_num_heads + h_kv]);
}
const float eg = expf(g_val);
#pragma unroll
for (int s = 0; s < LA_MAX_SLOTS; ++s) {
const unsigned long long i = lane + (unsigned long long)s * LA_WARP;
if (i < d_k) sc[s] *= eg;
}
}
}
const float vt = to_f(v[row * (kv_num_heads * d_v) + h_kv * d_v + j]);
const T* kt = k + row * (n_k_heads * d_k) + h_k * d_k;
if (needs_delta) {
// Step 2: retrieval r = Sáµ€ k_t (warp reduction over d_k)
float rpart = 0.0f;
#pragma unroll
for (int s = 0; s < LA_MAX_SLOTS; ++s) {
const unsigned long long i = lane + (unsigned long long)s * LA_WARP;
if (i < d_k) rpart += sc[s] * to_f(kt[i]);
}
const float r = la_warp_reduce_sum(rpart);
// Step 3: delta update S += k_t ⊗ (beta·(v_t − r))
float bt = beta_per_head ? to_f(beta[row * kv_num_heads + h_kv]) : to_f(beta[row]);
if (fuse_beta_sigmoid) bt = round_store<T>(la_sigmoid(bt));
const float dd = bt * (vt - r);
#pragma unroll
for (int s = 0; s < LA_MAX_SLOTS; ++s) {
const unsigned long long i = lane + (unsigned long long)s * LA_WARP;
if (i < d_k) sc[s] += to_f(kt[i]) * dd;
}
} else {
// linear / gated: S += k_t ⊗ v_t
#pragma unroll
for (int s = 0; s < LA_MAX_SLOTS; ++s) {
const unsigned long long i = lane + (unsigned long long)s * LA_WARP;
if (i < d_k) sc[s] += to_f(kt[i]) * vt;
}
}
// Step 4: readout o_t = scale · q_tᵀ S (warp reduction over d_k)
if (heads_per_group > 0) {
for (unsigned long long g = 0; g < heads_per_group; ++g) {
const unsigned long long h_q = h_kv * heads_per_group + g;
const T* qt = q + row * (q_num_heads * d_k) + h_q * d_k;
float opart = 0.0f;
#pragma unroll
for (int s = 0; s < LA_MAX_SLOTS; ++s) {
const unsigned long long i = lane + (unsigned long long)s * LA_WARP;
if (i < d_k) opart += to_f(qt[i]) * sc[s];
}
const float o = la_warp_reduce_sum(opart);
if (lane == 0) {
output[row * output_hidden + h_q * d_v + j] = from_f_val(o * scale, output);
}
}
} else {
// Inverse GQA: output slot is h_kv, query head h_kv·H_q/H_kv.
const unsigned long long h_q = h_kv * q_num_heads / kv_num_heads;
const T* qt = q + row * (q_num_heads * d_k) + h_q * d_k;
float opart = 0.0f;
#pragma unroll
for (int s = 0; s < LA_MAX_SLOTS; ++s) {
const unsigned long long i = lane + (unsigned long long)s * LA_WARP;
if (i < d_k) opart += to_f(qt[i]) * sc[s];
}
const float o = la_warp_reduce_sum(opart);
if (lane == 0) {
output[row * output_hidden + h_kv * d_v + j] = from_f_val(o * scale, output);
}
}
}
#pragma unroll
for (int s = 0; s < LA_MAX_SLOTS; ++s) {
const unsigned long long i = lane + (unsigned long long)s * LA_WARP;
if (i < d_k) present_state[sbase + i * d_v + j] = from_f_val(sc[s], present_state);
}
}
}
extern "C" __global__ void linear_attention_f32_coop(
const float* q, const float* k, const float* v,
const float* past_state, const float* decay, const float* beta,
const float* dt_bias, const float* neg_exp_A,
float* output, float* present_state,
unsigned long long batch, unsigned long long seq,
unsigned long long d_k, unsigned long long d_v,
unsigned long long q_num_heads, unsigned long long kv_num_heads,
unsigned long long n_k_heads, unsigned long long heads_per_group,
unsigned long long kv_per_k_head, unsigned long long output_hidden,
float scale, int needs_decay, int decay_per_key_dim,
int needs_delta, int beta_per_head,
int fuse_beta_sigmoid, int fuse_decay_softplus,
int fuse_neg_exp) {
linear_attention_core_coop<float>(
q, k, v, past_state, decay, beta, dt_bias, neg_exp_A, output,
present_state, batch, seq, d_k, d_v, q_num_heads, kv_num_heads, n_k_heads,
heads_per_group, kv_per_k_head, output_hidden, scale, needs_decay,
decay_per_key_dim, needs_delta, beta_per_head, fuse_beta_sigmoid,
fuse_decay_softplus, fuse_neg_exp);
}
extern "C" __global__ void linear_attention_f16_coop(
const __half* q, const __half* k, const __half* v,
const __half* past_state, const __half* decay, const __half* beta,
const __half* dt_bias, const __half* neg_exp_A,
__half* output, __half* present_state,
unsigned long long batch, unsigned long long seq,
unsigned long long d_k, unsigned long long d_v,
unsigned long long q_num_heads, unsigned long long kv_num_heads,
unsigned long long n_k_heads, unsigned long long heads_per_group,
unsigned long long kv_per_k_head, unsigned long long output_hidden,
float scale, int needs_decay, int decay_per_key_dim,
int needs_delta, int beta_per_head,
int fuse_beta_sigmoid, int fuse_decay_softplus,
int fuse_neg_exp) {
linear_attention_core_coop<__half>(
q, k, v, past_state, decay, beta, dt_bias, neg_exp_A, output,
present_state, batch, seq, d_k, d_v, q_num_heads, kv_num_heads, n_k_heads,
heads_per_group, kv_per_k_head, output_hidden, scale, needs_decay,
decay_per_key_dim, needs_delta, beta_per_head, fuse_beta_sigmoid,
fuse_decay_softplus, fuse_neg_exp);
}
extern "C" __global__ void linear_attention_bf16_coop(
const __nv_bfloat16* q, const __nv_bfloat16* k, const __nv_bfloat16* v,
const __nv_bfloat16* past_state, const __nv_bfloat16* decay,
const __nv_bfloat16* beta,
const __nv_bfloat16* dt_bias, const __nv_bfloat16* neg_exp_A,
__nv_bfloat16* output,
__nv_bfloat16* present_state,
unsigned long long batch, unsigned long long seq,
unsigned long long d_k, unsigned long long d_v,
unsigned long long q_num_heads, unsigned long long kv_num_heads,
unsigned long long n_k_heads, unsigned long long heads_per_group,
unsigned long long kv_per_k_head, unsigned long long output_hidden,
float scale, int needs_decay, int decay_per_key_dim,
int needs_delta, int beta_per_head,
int fuse_beta_sigmoid, int fuse_decay_softplus,
int fuse_neg_exp) {
linear_attention_core_coop<__nv_bfloat16>(
q, k, v, past_state, decay, beta, dt_bias, neg_exp_A, output,
present_state, batch, seq, d_k, d_v, q_num_heads, kv_num_heads, n_k_heads,
heads_per_group, kv_per_k_head, output_hidden, scale, needs_decay,
decay_per_key_dim, needs_delta, beta_per_head, fuse_beta_sigmoid,
fuse_decay_softplus, fuse_neg_exp);
}
"#;
#[derive(Clone, Copy, PartialEq, Eq)]
enum UpdateRule {
Linear,
Gated,
Delta,
GatedDelta,
}
impl UpdateRule {
fn parse(node: &Node) -> Result<Self> {
match node.attr("update_rule").and_then(|a| a.as_str()) {
Some("linear") => Ok(UpdateRule::Linear),
Some("gated") => Ok(UpdateRule::Gated),
Some("delta") => Ok(UpdateRule::Delta),
None | Some("gated_delta") => Ok(UpdateRule::GatedDelta),
Some(other) => Err(EpError::KernelFailed(format!(
"LinearAttention: update_rule must be one of linear, gated, delta, \
gated_delta; got {other:?}"
))),
}
}
fn needs_decay(self) -> bool {
matches!(self, UpdateRule::Gated | UpdateRule::GatedDelta)
}
fn needs_delta(self) -> bool {
matches!(self, UpdateRule::Delta | UpdateRule::GatedDelta)
}
}
fn read_heads(node: &Node, name: &str) -> Option<usize> {
node.attr(name)
.and_then(|a| a.as_int())
.and_then(|v| usize::try_from(v).ok())
.filter(|&v| v > 0)
}
pub(crate) fn unsupported_reason(node: &Node, input_dtypes: &[DataType]) -> Option<String> {
let q_num_heads = read_heads(node, "q_num_heads")?;
let kv_num_heads = read_heads(node, "kv_num_heads")?;
if UpdateRule::parse(node).is_err() {
return Some("LinearAttention: unrecognized update_rule".into());
}
if input_dtypes.len() < 3 {
return Some("LinearAttention: requires query, key, value inputs".into());
}
let dtype = input_dtypes[0];
if !matches!(
dtype,
DataType::Float32 | DataType::Float16 | DataType::BFloat16
) {
return Some(format!(
"LinearAttention dtype {dtype:?} (supported: Float32, Float16, BFloat16)"
));
}
if input_dtypes[..3.min(input_dtypes.len())]
.iter()
.any(|&d| d != dtype)
{
return Some("LinearAttention: query/key/value must share one dtype".into());
}
if q_num_heads >= kv_num_heads {
if !q_num_heads.is_multiple_of(kv_num_heads) {
return Some(format!(
"LinearAttention: q_num_heads {q_num_heads} not a multiple of kv_num_heads \
{kv_num_heads}"
));
}
} else if !kv_num_heads.is_multiple_of(q_num_heads) {
return Some(format!(
"LinearAttention: kv_num_heads {kv_num_heads} not a multiple of q_num_heads \
{q_num_heads} (inverse GQA)"
));
}
None
}
pub struct LinearAttentionFactory {
pub runtime: Arc<CudaRuntime>,
}
impl KernelFactory for LinearAttentionFactory {
fn create(&self, node: &Node, _input_shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>> {
let q_num_heads = read_heads(node, "q_num_heads").ok_or_else(|| {
EpError::KernelFailed(
"LinearAttention: `q_num_heads` must be a positive integer".into(),
)
})?;
let kv_num_heads = read_heads(node, "kv_num_heads").ok_or_else(|| {
EpError::KernelFailed(
"LinearAttention: `kv_num_heads` must be a positive integer".into(),
)
})?;
let update_rule = UpdateRule::parse(node)?;
let scale = match node.attr("scale").and_then(|a| a.as_float()) {
Some(s) if s != 0.0 => Some(s),
_ => None,
};
let fuse_beta_sigmoid = node
.attr(FUSE_BETA_SIGMOID_ATTR)
.and_then(|a| a.as_int())
.is_some_and(|v| v != 0);
let fuse_decay_softplus = node
.attr(FUSE_DECAY_SOFTPLUS_ATTR)
.and_then(|a| a.as_int())
.is_some_and(|v| v != 0);
let fuse_neg_exp = node
.attr(FUSE_NEG_EXP_ATTR)
.and_then(|a| a.as_int())
.is_some_and(|v| v != 0);
Ok(Box::new(LinearAttentionKernel {
runtime: self.runtime.clone(),
q_num_heads,
kv_num_heads,
update_rule,
scale,
fuse_beta_sigmoid,
fuse_decay_softplus,
fuse_neg_exp,
warp_coop: !linattn_warp_coop_disabled(),
}))
}
}
#[derive(Debug)]
struct LinearAttentionKernel {
runtime: Arc<CudaRuntime>,
q_num_heads: usize,
kv_num_heads: usize,
update_rule: UpdateRule,
scale: Option<f32>,
fuse_beta_sigmoid: bool,
fuse_decay_softplus: bool,
fuse_neg_exp: bool,
warp_coop: bool,
}
impl std::fmt::Debug for UpdateRule {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let name = match self {
UpdateRule::Linear => "linear",
UpdateRule::Gated => "gated",
UpdateRule::Delta => "delta",
UpdateRule::GatedDelta => "gated_delta",
};
f.write_str(name)
}
}
impl Kernel for LinearAttentionKernel {
fn execute(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
if inputs.len() < 3 || inputs.len() > 8 {
return Err(EpError::KernelFailed(format!(
"cuda_ep LinearAttention: expected 3..=8 inputs, got {}",
inputs.len()
)));
}
if outputs.len() != 2 {
return Err(EpError::KernelFailed(format!(
"cuda_ep LinearAttention: expected 2 outputs (output, present_state), got {}",
outputs.len()
)));
}
let dtype = inputs[0].dtype;
if !matches!(
dtype,
DataType::Float32 | DataType::Float16 | DataType::BFloat16
) {
return Err(not_implemented(format!(
"LinearAttention dtype {dtype:?} (supported: Float32, Float16, BFloat16)"
)));
}
let q = &inputs[0];
let k = &inputs[1];
let v = &inputs[2];
if q.shape.len() != 3 || k.shape.len() != 3 || v.shape.len() != 3 {
return Err(EpError::KernelFailed(format!(
"cuda_ep LinearAttention: query/key/value must be rank 3 [B, T, H·D], got {:?}, \
{:?}, {:?}",
q.shape, k.shape, v.shape
)));
}
let (batch, seq, q_hidden) = (q.shape[0], q.shape[1], q.shape[2]);
if k.shape[0] != batch || v.shape[0] != batch || k.shape[1] != seq || v.shape[1] != seq {
return Err(EpError::KernelFailed(
"cuda_ep LinearAttention: query/key/value batch and sequence dims must agree"
.into(),
));
}
let q_num_heads = self.q_num_heads;
let kv_num_heads = self.kv_num_heads;
if q_num_heads == 0 || kv_num_heads == 0 || !q_hidden.is_multiple_of(q_num_heads) {
return Err(EpError::KernelFailed(format!(
"cuda_ep LinearAttention: query hidden {q_hidden} not divisible by q_num_heads \
{q_num_heads}"
)));
}
let d_k = q_hidden / q_num_heads;
if d_k == 0 || d_k > MAX_D_K || !k.shape[2].is_multiple_of(d_k) {
return Err(not_implemented(format!(
"cuda_ep LinearAttention: d_k {d_k} unsupported (must be 1..={MAX_D_K} and divide \
key hidden {})",
k.shape[2]
)));
}
let n_k_heads = k.shape[2] / d_k;
if n_k_heads == 0 || !v.shape[2].is_multiple_of(kv_num_heads) {
return Err(EpError::KernelFailed(format!(
"cuda_ep LinearAttention: value hidden {} not divisible by kv_num_heads \
{kv_num_heads}",
v.shape[2]
)));
}
let d_v = v.shape[2] / kv_num_heads;
let heads_per_group = if q_num_heads >= kv_num_heads {
if !q_num_heads.is_multiple_of(kv_num_heads) {
return Err(EpError::KernelFailed(format!(
"cuda_ep LinearAttention: q_num_heads {q_num_heads} must be a multiple of \
kv_num_heads {kv_num_heads}"
)));
}
q_num_heads / kv_num_heads
} else {
if !kv_num_heads.is_multiple_of(q_num_heads) {
return Err(EpError::KernelFailed(format!(
"cuda_ep LinearAttention: kv_num_heads {kv_num_heads} must be a multiple of \
q_num_heads {q_num_heads} (inverse GQA)"
)));
}
0
};
if !kv_num_heads.is_multiple_of(n_k_heads) {
return Err(EpError::KernelFailed(format!(
"cuda_ep LinearAttention: kv_num_heads {kv_num_heads} must be a multiple of \
n_k_heads {n_k_heads}"
)));
}
let kv_per_k_head = kv_num_heads / n_k_heads;
let scale = self.scale.unwrap_or_else(|| 1.0 / (d_k as f32).sqrt());
let needs_decay = self.update_rule.needs_decay();
let needs_delta = self.update_rule.needs_delta();
let past_state = inputs.get(3);
let decay = inputs.get(4);
let beta = inputs.get(5);
let dt_bias = inputs.get(6);
let neg_exp_a = inputs.get(7);
if needs_decay && decay.is_none() {
return Err(EpError::KernelFailed(
"cuda_ep LinearAttention: decay input required for update_rule=gated/gated_delta"
.into(),
));
}
if needs_delta && beta.is_none() {
return Err(EpError::KernelFailed(
"cuda_ep LinearAttention: beta input required for update_rule=delta/gated_delta"
.into(),
));
}
let fuse_beta_sigmoid = self.fuse_beta_sigmoid && needs_delta;
let fuse_decay_softplus = self.fuse_decay_softplus && needs_decay;
let fuse_neg_exp = self.fuse_neg_exp && fuse_decay_softplus;
if fuse_decay_softplus && (dt_bias.is_none() || neg_exp_a.is_none()) {
return Err(EpError::KernelFailed(
"cuda_ep LinearAttention: folded decay softplus requires dt_bias and neg_exp_A \
inputs"
.into(),
));
}
let decay_per_key_dim = if needs_decay {
let s = decay.unwrap().shape;
if s.len() != 3 || s[0] != batch || s[1] != seq {
return Err(EpError::KernelFailed(format!(
"cuda_ep LinearAttention: decay must be [B={batch}, T={seq}, ...], got {s:?}"
)));
}
if s[2] == kv_num_heads * d_k {
true
} else if s[2] == kv_num_heads {
false
} else {
return Err(EpError::KernelFailed(format!(
"cuda_ep LinearAttention: decay last dim must be H_kv={kv_num_heads} or \
H_kv·d_k={}, got {}",
kv_num_heads * d_k,
s[2]
)));
}
} else {
false
};
let beta_per_head = if needs_delta {
let s = beta.unwrap().shape;
if s.len() != 3 || s[0] != batch || s[1] != seq {
return Err(EpError::KernelFailed(format!(
"cuda_ep LinearAttention: beta must be [B={batch}, T={seq}, ...], got {s:?}"
)));
}
if s[2] == kv_num_heads {
true
} else if s[2] == 1 {
false
} else {
return Err(EpError::KernelFailed(format!(
"cuda_ep LinearAttention: beta last dim must be H_kv={kv_num_heads} or 1, got {}",
s[2]
)));
}
} else {
false
};
if fuse_decay_softplus {
let want = if decay_per_key_dim {
kv_num_heads * d_k
} else {
kv_num_heads
};
for (name, view) in [
("dt_bias", dt_bias.unwrap()),
("neg_exp_A", neg_exp_a.unwrap()),
] {
let n: usize = view.shape.iter().product();
if n != want {
return Err(EpError::KernelFailed(format!(
"cuda_ep LinearAttention: folded {name} must have {want} elements, got \
{:?}",
view.shape
)));
}
}
}
if let Some(view) = past_state {
let s = view.shape;
if s.len() != 4 || s[0] != batch || s[1] != kv_num_heads || s[2] != d_k || s[3] != d_v {
return Err(EpError::KernelFailed(format!(
"cuda_ep LinearAttention: past_state must be [B={batch}, H_kv={kv_num_heads}, \
d_k={d_k}, d_v={d_v}], got {s:?}"
)));
}
}
let output_hidden = q_num_heads.max(kv_num_heads) * d_v;
if outputs[0].shape != [batch, seq, output_hidden] {
return Err(EpError::KernelFailed(format!(
"cuda_ep LinearAttention: output must be [B={batch}, T={seq}, {output_hidden}], \
got {:?}",
outputs[0].shape
)));
}
if outputs[1].shape != [batch, kv_num_heads, d_k, d_v] {
return Err(EpError::KernelFailed(format!(
"cuda_ep LinearAttention: present_state must be [B={batch}, H_kv={kv_num_heads}, \
d_k={d_k}, d_v={d_v}], got {:?}",
outputs[1].shape
)));
}
let all_inputs_ok = inputs
.iter()
.all(|input| input.dtype == dtype && input.is_contiguous());
if !all_inputs_ok
|| outputs[0].dtype != dtype
|| outputs[1].dtype != dtype
|| !outputs[0].is_contiguous()
|| !outputs[1].is_contiguous()
{
return Err(not_implemented(
"LinearAttention requires contiguous, uniform-dtype tensors",
));
}
let total = (batch * kv_num_heads * d_v) as u64;
if total == 0 || seq == 0 {
return Ok(());
}
let stem = match (dtype, self.warp_coop) {
(DataType::Float32, false) => "linear_attention_f32",
(DataType::Float16, false) => "linear_attention_f16",
(DataType::BFloat16, false) => "linear_attention_bf16",
(DataType::Float32, true) => "linear_attention_f32_coop",
(DataType::Float16, true) => "linear_attention_f16_coop",
(DataType::BFloat16, true) => "linear_attention_bf16_coop",
_ => unreachable!(),
};
if dtype != DataType::Float32 {
self.runtime.require_nvrtc_half_headers("LinearAttention")?;
}
let function = self
.runtime
.nvrtc_function("linear_attention_v3", SOURCE, stem)?;
let q_ptr = cuptr(q.data_ptr::<u8>() as *const c_void);
let k_ptr = cuptr(k.data_ptr::<u8>() as *const c_void);
let v_ptr = cuptr(v.data_ptr::<u8>() as *const c_void);
let past_ptr = past_state
.map(|t| cuptr(t.data_ptr::<u8>() as *const c_void))
.unwrap_or(0);
let decay_ptr = if needs_decay {
cuptr(decay.unwrap().data_ptr::<u8>() as *const c_void)
} else {
0
};
let beta_ptr = if needs_delta {
cuptr(beta.unwrap().data_ptr::<u8>() as *const c_void)
} else {
0
};
let dt_bias_ptr = if fuse_decay_softplus {
cuptr(dt_bias.unwrap().data_ptr::<u8>() as *const c_void)
} else {
0
};
let neg_exp_a_ptr = if fuse_decay_softplus {
cuptr(neg_exp_a.unwrap().data_ptr::<u8>() as *const c_void)
} else {
0
};
let output_ptr = cuptr(outputs[0].data_ptr_mut::<u8>() as *const c_void);
let present_ptr = cuptr(outputs[1].data_ptr_mut::<u8>() as *const c_void);
let batch = batch as u64;
let seq = seq as u64;
let d_k = d_k as u64;
let d_v = d_v as u64;
let q_num_heads = q_num_heads as u64;
let kv_num_heads = kv_num_heads as u64;
let n_k_heads = n_k_heads as u64;
let heads_per_group = heads_per_group as u64;
let kv_per_k_head = kv_per_k_head as u64;
let output_hidden = output_hidden as u64;
let needs_decay_i = i32::from(needs_decay);
let decay_per_key_dim_i = i32::from(decay_per_key_dim);
let needs_delta_i = i32::from(needs_delta);
let beta_per_head_i = i32::from(beta_per_head);
let fuse_beta_sigmoid_i = i32::from(fuse_beta_sigmoid);
let fuse_decay_softplus_i = i32::from(fuse_decay_softplus);
let fuse_neg_exp_i = i32::from(fuse_neg_exp);
let threads_needed = if self.warp_coop {
total.saturating_mul(u64::from(LA_WARP))
} else {
total
};
let grid = u32::try_from(threads_needed.div_ceil(BLOCK as u64))
.unwrap_or(u32::MAX)
.clamp(1, 65_535);
let stream = self.runtime.stream();
let mut builder = stream.launch_builder(&function);
builder
.arg(&q_ptr)
.arg(&k_ptr)
.arg(&v_ptr)
.arg(&past_ptr)
.arg(&decay_ptr)
.arg(&beta_ptr)
.arg(&dt_bias_ptr)
.arg(&neg_exp_a_ptr)
.arg(&output_ptr)
.arg(&present_ptr)
.arg(&batch)
.arg(&seq)
.arg(&d_k)
.arg(&d_v)
.arg(&q_num_heads)
.arg(&kv_num_heads)
.arg(&n_k_heads)
.arg(&heads_per_group)
.arg(&kv_per_k_head)
.arg(&output_hidden)
.arg(&scale)
.arg(&needs_decay_i)
.arg(&decay_per_key_dim_i)
.arg(&needs_delta_i)
.arg(&beta_per_head_i)
.arg(&fuse_beta_sigmoid_i)
.arg(&fuse_decay_softplus_i)
.arg(&fuse_neg_exp_i);
unsafe {
builder.launch(LaunchConfig {
grid_dim: (grid, 1, 1),
block_dim: (BLOCK, 1, 1),
shared_mem_bytes: 0,
})
}
.map_err(|error| driver_err("launch LinearAttention", error))?;
if self.runtime.is_capturing()? {
return Ok(());
}
self.runtime.synchronize()
}
fn supports_strided_input(&self, _idx: usize) -> bool {
false
}
}