use frink_core::cache::KvCache;
use frink_core::recurrent_state::RecurrentState;
use super::attn_block::KvStep;
use super::{Decoder, LayerWeights};
use crate::layer_shapes::AttnShape;
use crate::shortconv::window_from_history;
impl Decoder {
pub(crate) fn recurrent_block(
&self,
layer_idx: usize,
layer: &LayerWeights,
normed: &[f32],
rows: usize,
kv: KvStep<'_>,
) -> Vec<f32> {
let mut out = match self.config.layer_shape(layer_idx).attention {
AttnShape::ShortConv => self.shortconv_block(layer_idx, layer, normed, rows, kv),
AttnShape::Mamba1
| AttnShape::Mamba2
| AttnShape::Plamo2Ssm
| AttnShape::Gdn
| AttnShape::Lightning => self.ssm_block(layer_idx, layer, normed, rows, kv),
other => unreachable!("layer {layer_idx} is {other:?}, not a recurrent block"),
};
if let Some(post) = &layer.attn.post_attn_norm {
let hidden = post.len();
out = out
.chunks(hidden)
.flat_map(|row| {
frink_core::matmul::rms_norm(row, post, self.config.post_norm_eps())
})
.collect();
}
out
}
fn shortconv_block(
&self,
layer_idx: usize,
layer: &LayerWeights,
normed: &[f32],
rows: usize,
kv: KvStep<'_>,
) -> Vec<f32> {
let conv =
layer.attn.shortconv.as_ref().unwrap_or_else(|| {
panic!("layer {layer_idx} is ShortConv-shaped but has no weights")
});
let (l_cache, n_embd) = (conv.l_cache, conv.hidden_dim());
match kv {
KvStep::Decode(cache) | KvStep::Batched(cache) => {
conv.forward_rows(normed, rows, |bx| {
contiguous_step(cache, bx, l_cache, n_embd)
})
}
KvStep::Paged { cache, stores } => conv.forward_rows(normed, rows, |bx| {
{
let mut store = stores.write(layer_idx);
cache
.push(&mut store, bx, &[])
.expect("every caller reserves this row's pages before the stack runs");
}
let store = stores.read(layer_idx);
let table = cache.block_table();
let block = store.block_size();
window_from_history(l_cache, n_embd, cache.seq_len(), |i| {
store.k_row(table[i / block], i % block)
})
}),
}
}
fn ssm_block(
&self,
layer_idx: usize,
layer: &LayerWeights,
normed: &[f32],
rows: usize,
mut kv: KvStep<'_>,
) -> Vec<f32> {
let out = self.ssm_state_step(layer_idx, layer, normed, rows, kv.recurrent_slot());
match kv {
KvStep::Decode(cache) | KvStep::Batched(cache) => {
cache
.advance_len(rows)
.expect("unbounded/planned KvCache growth is infallible");
}
KvStep::Paged { cache, stores } => {
let mut store = stores.write(layer_idx);
for _ in 0..rows {
cache
.push(&mut store, &[], &[])
.expect("every caller reserves this row's pages before the stack runs");
}
}
}
out
}
pub(crate) fn ssm_state_step(
&self,
layer_idx: usize,
layer: &LayerWeights,
normed: &[f32],
rows: usize,
slot: &mut Option<RecurrentState>,
) -> Vec<f32> {
let block =
layer.attn.ssm.as_ref().unwrap_or_else(|| {
panic!("layer {layer_idx} runs a Mamba-2 block but has no weights")
});
let state = slot.get_or_insert_with(|| block.zero_state());
block.forward_rows(normed, rows, state, self.config.rms_norm_eps)
}
pub(crate) fn parallel_ssm_rows(
&self,
layer_idx: usize,
layer: &LayerWeights,
normed: &[f32],
rows: usize,
slot: &mut Option<RecurrentState>,
) -> Option<Vec<f32>> {
if layer.attn.ssm.is_none() || self.config.layer_shape(layer_idx).attention.is_recurrent() {
return None;
}
Some(self.ssm_state_step(layer_idx, layer, normed, rows, slot))
}
pub(crate) fn add_parallel_ssm(projected: &mut [f32], ssm: Option<Vec<f32>>) {
if let Some(ssm) = ssm {
assert_eq!(ssm.len(), projected.len());
for (p, s) in projected.iter_mut().zip(&ssm) {
*p += s;
}
}
}
}
fn contiguous_step(cache: &mut KvCache, bx: &[f32], l_cache: usize, n_embd: usize) -> Vec<f32> {
cache
.push(bx, &[])
.expect("unbounded/planned KvCache growth is infallible");
let rows = cache.rows();
window_from_history(l_cache, n_embd, rows, |i| {
&cache.k[i * n_embd..(i + 1) * n_embd]
})
}