fn silu(x: f32) -> f32 {
x / (1.0 + (-x).exp())
}
fn sigmoid(x: f32) -> f32 {
1.0 / (1.0 + (-x).exp())
}
fn matmul_a_by_bt(lhs: &[f32], rhs: &[f32], m: usize, k: usize, n: usize) -> Vec<f32> {
let mut out = vec![0.0f32; m * n];
for i in 0..m {
for j in 0..n {
let mut acc = 0.0f32;
for kk in 0..k {
acc += lhs[i * k + kk] * rhs[j * k + kk];
}
out[i * n + j] = acc;
}
}
out
}
#[derive(Debug, Clone)]
pub struct DenseFfnWeights {
pub gate: Vec<f32>,
pub up: Vec<f32>,
pub down: Vec<f32>,
}
#[derive(Debug, Clone, Copy)]
pub struct DenseFfnShape {
pub hidden_size: u32,
pub intermediate_size: u32,
}
pub fn dense_swiglu_cpu_ref(
x: &[f32],
weights: &DenseFfnWeights,
shape: DenseFfnShape,
) -> Vec<f32> {
let h = shape.hidden_size as usize;
let m = shape.intermediate_size as usize;
let seq_len = x.len() / h;
assert_eq!(x.len(), seq_len * h, "x shape mismatch");
assert_eq!(weights.gate.len(), m * h, "gate weight shape");
assert_eq!(weights.up.len(), m * h, "up weight shape");
assert_eq!(weights.down.len(), h * m, "down weight shape");
let a = matmul_a_by_bt(x, &weights.gate, seq_len, h, m);
let b = matmul_a_by_bt(x, &weights.up, seq_len, h, m);
let mut c = vec![0.0f32; seq_len * m];
for i in 0..(seq_len * m) {
c[i] = silu(a[i]) * b[i];
}
matmul_a_by_bt(&c, &weights.down, seq_len, m, h)
}
#[derive(Debug, Clone)]
pub struct MoeFfnWeights {
pub router: Vec<f32>,
pub expert_gate: Vec<f32>,
pub expert_up: Vec<f32>,
pub expert_down: Vec<f32>,
pub shared_gate_logit: Vec<f32>,
pub shared_gate: Vec<f32>, pub shared_up: Vec<f32>, pub shared_down: Vec<f32>, }
#[derive(Debug, Clone, Copy)]
pub struct MoeFfnShape {
pub hidden_size: u32,
pub num_experts: u32,
pub num_experts_per_tok: u32, pub moe_intermediate_size: u32,
pub shared_intermediate_size: u32,
}
pub fn moe_ffn_cpu_ref(x: &[f32], weights: &MoeFfnWeights, shape: MoeFfnShape) -> Vec<f32> {
let h = shape.hidden_size as usize;
let ne = shape.num_experts as usize;
let topk = shape.num_experts_per_tok as usize;
let m_moe = shape.moe_intermediate_size as usize;
let m_sh = shape.shared_intermediate_size as usize;
let seq_len = x.len() / h;
assert_eq!(x.len(), seq_len * h);
assert_eq!(weights.router.len(), ne * h);
assert_eq!(weights.expert_gate.len(), ne * m_moe * h);
assert_eq!(weights.expert_up.len(), ne * m_moe * h);
assert_eq!(weights.expert_down.len(), ne * h * m_moe);
assert_eq!(weights.shared_gate_logit.len(), h);
assert_eq!(weights.shared_gate.len(), m_sh * h);
assert_eq!(weights.shared_up.len(), m_sh * h);
assert_eq!(weights.shared_down.len(), h * m_sh);
assert!(topk <= ne, "top-k cannot exceed num_experts");
assert!(topk > 0, "top-k must be positive");
let mut output = vec![0.0f32; seq_len * h];
for t in 0..seq_len {
let x_t = &x[t * h..(t + 1) * h];
let mut logits = vec![0.0f32; ne];
for e in 0..ne {
let w_row = &weights.router[e * h..(e + 1) * h];
let mut acc = 0.0f32;
for i in 0..h {
acc += w_row[i] * x_t[i];
}
logits[e] = acc;
}
let max_l = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut exp = vec![0.0f32; ne];
let mut denom = 0.0f32;
for e in 0..ne {
exp[e] = (logits[e] - max_l).exp();
denom += exp[e];
}
let probs: Vec<f32> = exp.iter().map(|v| v / denom).collect();
let mut idx_sorted: Vec<usize> = (0..ne).collect();
idx_sorted.sort_by(|a, b| {
probs[*b]
.partial_cmp(&probs[*a])
.unwrap_or(std::cmp::Ordering::Equal)
});
let topk_idx: Vec<usize> = idx_sorted[..topk].to_vec();
let topk_probs: Vec<f32> = topk_idx.iter().map(|&i| probs[i]).collect();
let renorm_sum: f32 = topk_probs.iter().sum();
let topk_w: Vec<f32> = if renorm_sum > 1e-20 {
topk_probs.iter().map(|p| p / renorm_sum).collect()
} else {
vec![1.0 / topk as f32; topk]
};
let mut moe_out = vec![0.0f32; h];
for (w_i, &e_idx) in topk_w.iter().zip(topk_idx.iter()) {
let g_off = e_idx * m_moe * h;
let u_off = e_idx * m_moe * h;
let d_off = e_idx * h * m_moe;
let gate_w = &weights.expert_gate[g_off..g_off + m_moe * h];
let up_w = &weights.expert_up[u_off..u_off + m_moe * h];
let down_w = &weights.expert_down[d_off..d_off + h * m_moe];
let mut a = vec![0.0f32; m_moe];
for i in 0..m_moe {
let mut acc = 0.0f32;
for j in 0..h {
acc += gate_w[i * h + j] * x_t[j];
}
a[i] = acc;
}
let mut b = vec![0.0f32; m_moe];
for i in 0..m_moe {
let mut acc = 0.0f32;
for j in 0..h {
acc += up_w[i * h + j] * x_t[j];
}
b[i] = acc;
}
for i in 0..m_moe {
a[i] = silu(a[i]) * b[i];
}
let mut y = vec![0.0f32; h];
for i in 0..h {
let mut acc = 0.0f32;
for j in 0..m_moe {
acc += down_w[i * m_moe + j] * a[j];
}
y[i] = acc;
}
for i in 0..h {
moe_out[i] += w_i * y[i];
}
}
let shared_logit: f32 = weights
.shared_gate_logit
.iter()
.zip(x_t.iter())
.map(|(w, x)| w * x)
.sum();
let shared_gate_val = sigmoid(shared_logit);
let mut a_s = vec![0.0f32; m_sh];
for i in 0..m_sh {
let mut acc = 0.0f32;
for j in 0..h {
acc += weights.shared_gate[i * h + j] * x_t[j];
}
a_s[i] = acc;
}
let mut b_s = vec![0.0f32; m_sh];
for i in 0..m_sh {
let mut acc = 0.0f32;
for j in 0..h {
acc += weights.shared_up[i * h + j] * x_t[j];
}
b_s[i] = acc;
}
for i in 0..m_sh {
a_s[i] = silu(a_s[i]) * b_s[i];
}
let mut y_s = vec![0.0f32; h];
for i in 0..h {
let mut acc = 0.0f32;
for j in 0..m_sh {
acc += weights.shared_down[i * m_sh + j] * a_s[j];
}
y_s[i] = acc;
}
for i in 0..h {
output[t * h + i] = moe_out[i] + shared_gate_val * y_s[i];
}
}
output
}
#[cfg(test)]
mod tests {
use super::*;
fn mk_rand(seed: &mut u32, n: usize, scale: f32) -> Vec<f32> {
(0..n)
.map(|_| {
*seed = seed.wrapping_mul(1103515245).wrapping_add(12345);
((*seed as i32 as f32) / (i32::MAX as f32)) * scale
})
.collect()
}
#[test]
fn dense_swiglu_zero_weights_zero_output() {
let shape = DenseFfnShape {
hidden_size: 4,
intermediate_size: 8,
};
let weights = DenseFfnWeights {
gate: vec![0.0; 8 * 4],
up: vec![0.0; 8 * 4],
down: vec![0.0; 4 * 8],
};
let x: Vec<f32> = (0..16).map(|i| i as f32 * 0.1).collect(); let out = dense_swiglu_cpu_ref(&x, &weights, shape);
assert_eq!(out.len(), 16);
for v in &out {
assert!(v.abs() < 1e-7, "expected 0, got {}", v);
}
}
#[test]
fn dense_swiglu_matches_independent_recompute() {
let shape = DenseFfnShape {
hidden_size: 4,
intermediate_size: 8,
};
let mut seed = 0xAB_u32;
let weights = DenseFfnWeights {
gate: mk_rand(&mut seed, 8 * 4, 0.2),
up: mk_rand(&mut seed, 8 * 4, 0.2),
down: mk_rand(&mut seed, 4 * 8, 0.2),
};
let x = mk_rand(&mut seed, 8, 0.5);
let out = dense_swiglu_cpu_ref(&x, &weights, shape);
for t in 0..2 {
let x_t = &x[t * 4..(t + 1) * 4];
let mut a = [0.0f32; 8];
let mut b = [0.0f32; 8];
for i in 0..8 {
for j in 0..4 {
a[i] += weights.gate[i * 4 + j] * x_t[j];
b[i] += weights.up[i * 4 + j] * x_t[j];
}
}
for i in 0..8 {
a[i] = silu(a[i]) * b[i];
}
let mut y = [0.0f32; 4];
for i in 0..4 {
for j in 0..8 {
y[i] += weights.down[i * 8 + j] * a[j];
}
}
for j in 0..4 {
let d = (out[t * 4 + j] - y[j]).abs();
assert!(
d < 1e-5,
"t={}, j={}: got {}, want {}",
t,
j,
out[t * 4 + j],
y[j]
);
}
}
}
#[test]
fn dense_swiglu_deterministic() {
let shape = DenseFfnShape {
hidden_size: 4,
intermediate_size: 8,
};
let mut seed = 0xCC_u32;
let weights = DenseFfnWeights {
gate: mk_rand(&mut seed, 8 * 4, 0.1),
up: mk_rand(&mut seed, 8 * 4, 0.1),
down: mk_rand(&mut seed, 4 * 8, 0.1),
};
let x = vec![0.5; 4];
let o1 = dense_swiglu_cpu_ref(&x, &weights, shape);
let o2 = dense_swiglu_cpu_ref(&x, &weights, shape);
for i in 0..4 {
assert_eq!(o1[i].to_bits(), o2[i].to_bits());
}
}
#[test]
fn moe_4experts_routing_selects_top_k() {
let shape = MoeFfnShape {
hidden_size: 4,
num_experts: 4,
num_experts_per_tok: 2,
moe_intermediate_size: 4,
shared_intermediate_size: 4,
};
let mut seed = 0x100_u32;
let weights = MoeFfnWeights {
router: mk_rand(&mut seed, 4 * 4, 0.5),
expert_gate: mk_rand(&mut seed, 4 * 4 * 4, 0.1),
expert_up: mk_rand(&mut seed, 4 * 4 * 4, 0.1),
expert_down: mk_rand(&mut seed, 4 * 4 * 4, 0.1),
shared_gate_logit: mk_rand(&mut seed, 4, 0.1),
shared_gate: mk_rand(&mut seed, 4 * 4, 0.1),
shared_up: mk_rand(&mut seed, 4 * 4, 0.1),
shared_down: mk_rand(&mut seed, 4 * 4, 0.1),
};
let x: Vec<f32> = (0..4).map(|i| i as f32 * 0.1).collect();
let out = moe_ffn_cpu_ref(&x, &weights, shape);
assert_eq!(out.len(), 4);
for v in &out {
assert!(v.is_finite());
}
}
#[test]
fn moe_shared_expert_gate_controls_contribution() {
let shape = MoeFfnShape {
hidden_size: 4,
num_experts: 2,
num_experts_per_tok: 1,
moe_intermediate_size: 4,
shared_intermediate_size: 4,
};
let mut seed = 0x200_u32;
let base_weights = MoeFfnWeights {
router: mk_rand(&mut seed, 2 * 4, 0.5),
expert_gate: mk_rand(&mut seed, 2 * 4 * 4, 0.1),
expert_up: mk_rand(&mut seed, 2 * 4 * 4, 0.1),
expert_down: mk_rand(&mut seed, 2 * 4 * 4, 0.1),
shared_gate_logit: vec![0.0; 4], shared_gate: mk_rand(&mut seed, 4 * 4, 0.1),
shared_up: mk_rand(&mut seed, 4 * 4, 0.1),
shared_down: mk_rand(&mut seed, 4 * 4, 0.1),
};
let x: Vec<f32> = (0..4).map(|i| 0.1 * (i as f32 + 1.0)).collect();
let out_mid = moe_ffn_cpu_ref(&x, &base_weights, shape);
let mut weights_off = base_weights.clone();
weights_off.shared_gate_logit = vec![-1000.0; 4]; let out_off = moe_ffn_cpu_ref(&x, &weights_off, shape);
let mut weights_on = base_weights.clone();
weights_on.shared_gate_logit = vec![1000.0; 4];
let out_on = moe_ffn_cpu_ref(&x, &weights_on, shape);
for i in 0..4 {
let avg = 0.5 * (out_off[i] + out_on[i]);
let d = (out_mid[i] - avg).abs();
assert!(
d < 1e-3,
"gate linearity broken at {}: mid={}, avg_off_on={}, d={}",
i,
out_mid[i],
avg,
d
);
}
}
#[test]
fn moe_deterministic() {
let shape = MoeFfnShape {
hidden_size: 4,
num_experts: 3,
num_experts_per_tok: 2,
moe_intermediate_size: 4,
shared_intermediate_size: 4,
};
let mut seed = 0x300_u32;
let weights = MoeFfnWeights {
router: mk_rand(&mut seed, 3 * 4, 0.3),
expert_gate: mk_rand(&mut seed, 3 * 4 * 4, 0.1),
expert_up: mk_rand(&mut seed, 3 * 4 * 4, 0.1),
expert_down: mk_rand(&mut seed, 3 * 4 * 4, 0.1),
shared_gate_logit: mk_rand(&mut seed, 4, 0.1),
shared_gate: mk_rand(&mut seed, 4 * 4, 0.1),
shared_up: mk_rand(&mut seed, 4 * 4, 0.1),
shared_down: mk_rand(&mut seed, 4 * 4, 0.1),
};
let x: Vec<f32> = (0..4).map(|i| 0.1 * i as f32).collect();
let o1 = moe_ffn_cpu_ref(&x, &weights, shape);
let o2 = moe_ffn_cpu_ref(&x, &weights, shape);
for i in 0..4 {
assert_eq!(o1[i].to_bits(), o2[i].to_bits());
}
}
#[test]
fn moe_topk_all_experts_eq_softmax_weighted_sum() {
let shape = MoeFfnShape {
hidden_size: 2,
num_experts: 3,
num_experts_per_tok: 3, moe_intermediate_size: 2,
shared_intermediate_size: 2,
};
let mut seed = 0x400_u32;
let weights = MoeFfnWeights {
router: mk_rand(&mut seed, 3 * 2, 0.3),
expert_gate: mk_rand(&mut seed, 3 * 2 * 2, 0.2),
expert_up: mk_rand(&mut seed, 3 * 2 * 2, 0.2),
expert_down: mk_rand(&mut seed, 3 * 2 * 2, 0.2),
shared_gate_logit: vec![-1000.0; 2], shared_gate: vec![0.0; 2 * 2],
shared_up: vec![0.0; 2 * 2],
shared_down: vec![0.0; 2 * 2],
};
let x = vec![0.3, -0.2];
let out = moe_ffn_cpu_ref(&x, &weights, shape);
let mut logits = [0.0f32; 3];
for e in 0..3 {
for j in 0..2 {
logits[e] += weights.router[e * 2 + j] * x[j];
}
}
let max_l = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut exp = [0.0f32; 3];
let mut denom = 0.0f32;
for e in 0..3 {
exp[e] = (logits[e] - max_l).exp();
denom += exp[e];
}
let probs = [exp[0] / denom, exp[1] / denom, exp[2] / denom];
let mut expected = [0.0f32; 2];
for e in 0..3 {
let g_off = e * 2 * 2;
let u_off = e * 2 * 2;
let d_off = e * 2 * 2;
let mut a = [0.0f32; 2];
let mut b = [0.0f32; 2];
for i in 0..2 {
for j in 0..2 {
a[i] += weights.expert_gate[g_off + i * 2 + j] * x[j];
b[i] += weights.expert_up[u_off + i * 2 + j] * x[j];
}
}
for i in 0..2 {
a[i] = silu(a[i]) * b[i];
}
let mut y = [0.0f32; 2];
for i in 0..2 {
for j in 0..2 {
y[i] += weights.expert_down[d_off + i * 2 + j] * a[j];
}
}
for i in 0..2 {
expected[i] += probs[e] * y[i];
}
}
for i in 0..2 {
let d = (out[i] - expected[i]).abs();
assert!(
d < 1e-5,
"at {}: got {}, expected {}, d {}",
i,
out[i],
expected[i],
d
);
}
}
}