use ferrox_core::matmul::{rms_norm, rms_norm_per_head};
use super::{Decoder, LayerWeights};
impl Decoder {
pub(crate) fn apply_qk_norm(&self, x: &[f32], weight: &[f32]) -> Vec<f32> {
use crate::capability::QkNormStyle;
match self.config.qk_norm_style {
QkNormStyle::WholeVector => rms_norm(x, weight, self.config.rms_norm_eps),
QkNormStyle::PerHead => {
rms_norm_per_head(x, weight, self.config.head_dim, self.config.rms_norm_eps)
}
}
}
fn apply_qk_norms_batch(
&self,
layer: &LayerWeights,
q_batch: &mut [f32],
k_batch: &mut [f32],
q_width: usize,
kv_width: usize,
) {
if let Some(q_norm) = &layer.attn.q_norm {
for row in q_batch.chunks_mut(q_width) {
let normed = self.apply_qk_norm(row, q_norm);
row.copy_from_slice(&normed);
}
}
if let Some(k_norm) = &layer.attn.k_norm {
for row in k_batch.chunks_mut(kv_width) {
let normed = self.apply_qk_norm(row, k_norm);
row.copy_from_slice(&normed);
}
}
}
pub(crate) fn apply_qk_norms_pre_rope(
&self,
layer: &LayerWeights,
q_batch: &mut [f32],
k_batch: &mut [f32],
q_width: usize,
kv_width: usize,
) {
if !self.qk_norm_after_rope {
self.apply_qk_norms_batch(layer, q_batch, k_batch, q_width, kv_width);
}
}
pub(crate) fn apply_qk_norms_post_rope(
&self,
layer: &LayerWeights,
q_batch: &mut [f32],
k_batch: &mut [f32],
q_width: usize,
kv_width: usize,
) {
if self.qk_norm_after_rope {
self.apply_qk_norms_batch(layer, q_batch, k_batch, q_width, kv_width);
}
}
}
#[cfg(test)]
mod tests {
use crate::config::glm_5_2;
use crate::Decoder;
#[test]
fn the_two_hooks_are_exclusive_and_together_always_norm_exactly_once() {
for after in [false, true] {
let mut cfg = glm_5_2();
cfg.hidden_dim = 16;
cfg.n_heads = 4;
cfg.n_kv_heads = 2;
cfg.head_dim = 4;
cfg.moe.hidden_dim = 16;
cfg.moe.expert_ffn_dim = 8;
let mut decoder = Decoder::new_random_small(cfg, 1, 8);
decoder.qk_norm_after_rope = after;
let q_width = decoder.config.n_heads * decoder.config.head_dim;
let kv_width = decoder.config.n_kv_heads * decoder.config.head_dim;
decoder.layers[0].attn.q_norm = Some(vec![2.0; q_width]);
decoder.layers[0].attn.k_norm = Some(vec![3.0; kv_width]);
let q0: Vec<f32> = (0..q_width).map(|i| 0.5 + i as f32 * 0.25).collect();
let k0: Vec<f32> = (0..kv_width).map(|i| 1.0 - i as f32 * 0.1).collect();
let mut q = q0.clone();
let mut k = k0.clone();
let layer = &decoder.layers[0];
decoder.apply_qk_norms_pre_rope(layer, &mut q, &mut k, q_width, kv_width);
decoder.apply_qk_norms_post_rope(layer, &mut q, &mut k, q_width, kv_width);
let want_q = decoder.apply_qk_norm(&q0, &vec![2.0; q_width]);
let want_k = decoder.apply_qk_norm(&k0, &vec![3.0; kv_width]);
for (got, want) in q.iter().zip(want_q.iter()) {
assert!(
(got - want).abs() < 1e-6,
"after_rope={after}: Q normed {got} vs {want}"
);
}
for (got, want) in k.iter().zip(want_k.iter()) {
assert!(
(got - want).abs() < 1e-6,
"after_rope={after}: K normed {got} vs {want}"
);
}
assert!(
q.iter().zip(q0.iter()).any(|(a, b)| (a - b).abs() > 1e-3),
"after_rope={after}: the norm changed nothing"
);
}
}
}
#[cfg(all(test, feature = "metal"))]
mod metal_tests {
use super::*;
#[test]
fn metal_attention_refuses_a_layer_whose_qk_norm_runs_after_rope() {
let mut decoder = Decoder::new_random_small(crate::config::test_dense_fixture(), 1, 32);
let q_width = decoder.config.n_heads * decoder.config.head_dim;
let kv_width = decoder.config.n_kv_heads * decoder.config.head_dim;
decoder.layers[0].attn.q_norm = Some(vec![1.0; q_width]);
decoder.layers[0].attn.k_norm = Some(vec![1.0; kv_width]);
decoder.qk_norm_after_rope = false;
assert!(
decoder.layer_supports_metal_attn(&decoder.layers[0]),
"the fence below would prove nothing if this layer were ineligible anyway"
);
decoder.qk_norm_after_rope = true;
assert!(!decoder.layer_supports_metal_attn(&decoder.layers[0]));
}
}