use crate::Engine;
use crate::cache::{Cache, RecurLayer};
use crate::model::GpuTensor;
use cudarc::driver::{CudaSlice, LaunchConfig, PushKernelArg};
use memra_gguf::model_plan::KimiDeltaNetPlan;
use memra_gguf::source::TensorSource;
use std::sync::atomic::{AtomicU64, Ordering};
pub static KDA_FUSED6_DISPATCHES: AtomicU64 = AtomicU64::new(0);
pub static KDA_FUSED6_BF16_DISPATCHES: AtomicU64 = AtomicU64::new(0);
pub static KDA_FUSED6_Q8RP_DISPATCHES: AtomicU64 = AtomicU64::new(0);
pub const KDA_HEAD_DIM: usize = 128;
const KDA_MAX_CONV_KERNEL: usize = 8;
const KDA_L2_EPS: f32 = 1e-6;
pub struct KdaAttnLayer {
pub plan: KimiDeltaNetPlan,
pub wq: GpuTensor,
pub wk: GpuTensor,
pub wv: GpuTensor,
pub f_a: GpuTensor,
pub f_b: GpuTensor,
pub g_a: GpuTensor,
pub g_b: GpuTensor,
pub b_proj: GpuTensor,
pub wo: GpuTensor,
pub conv: CudaSlice<f32>,
pub a_log: GpuTensor,
pub dt_bias: GpuTensor,
pub o_norm: GpuTensor,
pub tp: Option<Box<crate::glm5_tp::Glm5TpKda>>,
}
impl KdaAttnLayer {
pub fn heads(&self) -> usize {
self.plan.num_heads as usize
}
pub fn head_dim(&self) -> usize {
self.plan.head_dim as usize
}
pub fn qkv(&self) -> usize {
self.heads() * self.head_dim()
}
pub fn conv_kernel(&self) -> usize {
self.plan.conv_kernel as usize
}
pub fn conv_width(&self) -> usize {
3 * self.qkv()
}
pub fn state_width(&self) -> usize {
self.heads() * self.head_dim() * self.head_dim()
}
pub fn load(
e: &Engine,
src: &dyn TensorSource,
il: u32,
plan: &KimiDeltaNetPlan,
) -> Result<Self, Box<dyn std::error::Error>> {
let heads = plan.num_heads as usize;
let head_dim = plan.head_dim as usize;
let kernel = plan.conv_kernel as usize;
if head_dim != KDA_HEAD_DIM {
return Err(format!(
"blk.{il}: KDA head_dim {head_dim} is not the {KDA_HEAD_DIM} the scan kernel is \
instantiated for; a new memra_kda_scan_s<N> instantiation is required before \
this geometry can serve"
)
.into());
}
if heads == 0 {
return Err(format!("blk.{il}: KDA num_heads must be positive").into());
}
if !(2..=KDA_MAX_CONV_KERNEL).contains(&kernel) {
return Err(format!(
"blk.{il}: KDA conv_kernel {kernel} outside the 2..={KDA_MAX_CONV_KERNEL} window \
the conv kernels hold in registers"
)
.into());
}
let p = |s: &str| format!("blk.{il}.{s}");
let load = |name: String| GpuTensor::load_from_source(e, src, &name);
let qkv = heads * head_dim;
let mut conv = e.zeros(3 * qkv * kernel)?;
for (plane, name) in [
"kda_q_conv1d.weight",
"kda_k_conv1d.weight",
"kda_v_conv1d.weight",
]
.into_iter()
.enumerate()
{
let w = load(p(name))?;
let src_data = w.float_data();
if src_data.len() != qkv * kernel {
return Err(format!(
"blk.{il}.{name}: {} elements, contract requires {}",
src_data.len(),
qkv * kernel
)
.into());
}
e.copy_into(&mut conv, plane * qkv * kernel, src_data, qkv * kernel)?;
}
Ok(Self {
plan: *plan,
wq: load(p("kda_q.weight"))?,
wk: load(p("kda_k.weight"))?,
wv: load(p("kda_v.weight"))?,
f_a: load(p("kda_f_a.weight"))?,
f_b: load(p("kda_f_b.weight"))?,
g_a: load(p("kda_g_a.weight"))?,
g_b: load(p("kda_g_b.weight"))?,
b_proj: load(p("kda_b.weight"))?,
wo: load(p("kda_out.weight"))?,
conv,
a_log: load(p("kda_a_log"))?,
dt_bias: load(p("kda_dt.bias"))?,
o_norm: load(p("kda_o_norm.weight"))?,
tp: None,
})
}
}
#[derive(Clone, Copy, PartialEq, Eq)]
pub(crate) enum ConvArm {
Prefill,
Decode,
}
pub struct KdaScanInputs {
pub q: CudaSlice<f32>,
pub k: CudaSlice<f32>,
pub v: CudaSlice<f32>,
pub g: CudaSlice<f32>,
pub beta: CudaSlice<f32>,
}
pub struct KdaRowsStash {
pub ring_snap: CudaSlice<f32>,
pub raws: [CudaSlice<f32>; 3],
pub scan: KdaScanInputs,
pub rows: usize,
}
pub(crate) enum KdaStash<'a> {
None,
Decode(&'a mut Option<KdaScanInputs>),
Rows(&'a mut Option<KdaRowsStash>),
}
#[allow(clippy::too_many_arguments)] fn kda_core(
e: &Engine,
la: &KdaAttnLayer,
x: &CudaSlice<f32>,
t: usize,
eps: f32,
ring: &mut CudaSlice<f32>,
state_in: &CudaSlice<f32>,
state_out: &mut CudaSlice<f32>,
arm: ConvArm,
stash: KdaStash<'_>,
scan_clock: Option<&mut u64>,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
if la.tp.is_some() {
return Err(format!(
"KDA layer is glm5-TP-sharded (MEMRA_GLM5_TP): the plain mixer path is unwired \
for a head shard — only the TP decode/prime walk may execute it (t={t}, arm \
{})",
if arm == ConvArm::Decode {
"decode"
} else {
"prefill"
}
)
.into());
}
let rows_exact = matches!(stash, KdaStash::Rows(_));
let gated = kda_core_gated(
e, la, x, t, eps, ring, state_in, state_out, arm, stash, scan_clock,
)?;
if rows_exact {
let y = e.matmul_rows_exact(&la.wo, &gated, t);
e.vws_recycle(gated);
y
} else {
e.matmul(&la.wo, &gated, t)
}
}
#[allow(clippy::too_many_arguments)] pub(crate) fn kda_core_gated(
e: &Engine,
la: &KdaAttnLayer,
x: &CudaSlice<f32>,
t: usize,
eps: f32,
ring: &mut CudaSlice<f32>,
state_in: &CudaSlice<f32>,
state_out: &mut CudaSlice<f32>,
arm: ConvArm,
stash: KdaStash<'_>,
mut scan_clock: Option<&mut u64>,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let heads = la.heads();
let head_dim = la.head_dim();
let qkv = la.qkv();
let kernel = la.conv_kernel();
let rows_exact = matches!(stash, KdaStash::Rows(_));
if rows_exact && arm != ConvArm::Prefill {
return Err("KDA rows stash requires the prefill conv arm".into());
}
if arm == ConvArm::Decode && t != 1 {
return Err(format!("KDA decode arm requires t == 1, got {t}").into());
}
if ring.len() < la.conv_width() * (kernel - 1) {
return Err(format!(
"KDA conv ring holds {} floats, layer needs {}",
ring.len(),
la.conv_width() * (kernel - 1)
)
.into());
}
if state_in.len() < la.state_width() || state_out.len() < la.state_width() {
return Err(format!(
"KDA recurrent state holds {}/{} floats, layer needs {}",
state_in.len(),
state_out.len(),
la.state_width()
)
.into());
}
let mut g6 = match e.kda_proj_fused6(la, x, t)? {
Some(outs) => outs,
None if rows_exact => {
[&la.wq, &la.wk, &la.wv, &la.f_a, &la.g_a, &la.b_proj]
.into_iter()
.map(|w| e.matmul_rows_exact(w, x, t))
.collect::<Result<Vec<_>, _>>()?
}
None => e.matmul_group(
&[&la.wq, &la.wk, &la.wv, &la.f_a, &la.g_a, &la.b_proj],
x,
t,
)?,
};
let beta_raw = g6.pop().unwrap(); let gate_down = g6.pop().unwrap(); let forget_down = g6.pop().unwrap(); let v_raw = g6.pop().unwrap(); let k_raw = g6.pop().unwrap();
let q_raw = g6.pop().unwrap();
let ring_snap = match &stash {
KdaStash::Rows(_) => {
let mut snap = e.vws_uninit(ring.len())?;
e.dtod_copy_into(ring, &mut snap, 0)?;
Some(snap)
}
_ => None,
};
let mut q_conv = if rows_exact {
e.vws_uninit(t * qkv)?
} else {
e.uninit(t * qkv)?
};
let mut k_conv = if rows_exact {
e.vws_uninit(t * qkv)?
} else {
e.uninit(t * qkv)?
};
let mut v_conv = if rows_exact {
e.vws_uninit(t * qkv)?
} else {
e.uninit(t * qkv)?
};
for (plane, (raw, out)) in [
(&q_raw, &mut q_conv),
(&k_raw, &mut k_conv),
(&v_raw, &mut v_conv),
]
.into_iter()
.enumerate()
{
match arm {
ConvArm::Prefill => e.kda_conv_silu(raw, &la.conv, ring, out, qkv, t, kernel, plane)?,
ConvArm::Decode => {
e.kda_conv_silu_decode(raw, ring, &la.conv, out, qkv, kernel, plane)?
}
}
}
if arm == ConvArm::Prefill {
for (plane, raw) in [&q_raw, &k_raw, &v_raw].into_iter().enumerate() {
e.kda_conv_ring_roll(raw, ring, qkv, t, kernel, plane)?;
}
}
let mut q_l2 = if rows_exact {
e.vws_uninit(t * qkv)?
} else {
e.uninit(t * qkv)?
};
let mut k_l2 = if rows_exact {
e.vws_uninit(t * qkv)?
} else {
e.uninit(t * qkv)?
};
e.l2_norm(&q_conv, &mut q_l2, head_dim, t * heads, KDA_L2_EPS)?;
e.l2_norm(&k_conv, &mut k_l2, head_dim, t * heads, KDA_L2_EPS)?;
if rows_exact {
e.vws_recycle(q_conv);
e.vws_recycle(k_conv);
}
let forget = if rows_exact {
e.matmul_rows_exact(&la.f_b, &forget_down, t)?
} else {
e.matmul(&la.f_b, &forget_down, t)?
};
let mut g_log = if rows_exact {
e.vws_uninit(t * qkv)?
} else {
e.uninit(t * qkv)?
};
e.kda_gate(
&forget,
la.dt_bias.float_data(),
la.a_log.float_data(),
&mut g_log,
qkv,
t,
head_dim,
la.plan.gate_lower_bound,
)?;
let mut beta = if rows_exact {
e.vws_uninit(t * heads)?
} else {
e.uninit(t * heads)?
};
e.sigmoid(&beta_raw, &mut beta, t * heads)?;
if rows_exact {
e.vws_recycle(forget_down);
e.vws_recycle(forget);
e.vws_recycle(beta_raw);
}
let scale = 1.0 / (head_dim as f32).sqrt();
let mut core = if rows_exact {
e.vws_uninit(t * qkv)?
} else {
e.uninit(t * qkv)?
};
let scan_t0 = scan_clock.as_ref().map(|_| {
let _ = e.stream().synchronize();
std::time::Instant::now()
});
e.kda_scan(
&q_l2, &k_l2, &v_conv, &g_log, &beta, state_in, state_out, &mut core, heads, t, scale,
)?;
if let (Some(ns), Some(t0)) = (scan_clock.take(), scan_t0) {
let _ = e.stream().synchronize();
*ns += t0.elapsed().as_nanos() as u64;
}
let gate = if rows_exact {
e.matmul_rows_exact(&la.g_b, &gate_down, t)?
} else {
e.matmul(&la.g_b, &gate_down, t)?
};
let mut gated = if rows_exact {
e.vws_uninit(t * qkv)?
} else {
e.uninit(t * qkv)?
};
e.kda_gated_rmsnorm(
&core,
la.o_norm.float_data(),
&gate,
&mut gated,
head_dim,
t * heads,
eps,
)?;
if rows_exact {
e.vws_recycle(gate_down);
e.vws_recycle(gate);
e.vws_recycle(core);
}
match stash {
KdaStash::None => {}
KdaStash::Decode(s) => {
*s = Some(KdaScanInputs {
q: q_l2,
k: k_l2,
v: v_conv,
g: g_log,
beta,
});
}
KdaStash::Rows(s) => {
if let Some(old) = s.take() {
e.vws_recycle(old.ring_snap);
for r in old.raws {
e.vws_recycle(r);
}
e.vws_recycle(old.scan.q);
e.vws_recycle(old.scan.k);
e.vws_recycle(old.scan.v);
e.vws_recycle(old.scan.g);
e.vws_recycle(old.scan.beta);
}
*s = Some(KdaRowsStash {
ring_snap: ring_snap.expect("rows arm snapshotted the ring above"),
raws: [q_raw, k_raw, v_raw],
scan: KdaScanInputs {
q: q_l2,
k: k_l2,
v: v_conv,
g: g_log,
beta,
},
rows: t,
});
}
}
Ok(gated)
}
pub fn kda_attn(
e: &Engine,
la: &KdaAttnLayer,
x: &CudaSlice<f32>,
t: usize,
eps: f32,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let mut ring = e.zeros(la.conv_width() * (la.conv_kernel() - 1))?;
let state_in = e.zeros(la.state_width())?;
let mut state_out = e.zeros(la.state_width())?;
kda_core(
e,
la,
x,
t,
eps,
&mut ring,
&state_in,
&mut state_out,
ConvArm::Prefill,
KdaStash::None,
None,
)
}
#[allow(clippy::too_many_arguments)] pub fn kda_attn_prime(
e: &Engine,
la: &KdaAttnLayer,
x: &CudaSlice<f32>,
t: usize,
eps: f32,
ring: &mut CudaSlice<f32>,
state_in: &CudaSlice<f32>,
state_out: &mut CudaSlice<f32>,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
kda_core(
e,
la,
x,
t,
eps,
ring,
state_in,
state_out,
ConvArm::Prefill,
KdaStash::None,
None,
)
}
pub fn kda_attn_decode(
e: &Engine,
la: &KdaAttnLayer,
x: &CudaSlice<f32>,
eps: f32,
ring: &mut CudaSlice<f32>,
state_in: &CudaSlice<f32>,
state_out: &mut CudaSlice<f32>,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
kda_core(
e,
la,
x,
1,
eps,
ring,
state_in,
state_out,
ConvArm::Decode,
KdaStash::None,
None,
)
}
#[allow(clippy::too_many_arguments)] fn kda_cached(
e: &Engine,
la: &KdaAttnLayer,
x: &CudaSlice<f32>,
t: usize,
eps: f32,
cache: &mut Cache,
il: usize,
arm: ConvArm,
stash: KdaStash<'_>,
scan_clock: Option<&mut u64>,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let rl = cache.recur[il].as_mut().ok_or_else(|| {
format!(
"blk.{il}: KDA layer has no recurrent state — the cache allocator saw a \
non-Recurrent StatePlan for a KDA layer"
)
})?;
let out = {
let RecurLayer {
conv_state,
ssm_state,
ssm_state_alt,
} = rl;
kda_core(
e,
la,
x,
t,
eps,
conv_state,
ssm_state,
ssm_state_alt,
arm,
stash,
scan_clock,
)?
};
std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
Ok(out)
}
pub fn kda_prime_cached(
e: &Engine,
la: &KdaAttnLayer,
x: &CudaSlice<f32>,
t: usize,
eps: f32,
cache: &mut Cache,
il: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
kda_cached(
e,
la,
x,
t,
eps,
cache,
il,
ConvArm::Prefill,
KdaStash::None,
None,
)
}
pub fn kda_decode_cached(
e: &Engine,
la: &KdaAttnLayer,
x: &CudaSlice<f32>,
eps: f32,
cache: &mut Cache,
il: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
kda_cached(
e,
la,
x,
1,
eps,
cache,
il,
ConvArm::Decode,
KdaStash::None,
None,
)
}
pub fn kda_decode_cached_stash(
e: &Engine,
la: &KdaAttnLayer,
x: &CudaSlice<f32>,
eps: f32,
cache: &mut Cache,
il: usize,
) -> Result<(CudaSlice<f32>, KdaScanInputs), Box<dyn std::error::Error>> {
let mut stash: Option<KdaScanInputs> = None;
let out = kda_cached(
e,
la,
x,
1,
eps,
cache,
il,
ConvArm::Decode,
KdaStash::Decode(&mut stash),
None,
)?;
let stash = stash.ok_or("kda_core returned without filling the requested scan stash")?;
Ok((out, stash))
}
#[allow(clippy::too_many_arguments)] pub fn kda_verify_rows_cached(
e: &Engine,
la: &KdaAttnLayer,
x: &CudaSlice<f32>,
t: usize,
eps: f32,
cache: &mut Cache,
il: usize,
scan_clock: Option<&mut u64>,
) -> Result<(CudaSlice<f32>, KdaRowsStash), Box<dyn std::error::Error>> {
let mut stash: Option<KdaRowsStash> = None;
let out = kda_cached(
e,
la,
x,
t,
eps,
cache,
il,
ConvArm::Prefill,
KdaStash::Rows(&mut stash),
scan_clock,
)?;
let stash = stash.ok_or("kda_core returned without filling the requested rows stash")?;
Ok((out, stash))
}
pub fn kda_verify_rollback_rows(
e: &Engine,
la: &KdaAttnLayer,
snap: &CudaSlice<f32>,
stash: &KdaRowsStash,
keep: usize,
cache: &mut Cache,
il: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let rl = cache.recur[il]
.as_mut()
.ok_or_else(|| format!("blk.{il}: KDA rows rollback on a layer with no recurrent state"))?;
kda_verify_rollback_rows_on(e, la, snap, stash, keep, rl, il)
}
pub fn kda_verify_rollback_rows_on(
e: &Engine,
la: &KdaAttnLayer,
snap: &CudaSlice<f32>,
stash: &KdaRowsStash,
keep: usize,
rl: &mut RecurLayer,
il: usize,
) -> Result<(), Box<dyn std::error::Error>> {
if keep == 0 || keep >= stash.rows {
return Err(format!(
"blk.{il}: KDA rows rollback keep={keep} outside 1..{} (full accept keeps the \
resident state and never replays)",
stash.rows
)
.into());
}
let qkv = la.qkv();
let kernel = la.conv_kernel();
let heads = la.heads();
let scale = 1.0 / (la.head_dim() as f32).sqrt();
e.copy_into(
&mut rl.conv_state,
0,
&stash.ring_snap,
stash.ring_snap.len(),
)?;
for (plane, raw) in stash.raws.iter().enumerate() {
e.kda_conv_ring_roll(raw, &mut rl.conv_state, qkv, keep, kernel, plane)?;
}
let mut o = e.uninit(keep * qkv)?;
{
let RecurLayer {
ssm_state: _,
ssm_state_alt,
..
} = rl;
e.kda_scan(
&stash.scan.q,
&stash.scan.k,
&stash.scan.v,
&stash.scan.g,
&stash.scan.beta,
snap,
ssm_state_alt,
&mut o,
heads,
keep,
scale,
)?;
}
std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
Ok(())
}
pub fn kda_scan_replay(
e: &Engine,
la: &KdaAttnLayer,
snap: &CudaSlice<f32>,
inputs: &[KdaScanInputs],
cache: &mut Cache,
il: usize,
) -> Result<(), Box<dyn std::error::Error>> {
if inputs.is_empty() {
return Err(format!(
"blk.{il}: KDA replay needs at least one stashed row (rollback keep >= 1; a \
restore TO the snapshot itself is a different contract)"
)
.into());
}
if la.tp.is_some() {
return Err(format!(
"blk.{il}: KDA scan replay (the PER-ROW rollback seam) is unwired for a \
glm5-TP-sharded layer — the spec x TP composition requires the BATCHED \
verify walk, whose rollback rides kda_verify_rollback_rows_on per rank"
)
.into());
}
let heads = la.heads();
let scale = 1.0 / (la.head_dim() as f32).sqrt();
let qkv = la.qkv();
let rl = cache.recur[il]
.as_mut()
.ok_or_else(|| format!("blk.{il}: KDA replay on a layer with no recurrent state"))?;
let mut o = e.uninit(qkv)?; for (r, inp) in inputs.iter().enumerate() {
{
let RecurLayer {
ssm_state,
ssm_state_alt,
..
} = rl;
let state_in: &CudaSlice<f32> = if r == 0 { snap } else { ssm_state };
e.kda_scan(
&inp.q,
&inp.k,
&inp.v,
&inp.g,
&inp.beta,
state_in,
ssm_state_alt,
&mut o,
heads,
1,
scale,
)?;
}
std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
}
Ok(())
}
impl Engine {
#[allow(clippy::too_many_arguments)]
pub fn kda_conv_silu(
&self,
x_tm: &CudaSlice<f32>,
w: &CudaSlice<f32>,
ring: &CudaSlice<f32>,
y_tm: &mut CudaSlice<f32>,
qkv: usize,
t: usize,
kernel: usize,
plane: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("memra_kda_conv_silu_f32");
let cfg = LaunchConfig {
grid_dim: (qkv.div_ceil(256) as u32, t as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (n, tt, k, p) = (qkv as i32, t as i32, kernel as i32, plane as i32);
let stream = self.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(x_tm)
.arg(w)
.arg(ring)
.arg(&mut *y_tm)
.arg(&n)
.arg(&tt)
.arg(&k)
.arg(&p);
unsafe { b.launch(cfg)? };
Ok(())
}
pub fn kda_conv_ring_roll(
&self,
x_tm: &CudaSlice<f32>,
ring: &mut CudaSlice<f32>,
qkv: usize,
t: usize,
kernel: usize,
plane: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("memra_kda_conv_ring_roll_f32");
let cfg = LaunchConfig {
grid_dim: (qkv.div_ceil(256) as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (n, tt, k, p) = (qkv as i32, t as i32, kernel as i32, plane as i32);
let stream = self.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(x_tm).arg(&mut *ring).arg(&n).arg(&tt).arg(&k).arg(&p);
unsafe { b.launch(cfg)? };
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn kda_conv_silu_decode(
&self,
x_new: &CudaSlice<f32>,
ring: &mut CudaSlice<f32>,
w: &CudaSlice<f32>,
y: &mut CudaSlice<f32>,
qkv: usize,
kernel: usize,
plane: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("memra_kda_conv_silu_decode_f32");
let cfg = LaunchConfig {
grid_dim: (qkv.div_ceil(256) as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (n, k, p) = (qkv as i32, kernel as i32, plane as i32);
let stream = self.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(x_new)
.arg(&mut *ring)
.arg(w)
.arg(&mut *y)
.arg(&n)
.arg(&k)
.arg(&p);
unsafe { b.launch(cfg)? };
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn kda_gate(
&self,
forget: &CudaSlice<f32>,
dt_bias: &CudaSlice<f32>,
a_log: &CudaSlice<f32>,
g: &mut CudaSlice<f32>,
qkv: usize,
t: usize,
head_dim: usize,
lower_bound: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("memra_kda_gate_f32");
let cfg = LaunchConfig {
grid_dim: (qkv.div_ceil(256) as u32, t as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (n, tt, hd, lb) = (qkv as i32, t as i32, head_dim as i32, lower_bound);
let stream = self.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(forget)
.arg(dt_bias)
.arg(a_log)
.arg(&mut *g)
.arg(&n)
.arg(&tt)
.arg(&hd)
.arg(&lb);
unsafe { b.launch(cfg)? };
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn kda_scan(
&self,
q: &CudaSlice<f32>,
k: &CudaSlice<f32>,
v: &CudaSlice<f32>,
g: &CudaSlice<f32>,
beta: &CudaSlice<f32>,
state_in: &CudaSlice<f32>,
state_out: &mut CudaSlice<f32>,
o: &mut CudaSlice<f32>,
heads: usize,
t: usize,
scale: f32,
) -> Result<(), Box<dyn std::error::Error>> {
const COLS_PER_BLOCK: u32 = 4;
let f = self.func("memra_kda_scan_s128");
let cfg = LaunchConfig {
grid_dim: (
heads as u32,
1,
(KDA_HEAD_DIM as u32).div_ceil(COLS_PER_BLOCK),
),
block_dim: (32, COLS_PER_BLOCK, 1),
shared_mem_bytes: 0,
};
let (h, tt, s) = (heads as i32, t as i32, scale);
let stream = self.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(v)
.arg(g)
.arg(beta)
.arg(state_in)
.arg(&mut *state_out)
.arg(&mut *o)
.arg(&h)
.arg(&tt)
.arg(&s);
unsafe { b.launch(cfg)? };
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn kda_gated_rmsnorm(
&self,
core: &CudaSlice<f32>,
w: &CudaSlice<f32>,
gate: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("memra_kda_gated_rmsnorm_f32");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (nc, ep) = (ncols as i32, eps);
let stream = self.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(core)
.arg(w)
.arg(gate)
.arg(&mut *dst)
.arg(&nc)
.arg(&ep);
unsafe { b.launch(cfg)? };
Ok(())
}
pub fn kda_proj_fused6(
&self,
la: &KdaAttnLayer,
x: &CudaSlice<f32>,
t: usize,
) -> Result<Option<Vec<CudaSlice<f32>>>, Box<dyn std::error::Error>> {
if std::env::var("MEMRA_KDA_FUSED_PROJ").as_deref() != Ok("1") {
return Ok(None);
}
if la.tp.is_some() {
static TP_F6_DECLINE: std::sync::Once = std::sync::Once::new();
TP_F6_DECLINE.call_once(|| {
eprintln!(
"[kda-fused-proj] DECLINED on a glm5-TP head shard: the door is gated \
on full-width projections (the load preflight refuses the pair; this \
is the per-call twin for a post-load flag set)"
);
});
return Ok(None);
}
if !(1..=15).contains(&t) {
return Ok(None);
}
let f32w = |w: &GpuTensor| -> Option<usize> {
match w {
GpuTensor::Float { .. } => Some(w.in_features()),
_ => None,
}
};
let (Some(in_fa), Some(in_ga), Some(in_b)) =
(f32w(&la.f_a), f32w(&la.g_a), f32w(&la.b_proj))
else {
return Ok(None);
};
let bf16 = |w: &GpuTensor| -> Option<usize> {
match w {
GpuTensor::FloatBf16 { .. } => Some(w.in_features()),
_ => None,
}
};
if let (Some(in_q), Some(in_k), Some(in_v)) = (bf16(&la.wq), bf16(&la.wk), bf16(&la.wv)) {
if crate::glm5_w8_on() && !(crate::step_tp_w8_on() && crate::w8_hybrid_on()) {
if !Self::bf16_mmv_on() || !crate::b200_gemv_v2_on() {
return Ok(None);
}
let in_f = in_q;
if [in_k, in_v, in_fa, in_ga, in_b].iter().any(|&i| i != in_f)
|| !in_f.is_multiple_of(128)
|| x.len() < t * in_f
|| Engine::q8_v2_smem_bytes(in_f) > 48 * 1024
{
return Ok(None);
}
let dims = [
la.wq.out_features(),
la.wk.out_features(),
la.wv.out_features(),
la.f_a.out_features(),
la.g_a.out_features(),
la.b_proj.out_features(),
];
let (
GpuTensor::FloatBf16 { data: bq, .. },
GpuTensor::FloatBf16 { data: bk, .. },
GpuTensor::FloatBf16 { data: bv, .. },
) = (&la.wq, &la.wk, &la.wv)
else {
unreachable!("bf16() above only admits FloatBf16");
};
let (
GpuTensor::Float { data: wfa, .. },
GpuTensor::Float { data: wga, .. },
GpuTensor::Float { data: wb, .. },
) = (&la.f_a, &la.g_a, &la.b_proj)
else {
unreachable!("f32w() above only admits Float");
};
let mut outs = [
self.uninit(t * dims[0])?,
self.uninit(t * dims[1])?,
self.uninit(t * dims[2])?,
self.uninit(t * dims[3])?,
self.uninit(t * dims[4])?,
self.uninit(t * dims[5])?,
];
self.kda_proj_fused6_q8rp_raw(
bq, bk, bv, wfa, wga, wb, x, &mut outs, in_f, dims, t,
)?;
if KDA_FUSED6_Q8RP_DISPATCHES.fetch_add(1, Ordering::Relaxed) == 0 {
eprintln!(
"[kda-fused6] engaged arm=q8rp_v2 in_f={in_f} out={dims:?} t={t} (one \
launch replaces the six W8-mirror projections and their six redundant \
activation quantizes; MEMRA_KDA_FUSED_PROJ=1 MEMRA_B200_GEMV_V2=1)"
);
}
return Ok(Some(outs.into_iter().collect()));
}
if !Self::bf16_mmv_on()
|| (crate::step_tp_w8_on() && crate::w8_hybrid_on())
|| crate::b200_bf16_gemv_lt_on()
{
return Ok(None);
}
let in_f = in_q;
if [in_k, in_v, in_fa, in_ga, in_b].iter().any(|&i| i != in_f)
|| !in_f.is_multiple_of(128)
|| x.len() < t * in_f
{
return Ok(None);
}
let dims = [
la.wq.out_features(),
la.wk.out_features(),
la.wv.out_features(),
la.f_a.out_features(),
la.g_a.out_features(),
la.b_proj.out_features(),
];
let (
GpuTensor::FloatBf16 { data: bq, .. },
GpuTensor::FloatBf16 { data: bk, .. },
GpuTensor::FloatBf16 { data: bv, .. },
) = (&la.wq, &la.wk, &la.wv)
else {
unreachable!("bf16() above only admits FloatBf16");
};
let (
GpuTensor::Float { data: wfa, .. },
GpuTensor::Float { data: wga, .. },
GpuTensor::Float { data: wb, .. },
) = (&la.f_a, &la.g_a, &la.b_proj)
else {
unreachable!("f32w() above only admits Float");
};
let mut outs = [
self.uninit(t * dims[0])?,
self.uninit(t * dims[1])?,
self.uninit(t * dims[2])?,
self.uninit(t * dims[3])?,
self.uninit(t * dims[4])?,
self.uninit(t * dims[5])?,
];
self.kda_proj_fused6_bf16_raw(bq, bk, bv, wfa, wga, wb, x, &mut outs, in_f, dims, t)?;
if KDA_FUSED6_BF16_DISPATCHES.fetch_add(1, Ordering::Relaxed) == 0 {
eprintln!(
"[kda-fused6] engaged arm=bf16 in_f={in_f} out={dims:?} t={t} (one launch \
replaces the six-projection group on the bf16-resident serving recipe; \
MEMRA_KDA_FUSED_PROJ=1)"
);
}
return Ok(Some(outs.into_iter().collect()));
}
if std::env::var("MEMRA_FAST").as_deref() == Ok("0")
|| !self.mmvq_supports(crate::QT_Q8_0)
|| (t >= 2 && std::env::var("MEMRA_NO_BATCHED").is_ok())
|| (t >= 5 && !Self::b8_enabled())
{
return Ok(None);
}
let q8 = |w: &GpuTensor| -> Option<(usize, usize)> {
match w {
GpuTensor::Quant {
qtype: crate::QT_Q8_0,
row_bytes,
scale,
rp: false,
rp4: None,
..
} if *scale == 1.0 => Some((w.in_features(), *row_bytes)),
_ => None,
}
};
let (Some((in_q, rb_q)), Some((in_k, rb_k)), Some((in_v, rb_v))) =
(q8(&la.wq), q8(&la.wk), q8(&la.wv))
else {
return Ok(None);
};
let in_f = in_q;
if [in_k, in_v, in_fa, in_ga, in_b].iter().any(|&i| i != in_f)
|| rb_k != rb_q
|| rb_v != rb_q
|| !in_f.is_multiple_of(128)
|| x.len() < t * in_f
{
return Ok(None);
}
let dims = [
la.wq.out_features(),
la.wk.out_features(),
la.wv.out_features(),
la.f_a.out_features(),
la.g_a.out_features(),
la.b_proj.out_features(),
];
let (
GpuTensor::Quant { bytes: bq, .. },
GpuTensor::Quant { bytes: bk, .. },
GpuTensor::Quant { bytes: bv, .. },
) = (&la.wq, &la.wk, &la.wv)
else {
unreachable!("q8() above only admits Quant");
};
let (
GpuTensor::Float { data: wfa, .. },
GpuTensor::Float { data: wga, .. },
GpuTensor::Float { data: wb, .. },
) = (&la.f_a, &la.g_a, &la.b_proj)
else {
unreachable!("f32w() above only admits Float");
};
let (aq, ad) = self.quantize_q8_1(x, t, in_f)?;
let mut outs = [
self.uninit(t * dims[0])?,
self.uninit(t * dims[1])?,
self.uninit(t * dims[2])?,
self.uninit(t * dims[3])?,
self.uninit(t * dims[4])?,
self.uninit(t * dims[5])?,
];
self.kda_proj_fused6_raw(
bq, bk, bv, wfa, wga, wb, &aq, &ad, x, &mut outs, in_f, dims, t, rb_q,
)?;
if KDA_FUSED6_DISPATCHES.fetch_add(1, Ordering::Relaxed) == 0 {
eprintln!(
"[kda-fused6] engaged in_f={in_f} out={dims:?} t={t} (one launch replaces the \
six-projection group; MEMRA_KDA_FUSED_PROJ=1)"
);
}
Ok(Some(outs.into_iter().collect()))
}
#[allow(clippy::too_many_arguments)] pub fn kda_proj_fused6_raw(
&self,
wq: &CudaSlice<u8>,
wk: &CudaSlice<u8>,
wv: &CudaSlice<u8>,
wfa: &CudaSlice<f32>,
wga: &CudaSlice<f32>,
wb: &CudaSlice<f32>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
x: &CudaSlice<f32>,
outs: &mut [CudaSlice<f32>; 6],
in_f: usize,
dims: [usize; 6],
t: usize,
row_bytes: usize,
) -> Result<(), Box<dyn std::error::Error>> {
const ROWS_PER_BLOCK: usize = 4; if t == 0
|| !in_f.is_multiple_of(128)
|| x.len() < t * in_f
|| aq.len() < t * in_f
|| ad.len() < t * (in_f / 32)
{
return Err("kda_proj_fused6 geometry".into());
}
for (i, (w, want_rows)) in [(wq, dims[0]), (wk, dims[1]), (wv, dims[2])]
.into_iter()
.enumerate()
{
if w.len() < want_rows * row_bytes {
return Err(format!(
"kda_proj_fused6: q8 weight {i} holds {} bytes, needs {}",
w.len(),
want_rows * row_bytes
)
.into());
}
}
for (i, (w, want_rows)) in [(wfa, dims[3]), (wga, dims[4]), (wb, dims[5])]
.into_iter()
.enumerate()
{
if w.len() < want_rows * in_f {
return Err(format!(
"kda_proj_fused6: f32 weight {} holds {} floats, needs {}",
i + 3,
w.len(),
want_rows * in_f
)
.into());
}
}
for (i, (o, want)) in outs.iter().zip(dims).enumerate() {
if o.len() < t * want {
return Err(format!("kda_proj_fused6: output {i} too small").into());
}
}
let blocks: usize = dims.iter().map(|d| d.div_ceil(ROWS_PER_BLOCK)).sum();
let f = self.func("qmatvec_kda6_q8f32_mmvq");
let cfg = LaunchConfig {
grid_dim: (blocks as u32, t as u32, 1),
block_dim: (32, ROWS_PER_BLOCK as u32, 1),
shared_mem_bytes: 0,
};
let inf = in_f as i32;
let d = dims.map(|v| v as i32);
let (mi, rb) = (t as i32, row_bytes as i64);
let [o0, o1, o2, o3, o4, o5] = outs;
let stream = self.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(wq)
.arg(wk)
.arg(wv)
.arg(wfa)
.arg(wga)
.arg(wb)
.arg(aq)
.arg(ad)
.arg(x)
.arg(&mut *o0)
.arg(&mut *o1)
.arg(&mut *o2)
.arg(&mut *o3)
.arg(&mut *o4)
.arg(&mut *o5)
.arg(&inf)
.arg(&d[0])
.arg(&d[1])
.arg(&d[2])
.arg(&d[3])
.arg(&d[4])
.arg(&d[5])
.arg(&mi)
.arg(&rb);
unsafe { b.launch(cfg)? };
Ok(())
}
#[allow(clippy::too_many_arguments)] pub fn kda_proj_fused6_bf16_raw(
&self,
wq: &CudaSlice<u8>,
wk: &CudaSlice<u8>,
wv: &CudaSlice<u8>,
wfa: &CudaSlice<f32>,
wga: &CudaSlice<f32>,
wb: &CudaSlice<f32>,
x: &CudaSlice<f32>,
outs: &mut [CudaSlice<f32>; 6],
in_f: usize,
dims: [usize; 6],
t: usize,
) -> Result<(), Box<dyn std::error::Error>> {
self.kda_proj_fused6_bf16_arm_raw(
wq,
wk,
wv,
wfa,
wga,
wb,
x,
outs,
in_f,
dims,
t,
crate::b200_gemv_v2_level(),
)
}
#[allow(clippy::too_many_arguments)] pub fn kda_proj_fused6_bf16_arm_raw(
&self,
wq: &CudaSlice<u8>,
wk: &CudaSlice<u8>,
wv: &CudaSlice<u8>,
wfa: &CudaSlice<f32>,
wga: &CudaSlice<f32>,
wb: &CudaSlice<f32>,
x: &CudaSlice<f32>,
outs: &mut [CudaSlice<f32>; 6],
in_f: usize,
dims: [usize; 6],
t: usize,
arm: u8,
) -> Result<(), Box<dyn std::error::Error>> {
if t == 0 || !in_f.is_multiple_of(128) || x.len() < t * in_f {
return Err("kda_proj_fused6_bf16 geometry".into());
}
for (i, (w, want_rows)) in [(wq, dims[0]), (wk, dims[1]), (wv, dims[2])]
.into_iter()
.enumerate()
{
if w.len() < want_rows * in_f * 2 {
return Err(format!(
"kda_proj_fused6_bf16: bf16 weight {i} holds {} bytes, needs {}",
w.len(),
want_rows * in_f * 2
)
.into());
}
}
for (i, (w, want_rows)) in [(wfa, dims[3]), (wga, dims[4]), (wb, dims[5])]
.into_iter()
.enumerate()
{
if w.len() < want_rows * in_f {
return Err(format!(
"kda_proj_fused6_bf16: f32 weight {} holds {} floats, needs {}",
i + 3,
w.len(),
want_rows * in_f
)
.into());
}
}
for (i, (o, want)) in outs.iter().zip(dims).enumerate() {
if o.len() < t * want {
return Err(format!("kda_proj_fused6_bf16: output {i} too small").into());
}
}
let arm = if arm >= 2 && !crate::gemv_v3_fits() {
1
} else {
arm
};
let nb = crate::mmv_block();
let rpb = if arm >= 1 { crate::GEMV_V2_ROWS } else { 4 };
let blocks: usize = dims.iter().map(|d| d.div_ceil(rpb)).sum();
let f = self.func(match arm {
0 => "qmatvec_kda6_bf16f32",
1 => "qmatvec_kda6_bf16f32_v2",
_ => "qmatvec_kda6_bf16f32_v3",
});
let cfg = LaunchConfig {
grid_dim: (blocks as u32, t as u32, 1),
block_dim: (nb, 1, 1),
shared_mem_bytes: match arm {
0 => 0,
1 => (crate::GEMV_V2_ROWS as u32) * nb * 4,
_ => crate::gemv_v3_smem_bytes(nb as usize) as u32,
},
};
let inf = in_f as i32;
let d = dims.map(|v| v as i32);
let mi = t as i32;
let [o0, o1, o2, o3, o4, o5] = outs;
let stream = self.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(wq)
.arg(wk)
.arg(wv)
.arg(wfa)
.arg(wga)
.arg(wb)
.arg(x)
.arg(&mut *o0)
.arg(&mut *o1)
.arg(&mut *o2)
.arg(&mut *o3)
.arg(&mut *o4)
.arg(&mut *o5)
.arg(&inf)
.arg(&d[0])
.arg(&d[1])
.arg(&d[2])
.arg(&d[3])
.arg(&d[4])
.arg(&d[5])
.arg(&mi);
unsafe { b.launch(cfg)? };
Ok(())
}
}