use ferrox_core::weight_matrix::WeightMatrix;
pub fn grouped_output_projection(
attn_out: &[f32],
group_down: &[WeightMatrix],
wo_b: &WeightMatrix,
) -> Vec<f32> {
let n_groups = group_down.len();
assert!(n_groups > 0, "must have at least one output group");
assert_eq!(
attn_out.len() % n_groups,
0,
"attention output width must split evenly across groups"
);
let o_group_dim = attn_out.len() / n_groups;
let mut combined = Vec::new();
for (g, down) in group_down.iter().enumerate() {
assert_eq!(
down.cols(),
o_group_dim,
"group {g}'s down-projection must accept exactly one group's slice width"
);
let slice = &attn_out[g * o_group_dim..(g + 1) * o_group_dim];
combined.extend(down.apply(slice));
}
wo_b.apply(&combined)
}
#[cfg(test)]
mod tests {
use super::*;
use ferrox_core::tensor::Tensor;
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]))
}
#[test]
fn single_group_matches_a_plain_two_matrix_low_rank_projection() {
let attn_out = vec![1.0, 2.0, 3.0, 4.0];
let down_data = [0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8];
let down = wm(&down_data, 2, 4); let down_again = wm(&down_data, 2, 4);
let wo_b = wm(&[1.0, 0.0, 0.0, 1.0], 2, 2);
let out = grouped_output_projection(&attn_out, &[down], &wo_b);
let expected = wo_b.apply(&down_again.apply(&attn_out));
assert_eq!(out.len(), expected.len());
for (a, b) in out.iter().zip(expected.iter()) {
assert!((a - b).abs() < 1e-6);
}
}
#[test]
fn groups_are_independent_changing_one_groups_slice_only_affects_its_own_contribution() {
let attn_out = vec![1.0, 2.0, 100.0, -50.0]; let down0 = wm(&[0.5, -0.5], 1, 2); let down1_zero = wm(&[0.0, 0.0], 1, 2);
let wo_b = wm(&[2.0, 3.0], 1, 2);
let out = grouped_output_projection(&attn_out, &[down0, down1_zero], &wo_b);
assert!((out[0] - (-1.0)).abs() < 1e-5, "out[0]={}", out[0]);
}
#[test]
fn group_order_is_preserved_in_the_concatenation_fed_to_wo_b() {
let attn_out = vec![10.0, 20.0]; let down0 = wm(&[1.0], 1, 1); let down1 = wm(&[1.0], 1, 1); let wo_b = wm(&[1.0, 0.0], 1, 2);
let out = grouped_output_projection(&attn_out, &[down0, down1], &wo_b);
assert!(
(out[0] - 10.0).abs() < 1e-5,
"expected group 0's value first, out[0]={}",
out[0]
);
}
#[test]
#[should_panic(expected = "split evenly")]
fn mismatched_group_count_panics_rather_than_silently_truncating() {
let attn_out = vec![1.0, 2.0, 3.0]; let down = wm(&[1.0, 1.0, 1.0], 1, 3);
let down2 = wm(&[1.0, 1.0, 1.0], 1, 3);
let wo_b = wm(&[1.0, 1.0], 1, 2);
let _ = grouped_output_projection(&attn_out, &[down, down2], &wo_b);
}
}