use ferrox_core::matmul::{rms_norm, silu};
use ferrox_core::weight_matrix::WeightMatrix;
#[derive(Debug, Clone, Copy)]
pub struct GdnConfig {
pub hidden_dim: usize,
pub num_v_heads: usize,
pub head_dim: usize,
pub conv_kernel_size: usize,
pub rms_norm_eps: f32,
}
impl GdnConfig {
pub fn qkv_dim(&self) -> usize {
3 * self.num_v_heads * self.head_dim
}
pub fn v_dim(&self) -> usize {
self.num_v_heads * self.head_dim
}
}
pub struct GdnWeights {
pub attn_qkv: WeightMatrix, pub attn_gate: WeightMatrix, pub ssm_conv1d: Vec<f32>,
pub ssm_dt: Vec<f32>, pub ssm_a: Vec<f32>, pub ssm_beta: WeightMatrix, pub ssm_alpha: WeightMatrix, pub ssm_norm: Vec<f32>, pub ssm_out: WeightMatrix, }
pub struct GdnState {
conv_hist: Vec<f32>,
recurrent: Vec<f32>,
}
impl GdnState {
pub fn new(cfg: &GdnConfig) -> Self {
Self {
conv_hist: Vec::new(),
recurrent: vec![0f32; cfg.num_v_heads * cfg.head_dim * cfg.head_dim],
}
}
}
fn softplus(x: f32) -> f32 {
if x > 20.0 {
x
} else {
(1.0 + x.exp()).ln()
}
}
fn sigmoid(x: f32) -> f32 {
1.0 / (1.0 + (-x).exp())
}
fn l2_normalize(v: &mut [f32], eps: f32) {
let norm_sq: f32 = v.iter().map(|x| x * x).sum();
let scale = 1.0 / (norm_sq + eps).sqrt();
for x in v.iter_mut() {
*x *= scale;
}
}
fn causal_conv_step(
weight: &[f32],
history: &mut Vec<f32>,
current: &[f32],
kernel_size: usize,
dim: usize,
) -> Vec<f32> {
let hist_len = history.len() / dim.max(1);
let missing = (kernel_size - 1).saturating_sub(hist_len);
let mut y = vec![0f32; dim];
for j in 0..kernel_size {
if j < missing {
continue;
}
let src: &[f32] = if j == kernel_size - 1 {
current
} else {
let hist_idx = j - missing;
&history[hist_idx * dim..(hist_idx + 1) * dim]
};
for d in 0..dim {
y[d] += weight[d * kernel_size + j] * src[d];
}
}
for v in y.iter_mut() {
*v = silu(*v);
}
history.extend_from_slice(current);
let max_hist_len = (kernel_size - 1) * dim;
if history.len() > max_hist_len {
let excess = history.len() - max_hist_len;
history.drain(0..excess);
}
y
}
pub fn gdn_forward_token(
weights: &GdnWeights,
cfg: &GdnConfig,
hidden: &[f32],
state: &mut GdnState,
) -> Vec<f32> {
assert_eq!(hidden.len(), cfg.hidden_dim);
let qkv_dim = cfg.qkv_dim();
let v_dim = cfg.v_dim();
let head_dim = cfg.head_dim;
let n_heads = cfg.num_v_heads;
let qkv_lin = weights.attn_qkv.apply(hidden);
let z = weights.attn_gate.apply(hidden);
let beta_raw = weights.ssm_beta.apply(hidden);
let alpha_raw = weights.ssm_alpha.apply(hidden);
let qkv = causal_conv_step(
&weights.ssm_conv1d,
&mut state.conv_hist,
&qkv_lin,
cfg.conv_kernel_size,
qkv_dim,
);
let qk_dim = n_heads * head_dim;
let (q_all, rest) = qkv.split_at(qk_dim);
let (k_all, v_all) = rest.split_at(qk_dim);
let scale = 1.0 / (head_dim as f32).sqrt();
let mut y_flat = vec![0f32; v_dim];
#[allow(clippy::needless_range_loop)]
for h in 0..n_heads {
let base = h * head_dim;
let mut q_h = q_all[base..base + head_dim].to_vec();
let mut k_h = k_all[base..base + head_dim].to_vec();
let v_h = &v_all[base..base + head_dim];
l2_normalize(&mut q_h, 1e-6);
l2_normalize(&mut k_h, 1e-6);
for x in q_h.iter_mut() {
*x *= scale;
}
let gate = softplus(alpha_raw[h] + weights.ssm_dt[h]) * weights.ssm_a[h];
let decay = gate.exp();
let beta = sigmoid(beta_raw[h]);
let s_base = h * head_dim * head_dim;
let s = &mut state.recurrent[s_base..s_base + head_dim * head_dim];
for cell in s.iter_mut() {
*cell *= decay;
}
let mut kv_mem = vec![0f32; head_dim];
for v_idx in 0..head_dim {
let mut acc = 0f32;
for k_idx in 0..head_dim {
acc += s[v_idx * head_dim + k_idx] * k_h[k_idx];
}
kv_mem[v_idx] = acc;
}
for v_idx in 0..head_dim {
let delta = (v_h[v_idx] - kv_mem[v_idx]) * beta;
for k_idx in 0..head_dim {
s[v_idx * head_dim + k_idx] += delta * k_h[k_idx];
}
}
for v_idx in 0..head_dim {
let mut acc = 0f32;
for k_idx in 0..head_dim {
acc += s[v_idx * head_dim + k_idx] * q_h[k_idx];
}
y_flat[base + v_idx] = acc;
}
}
let mut gated = vec![0f32; v_dim];
for h in 0..n_heads {
let base = h * head_dim;
let normed = rms_norm(
&y_flat[base..base + head_dim],
&weights.ssm_norm,
cfg.rms_norm_eps,
);
for i in 0..head_dim {
gated[base + i] = silu(z[base + i]) * normed[i];
}
}
weights.ssm_out.apply(&gated)
}
#[cfg(test)]
mod tests {
use super::*;
use ferrox_core::tensor::Tensor;
const HIDDEN: usize = 4;
const N_HEADS: usize = 2;
const HEAD_DIM: usize = 2;
const CONV_K: usize = 2;
const QKV_DIM: usize = 3 * N_HEADS * HEAD_DIM; const V_DIM: usize = N_HEADS * HEAD_DIM;
fn wm(data: &[f32], rows: usize, cols: usize) -> WeightMatrix {
assert_eq!(data.len(), rows * cols);
WeightMatrix::F32(Tensor::new(data.to_vec(), vec![rows, cols]))
}
fn cfg() -> GdnConfig {
GdnConfig {
hidden_dim: HIDDEN,
num_v_heads: N_HEADS,
head_dim: HEAD_DIM,
conv_kernel_size: CONV_K,
rms_norm_eps: 1e-5,
}
}
fn make_weights() -> GdnWeights {
let mut qkv = Vec::with_capacity(QKV_DIM * HIDDEN);
for i in 0..QKV_DIM * HIDDEN {
qkv.push(((i % 7) as f32 - 3.0) * 0.1);
}
let mut gate = Vec::with_capacity(V_DIM * HIDDEN);
for i in 0..V_DIM * HIDDEN {
gate.push(((i % 5) as f32 - 2.0) * 0.08);
}
let mut conv = Vec::with_capacity(QKV_DIM * CONV_K);
for i in 0..QKV_DIM * CONV_K {
conv.push(if i % CONV_K == CONV_K - 1 { 1.0 } else { 0.1 });
}
let mut beta = Vec::with_capacity(N_HEADS * HIDDEN);
let mut alpha = Vec::with_capacity(N_HEADS * HIDDEN);
for i in 0..N_HEADS * HIDDEN {
beta.push(((i % 3) as f32 - 1.0) * 0.2);
alpha.push(((i % 4) as f32 - 1.5) * 0.15);
}
let mut out = Vec::with_capacity(HIDDEN * V_DIM);
for i in 0..HIDDEN * V_DIM {
out.push(((i % 6) as f32 - 2.5) * 0.12);
}
GdnWeights {
attn_qkv: wm(&qkv, QKV_DIM, HIDDEN),
attn_gate: wm(&gate, V_DIM, HIDDEN),
ssm_conv1d: conv,
ssm_dt: vec![0.1, -0.05],
ssm_a: vec![-0.5, -0.75],
ssm_beta: wm(&beta, N_HEADS, HIDDEN),
ssm_alpha: wm(&alpha, N_HEADS, HIDDEN),
ssm_norm: vec![1.0, 1.0],
ssm_out: wm(&out, HIDDEN, V_DIM),
}
}
#[test]
fn gdn_forward_token_tiny_dims_finite_and_shaped() {
let weights = make_weights();
let cfg = cfg();
let mut state = GdnState::new(&cfg);
let hidden = [0.2f32, -0.1, 0.3, -0.4];
let out0 = gdn_forward_token(&weights, &cfg, &hidden, &mut state);
assert_eq!(out0.len(), HIDDEN);
assert!(out0.iter().all(|x| x.is_finite()));
let out1 = gdn_forward_token(&weights, &cfg, &hidden, &mut state);
assert_eq!(out1.len(), HIDDEN);
assert!(out1.iter().all(|x| x.is_finite()));
assert!(
out0.iter()
.zip(out1.iter())
.any(|(a, b)| (a - b).abs() > 1e-6),
"recurrent state should change the second token"
);
}
#[test]
fn softplus_matches_closed_form_at_zero() {
assert!((softplus(0.0) - (2.0f32).ln()).abs() < 1e-6);
}
}