#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct DiffAttnConfig {
pub num_heads: usize,
pub num_kv_heads: usize,
pub head_dim: usize,
pub layer_depth: usize,
}
impl DiffAttnConfig {
#[inline]
pub fn lambda_init(&self) -> f32 {
0.8 - 0.6 * (-0.3 * self.layer_depth as f32).exp()
}
#[inline]
pub fn q_heads_packed(&self) -> usize {
2 * self.num_heads
}
#[inline]
pub fn kv_heads_packed(&self) -> usize {
2 * self.num_kv_heads
}
#[inline]
pub fn n_rep(&self) -> usize {
assert!(self.num_kv_heads > 0, "num_kv_heads must be > 0");
assert_eq!(
self.num_heads % self.num_kv_heads,
0,
"num_heads ({}) must be divisible by num_kv_heads ({})",
self.num_heads,
self.num_kv_heads
);
self.num_heads / self.num_kv_heads
}
#[inline]
pub fn q_dim(&self) -> usize {
self.q_heads_packed() * self.head_dim
}
#[inline]
pub fn k_dim(&self) -> usize {
self.kv_heads_packed() * self.head_dim
}
#[inline]
pub fn v_dim(&self) -> usize {
self.num_kv_heads * 2 * self.head_dim
}
#[inline]
pub fn out_dim(&self) -> usize {
self.num_heads * 2 * self.head_dim
}
}
pub struct DiffLambdaParams {
pub lambda_q1: Vec<f32>,
pub lambda_k1: Vec<f32>,
pub lambda_q2: Vec<f32>,
pub lambda_k2: Vec<f32>,
}
#[inline]
pub fn compute_lambda_full(params: &DiffLambdaParams, lambda_init: f32) -> f32 {
let len = params.lambda_q1.len();
assert_eq!(params.lambda_k1.len(), len, "lambda_k1 length != lambda_q1");
assert_eq!(params.lambda_q2.len(), len, "lambda_q2 length != lambda_q1");
assert_eq!(params.lambda_k2.len(), len, "lambda_k2 length != lambda_q1");
let dot1: f32 = params
.lambda_q1
.iter()
.zip(params.lambda_k1.iter())
.map(|(&q, &k)| q * k)
.sum();
let dot2: f32 = params
.lambda_q2
.iter()
.zip(params.lambda_k2.iter())
.map(|(&q, &k)| q * k)
.sum();
dot1.exp() - dot2.exp() + lambda_init
}
#[derive(Default, Clone, Debug)]
pub struct DiffAttnScratch {
scores1: Vec<f32>,
scores2: Vec<f32>,
v_head_t: Vec<f32>,
k_packed: Vec<f32>,
q_packed: Vec<f32>,
context: Vec<f32>,
}
impl DiffAttnScratch {
pub fn reserve_for(&mut self, seq_len: usize, cfg: &DiffAttnConfig) {
let n_pairs = cfg.num_heads;
let v_head_dim = 2 * cfg.head_dim; self.scores1.resize(n_pairs * seq_len * seq_len, 0.0_f32);
self.scores2.resize(n_pairs * seq_len * seq_len, 0.0_f32);
self.v_head_t.resize(v_head_dim * seq_len, 0.0_f32);
self.k_packed.resize(seq_len * cfg.head_dim, 0.0_f32);
self.q_packed.resize(seq_len * cfg.head_dim, 0.0_f32);
self.context.resize(n_pairs * seq_len * v_head_dim, 0.0_f32);
}
}
const MASK_VALUE: f32 = -10_000.0_f32;
#[allow(clippy::too_many_arguments)]
pub fn apply_differential_attention(
q_buf: &[f32],
k_buf: &[f32],
v_buf: &[f32],
lambda_params: &DiffLambdaParams,
subln_weight: &[f32],
subln_eps: f32,
attn_out: &mut [f32],
seq_len: usize,
cfg: &DiffAttnConfig,
scratch: &mut DiffAttnScratch,
) {
use crate::forward::cpu::matmul_bt;
assert!(cfg.num_heads > 0, "num_heads must be > 0");
assert!(cfg.num_kv_heads > 0, "num_kv_heads must be > 0");
assert!(cfg.head_dim > 0, "head_dim must be > 0");
assert_eq!(
cfg.num_heads % cfg.num_kv_heads,
0,
"num_heads ({}) must be divisible by num_kv_heads ({})",
cfg.num_heads,
cfg.num_kv_heads
);
assert_eq!(
q_buf.len(),
seq_len * cfg.q_dim(),
"q_buf length mismatch: expected {} got {}",
seq_len * cfg.q_dim(),
q_buf.len()
);
assert_eq!(
k_buf.len(),
seq_len * cfg.k_dim(),
"k_buf length mismatch: expected {} got {}",
seq_len * cfg.k_dim(),
k_buf.len()
);
assert_eq!(
v_buf.len(),
seq_len * cfg.v_dim(),
"v_buf length mismatch: expected {} got {}",
seq_len * cfg.v_dim(),
v_buf.len()
);
assert_eq!(
subln_weight.len(),
2 * cfg.head_dim,
"subln_weight length must be 2*head_dim"
);
assert_eq!(
attn_out.len(),
seq_len * cfg.out_dim(),
"attn_out length mismatch: expected {} got {}",
seq_len * cfg.out_dim(),
attn_out.len()
);
for (name, v) in [
("lambda_q1", &lambda_params.lambda_q1),
("lambda_k1", &lambda_params.lambda_k1),
("lambda_q2", &lambda_params.lambda_q2),
("lambda_k2", &lambda_params.lambda_k2),
] {
assert_eq!(
v.len(),
cfg.head_dim,
"{name} length must equal head_dim ({}), got {}",
cfg.head_dim,
v.len()
);
}
if seq_len == 0 {
return;
}
scratch.reserve_for(seq_len, cfg);
let lambda_init = cfg.lambda_init();
let lambda_full = compute_lambda_full(lambda_params, lambda_init);
let scale = (cfg.head_dim as f32).powf(-0.5);
let head_dim = cfg.head_dim;
let v_head_dim = 2 * head_dim;
let n_rep = cfg.n_rep();
let q_row_stride = cfg.q_dim(); let k_row_stride = cfg.k_dim(); let v_row_stride = cfg.v_dim(); let out_row_stride = cfg.out_dim();
for kv_h in 0..cfg.num_kv_heads {
let v_head_offset = kv_h * v_head_dim;
let v_t = &mut scratch.v_head_t[..v_head_dim * seq_len];
for pos in 0..seq_len {
let src_off = pos * v_row_stride + v_head_offset;
let v_row = &v_buf[src_off..src_off + v_head_dim];
for d in 0..v_head_dim {
v_t[d * seq_len + pos] = v_row[d];
}
}
let pair_start = kv_h * n_rep;
let pair_end = pair_start + n_rep;
for pair_h in pair_start..pair_end {
let q_h0 = 2 * pair_h; let q_h1 = 2 * pair_h + 1; let k_h0 = 2 * kv_h;
let k_h1 = 2 * kv_h + 1;
for (q_ph, k_ph, scores_target) in [
(q_h0, k_h0, &mut scratch.scores1),
(q_h1, k_h1, &mut scratch.scores2),
] {
let q_packed = &mut scratch.q_packed[..seq_len * head_dim];
for pos in 0..seq_len {
let src_off = pos * q_row_stride + q_ph * head_dim;
q_packed[pos * head_dim..pos * head_dim + head_dim]
.copy_from_slice(&q_buf[src_off..src_off + head_dim]);
}
let k_packed = &mut scratch.k_packed[..seq_len * head_dim];
for pos in 0..seq_len {
let src_off = pos * k_row_stride + k_ph * head_dim;
k_packed[pos * head_dim..pos * head_dim + head_dim]
.copy_from_slice(&k_buf[src_off..src_off + head_dim]);
}
let score_off = pair_h * seq_len * seq_len;
let score_slice = &mut scores_target[score_off..score_off + seq_len * seq_len];
matmul_bt(q_packed, k_packed, score_slice, seq_len, head_dim, seq_len);
for qi in 0..seq_len {
let row = &mut score_slice[qi * seq_len..(qi + 1) * seq_len];
for (ki, v) in row.iter_mut().enumerate() {
if ki <= qi {
*v *= scale;
} else {
*v = MASK_VALUE;
}
}
}
for qi in 0..seq_len {
let row = &mut score_slice[qi * seq_len..(qi + 1) * seq_len];
let valid = qi + 1;
let max_val = row[..valid]
.iter()
.copied()
.fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0_f32;
for v in &mut row[..valid] {
*v = (*v - max_val).exp();
sum += *v;
}
if sum > 0.0 {
let inv = 1.0 / sum;
for v in &mut row[..valid] {
*v *= inv;
}
}
row[valid..].fill(0.0);
}
}
let s1_off = pair_h * seq_len * seq_len;
let s2_off = pair_h * seq_len * seq_len;
for i in 0..(seq_len * seq_len) {
scratch.scores1[s1_off + i] -= lambda_full * scratch.scores2[s2_off + i];
}
let ctx_off = pair_h * seq_len * v_head_dim;
let ctx_slice = &mut scratch.context[ctx_off..ctx_off + seq_len * v_head_dim];
let diff_scores = &scratch.scores1[s1_off..s1_off + seq_len * seq_len];
matmul_bt(
diff_scores,
&scratch.v_head_t[..v_head_dim * seq_len],
ctx_slice,
seq_len,
seq_len,
v_head_dim,
);
{
use crate::forward::cpu::rms_norm;
rms_norm(ctx_slice, subln_weight, v_head_dim, subln_eps);
}
let scale_factor = 1.0 - lambda_init;
for pos in 0..seq_len {
let src_off = pos * v_head_dim;
let dst_off = pos * out_row_stride + pair_h * v_head_dim;
let src = &ctx_slice[src_off..src_off + v_head_dim];
let dst = &mut attn_out[dst_off..dst_off + v_head_dim];
for (d, s) in dst.iter_mut().zip(src.iter()) {
*d = s * scale_factor;
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn det_data(len: usize, seed: u64) -> Vec<f32> {
let mut state = seed.wrapping_add(0x9e37_79b9_7f4a_7c15);
let mut out = Vec::with_capacity(len);
for _ in 0..len {
state ^= state << 7;
state ^= state >> 9;
state = state.wrapping_mul(0x2545_f491_4f6c_dd1d);
let mantissa = ((state >> 41) as u32) & 0x007f_ffff;
let x = f32::from_bits(0x3f80_0000 | mantissa) - 1.5;
out.push(x);
}
out
}
#[allow(dead_code)]
fn max_abs_diff(a: &[f32], b: &[f32]) -> f32 {
assert_eq!(a.len(), b.len());
a.iter()
.zip(b.iter())
.map(|(x, y)| (x - y).abs())
.fold(0.0_f32, f32::max)
}
fn make_cfg(
num_heads: usize,
num_kv_heads: usize,
head_dim: usize,
depth: usize,
) -> DiffAttnConfig {
DiffAttnConfig {
num_heads,
num_kv_heads,
head_dim,
layer_depth: depth,
}
}
fn zero_lambda_params(head_dim: usize) -> DiffLambdaParams {
DiffLambdaParams {
lambda_q1: vec![0.0; head_dim],
lambda_k1: vec![0.0; head_dim],
lambda_q2: vec![0.0; head_dim],
lambda_k2: vec![0.0; head_dim],
}
}
#[test]
fn test_lambda_init_schedule() {
let cfg0 = make_cfg(2, 2, 4, 0);
let li0 = cfg0.lambda_init();
assert!(
(li0 - 0.2).abs() < 1e-6,
"lambda_init(0) expected 0.2, got {li0}"
);
let prev_cfg = make_cfg(2, 2, 4, 5);
let li5 = prev_cfg.lambda_init();
let cfg10 = make_cfg(2, 2, 4, 10);
let li10 = cfg10.lambda_init();
assert!(li5 > li0, "lambda_init should increase with depth");
assert!(li10 > li5, "lambda_init should increase with depth");
assert!(li10 < 0.8, "lambda_init should approach but stay below 0.8");
let mut prev = cfg0.lambda_init();
for d in 1..=50 {
let cur = make_cfg(2, 2, 4, d).lambda_init();
assert!(
cur > prev,
"lambda_init not monotone at depth {d}: cur={cur} prev={prev}"
);
prev = cur;
}
}
#[test]
fn test_compute_lambda_full() {
let params = zero_lambda_params(4);
let lambda_init = 0.3;
let lf = compute_lambda_full(¶ms, lambda_init);
assert!(
(lf - lambda_init).abs() < 1e-6,
"all-zero params: lambda_full expected {lambda_init}, got {lf}"
);
let params2 = DiffLambdaParams {
lambda_q1: vec![1.0, 0.0, 0.0, 0.0],
lambda_k1: vec![2.0, 0.0, 0.0, 0.0],
lambda_q2: vec![0.0, 1.0, 0.0, 0.0],
lambda_k2: vec![0.0, 1.0, 0.0, 0.0],
};
let expected = 2.0f32.exp() - 1.0f32.exp() + 0.3;
let got = compute_lambda_full(¶ms2, 0.3);
assert!(
(got - expected).abs() < 1e-5,
"hand-computed lambda_full: expected {expected}, got {got}"
);
}
#[test]
fn test_diff_attn_shapes() {
let cfg = make_cfg(2, 2, 4, 0);
let seq_len = 5;
let q = det_data(seq_len * cfg.q_dim(), 1);
let k = det_data(seq_len * cfg.k_dim(), 2);
let v = det_data(seq_len * cfg.v_dim(), 3);
let params = zero_lambda_params(cfg.head_dim);
let subln_w = vec![1.0f32; 2 * cfg.head_dim];
let mut out = vec![0.0f32; seq_len * cfg.out_dim()];
let mut scratch = DiffAttnScratch::default();
apply_differential_attention(
&q,
&k,
&v,
¶ms,
&subln_w,
1e-6,
&mut out,
seq_len,
&cfg,
&mut scratch,
);
assert_eq!(
out.len(),
seq_len * cfg.num_heads * 2 * cfg.head_dim,
"output length must be seq_len * num_heads * 2 * head_dim"
);
}
#[test]
fn test_diff_attn_single_head_single_token() {
let head_dim = 2_usize;
let cfg = make_cfg(1, 1, head_dim, 3); let lambda_init = cfg.lambda_init();
let seq_len = 1_usize;
let q = vec![1.0f32, 0.0, 1.0, 0.0]; let k = vec![1.0f32, 0.0, 1.0, 0.0]; let v = vec![2.0f32, 3.0, 0.0, 0.0];
let params = zero_lambda_params(head_dim); let subln_w = vec![1.0f32; 2 * head_dim]; let mut out = vec![0.0f32; seq_len * cfg.out_dim()]; let mut scratch = DiffAttnScratch::default();
apply_differential_attention(
&q,
&k,
&v,
¶ms,
&subln_w,
1e-6,
&mut out,
seq_len,
&cfg,
&mut scratch,
);
assert!(out[0].is_finite(), "output must be finite");
let scale_factor = 1.0 - lambda_init;
assert!(
scale_factor > 0.0,
"scale factor (1-lambda_init) must be positive"
);
}
#[test]
fn test_diff_attn_causal_masking() {
let head_dim = 2_usize;
let cfg = make_cfg(1, 1, head_dim, 0);
let seq_len = 3_usize;
let v_head_dim = 2 * head_dim;
let q = det_data(seq_len * cfg.q_dim(), 101);
let k = det_data(seq_len * cfg.k_dim(), 202);
let params = zero_lambda_params(head_dim);
let subln_w = vec![1.0f32; 2 * head_dim];
let v_base = det_data(seq_len * cfg.v_dim(), 303);
let mut v_perturbed = v_base.clone();
for pos in 1..seq_len {
for d in 0..v_head_dim {
v_perturbed[pos * v_head_dim + d] += 12_345.0;
}
}
let run = |v: &[f32]| {
let mut out = vec![0.0f32; seq_len * cfg.out_dim()];
let mut scratch = DiffAttnScratch::default();
apply_differential_attention(
&q,
&k,
v,
¶ms,
&subln_w,
1e-6,
&mut out,
seq_len,
&cfg,
&mut scratch,
);
out
};
let out_base = run(&v_base);
let out_perturbed = run(&v_perturbed);
for d in 0..v_head_dim {
assert_eq!(
out_base[d].to_bits(),
out_perturbed[d].to_bits(),
"position 0 changed when only future V changed — causal mask leak at dim {d}"
);
}
let pos2_changed = (0..v_head_dim).any(|d| {
let off = 2 * cfg.out_dim() + d;
out_base[off] != out_perturbed[off]
});
assert!(
pos2_changed,
"position 2 should be affected by the future-V perturbation"
);
}
#[test]
fn test_lambda_one_zeroes_identical_maps() {
let head_dim = 4_usize;
let cfg = make_cfg(1, 1, head_dim, 0); let seq_len = 3_usize;
let lambda_init = cfg.lambda_init();
let mut lambda_q1 = vec![0.0f32; head_dim];
let mut lambda_k1 = vec![0.0f32; head_dim];
lambda_q1[0] = (2.0 - lambda_init).ln();
lambda_k1[0] = 1.0;
let params = DiffLambdaParams {
lambda_q1,
lambda_k1,
lambda_q2: vec![0.0f32; head_dim],
lambda_k2: vec![0.0f32; head_dim],
};
let lambda_full = compute_lambda_full(¶ms, lambda_init);
assert!(
(lambda_full - 1.0).abs() < 1e-5,
"test setup error: lambda_full should be ≈1.0, got {lambda_full}"
);
let mut q = vec![0.0f32; seq_len * cfg.q_dim()];
let mut k = vec![0.0f32; seq_len * cfg.k_dim()];
for pos in 0..seq_len {
let base = pos * cfg.q_dim();
for d in 0..head_dim {
q[base + d] = det_data(1, (pos * head_dim + d) as u64)[0];
q[base + head_dim + d] = q[base + d]; }
let kbase = pos * cfg.k_dim();
for d in 0..head_dim {
k[kbase + d] = det_data(1, (pos * head_dim + d + 1000) as u64)[0];
k[kbase + head_dim + d] = k[kbase + d]; }
}
let v = det_data(seq_len * cfg.v_dim(), 77);
let subln_w = vec![1.0f32; 2 * head_dim];
let mut out = vec![0.0f32; seq_len * cfg.out_dim()];
let mut scratch = DiffAttnScratch::default();
apply_differential_attention(
&q,
&k,
&v,
¶ms,
&subln_w,
1e-6,
&mut out,
seq_len,
&cfg,
&mut scratch,
);
let max_abs = out.iter().copied().fold(0.0f32, |m, x| m.max(x.abs()));
assert!(
max_abs < 1e-3,
"identical maps with lambda_full≈1 must yield ~0 output, got max_abs={max_abs}"
);
}
#[test]
#[should_panic(expected = "num_heads must be > 0")]
fn test_zero_num_heads_panics() {
let cfg = make_cfg(0, 1, 4, 0);
let params = zero_lambda_params(4);
let mut out: Vec<f32> = vec![];
let mut scratch = DiffAttnScratch::default();
apply_differential_attention(
&[],
&[],
&[],
¶ms,
&[],
1e-6,
&mut out,
1,
&cfg,
&mut scratch,
);
}
#[test]
#[should_panic(expected = "num_kv_heads must be > 0")]
fn test_zero_num_kv_heads_panics() {
let cfg = make_cfg(2, 0, 4, 0);
let params = zero_lambda_params(4);
let mut out: Vec<f32> = vec![];
let mut scratch = DiffAttnScratch::default();
apply_differential_attention(
&[],
&[],
&[],
¶ms,
&[],
1e-6,
&mut out,
1,
&cfg,
&mut scratch,
);
}
#[test]
#[should_panic(expected = "head_dim must be > 0")]
fn test_zero_head_dim_panics() {
let cfg = make_cfg(2, 2, 0, 0);
let params = zero_lambda_params(0);
let mut out: Vec<f32> = vec![];
let mut scratch = DiffAttnScratch::default();
apply_differential_attention(
&[],
&[],
&[],
¶ms,
&[],
1e-6,
&mut out,
1,
&cfg,
&mut scratch,
);
}
#[test]
#[should_panic(expected = "num_kv_heads must be > 0")]
fn test_n_rep_zero_num_kv_heads_panics() {
let _ = make_cfg(4, 0, 8, 0).n_rep();
}
#[test]
#[should_panic(expected = "must be divisible")]
fn test_n_rep_non_divisible_panics() {
let _ = make_cfg(5, 2, 8, 0).n_rep();
}
}