use crate::error::WhisperResult;
use crate::model::lfm2::gqa::{GqaConfig, GroupedQueryAttention};
use crate::model::lfm2::layer::LayerNormNoBias;
use crate::model::lfm2::mlp::{MlpActivation, MlpConfig, MlpFfn};
use crate::model::lfm2::rope::RotaryEmbedding;
#[derive(Debug, Clone)]
pub struct MoonshineEncoderBlock {
pub ln1: LayerNormNoBias,
pub self_attn: GroupedQueryAttention,
pub ln2: LayerNormNoBias,
pub ffn: MlpFfn,
}
impl MoonshineEncoderBlock {
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 gqa_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),
};
let mlp_config = MlpConfig {
hidden_size: d_model,
intermediate_size,
bias: false,
activation: MlpActivation::Gelu,
};
Ok(Self {
ln1: LayerNormNoBias::new(d_model),
self_attn: GroupedQueryAttention::new(gqa_config)?,
ln2: LayerNormNoBias::new(d_model),
ffn: MlpFfn::new(mlp_config)?,
})
}
pub fn forward(
&self,
x: &[f32],
seq_len: usize,
rope: &RotaryEmbedding,
) -> WhisperResult<Vec<f32>> {
let normed = self.ln1.forward(x, seq_len)?;
let attn_out = self
.self_attn
.forward_with_rope(&normed, seq_len, Some(rope))?;
let mut residual = add_vectors(x, &attn_out);
let normed2 = self.ln2.forward(&residual, seq_len)?;
let ffn_out = self.ffn.forward(&normed2, seq_len)?;
add_vectors_inplace(&mut residual, &ffn_out);
Ok(residual)
}
pub fn forward_probed(
&self,
x: &[f32],
seq_len: usize,
rope: &RotaryEmbedding,
block_idx: usize,
probe: &mut crate::probe::ActivationProbe,
) -> WhisperResult<Vec<f32>> {
let d_model = self.ln1.weight.len();
let prefix = format!("encoder.block_{block_idx}");
let normed = self.ln1.forward(x, seq_len)?;
probe.record(&format!("{prefix}.ln1_out"), &normed, &[seq_len, d_model]);
let attn_out = self
.self_attn
.forward_with_rope(&normed, seq_len, Some(rope))?;
probe.record(
&format!("{prefix}.self_attn_out"),
&attn_out,
&[seq_len, d_model],
);
let mut residual = add_vectors(x, &attn_out);
probe.record(
&format!("{prefix}.residual_1"),
&residual,
&[seq_len, d_model],
);
let normed2 = self.ln2.forward(&residual, seq_len)?;
probe.record(&format!("{prefix}.ln2_out"), &normed2, &[seq_len, d_model]);
let ffn_out = self.ffn.forward(&normed2, seq_len)?;
probe.record(&format!("{prefix}.ffn_out"), &ffn_out, &[seq_len, d_model]);
add_vectors_inplace(&mut residual, &ffn_out);
probe.record(
&format!("{prefix}.residual_2"),
&residual,
&[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_encoder_block_new() {
let block = MoonshineEncoderBlock::new(288, 8, 8, 1152);
assert!(block.is_ok());
}
#[test]
fn test_moonshine_encoder_block_forward_shape() {
let block = MoonshineEncoderBlock::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 seq_len = 7; let d_model = 288;
let input = vec![0.1_f32; seq_len * d_model];
let output = block.forward(&input, seq_len, &rope).expect("forward");
assert_eq!(output.len(), seq_len * d_model);
}
#[test]
fn test_moonshine_encoder_block_residual() {
let block = MoonshineEncoderBlock::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 seq_len = 4;
let d_model = 288;
let input = vec![1.0_f32; seq_len * d_model];
let output = block.forward(&input, seq_len, &rope).expect("forward");
assert_eq!(output.len(), input.len());
assert!(output.iter().all(|v| v.is_finite()));
}
#[test]
fn test_add_vectors() {
let a = vec![1.0, 2.0, 3.0];
let b = vec![4.0, 5.0, 6.0];
let c = add_vectors(&a, &b);
assert_eq!(c, vec![5.0, 7.0, 9.0]);
}
}