use ferrox_core::attention::{
apply_rope, apply_rope_interleaved, apply_rope_interleaved_with_freq_factors,
apply_rope_with_freq_factors,
};
use super::Decoder;
impl Decoder {
pub(crate) fn apply_rope_head_theta(
&self,
slice: &mut [f32],
pos: usize,
theta: f32,
freq_factors: Option<&[f32]>,
) {
use crate::config::RopeLayout;
let slice = match self.config.rope_dim {
Some(rot) if rot < slice.len() => &mut slice[..rot],
_ => slice,
};
match (self.config.rope_layout, freq_factors) {
(RopeLayout::Norm, Some(freq_factors)) => {
apply_rope_interleaved_with_freq_factors(slice, pos, theta, freq_factors)
}
(RopeLayout::Norm, None) => apply_rope_interleaved(slice, pos, theta),
(RopeLayout::Neox, Some(freq_factors)) => {
apply_rope_with_freq_factors(slice, pos, theta, freq_factors)
}
(RopeLayout::Neox, None) => apply_rope(slice, pos, theta),
}
}
pub(crate) fn apply_rope_head_layer(&self, slice: &mut [f32], pos: usize, layer_idx: usize) {
let (theta, freq_factors) = self.config.layer_rope(layer_idx);
self.apply_rope_head_theta(slice, pos, theta, freq_factors)
}
#[inline]
pub(crate) fn apply_rope_attn_factor(&self, q: &mut [f32], k: &mut [f32]) {
let m = self.config.rope_attn_factor;
if m == 1.0 {
return;
}
let head_dim = self.config.head_dim;
let rot = self.config.rope_dim.unwrap_or(head_dim).min(head_dim);
for buf in [q, k] {
for head in buf.chunks_mut(head_dim) {
let n = rot.min(head.len());
for v in head[..n].iter_mut() {
*v *= m;
}
}
}
}
#[inline]
pub(crate) fn apply_attention_scale(&self, q: &mut [f32]) {
let Some(scale) = self.config.attention_scale else {
return;
};
let compensate = scale * (self.config.head_dim as f32).sqrt();
for v in q.iter_mut() {
*v *= compensate;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn partial_rotary_leaves_the_tail_untouched() {
let mut cfg = crate::config::test_dense_fixture();
cfg.head_dim = 8;
cfg.rope_layout = crate::config::RopeLayout::Neox;
cfg.rope_freqs = None;
cfg.rope_dim = Some(4);
let decoder = Decoder::new_random_small(cfg, 1, 32);
let mut head: Vec<f32> = (0..8).map(|i| 1.0 + i as f32).collect();
let before = head.clone();
decoder.apply_rope_head_theta(&mut head, 3, 10000.0, None);
assert_eq!(
&head[4..],
&before[4..],
"dims at or past rope_dim must not rotate"
);
assert!(
head[..4] != before[..4],
"dims below rope_dim must rotate at a non-zero position"
);
}
#[test]
fn attn_factor_scales_only_the_rotated_channels() {
let mut cfg = crate::config::test_dense_fixture();
cfg.head_dim = 8;
cfg.n_heads = 2;
cfg.n_kv_heads = 2;
cfg.rope_dim = Some(4);
cfg.rope_attn_factor = 2.0;
let decoder = Decoder::new_random_small(cfg, 1, 32);
let mut q: Vec<f32> = (0..16).map(|i| 1.0 + i as f32).collect();
let mut k: Vec<f32> = (0..16).map(|i| 1.0 + i as f32).collect();
let before = q.clone();
decoder.apply_rope_attn_factor(&mut q, &mut k);
for h in 0..2 {
let base = h * 8;
for i in 0..4 {
assert_eq!(
q[base + i],
before[base + i] * 2.0,
"rotated channel {i} of head {h} must be scaled"
);
}
for i in 4..8 {
assert_eq!(
q[base + i],
before[base + i],
"pass-through channel {i} of head {h} must be untouched"
);
}
}
assert_eq!(q, k, "q and k take the same magnitude scale");
}
#[test]
fn attn_factor_scales_the_whole_head_without_partial_rotary() {
let mut cfg = crate::config::test_dense_fixture();
cfg.head_dim = 8;
cfg.n_heads = 1;
cfg.n_kv_heads = 1;
cfg.rope_dim = None;
cfg.rope_attn_factor = 3.0;
let decoder = Decoder::new_random_small(cfg, 1, 32);
let mut q: Vec<f32> = (0..8).map(|i| 1.0 + i as f32).collect();
let mut k = q.clone();
let before = q.clone();
decoder.apply_rope_attn_factor(&mut q, &mut k);
for i in 0..8 {
assert_eq!(q[i], before[i] * 3.0);
}
}
#[test]
fn full_rotary_still_rotates_the_whole_head() {
let mut cfg = crate::config::test_dense_fixture();
cfg.head_dim = 8;
cfg.rope_layout = crate::config::RopeLayout::Neox;
cfg.rope_freqs = None;
cfg.rope_dim = None;
let decoder = Decoder::new_random_small(cfg, 1, 32);
let mut head: Vec<f32> = (0..8).map(|i| 1.0 + i as f32).collect();
let before = head.clone();
decoder.apply_rope_head_theta(&mut head, 3, 10000.0, None);
assert!(head[4..] != before[4..]);
}
#[test]
fn rope_attn_factor_scales_q_and_k_only() {
let mut cfg = crate::config::test_dense_fixture();
cfg.rope_attn_factor = 2.0;
let decoder = Decoder::new_random_small(cfg, 1, 32);
let mut q = vec![1.0f32, -2.0, 3.0];
let mut k = vec![0.5f32, 4.0];
decoder.apply_rope_attn_factor(&mut q, &mut k);
assert_eq!(q, vec![2.0, -4.0, 6.0]);
assert_eq!(k, vec![1.0, 8.0]);
let mut cfg = crate::config::test_dense_fixture();
cfg.rope_attn_factor = 1.0;
let decoder = Decoder::new_random_small(cfg, 1, 32);
let mut q = vec![1.0f32, -2.0];
let mut k = vec![3.0f32];
decoder.apply_rope_attn_factor(&mut q, &mut k);
assert_eq!(q, vec![1.0, -2.0]);
assert_eq!(k, vec![3.0]);
}
#[test]
fn a_sliding_layer_ropes_unscaled_where_a_full_attention_layer_ropes_scaled() {
use crate::config::{RopeFreqs, RopeLayout};
const FACTOR: f32 = 8.0;
let mut cfg = crate::config::test_dense_fixture();
cfg.head_dim = 8;
cfg.rope_layout = RopeLayout::Norm;
cfg.rope_theta = 1_000_000.0;
cfg.rope_theta_swa = Some(10_000.0);
cfg.sliding_window = Some(4);
cfg.swa_pattern = Some(2);
cfg.swa_dense_first = false;
cfg.rope_freqs = Some(RopeFreqs {
full: vec![FACTOR; 4],
swa: Some(vec![1.0; 4]),
});
assert_eq!(cfg.layer_sliding_window(0), Some(4), "layer 0 must slide");
assert_eq!(cfg.layer_sliding_window(1), None, "layer 1 must be dense");
let decoder = Decoder::new_random_small(cfg, 2, 32);
let source: Vec<f32> = (0..8).map(|i| 1.0 + i as f32).collect();
let pos = 7;
let mut sliding = source.clone();
decoder.apply_rope_head_layer(&mut sliding, pos, 0);
let mut want_sliding = source.clone();
apply_rope_interleaved(&mut want_sliding, pos, 10_000.0);
assert_eq!(
sliding, want_sliding,
"a Gemma-3 sliding layer rotates at the raw position, base 10000"
);
let mut full = source.clone();
decoder.apply_rope_head_layer(&mut full, pos, 1);
let mut want_full = source.clone();
apply_rope_interleaved_with_freq_factors(&mut want_full, pos, 1_000_000.0, &[FACTOR; 4]);
assert_eq!(
full, want_full,
"a Gemma-3 full-attention layer rotates at p/8, base 1e6"
);
assert_ne!(
sliding, full,
"if these agree the test proves nothing about per-layer RoPE"
);
let mut was_wrong = source.clone();
apply_rope_interleaved_with_freq_factors(&mut was_wrong, pos, 10_000.0, &[FACTOR; 4]);
assert_ne!(
sliding, was_wrong,
"the sliding layer must not carry the full-attention scale"
);
}
#[test]
fn an_inheriting_architecture_ropes_both_kinds_of_layer_with_one_set() {
use crate::config::{RopeFreqs, RopeLayout};
let mut cfg = crate::config::test_dense_fixture();
cfg.head_dim = 8;
cfg.rope_layout = RopeLayout::Norm;
cfg.sliding_window = Some(4);
cfg.swa_pattern = Some(2);
cfg.rope_freqs = Some(RopeFreqs {
full: vec![8.0; 4],
swa: None,
});
assert!(
!cfg.rope_freqs_vary_by_layer(),
"an inheriting model resolves to ONE divisor set for every layer"
);
let decoder = Decoder::new_random_small(cfg, 2, 32);
let source: Vec<f32> = (0..8).map(|i| 1.0 + i as f32).collect();
let mut sliding = source.clone();
let mut full = source.clone();
decoder.apply_rope_head_layer(&mut sliding, 7, 0);
decoder.apply_rope_head_layer(&mut full, 7, 1);
assert_eq!(
sliding, full,
"both layers share one base and one divisor set here"
);
}
}