use crate::error::WhisperResult;
use crate::model::lfm2::gqa::{GqaConfig, GroupedQueryAttention};
use crate::model::lfm2::layer::LayerNormNoBias;
use crate::model::lfm2::mlp::GatedMlpFfn;
use crate::model::lfm2::rope::RotaryEmbedding;
use crate::model::LayerKVCache;
#[derive(Debug, Clone)]
pub struct MoonshineDecoderBlock {
pub ln1: LayerNormNoBias,
pub self_attn: GroupedQueryAttention,
pub ln_cross: LayerNormNoBias,
pub cross_attn: GroupedQueryAttention,
pub ln2: LayerNormNoBias,
pub ffn: GatedMlpFfn,
}
impl MoonshineDecoderBlock {
pub fn new(
d_model: usize,
n_q_heads: usize,
n_kv_heads: usize,
intermediate_size: usize,
) -> WhisperResult<Self> {
let head_dim = d_model / n_q_heads;
let self_attn_config = GqaConfig {
hidden_size: d_model,
num_q_heads: n_q_heads,
num_kv_heads: n_kv_heads,
head_dim,
causal: true,
dropout: 0.0,
pad_head_dim_to: Some(8),
};
let cross_attn_config = GqaConfig {
hidden_size: d_model,
num_q_heads: n_q_heads,
num_kv_heads: n_kv_heads,
head_dim,
causal: false,
dropout: 0.0,
pad_head_dim_to: Some(8),
};
Ok(Self {
ln1: LayerNormNoBias::new(d_model),
self_attn: GroupedQueryAttention::new(self_attn_config)?,
ln_cross: LayerNormNoBias::new(d_model),
cross_attn: GroupedQueryAttention::new(cross_attn_config)?,
ln2: LayerNormNoBias::new(d_model),
ffn: GatedMlpFfn::new(d_model, intermediate_size)?,
})
}
#[allow(clippy::too_many_arguments)]
pub fn forward_cached(
&self,
x: &[f32],
encoder_out: &[f32],
enc_seq_len: usize,
position: usize,
rope: &RotaryEmbedding,
self_attn_cache: &mut LayerKVCache,
cross_attn_cache: &mut LayerKVCache,
cross_attn_cached: bool,
) -> WhisperResult<Vec<f32>> {
let d_model = self.self_attn.config.hidden_size;
let normed = self.ln1.forward(x, 1)?;
let (q, k_new, v_new) = self
.self_attn
.project_qkv_single(&normed, Some(rope), position)?;
self_attn_cache.append(&k_new, &v_new)?;
let k_full = self_attn_cache.get_key();
let v_full = self_attn_cache.get_value();
let cache_len = self_attn_cache.len();
let attn_out = self
.self_attn
.attention_cached(&q, k_full, v_full, cache_len)?;
let attn_out = self.self_attn.output_projection(&attn_out);
let mut residual: Vec<f32> = x.iter().zip(attn_out.iter()).map(|(a, b)| a + b).collect();
let normed_cross = self.ln_cross.forward(&residual, 1)?;
let cross_out = if !cross_attn_cached || cross_attn_cache.is_empty() {
let (k_enc, v_enc) = self.cross_attn.project_kv(encoder_out, enc_seq_len);
cross_attn_cache.append(&k_enc, &v_enc)?;
let q_cross = self.cross_attn.project_q(&normed_cross);
let attn_out =
self.cross_attn
.attention_cached(&q_cross, &k_enc, &v_enc, enc_seq_len)?;
self.cross_attn.output_projection(&attn_out)
} else {
let k_cached = cross_attn_cache.get_key();
let v_cached = cross_attn_cache.get_value();
let cached_enc_len = cross_attn_cache.len();
let q_cross = self.cross_attn.project_q(&normed_cross);
let attn_out =
self.cross_attn
.attention_cached(&q_cross, k_cached, v_cached, cached_enc_len)?;
self.cross_attn.output_projection(&attn_out)
};
add_vectors_inplace(&mut residual, &cross_out);
let normed2 = self.ln2.forward(&residual, 1)?;
let ffn_out = self.ffn.forward(&normed2, 1)?;
add_vectors_inplace(&mut residual, &ffn_out);
debug_assert_eq!(residual.len(), d_model);
Ok(residual)
}
pub fn forward(
&self,
x: &[f32],
encoder_out: &[f32],
dec_seq_len: usize,
enc_seq_len: usize,
rope: &RotaryEmbedding,
) -> WhisperResult<Vec<f32>> {
let normed = self.ln1.forward(x, dec_seq_len)?;
let self_attn_out = self
.self_attn
.forward_with_rope(&normed, dec_seq_len, Some(rope))?;
let mut residual = add_vectors(x, &self_attn_out);
let normed_cross = self.ln_cross.forward(&residual, dec_seq_len)?;
let cross_attn_out = self.cross_attn.forward_cross_attention(
&normed_cross,
encoder_out,
dec_seq_len,
enc_seq_len,
)?;
add_vectors_inplace(&mut residual, &cross_attn_out);
let normed2 = self.ln2.forward(&residual, dec_seq_len)?;
let ffn_out = self.ffn.forward(&normed2, dec_seq_len)?;
add_vectors_inplace(&mut residual, &ffn_out);
Ok(residual)
}
#[allow(clippy::too_many_arguments)]
pub fn forward_probed(
&self,
x: &[f32],
encoder_out: &[f32],
dec_seq_len: usize,
enc_seq_len: usize,
rope: &RotaryEmbedding,
block_idx: usize,
probe: &mut crate::probe::ActivationProbe,
) -> WhisperResult<Vec<f32>> {
let d_model = self.self_attn.config.hidden_size;
let prefix = format!("decoder.block_{block_idx}");
let normed = self.ln1.forward(x, dec_seq_len)?;
probe.record(
&format!("{prefix}.ln1_out"),
&normed,
&[dec_seq_len, d_model],
);
let self_attn_out = self
.self_attn
.forward_with_rope(&normed, dec_seq_len, Some(rope))?;
probe.record(
&format!("{prefix}.self_attn_out"),
&self_attn_out,
&[dec_seq_len, d_model],
);
let mut residual = add_vectors(x, &self_attn_out);
probe.record(
&format!("{prefix}.residual_1"),
&residual,
&[dec_seq_len, d_model],
);
let normed_cross = self.ln_cross.forward(&residual, dec_seq_len)?;
probe.record(
&format!("{prefix}.ln_cross_out"),
&normed_cross,
&[dec_seq_len, d_model],
);
let cross_attn_out = self.cross_attn.forward_cross_attention(
&normed_cross,
encoder_out,
dec_seq_len,
enc_seq_len,
)?;
probe.record(
&format!("{prefix}.cross_attn_out"),
&cross_attn_out,
&[dec_seq_len, d_model],
);
add_vectors_inplace(&mut residual, &cross_attn_out);
probe.record(
&format!("{prefix}.residual_2"),
&residual,
&[dec_seq_len, d_model],
);
let normed2 = self.ln2.forward(&residual, dec_seq_len)?;
probe.record(
&format!("{prefix}.ln2_out"),
&normed2,
&[dec_seq_len, d_model],
);
let ffn_out = self.ffn.forward(&normed2, dec_seq_len)?;
probe.record(
&format!("{prefix}.ffn_out"),
&ffn_out,
&[dec_seq_len, d_model],
);
add_vectors_inplace(&mut residual, &ffn_out);
probe.record(
&format!("{prefix}.residual_3"),
&residual,
&[dec_seq_len, d_model],
);
Ok(residual)
}
}
fn add_vectors(a: &[f32], b: &[f32]) -> Vec<f32> {
a.iter().zip(b.iter()).map(|(x, y)| x + y).collect()
}
fn add_vectors_inplace(a: &mut [f32], b: &[f32]) {
for (x, y) in a.iter_mut().zip(b.iter()) {
*x += y;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_moonshine_decoder_block_new() {
let block = MoonshineDecoderBlock::new(288, 8, 8, 1152);
assert!(block.is_ok());
}
#[test]
fn test_moonshine_decoder_block_forward_shape() {
let block = MoonshineDecoderBlock::new(288, 8, 8, 1152).expect("block creation");
let rope = RotaryEmbedding::new(crate::model::lfm2::rope::RopeConfig {
head_dim: 40, base: 10000.0,
max_seq_len: 2048,
rotary_dim: Some(32),
})
.expect("rope creation");
let d_model = 288;
let dec_seq_len = 3;
let enc_seq_len = 7;
let decoder_input = vec![0.1_f32; dec_seq_len * d_model];
let encoder_output = vec![0.2_f32; enc_seq_len * d_model];
let output = block
.forward(
&decoder_input,
&encoder_output,
dec_seq_len,
enc_seq_len,
&rope,
)
.expect("forward");
assert_eq!(output.len(), dec_seq_len * d_model);
}
#[test]
fn test_moonshine_decoder_block_forward_cached_shape() {
let block = MoonshineDecoderBlock::new(288, 8, 8, 1152).expect("block creation");
let rope = RotaryEmbedding::new(crate::model::lfm2::rope::RopeConfig {
head_dim: 40, base: 10000.0,
max_seq_len: 2048,
rotary_dim: Some(32),
})
.expect("rope creation");
let d_model = 288;
let padded_kv_dim = 8 * 40;
let enc_seq_len = 7;
let max_tokens = 100;
let encoder_output = vec![0.2_f32; enc_seq_len * d_model];
let mut self_cache = LayerKVCache::new(padded_kv_dim, max_tokens);
let mut cross_cache = LayerKVCache::new(padded_kv_dim, max_tokens);
for pos in 0..3 {
let x = vec![0.1_f32; d_model];
let out = block
.forward_cached(
&x,
&encoder_output,
enc_seq_len,
pos,
&rope,
&mut self_cache,
&mut cross_cache,
pos > 0,
)
.expect("forward_cached");
assert_eq!(out.len(), d_model);
assert!(out.iter().all(|v| v.is_finite()));
}
assert_eq!(self_cache.len(), 3);
assert_eq!(cross_cache.len(), enc_seq_len);
}
#[test]
fn test_moonshine_decoder_block_finite_output() {
let block = MoonshineDecoderBlock::new(288, 8, 8, 1152).expect("block creation");
let rope = RotaryEmbedding::new(crate::model::lfm2::rope::RopeConfig {
head_dim: 40,
base: 10000.0,
max_seq_len: 2048,
rotary_dim: Some(32),
})
.expect("rope creation");
let d_model = 288;
let dec_seq_len = 1;
let enc_seq_len = 5;
let decoder_input = vec![1.0_f32; dec_seq_len * d_model];
let encoder_output = vec![0.5_f32; enc_seq_len * d_model];
let output = block
.forward(
&decoder_input,
&encoder_output,
dec_seq_len,
enc_seq_len,
&rope,
)
.expect("forward");
assert!(output.iter().all(|v| v.is_finite()));
}
}