use crate::forward::ShaderModuleTuned as _;
pub fn l2norm(x: &mut [f32], eps: f32) {
let ss: f32 = x.iter().map(|v| v * v).sum();
let inv = 1.0 / (ss + eps).sqrt();
for v in x.iter_mut() {
*v *= inv;
}
}
pub fn softplus(x: f32) -> f32 {
if x > 20.0 {
x
} else if x < -20.0 {
x.exp()
} else {
x.exp().ln_1p()
}
}
pub fn delta_step(
state: &mut [f32],
q: &[f32],
k: &[f32],
v: &[f32],
g: f32,
beta: f32,
) -> Vec<f32> {
let (dk, dv) = (q.len(), v.len());
debug_assert_eq!(state.len(), dk * dv);
let decay = g.exp();
let mut kv_mem = vec![0f32; dv];
for i in 0..dk {
for j in 0..dv {
let s = state[i * dv + j] * decay;
state[i * dv + j] = s;
kv_mem[j] += k[i] * s;
}
}
let delta: Vec<f32> = (0..dv).map(|j| (v[j] - kv_mem[j]) * beta).collect();
let mut out = vec![0f32; dv];
for i in 0..dk {
for j in 0..dv {
let s = state[i * dv + j] + k[i] * delta[j];
state[i * dv + j] = s;
out[j] += q[i] * s;
}
}
out
}
pub fn silu(x: f32) -> f32 {
x / (1.0 + (-x).exp())
}
pub struct DeltaNetRef {
pub nk: usize,
pub nv: usize,
pub dk: usize,
pub dv: usize,
pub kernel: usize,
pub eps: f32,
pub w_qkv: Vec<f32>, pub w_z: Vec<f32>, pub w_b: Vec<f32>, pub w_a: Vec<f32>, pub conv_w: Vec<f32>, pub a_log: Vec<f32>, pub dt_bias: Vec<f32>, pub norm_w: Vec<f32>, pub w_out: Vec<f32>, }
pub struct DeltaNetState {
pub conv: Vec<f32>,
pub s: Vec<f32>,
}
impl DeltaNetRef {
pub fn conv_dim(&self) -> usize {
2 * self.nk * self.dk + self.nv * self.dv
}
pub fn fresh_state(&self) -> DeltaNetState {
DeltaNetState {
conv: vec![0.0; self.conv_dim() * self.kernel],
s: vec![0.0; self.nv * self.dk * self.dv],
}
}
pub fn step(&self, st: &mut DeltaNetState, x: &[f32], hidden: usize) -> Vec<f32> {
let (nv, dv) = (self.nv, self.dv);
let matvec = |w: &[f32], rows: usize| -> Vec<f32> {
(0..rows)
.map(|r| (0..hidden).map(|c| w[r * hidden + c] * x[c]).sum())
.collect()
};
let mixed = matvec(&self.w_qkv, self.conv_dim());
let z = matvec(&self.w_z, nv * dv);
let b = matvec(&self.w_b, nv);
let a = matvec(&self.w_a, nv);
let core = self.core(st, &mixed, &z, &b, &a);
(0..hidden)
.map(|r| {
(0..nv * dv)
.map(|c| self.w_out[r * nv * dv + c] * core[c])
.sum()
})
.collect()
}
pub fn core(
&self,
st: &mut DeltaNetState,
mixed: &[f32],
z: &[f32],
b: &[f32],
a: &[f32],
) -> Vec<f32> {
let (nk, nv, dk, dv, kn) = (self.nk, self.nv, self.dk, self.dv, self.kernel);
let conv_dim = self.conv_dim();
let mut conved = vec![0f32; conv_dim];
for c in 0..conv_dim {
let ring = &mut st.conv[c * kn..(c + 1) * kn];
ring.copy_within(1.., 0);
ring[kn - 1] = mixed[c];
let acc: f32 = ring
.iter()
.zip(&self.conv_w[c * kn..(c + 1) * kn])
.map(|(r, w)| r * w)
.sum();
conved[c] = silu(acc);
}
let (qs, rest) = conved.split_at(nk * dk);
let (ks, vs) = rest.split_at(nk * dk);
let rep = nv / nk;
let scale = 1.0 / (dk as f32).sqrt();
let mut core = vec![0f32; nv * dv];
for h in 0..nv {
let kh = h / rep; let mut q = qs[kh * dk..(kh + 1) * dk].to_vec();
let mut k = ks[kh * dk..(kh + 1) * dk].to_vec();
l2norm(&mut q, 1e-6);
l2norm(&mut k, 1e-6);
for qv in q.iter_mut() {
*qv *= scale;
}
let beta = 1.0 / (1.0 + (-b[h]).exp());
let g = -self.a_log[h].exp() * softplus(a[h] + self.dt_bias[h]);
let out = delta_step(
&mut st.s[h * dk * dv..(h + 1) * dk * dv],
&q,
&k,
&vs[h * dv..(h + 1) * dv],
g,
beta,
);
let ms = out.iter().map(|v| v * v).sum::<f32>() / dv as f32;
let inv = 1.0 / (ms + self.eps).sqrt();
for j in 0..dv {
core[h * dv + j] = out[j] * inv * self.norm_w[j] * silu(z[h * dv + j]);
}
}
core
}
}
impl DeltaNetRef {
pub fn core_chunk(
&self,
st: &mut DeltaNetState,
mixed_c: &[Vec<f32>],
z_c: &[Vec<f32>],
b_c: &[Vec<f32>],
a_c: &[Vec<f32>],
) -> Vec<Vec<f32>> {
let (nk, nv, dk, dv, kn) = (self.nk, self.nv, self.dk, self.dv, self.kernel);
let conv_dim = self.conv_dim();
let cc = mixed_c.len();
assert!(cc >= 1 && z_c.len() == cc && b_c.len() == cc && a_c.len() == cc);
let conved_c: Vec<Vec<f32>> = mixed_c
.iter()
.map(|mixed| {
let mut conved = vec![0f32; conv_dim];
for c in 0..conv_dim {
let ring = &mut st.conv[c * kn..(c + 1) * kn];
ring.copy_within(1.., 0);
ring[kn - 1] = mixed[c];
let acc: f32 = ring
.iter()
.zip(&self.conv_w[c * kn..(c + 1) * kn])
.map(|(r, w)| r * w)
.sum();
conved[c] = silu(acc);
}
conved
})
.collect();
let rep = nv / nk;
let scale = 1.0 / (dk as f32).sqrt();
let mut core_c = vec![vec![0f32; nv * dv]; cc];
for h in 0..nv {
let kh = h / rep;
let mut qs = Vec::with_capacity(cc);
let mut ks = Vec::with_capacity(cc);
let mut vs = Vec::with_capacity(cc);
let mut betas = Vec::with_capacity(cc);
let mut lg = Vec::with_capacity(cc); let mut lg_run = 0f32;
for r in 0..cc {
let conved = &conved_c[r];
let (qsl, rest) = conved.split_at(nk * dk);
let (ksl, vsl) = rest.split_at(nk * dk);
let mut q = qsl[kh * dk..(kh + 1) * dk].to_vec();
let mut k = ksl[kh * dk..(kh + 1) * dk].to_vec();
l2norm(&mut q, 1e-6);
l2norm(&mut k, 1e-6);
for qv in q.iter_mut() {
*qv *= scale;
}
qs.push(q);
ks.push(k);
vs.push(vsl[h * dv..(h + 1) * dv].to_vec());
betas.push(1.0 / (1.0 + (-b_c[r][h]).exp()));
let g = -self.a_log[h].exp() * softplus(a_c[r][h] + self.dt_bias[h]);
lg_run += g;
lg.push(lg_run);
}
let s0 = st.s[h * dk * dv..(h + 1) * dk * dv].to_vec();
let s0t = |x: &[f32]| -> Vec<f32> {
let mut out = vec![0f32; dv];
for i in 0..dk {
let xi = x[i];
for j in 0..dv {
out[j] += xi * s0[i * dv + j];
}
}
out
};
let mut us: Vec<Vec<f32>> = Vec::with_capacity(cc);
for r in 0..cc {
let gr = lg[r].exp();
let s0k = s0t(&ks[r]);
let mut u = vec![0f32; dv];
for j in 0..dv {
u[j] = vs[r][j] - gr * s0k[j];
}
for i in 0..r {
let kik: f32 = ks[i].iter().zip(&ks[r]).map(|(a, b)| a * b).sum();
let ratio = (lg[r] - lg[i]).exp();
let w = ratio * kik;
for j in 0..dv {
u[j] -= w * us[i][j];
}
}
for uj in u.iter_mut() {
*uj *= betas[r];
}
us.push(u);
}
for r in 0..cc {
let gr = lg[r].exp();
let s0q = s0t(&qs[r]);
let mut o: Vec<f32> = (0..dv).map(|j| gr * s0q[j]).collect();
for i in 0..=r {
let kiq: f32 = ks[i].iter().zip(&qs[r]).map(|(a, b)| a * b).sum();
let w = (lg[r] - lg[i]).exp() * kiq;
for j in 0..dv {
o[j] += w * us[i][j];
}
}
let ms = o.iter().map(|v| v * v).sum::<f32>() / dv as f32;
let inv = 1.0 / (ms + self.eps).sqrt();
for j in 0..dv {
core_c[r][h * dv + j] = o[j] * inv * self.norm_w[j] * silu(z_c[r][h * dv + j]);
}
}
let g_c = lg[cc - 1].exp();
let sh = &mut st.s[h * dk * dv..(h + 1) * dk * dv];
for (i, s) in sh.iter_mut().enumerate() {
*s = s0[i] * g_c;
}
for i in 0..cc {
let ratio = (lg[cc - 1] - lg[i]).exp();
for d in 0..dk {
let kd = ks[i][d] * ratio;
for j in 0..dv {
sh[d * dv + j] += kd * us[i][j];
}
}
}
}
core_c
}
}
const DN_CONV: &str = r#"
@group(0) @binding(0) var<storage, read> mixed: array<f32>; // [conv_dim] at dims.z
@group(0) @binding(1) var<storage, read> cw: array<f32>; // [conv_dim, K] taps oldest-first
@group(0) @binding(2) var<storage, read_write> ring: array<f32>; // [conv_dim, K] newest-last
@group(0) @binding(3) var<storage, read_write> conved: array<f32>; // [conv_dim] at dims.w
@group(0) @binding(4) var<uniform> dims: vec4<u32>;
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
let c = gid.x;
if (c >= dims.x) { return; }
let kk = dims.y;
let base = c * kk;
var acc = 0.0;
for (var t = 0u; t + 1u < kk; t = t + 1u) {
let nxt = ring[base + t + 1u];
ring[base + t] = nxt;
acc = acc + nxt * cw[base + t];
}
let xn = mixed[dims.z + c];
ring[base + kk - 1u] = xn;
acc = acc + xn * cw[base + kk - 1u];
conved[dims.w + c] = acc / (1.0 + exp(-acc));
}
"#;
const DN_CHUNK: &str = r#"
@group(0) @binding(0) var<storage, read> conved: array<f32>; // [C, conv_dim]
@group(0) @binding(1) var<storage, read> zin: array<f32>; // [C, nv·dv]
@group(0) @binding(2) var<storage, read> ab: array<f32>; // [C, 2·nv] (b | a per pos)
@group(0) @binding(3) var<storage, read> gp: array<f32>; // [A_log(nv) | dt_bias(nv) | norm_w(dv)]
@group(0) @binding(4) var<storage, read_write> st: array<f32>; // [nv, dk, dv]
@group(0) @binding(5) var<storage, read_write> core_c: array<f32>; // [C, nv·dv]
@group(0) @binding(6) var<uniform> dims: vec4<u32>; // (nv, nk, dk, dv)
@group(0) @binding(7) var<uniform> dims2: vec4<u32>; // (C, conv_dim, _, _)
@group(0) @binding(8) var<uniform> epsm: vec4<f32>;
const CMAX = 16u;
var<workgroup> ksh: array<f32, 2048>; // [C, dk] normalized k
var<workgroup> qsh: array<f32, 2048>; // [C, dk] normalized+scaled q
var<workgroup> ush: array<f32, 2048>; // [C, dv] pseudo-values
var<workgroup> amat: array<f32, 256>; // [C, C] decay-ratio'd KKᵀ (strict lower)
var<workgroup> qkm: array<f32, 256>; // [C, C] decay-ratio'd QKᵀ (lower incl. diag)
var<workgroup> osh: array<f32, 128>; // one position's o (norm phase scratch)
var<workgroup> lgs: array<f32, 16>; // cumulative log-decay per position
var<workgroup> bets: array<f32, 16>; // β per position
var<workgroup> oinv: f32;
fn softplus(x: f32) -> f32 {
if (x > 20.0) { return x; }
if (x < -20.0) { return exp(x); }
return log(1.0 + exp(x));
}
@compute @workgroup_size(128)
fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
let h = wid.x;
let nv = dims.x; let nk = dims.y; let dk = dims.z; let dv = dims.w;
let cc = dims2.x; let cdim = dims2.y;
let kh = h / (nv / nk);
let j = lid.x;
// P1a: stage raw k/q for every position (coalesced lanes over dk).
for (var r = 0u; r < cc; r = r + 1u) {
for (var i = j; i < dk; i = i + 128u) {
qsh[r * dk + i] = conved[r * cdim + kh * dk + i];
ksh[r * dk + i] = conved[r * cdim + nk * dk + kh * dk + i];
}
}
workgroupBarrier();
// P1b: thread 0 — per-position l2 norms (sequential sums, the CPU reference's order),
// β, and the cumulative log-decay.
if (j == 0u) {
let scale = 1.0 / sqrt(f32(dk));
var lg_run = 0.0;
for (var r = 0u; r < cc; r = r + 1u) {
var sq = 0.0;
var sk = 0.0;
for (var i = 0u; i < dk; i = i + 1u) { let v = qsh[r * dk + i]; sq = sq + v * v; }
for (var i = 0u; i < dk; i = i + 1u) { let v = ksh[r * dk + i]; sk = sk + v * v; }
let invq = 1.0 / sqrt(sq + 1e-6);
let invk = 1.0 / sqrt(sk + 1e-6);
for (var i = 0u; i < dk; i = i + 1u) {
qsh[r * dk + i] = qsh[r * dk + i] * invq * scale;
ksh[r * dk + i] = ksh[r * dk + i] * invk;
}
bets[r] = 1.0 / (1.0 + exp(-ab[r * 2u * nv + h]));
let g = -exp(gp[h]) * softplus(ab[r * 2u * nv + nv + h] + gp[nv + h]);
lg_run = lg_run + g;
lgs[r] = lg_run;
}
}
workgroupBarrier();
// P2: pair matrices — A (i<r) and QK (i≤r), one (r,i) pair per lane sweep.
for (var p = j; p < cc * cc; p = p + 128u) {
let r = p / cc;
let i = p % cc;
if (i <= r) {
var kk_d = 0.0;
var qk_d = 0.0;
for (var d = 0u; d < dk; d = d + 1u) {
kk_d = kk_d + ksh[i * dk + d] * ksh[r * dk + d];
qk_d = qk_d + ksh[i * dk + d] * qsh[r * dk + d];
}
let ratio = exp(lgs[r] - lgs[i]);
if (i < r) { amat[r * CMAX + i] = ratio * kk_d; }
qkm[r * CMAX + i] = ratio * qk_d;
}
}
workgroupBarrier();
// P3: one state sweep — per lane j (a dv column), C accumulators for k and q.
var acck: array<f32, CMAX>;
var accq: array<f32, CMAX>;
for (var r = 0u; r < cc; r = r + 1u) { acck[r] = 0.0; accq[r] = 0.0; }
if (j < dv) {
let sbase = h * dk * dv + j;
for (var i = 0u; i < dk; i = i + 1u) {
let s = st[sbase + i * dv];
for (var r = 0u; r < cc; r = r + 1u) {
acck[r] = acck[r] + ksh[r * dk + i] * s;
accq[r] = accq[r] + qsh[r * dk + i] * s;
}
}
}
// P4: pseudo-value solve, sequential in r, parallel in j.
for (var r = 0u; r < cc; r = r + 1u) {
if (j < dv) {
let gr = exp(lgs[r]);
var u = conved[r * cdim + 2u * nk * dk + h * dv + j] - gr * acck[r];
for (var i = 0u; i < r; i = i + 1u) {
u = u - amat[r * CMAX + i] * ush[i * dv + j];
}
ush[r * dv + j] = u * bets[r];
}
workgroupBarrier();
}
// P5: outputs + gated RMSNorm per position (thread-0 reduction, DN_STEP's structure).
for (var r = 0u; r < cc; r = r + 1u) {
if (j < dv) {
var o = exp(lgs[r]) * accq[r];
for (var i = 0u; i <= r; i = i + 1u) {
o = o + qkm[r * CMAX + i] * ush[i * dv + j];
}
osh[j] = o;
}
workgroupBarrier();
if (j == 0u) {
var ms = 0.0;
for (var jj = 0u; jj < dv; jj = jj + 1u) { ms = ms + osh[jj] * osh[jj]; }
oinv = 1.0 / sqrt(ms / f32(dv) + epsm.x);
}
workgroupBarrier();
if (j < dv) {
let zz = zin[r * nv * dv + h * dv + j];
core_c[r * nv * dv + h * dv + j] =
osh[j] * oinv * gp[2u * nv + j] * (zz / (1.0 + exp(-zz)));
}
workgroupBarrier();
}
// P6: boundary state update — S_C = Γ_C·S₀ + Σ_r (Γ_C/Γ_r)·k_r·u_rᵀ, one write per element.
if (j < dv) {
let sbase = h * dk * dv + j;
let lgc = lgs[cc - 1u];
let gc = exp(lgc);
for (var i = 0u; i < dk; i = i + 1u) {
var s = st[sbase + i * dv] * gc;
for (var r = 0u; r < cc; r = r + 1u) {
s = s + exp(lgc - lgs[r]) * ksh[r * dk + i] * ush[r * dv + j];
}
st[sbase + i * dv] = s;
}
}
}
"#;
const DN_STEP: &str = r#"
@group(0) @binding(0) var<storage, read> conved: array<f32>; // [q(nk·dk) | k(nk·dk) | v(nv·dv)]
@group(0) @binding(1) var<storage, read> zin: array<f32>; // [nv·dv]
@group(0) @binding(2) var<storage, read> ab: array<f32>; // [b(nv) | a(nv)]
@group(0) @binding(3) var<storage, read> gp: array<f32>; // [A_log(nv) | dt_bias(nv) | norm_w(dv)]
@group(0) @binding(4) var<storage, read_write> st: array<f32>; // [nv, dk, dv]
@group(0) @binding(5) var<storage, read_write> core: array<f32>; // [nv·dv]
@group(0) @binding(6) var<uniform> dims: vec4<u32>;
@group(0) @binding(7) var<uniform> epsm: vec4<f32>;
var<workgroup> qsh: array<f32, 256>;
var<workgroup> ksh: array<f32, 256>;
var<workgroup> osh: array<f32, 128>;
var<workgroup> scal: array<f32, 3>; // decay, beta, gated-norm inv
fn softplus(x: f32) -> f32 {
if (x > 20.0) { return x; }
if (x < -20.0) { return exp(x); }
return log(1.0 + exp(x));
}
@compute @workgroup_size(128)
fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
let h = wid.x;
let nv = dims.x; let nk = dims.y; let dk = dims.z; let dv = dims.w;
let kh = h / (nv / nk);
let j = lid.x;
for (var i = j; i < dk; i = i + 128u) {
qsh[i] = conved[kh * dk + i];
ksh[i] = conved[nk * dk + kh * dk + i];
}
workgroupBarrier();
if (j == 0u) {
var sq = 0.0;
var sk = 0.0;
for (var i = 0u; i < dk; i = i + 1u) { sq = sq + qsh[i] * qsh[i]; }
for (var i = 0u; i < dk; i = i + 1u) { sk = sk + ksh[i] * ksh[i]; }
let invq = 1.0 / sqrt(sq + 1e-6);
let invk = 1.0 / sqrt(sk + 1e-6);
let scale = 1.0 / sqrt(f32(dk));
for (var i = 0u; i < dk; i = i + 1u) {
qsh[i] = qsh[i] * invq * scale; // two mults, matching the CPU reference's order
ksh[i] = ksh[i] * invk;
}
let g = -exp(gp[h]) * softplus(ab[nv + h] + gp[nv + h]);
scal[0] = exp(g);
scal[1] = 1.0 / (1.0 + exp(-ab[h]));
}
workgroupBarrier();
if (j < dv) {
let sbase = h * dk * dv + j;
let decay = scal[0];
// Pass 1: kv_mem from the DECAYED state (read-only — pass 2 recomputes the identical
// product, so nothing needs to be stored between passes).
var kv = 0.0;
for (var i = 0u; i < dk; i = i + 1u) {
kv = kv + ksh[i] * (st[sbase + i * dv] * decay);
}
let delta = (conved[2u * nk * dk + h * dv + j] - kv) * scal[1];
var o = 0.0;
for (var i = 0u; i < dk; i = i + 1u) {
let s = st[sbase + i * dv] * decay + ksh[i] * delta;
st[sbase + i * dv] = s;
o = o + qsh[i] * s;
}
osh[j] = o;
}
workgroupBarrier();
if (j == 0u) {
var ms = 0.0;
for (var jj = 0u; jj < dv; jj = jj + 1u) { ms = ms + osh[jj] * osh[jj]; }
scal[2] = 1.0 / sqrt(ms / f32(dv) + epsm.x);
}
workgroupBarrier();
if (j < dv) {
let zz = zin[h * dv + j];
core[h * dv + j] = osh[j] * scal[2] * gp[2u * nv + j] * (zz / (1.0 + exp(-zz)));
}
}
"#;
pub struct DeltaNetGpu {
conv_pl: wgpu::ComputePipeline,
step_pl: wgpu::ComputePipeline,
conv_bg: wgpu::BindGroup,
step_bg: wgpu::BindGroup,
mixed: wgpu::Buffer,
zin: wgpu::Buffer,
ab: wgpu::Buffer,
ring: wgpu::Buffer,
st: wgpu::Buffer,
core: wgpu::Buffer,
chunk_pl: wgpu::ComputePipeline,
chunk_bg: wgpu::BindGroup,
conv_c_bgs: Vec<wgpu::BindGroup>,
mixed_c: wgpu::Buffer,
zin_c: wgpu::Buffer,
ab_c: wgpu::Buffer,
core_c: wgpu::Buffer,
dims2: wgpu::Buffer,
nv: u32,
nk: u32,
dk: u32,
dv: u32,
kernel: u32,
}
pub const DN_CHUNK_CMAX: usize = 16;
impl DeltaNetGpu {
#[allow(clippy::too_many_arguments)]
pub fn new(
ctx: &crate::GpuCtx,
nk: usize,
nv: usize,
dk: usize,
dv: usize,
kernel: usize,
eps: f32,
conv_w: &[f32],
a_log: &[f32],
dt_bias: &[f32],
norm_w: &[f32],
) -> anyhow::Result<Self> {
anyhow::ensure!(
dk <= 256 && dv <= 128,
"DN_STEP geometry: dk ≤ 256, dv ≤ 128"
);
anyhow::ensure!(
nv.is_multiple_of(nk),
"nv must be a multiple of nk (repeat_interleave)"
);
let conv_dim = 2 * nk * dk + nv * dv;
anyhow::ensure!(conv_w.len() == conv_dim * kernel, "conv_w shape");
anyhow::ensure!(a_log.len() == nv && dt_bias.len() == nv && norm_w.len() == dv);
let module = |label: &str, src: &str| {
ctx.device
.shader_module_tuned(wgpu::ShaderModuleDescriptor {
label: Some(label),
source: wgpu::ShaderSource::Wgsl(src.into()),
})
};
let pipeline = |label: &str, m: &wgpu::ShaderModule| {
ctx.device
.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
label: Some(label),
layout: None,
module: m,
entry_point: Some("main"),
compilation_options: wgpu::PipelineCompilationOptions::default(),
cache: None,
})
};
let conv_pl = pipeline("dn_conv", &module("dn_conv", DN_CONV));
let step_pl = pipeline("dn_step", &module("dn_step", DN_STEP));
let mixed = ctx.empty(conv_dim);
let zin = ctx.empty(nv * dv);
let ab = ctx.empty(2 * nv);
let cw = ctx.storage(conv_w);
let ring = ctx.storage(&vec![0f32; conv_dim * kernel]);
let conved = ctx.empty(conv_dim);
let gp: Vec<f32> = a_log.iter().chain(dt_bias).chain(norm_w).copied().collect();
let gp = ctx.storage(&gp);
let st = ctx.storage(&vec![0f32; nv * dk * dv]);
let core = ctx.empty(nv * dv);
let uni = |vals: [u32; 4]| {
use wgpu::util::DeviceExt;
ctx.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("dims"),
contents: bytemuck::cast_slice(&vals),
usage: wgpu::BufferUsages::UNIFORM,
})
};
let conv_dims = uni([conv_dim as u32, kernel as u32, 0, 0]);
let step_dims = uni([nv as u32, nk as u32, dk as u32, dv as u32]);
let epsm = {
use wgpu::util::DeviceExt;
ctx.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("epsm"),
contents: bytemuck::cast_slice(&[eps, 0.0, 0.0, 0.0]),
usage: wgpu::BufferUsages::UNIFORM,
})
};
let bind = |pl: &wgpu::ComputePipeline, entries: &[&wgpu::Buffer]| {
let e: Vec<wgpu::BindGroupEntry> = entries
.iter()
.enumerate()
.map(|(i, b)| wgpu::BindGroupEntry {
binding: i as u32,
resource: b.as_entire_binding(),
})
.collect();
ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pl.get_bind_group_layout(0),
entries: &e,
})
};
let conv_bg = bind(&conv_pl, &[&mixed, &cw, &ring, &conved, &conv_dims]);
let step_bg = bind(
&step_pl,
&[&conved, &zin, &ab, &gp, &st, &core, &step_dims, &epsm],
);
anyhow::ensure!(
dk <= 128,
"DN_CHUNK geometry: dk ≤ 128 (shared k/q staging)"
);
let cmax = DN_CHUNK_CMAX;
let chunk_pl = pipeline("dn_chunk", &module("dn_chunk", DN_CHUNK));
let mixed_c = ctx.empty(cmax * conv_dim);
let conved_c = ctx.empty(cmax * conv_dim);
let zin_c = ctx.empty(cmax * nv * dv);
let ab_c = ctx.empty(cmax * 2 * nv);
let core_c = ctx.empty(cmax * nv * dv);
let dims2 = {
use wgpu::util::DeviceExt;
ctx.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("dn_chunk_dims2"),
contents: bytemuck::cast_slice(&[0u32, conv_dim as u32, 0, 0]),
usage: wgpu::BufferUsages::UNIFORM | wgpu::BufferUsages::COPY_DST,
})
};
let conv_c_bgs: Vec<wgpu::BindGroup> = (0..cmax)
.map(|r| {
let off = (r * conv_dim) as u32;
let d = uni([conv_dim as u32, kernel as u32, off, off]);
bind(&conv_pl, &[&mixed_c, &cw, &ring, &conved_c, &d])
})
.collect();
let chunk_bg = bind(
&chunk_pl,
&[
&conved_c, &zin_c, &ab_c, &gp, &st, &core_c, &step_dims, &dims2, &epsm,
],
);
Ok(Self {
conv_pl,
step_pl,
conv_bg,
step_bg,
mixed,
zin,
ab,
ring,
st,
core,
chunk_pl,
chunk_bg,
conv_c_bgs,
mixed_c,
zin_c,
ab_c,
core_c,
dims2,
nv: nv as u32,
nk: nk as u32,
dk: dk as u32,
dv: dv as u32,
kernel: kernel as u32,
})
}
pub fn chunk(
&self,
ctx: &crate::GpuCtx,
mixed_c: &[Vec<f32>],
z_c: &[Vec<f32>],
b_c: &[Vec<f32>],
a_c: &[Vec<f32>],
) -> anyhow::Result<Vec<Vec<f32>>> {
let (nv, dv) = (self.nv as usize, self.dv as usize);
let conv_dim = (2 * self.nk * self.dk) as usize + nv * dv;
let cc = mixed_c.len();
anyhow::ensure!((1..=DN_CHUNK_CMAX).contains(&cc), "chunk size 1..=CMAX");
anyhow::ensure!(z_c.len() == cc && b_c.len() == cc && a_c.len() == cc);
let mut mx = Vec::with_capacity(cc * conv_dim);
let mut zz = Vec::with_capacity(cc * nv * dv);
let mut ba = Vec::with_capacity(cc * 2 * nv);
for r in 0..cc {
anyhow::ensure!(mixed_c[r].len() == conv_dim, "mixed len");
anyhow::ensure!(z_c[r].len() == nv * dv && b_c[r].len() == nv && a_c[r].len() == nv);
mx.extend_from_slice(&mixed_c[r]);
zz.extend_from_slice(&z_c[r]);
ba.extend_from_slice(&b_c[r]);
ba.extend_from_slice(&a_c[r]);
}
ctx.queue
.write_buffer(&self.mixed_c, 0, bytemuck::cast_slice(&mx));
ctx.queue
.write_buffer(&self.zin_c, 0, bytemuck::cast_slice(&zz));
ctx.queue
.write_buffer(&self.ab_c, 0, bytemuck::cast_slice(&ba));
ctx.queue.write_buffer(
&self.dims2,
0,
bytemuck::cast_slice(&[cc as u32, conv_dim as u32, 0, 0]),
);
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
{
let mut pass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
for bg in self.conv_c_bgs.iter().take(cc) {
pass.set_pipeline(&self.conv_pl);
pass.set_bind_group(0, bg, &[]);
pass.dispatch_workgroups((conv_dim as u32).div_ceil(64), 1, 1);
}
pass.set_pipeline(&self.chunk_pl);
pass.set_bind_group(0, &self.chunk_bg, &[]);
pass.dispatch_workgroups(self.nv, 1, 1);
}
ctx.queue.submit([enc.finish()]);
let flat = ctx.read(&self.core_c, cc * nv * dv)?;
Ok(flat.chunks(nv * dv).map(|c| c.to_vec()).collect())
}
pub fn reset(&self, ctx: &crate::GpuCtx) {
let conv_dim = (2 * self.nk * self.dk + self.nv * self.dv) as usize;
let zring = vec![0f32; conv_dim * self.kernel as usize];
let zst = vec![0f32; (self.nv * self.dk * self.dv) as usize];
ctx.queue
.write_buffer(&self.ring, 0, bytemuck::cast_slice(&zring));
ctx.queue
.write_buffer(&self.st, 0, bytemuck::cast_slice(&zst));
}
pub fn step(
&self,
ctx: &crate::GpuCtx,
mixed: &[f32],
z: &[f32],
b: &[f32],
a: &[f32],
) -> anyhow::Result<Vec<f32>> {
let conv_dim = (2 * self.nk * self.dk + self.nv * self.dv) as usize;
anyhow::ensure!(mixed.len() == conv_dim, "mixed len");
anyhow::ensure!(z.len() == (self.nv * self.dv) as usize, "z len");
anyhow::ensure!(b.len() == self.nv as usize && a.len() == self.nv as usize);
ctx.queue
.write_buffer(&self.mixed, 0, bytemuck::cast_slice(mixed));
ctx.queue
.write_buffer(&self.zin, 0, bytemuck::cast_slice(z));
let ba: Vec<f32> = b.iter().chain(a).copied().collect();
ctx.queue
.write_buffer(&self.ab, 0, bytemuck::cast_slice(&ba));
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
{
let mut pass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
pass.set_pipeline(&self.conv_pl);
pass.set_bind_group(0, &self.conv_bg, &[]);
pass.dispatch_workgroups((conv_dim as u32).div_ceil(64), 1, 1);
pass.set_pipeline(&self.step_pl);
pass.set_bind_group(0, &self.step_bg, &[]);
pass.dispatch_workgroups(self.nv, 1, 1);
}
ctx.queue.submit([enc.finish()]);
ctx.read(&self.core, (self.nv * self.dv) as usize)
}
pub fn read_state(&self, ctx: &crate::GpuCtx) -> anyhow::Result<Vec<f32>> {
ctx.read(&self.st, (self.nv * self.dk * self.dv) as usize)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn delta_step_matches_hand_computed_case() {
let mut s = vec![1.0, 0.0, 0.0, 2.0];
let q = [0.6f32, 0.8]; let k = [1.0f32, 0.0];
let v = [3.0f32, 4.0];
let g = 0.0f32; let beta = 0.5f32;
let o = delta_step(&mut s, &q, &k, &v, g, beta);
assert!(
(o[0] - 1.2).abs() < 1e-6 && (o[1] - 2.8).abs() < 1e-6,
"{o:?}"
);
assert_eq!(s, vec![2.0, 2.0, 0.0, 2.0]);
}
#[test]
fn decay_halves_state_before_everything_else() {
let mut s = vec![2.0f32, 0.0, 0.0, 0.0];
let o = delta_step(
&mut s,
&[1.0, 0.0],
&[1.0, 0.0],
&[1.0, 0.0],
0.5f32.ln(),
1.0,
);
assert!((s[0] - 1.0).abs() < 1e-6 && (o[0] - 1.0).abs() < 1e-6);
}
#[test]
fn beta_one_reproduces_pure_delta_rule_matrix_form() {
let (dk, dv) = (4usize, 3usize);
let mut s: Vec<f32> = (0..dk * dv)
.map(|i| ((i * 37 % 11) as f32 - 5.0) * 0.1)
.collect();
let s0 = s.clone();
let mut k = vec![0.3f32, -0.5, 0.7, 0.2];
l2norm(&mut k, 1e-6);
let v = [0.9f32, -0.2, 0.4];
let q = vec![0.25f32; dk];
let _ = delta_step(&mut s, &q, &k, &v, 0.0, 1.0);
for i in 0..dk {
for j in 0..dv {
let kts: f32 = (0..dk).map(|t| k[t] * s0[t * dv + j]).sum();
let expect = s0[i * dv + j] - k[i] * kts + k[i] * v[j];
assert!(
(s[i * dv + j] - expect).abs() < 1e-5,
"S[{i},{j}] = {} vs {expect}",
s[i * dv + j]
);
}
}
}
#[test]
fn state_norm_is_contractive_under_decay_and_bounded_beta() {
let (dk, dv) = (8usize, 8);
let mut s = vec![0f32; dk * dv];
let mut seed = 42u64;
let mut rng = move || {
seed ^= seed << 13;
seed ^= seed >> 7;
seed ^= seed << 17;
((seed >> 40) as u32 as f32 / (1u32 << 24) as f32) - 0.5
};
let mut peak = 0f32;
for _ in 0..200 {
let mut q: Vec<f32> = (0..dk).map(|_| rng()).collect();
let mut k: Vec<f32> = (0..dk).map(|_| rng()).collect();
let v: Vec<f32> = (0..dv).map(|_| rng()).collect();
l2norm(&mut q, 1e-6);
l2norm(&mut k, 1e-6);
let g = -softplus(rng() * 4.0); let beta = 1.0 / (1.0 + (-rng() * 4.0).exp());
let _ = delta_step(&mut s, &q, &k, &v, g, beta);
let n: f32 = s.iter().map(|x| x * x).sum::<f32>().sqrt();
peak = peak.max(n);
assert!(n.is_finite());
}
assert!(peak < 20.0, "state norm diverged: {peak}");
}
#[test]
fn full_layer_step_carries_conv_ring_and_recurrent_state() {
let (hidden, nk, nv, dk, dv, kn) = (16usize, 2usize, 4usize, 4usize, 4usize, 4usize);
let mut seed = 7u64;
let mut rng = move || {
seed ^= seed << 13;
seed ^= seed >> 7;
seed ^= seed << 17;
((seed >> 40) as u32 as f32 / (1u32 << 24) as f32) - 0.5
};
let conv_dim = 2 * nk * dk + nv * dv;
let r = DeltaNetRef {
nk,
nv,
dk,
dv,
kernel: kn,
eps: 1e-6,
w_qkv: (0..conv_dim * hidden).map(|_| rng() * 0.3).collect(),
w_z: (0..nv * dv * hidden).map(|_| rng() * 0.3).collect(),
w_b: (0..nv * hidden).map(|_| rng() * 0.3).collect(),
w_a: (0..nv * hidden).map(|_| rng() * 0.3).collect(),
conv_w: (0..conv_dim * kn).map(|_| rng() * 0.4).collect(),
a_log: (0..nv).map(|_| rng().abs() + 0.1).collect(),
dt_bias: vec![1.0; nv],
norm_w: (0..dv).map(|_| 1.0 + rng() * 0.1).collect(),
w_out: (0..hidden * nv * dv).map(|_| rng() * 0.3).collect(),
};
let x: Vec<f32> = (0..hidden).map(|_| rng()).collect();
let mut st = r.fresh_state();
let y1 = r.step(&mut st, &x, hidden);
let y2 = r.step(&mut st, &x, hidden);
assert!(
y1.iter().zip(&y2).any(|(a, b)| (a - b).abs() > 1e-6),
"state must advance"
);
let mut st2 = r.fresh_state();
let y1b = r.step(&mut st2, &x, hidden);
for (a, b) in y1.iter().zip(&y1b) {
assert!((a - b).abs() < 1e-7, "fresh-state replay must be exact");
}
}
}