use crate::attention::apply_rope;
use crate::matmul::rms_norm;
pub fn channel_gated_pool(kv_block: &[Vec<f32>], score_block: &[Vec<f32>]) -> Vec<f32> {
assert_eq!(
kv_block.len(),
score_block.len(),
"kv_block and score_block must have the same number of raw positions"
);
assert!(
!kv_block.is_empty(),
"a compression block must have at least one raw position"
);
let n_embd_head = kv_block[0].len();
for (kv, score) in kv_block.iter().zip(score_block.iter()) {
assert_eq!(kv.len(), n_embd_head);
assert_eq!(score.len(), n_embd_head);
}
let mut out = vec![0f32; n_embd_head];
for c in 0..n_embd_head {
let max = score_block
.iter()
.map(|s| s[c])
.fold(f32::NEG_INFINITY, f32::max);
let mut weights: Vec<f32> = score_block.iter().map(|s| (s[c] - max).exp()).collect();
let sum: f32 = weights.iter().sum();
for w in weights.iter_mut() {
*w /= sum;
}
out[c] = kv_block
.iter()
.zip(weights.iter())
.map(|(kv, &w)| kv[c] * w)
.sum();
}
out
}
pub fn compress_block(
kv_block: &[Vec<f32>],
score_block: &[Vec<f32>],
norm_weight: &[f32],
rms_eps: f32,
n_embd_head_rope: usize,
block_position: usize,
compress_rope_theta: f32,
) -> Vec<f32> {
let pooled = channel_gated_pool(kv_block, score_block);
let n_embd_head = pooled.len();
assert!(
n_embd_head_rope <= n_embd_head,
"rope slice cannot exceed the compressed vector's width"
);
assert_eq!(norm_weight.len(), n_embd_head);
let mut normed = rms_norm(&pooled, norm_weight, rms_eps);
let rope_start = n_embd_head - n_embd_head_rope;
apply_rope(
&mut normed[rope_start..],
block_position,
compress_rope_theta,
);
normed
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn channel_gated_pool_with_one_hot_scores_selects_that_positions_value() {
let kv_block = vec![vec![1.0, 2.0], vec![9.0, -3.0], vec![100.0, 200.0]];
let score_block = vec![vec![0.0, 0.0], vec![50.0, 50.0], vec![0.0, 0.0]];
let out = channel_gated_pool(&kv_block, &score_block);
assert_eq!(out.len(), 2);
assert!((out[0] - 9.0).abs() < 1e-3, "out[0]={}", out[0]);
assert!((out[1] - (-3.0)).abs() < 1e-3, "out[1]={}", out[1]);
}
#[test]
fn channel_gated_pool_gates_each_channel_independently() {
let kv_block = vec![vec![10.0, 20.0], vec![30.0, 40.0]];
let score_block = vec![vec![50.0, -50.0], vec![-50.0, 50.0]];
let out = channel_gated_pool(&kv_block, &score_block);
assert!((out[0] - 10.0).abs() < 1e-3, "out[0]={}", out[0]);
assert!((out[1] - 40.0).abs() < 1e-3, "out[1]={}", out[1]);
}
#[test]
fn channel_gated_pool_with_uniform_scores_is_a_plain_average() {
let kv_block = vec![vec![2.0, 4.0], vec![6.0, 8.0], vec![10.0, 12.0]];
let score_block = vec![vec![0.0, 0.0]; 3];
let out = channel_gated_pool(&kv_block, &score_block);
assert!((out[0] - 6.0).abs() < 1e-5, "out[0]={}", out[0]);
assert!((out[1] - 8.0).abs() < 1e-5, "out[1]={}", out[1]);
}
#[test]
fn channel_gated_pool_single_position_block_is_identity() {
let kv_block = vec![vec![1.5, -2.5, 3.5]];
let score_block = vec![vec![7.0, -3.0, 0.0]];
let out = channel_gated_pool(&kv_block, &score_block);
assert_eq!(out, vec![1.5, -2.5, 3.5]);
}
#[test]
fn compress_block_leaves_nope_slice_untouched_by_rope() {
let kv_block = vec![vec![1.0, 2.0, 3.0, 4.0]];
let score_block = vec![vec![0.0, 0.0, 0.0, 0.0]];
let norm_weight = vec![1.0, 1.0, 1.0, 1.0];
let out = compress_block(&kv_block, &score_block, &norm_weight, 1e-5, 2, 0, 10000.0);
let plain = rms_norm(&[1.0, 2.0, 3.0, 4.0], &norm_weight, 1e-5);
for (a, b) in out.iter().zip(plain.iter()) {
assert!((a - b).abs() < 1e-5, "a={a} b={b}");
}
}
#[test]
fn compress_block_rotates_only_the_rope_slice_at_nonzero_position() {
let kv_block = vec![vec![1.0, 2.0, 3.0, 4.0]];
let score_block = vec![vec![0.0, 0.0, 0.0, 0.0]];
let norm_weight = vec![1.0, 1.0, 1.0, 1.0];
let out = compress_block(&kv_block, &score_block, &norm_weight, 1e-5, 2, 5, 10000.0);
let plain = rms_norm(&[1.0, 2.0, 3.0, 4.0], &norm_weight, 1e-5);
assert!((out[0] - plain[0]).abs() < 1e-5);
assert!((out[1] - plain[1]).abs() < 1e-5);
let differs = (out[2] - plain[2]).abs() > 1e-4 || (out[3] - plain[3]).abs() > 1e-4;
assert!(differs, "rope slice must be rotated at a nonzero position");
}
#[test]
fn compress_block_rope_preserves_rope_slice_norm() {
let kv_block = vec![vec![0.3, -0.7, 1.1, -1.3, 0.2, 0.5]];
let score_block = vec![vec![0.0; 6]];
let norm_weight = vec![1.0; 6];
let out = compress_block(&kv_block, &score_block, &norm_weight, 1e-5, 4, 3, 10000.0);
let plain = rms_norm(&[0.3, -0.7, 1.1, -1.3, 0.2, 0.5], &norm_weight, 1e-5);
let rope_norm_before: f32 = plain[2..].iter().map(|v| v * v).sum::<f32>().sqrt();
let rope_norm_after: f32 = out[2..].iter().map(|v| v * v).sum::<f32>().sqrt();
assert!((rope_norm_before - rope_norm_after).abs() < 1e-4);
}
}