use crate::clamp_kqv::clamp_in_place;
use super::{Decoder, LayerWeights};
impl Decoder {
#[allow(clippy::too_many_arguments)]
pub(crate) fn apply_qkv_bias_and_clamp(
&self,
layer: &LayerWeights,
q: &mut [f32],
k: &mut [f32],
v: &mut [f32],
q_width: usize,
kv_width: usize,
v_width: usize,
) {
let add_bias = |x: &mut [f32], bias: Option<&Vec<f32>>, width: usize| {
if let Some(bias) = bias {
debug_assert_eq!(bias.len(), width);
for row in x.chunks_mut(width) {
for (x, b) in row.iter_mut().zip(bias.iter()) {
*x += b;
}
}
}
};
add_bias(q, layer.attn.q_bias.as_ref(), q_width);
add_bias(k, layer.attn.k_bias.as_ref(), kv_width);
add_bias(v, layer.attn.v_bias.as_ref(), v_width);
if let Some(c) = self.config.clamp_kqv {
clamp_in_place(q, c);
clamp_in_place(k, c);
clamp_in_place(v, c);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_clamp_follows_the_bias() {
let mut cfg = crate::config::test_dense_fixture();
cfg.clamp_kqv = Some(8.0);
let mut d = Decoder::new_random_small(cfg, 1, 32);
let q_width = d.config.n_heads * d.config.head_dim;
let kv_width = d.config.n_kv_heads * d.config.head_dim;
d.layers[0].attn.q_bias = Some(vec![1.0; q_width]);
d.layers[0].attn.k_bias = Some(vec![-1.0; kv_width]);
let mut q = vec![7.5f32; q_width];
let mut k = vec![-7.5f32; kv_width];
let mut v = vec![100.0f32; kv_width];
let layer = &d.layers[0];
d.apply_qkv_bias_and_clamp(layer, &mut q, &mut k, &mut v, q_width, kv_width, kv_width);
assert!(q.iter().all(|&x| x == 8.0), "{q:?}");
assert!(k.iter().all(|&x| x == -8.0), "{k:?}");
assert!(
v.iter().all(|&x| x == 8.0),
"V has no bias and is still clamped: {v:?}"
);
}
#[test]
fn no_clamp_means_the_bias_alone() {
let mut d = Decoder::new_random_small(crate::config::test_dense_fixture(), 1, 32);
assert_eq!(d.config.clamp_kqv, None);
let q_width = d.config.n_heads * d.config.head_dim;
let kv_width = d.config.n_kv_heads * d.config.head_dim;
d.layers[0].attn.v_bias = Some(vec![0.5; kv_width]);
let mut q = vec![100.0f32; 2 * q_width];
let mut k = vec![-100.0f32; 2 * kv_width];
let mut v = vec![100.0f32; 2 * kv_width];
let layer = &d.layers[0];
d.apply_qkv_bias_and_clamp(layer, &mut q, &mut k, &mut v, q_width, kv_width, kv_width);
assert!(q.iter().all(|&x| x == 100.0));
assert!(k.iter().all(|&x| x == -100.0));
assert!(v.iter().all(|&x| x == 100.5), "{v:?}");
}
}