use frink_core::attention::{
apply_rope_interleaved, apply_rope_interleaved_with_freq_factors, causal_mla_attention_scaled,
};
use frink_core::matmul::rms_norm;
use frink_core::mla_absorbed::causal_mla_absorbed_attention;
use frink_core::weight_matrix::WeightMatrix;
use crate::config::MlaConfig;
pub use crate::mla_q_proj::MlaQProj;
use crate::mla_yarn::{kq_scale, MlaYarn};
pub struct MlaAttnWeights {
pub q: MlaQProj,
pub kv_a_proj_with_mqa: WeightMatrix, pub kv_a_layernorm: Vec<f32>, pub kv_b: MlaKvB,
pub o_proj: WeightMatrix, pub g_proj: Option<WeightMatrix>, }
pub enum MlaKvB {
Combined(WeightMatrix),
Split {
k_b: Vec<WeightMatrix>,
v_b: Vec<WeightMatrix>,
},
}
impl MlaKvB {
pub fn cached_positions(&self, cfg: &MlaConfig, k_cache_len: usize) -> usize {
match self {
MlaKvB::Combined(_) => {
k_cache_len / (cfg.num_heads * (cfg.qk_nope_head_dim + cfg.qk_rope_head_dim))
}
MlaKvB::Split { .. } => k_cache_len / (cfg.kv_lora_rank + cfg.qk_rope_head_dim),
}
}
}
#[allow(clippy::too_many_arguments)]
pub fn mla_forward_token(
weights: &MlaAttnWeights,
cfg: &MlaConfig,
yarn: Option<&MlaYarn>,
hidden: &[f32],
rms_norm_eps: f32,
k_cache: &mut Vec<f32>,
v_cache: &mut Vec<f32>,
) -> Vec<f32> {
let q_head_dim = cfg.qk_nope_head_dim + cfg.qk_rope_head_dim;
let pos = weights.kv_b.cached_positions(cfg, k_cache.len());
let mut query = weights.q.apply(hidden, rms_norm_eps); if let Some(rope) = &cfg.rope {
for h in 0..cfg.num_heads {
let q_rot_h = &mut query[h * q_head_dim + cfg.qk_nope_head_dim..(h + 1) * q_head_dim];
rotate_pe(q_rot_h, pos, rope.theta, yarn);
}
}
let compressed_kv = weights.kv_a_proj_with_mqa.apply(hidden);
let (k_pass_c, k_rot_raw) = compressed_kv.split_at(cfg.kv_lora_rank);
let mut k_rot = k_rot_raw.to_vec();
if let Some(rope) = &cfg.rope {
rotate_pe(&mut k_rot, pos, rope.theta, yarn);
}
let k_pass_c_normed = rms_norm(k_pass_c, &weights.kv_a_layernorm, rms_norm_eps);
let scale = kq_scale(yarn, q_head_dim);
let attn_out = match &weights.kv_b {
MlaKvB::Split { k_b, v_b } => absorbed_attention(
cfg,
k_b,
v_b,
&query,
&k_pass_c_normed,
&k_rot,
k_cache,
scale,
),
MlaKvB::Combined(kv_b_proj) => naive_attention(
cfg,
kv_b_proj,
&query,
&k_pass_c_normed,
&k_rot,
k_cache,
v_cache,
scale,
),
};
let gated = match &weights.g_proj {
Some(g_proj) => {
let g = g_proj.apply(hidden);
attn_out
.iter()
.zip(g.iter())
.map(|(a, g)| a * (1.0 / (1.0 + (-g).exp())))
.collect::<Vec<f32>>()
}
None => attn_out,
};
weights.o_proj.apply(&gated)
}
fn rotate_pe(slice: &mut [f32], pos: usize, theta: f32, yarn: Option<&MlaYarn>) {
match yarn {
None => apply_rope_interleaved(slice, pos, theta),
Some(y) => {
apply_rope_interleaved_with_freq_factors(slice, pos, theta, &y.freq_factors);
if y.pe_magnitude != 1.0 {
for v in slice.iter_mut() {
*v *= y.pe_magnitude;
}
}
}
}
}
#[allow(clippy::too_many_arguments)]
fn absorbed_attention(
cfg: &MlaConfig,
k_b: &[WeightMatrix],
v_b: &[WeightMatrix],
query: &[f32],
k_pass_c_normed: &[f32],
k_rot: &[f32],
k_cache: &mut Vec<f32>,
scale: f32,
) -> Vec<f32> {
let q_head_dim = cfg.qk_nope_head_dim + cfg.qk_rope_head_dim;
let width = cfg.kv_lora_rank + cfg.qk_rope_head_dim;
let mut q_abs = vec![0f32; cfg.num_heads * width];
for h in 0..cfg.num_heads {
let q_h = &query[h * q_head_dim..(h + 1) * q_head_dim];
let absorbed = k_b[h].apply(&q_h[..cfg.qk_nope_head_dim]);
q_abs[h * width..h * width + cfg.kv_lora_rank].copy_from_slice(&absorbed);
q_abs[h * width + cfg.kv_lora_rank..(h + 1) * width]
.copy_from_slice(&q_h[cfg.qk_nope_head_dim..]);
}
k_cache.extend_from_slice(k_pass_c_normed);
k_cache.extend_from_slice(k_rot);
let seq_len = k_cache.len() / width;
let out_lat = causal_mla_absorbed_attention(
&q_abs,
k_cache,
cfg.num_heads,
cfg.kv_lora_rank,
cfg.qk_rope_head_dim,
seq_len,
scale,
);
let mut out = vec![0f32; cfg.num_heads * cfg.v_head_dim];
for h in 0..cfg.num_heads {
let v_h = v_b[h].apply(&out_lat[h * cfg.kv_lora_rank..(h + 1) * cfg.kv_lora_rank]);
out[h * cfg.v_head_dim..(h + 1) * cfg.v_head_dim].copy_from_slice(&v_h);
}
out
}
#[allow(clippy::too_many_arguments)]
fn naive_attention(
cfg: &MlaConfig,
kv_b_proj: &WeightMatrix,
query: &[f32],
k_pass_c_normed: &[f32],
k_rot: &[f32],
k_cache: &mut Vec<f32>,
v_cache: &mut Vec<f32>,
scale: f32,
) -> Vec<f32> {
let q_head_dim = cfg.qk_nope_head_dim + cfg.qk_rope_head_dim;
let k_pass_full = kv_b_proj.apply(k_pass_c_normed);
let mut key_step = vec![0f32; cfg.num_heads * q_head_dim];
let mut value_step = vec![0f32; cfg.num_heads * cfg.v_head_dim];
let kpf_stride = cfg.qk_nope_head_dim + cfg.v_head_dim;
for h in 0..cfg.num_heads {
let k_pass = &k_pass_full[h * kpf_stride..h * kpf_stride + cfg.qk_nope_head_dim];
let v_h = &k_pass_full[h * kpf_stride + cfg.qk_nope_head_dim..(h + 1) * kpf_stride];
let key_h = &mut key_step[h * q_head_dim..(h + 1) * q_head_dim];
key_h[..cfg.qk_nope_head_dim].copy_from_slice(k_pass);
key_h[cfg.qk_nope_head_dim..].copy_from_slice(k_rot);
value_step[h * cfg.v_head_dim..(h + 1) * cfg.v_head_dim].copy_from_slice(v_h);
}
k_cache.extend_from_slice(&key_step);
v_cache.extend_from_slice(&value_step);
let seq_len = k_cache.len() / (cfg.num_heads * q_head_dim);
causal_mla_attention_scaled(
query,
k_cache,
v_cache,
cfg.num_heads,
q_head_dim,
cfg.v_head_dim,
seq_len,
scale,
)
}
#[cfg(test)]
mod tests;