use crate::error::{WhisperError, WhisperResult};
use crate::model::Conv1d;
#[derive(Debug, Clone)]
pub struct GroupNorm {
pub weight: Vec<f32>,
pub bias: Vec<f32>,
pub num_groups: usize,
pub num_channels: usize,
pub eps: f32,
}
impl GroupNorm {
#[must_use]
pub fn new(num_groups: usize, num_channels: usize) -> Self {
Self {
weight: vec![1.0; num_channels],
bias: vec![0.0; num_channels],
num_groups,
num_channels,
eps: 1e-5,
}
}
pub fn forward(&self, input: &[f32], seq_len: usize) -> WhisperResult<Vec<f32>> {
let expected = seq_len * self.num_channels;
if input.len() != expected {
return Err(WhisperError::Model(format!(
"GroupNorm input length {} != seq_len({}) * channels({})",
input.len(),
seq_len,
self.num_channels
)));
}
let channels_per_group = self.num_channels / self.num_groups;
if seq_len == 0 {
return Ok(Vec::new());
}
let mut output = vec![0.0_f32; input.len()];
for g in 0..self.num_groups {
let ch_start = g * channels_per_group;
let ch_end = ch_start + channels_per_group;
let count = (seq_len * channels_per_group) as f32;
let mut sum = 0.0_f32;
for s in 0..seq_len {
let row_start = s * self.num_channels;
for c in ch_start..ch_end {
sum += input[row_start + c];
}
}
let mean = sum / count;
let mut var_sum = 0.0_f32;
for s in 0..seq_len {
let row_start = s * self.num_channels;
for c in ch_start..ch_end {
let diff = input[row_start + c] - mean;
var_sum += diff * diff;
}
}
let variance = var_sum / count;
let inv_std = 1.0 / (variance + self.eps).sqrt();
for s in 0..seq_len {
let row_start = s * self.num_channels;
for c in ch_start..ch_end {
let normalized = (input[row_start + c] - mean) * inv_std;
output[row_start + c] = normalized * self.weight[c] + self.bias[c];
}
}
}
Ok(output)
}
}
#[derive(Debug, Clone)]
pub struct ConvStem {
pub conv1: Conv1d,
pub groupnorm: GroupNorm,
pub conv2: Conv1d,
pub conv3: Conv1d,
pub intermediate_channels: usize,
pub d_model: usize,
}
pub const CONV_STEM_TOTAL_STRIDE: usize = 384;
impl ConvStem {
#[must_use]
pub fn new(d_model: usize) -> Self {
Self {
conv1: Conv1d::new(1, d_model, 127, 64, 0),
groupnorm: GroupNorm::new(1, d_model),
conv2: Conv1d::new(d_model, 2 * d_model, 7, 3, 0),
conv3: Conv1d::new(2 * d_model, d_model, 3, 2, 0),
intermediate_channels: 2 * d_model,
d_model,
}
}
pub fn forward(&self, audio: &[f32]) -> WhisperResult<Vec<f32>> {
if audio.is_empty() {
return Err(WhisperError::Audio("empty audio input".into()));
}
let x = self.conv1.forward(audio)?;
let x = crate::simd::tanh_activation(&x);
let seq_len = x.len() / self.d_model;
let x = self.groupnorm.forward(&x, seq_len)?;
let x = self.conv2.forward(&x)?;
let x = crate::simd::gelu(&x);
let x = self.conv3.forward(&x)?;
let x = crate::simd::gelu(&x);
Ok(x)
}
pub fn forward_probed(
&self,
audio: &[f32],
probe: &mut crate::probe::ActivationProbe,
) -> crate::error::WhisperResult<Vec<f32>> {
if audio.is_empty() {
return Err(WhisperError::Audio("empty audio input".into()));
}
let x = self.conv1.forward(audio)?;
let seq_len = x.len() / self.d_model;
probe.record("conv_stem.conv1_out", &x, &[seq_len, self.d_model]);
let x = crate::simd::tanh_activation(&x);
probe.record("conv_stem.tanh_out", &x, &[seq_len, self.d_model]);
let x = self.groupnorm.forward(&x, seq_len)?;
probe.record("conv_stem.groupnorm_out", &x, &[seq_len, self.d_model]);
let x = self.conv2.forward(&x)?;
let seq_len2 = x.len() / self.intermediate_channels;
probe.record(
"conv_stem.conv2_out",
&x,
&[seq_len2, self.intermediate_channels],
);
let x = crate::simd::gelu(&x);
let x = self.conv3.forward(&x)?;
let seq_len3 = x.len() / self.d_model;
probe.record("conv_stem.conv3_out", &x, &[seq_len3, self.d_model]);
let x = crate::simd::gelu(&x);
probe.record("conv_stem.gelu3_out", &x, &[seq_len3, self.d_model]);
Ok(x)
}
#[must_use]
pub fn output_frames(audio_samples: usize) -> usize {
if audio_samples == 0 {
return 0;
}
let after_conv1 = if audio_samples >= 127 {
(audio_samples - 127) / 64 + 1
} else {
0
};
let after_conv2 = if after_conv1 >= 7 {
(after_conv1 - 7) / 3 + 1
} else {
0
};
if after_conv2 >= 3 {
(after_conv2 - 3) / 2 + 1
} else {
0
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_conv_stem_new() {
let stem = ConvStem::new(288);
assert_eq!(stem.conv1.in_channels, 1);
assert_eq!(stem.conv1.out_channels, 288);
assert_eq!(stem.conv1.kernel_size, 127);
assert_eq!(stem.conv1.stride, 64);
assert_eq!(stem.conv2.in_channels, 288);
assert_eq!(stem.conv2.out_channels, 576);
assert_eq!(stem.conv2.kernel_size, 7);
assert_eq!(stem.conv2.stride, 3);
assert_eq!(stem.conv3.in_channels, 576);
assert_eq!(stem.conv3.out_channels, 288);
assert_eq!(stem.conv3.kernel_size, 3);
assert_eq!(stem.conv3.stride, 2);
assert_eq!(stem.d_model, 288);
assert_eq!(stem.groupnorm.num_channels, 288);
}
#[test]
fn test_conv_stem_output_frames() {
let frames = ConvStem::output_frames(24_000);
assert!(
frames > 0,
"1.5s audio should produce >0 frames, got {frames}"
);
assert!(
frames > 30 && frames < 100,
"1.5s audio should produce 30-100 frames, got {frames}"
);
let frames_30s = ConvStem::output_frames(480_000);
assert!(
frames_30s > 500,
"30s audio should produce >500 frames, got {frames_30s}"
);
let ratio = frames_30s as f32 / frames as f32;
assert!(
(ratio - 20.0).abs() < 5.0,
"30s/1.5s frame ratio should be ~20x, got {ratio:.1}x"
);
}
#[test]
fn test_conv_stem_output_frames_empty() {
assert_eq!(ConvStem::output_frames(0), 0);
}
#[test]
fn test_conv_stem_output_frames_short() {
assert_eq!(ConvStem::output_frames(0), 0);
assert_eq!(ConvStem::output_frames(126), 0);
}
#[test]
fn test_conv_stem_forward_produces_output() {
let stem = ConvStem::new(8);
let audio = vec![0.0_f32; 32_000];
let output = stem.forward(&audio).expect("forward should succeed");
let expected_frames = ConvStem::output_frames(32_000);
assert_eq!(
output.len(),
expected_frames * 8,
"output should be frames × d_model"
);
}
#[test]
fn test_conv_stem_forward_empty_errors() {
let stem = ConvStem::new(8);
let result = stem.forward(&[]);
assert!(result.is_err());
}
#[test]
fn test_conv_stem_total_stride() {
assert_eq!(CONV_STEM_TOTAL_STRIDE, 64 * 3 * 2);
}
#[test]
fn test_output_frames_boundary_conditions() {
assert_eq!(ConvStem::output_frames(0), 0);
assert_eq!(ConvStem::output_frames(126), 0);
assert_eq!(ConvStem::output_frames(127), 0);
assert_eq!(ConvStem::output_frames(128), 0);
let first_nonzero = ConvStem::output_frames(895);
assert_eq!(
first_nonzero, 1,
"895 samples should produce exactly 1 frame"
);
assert_eq!(
ConvStem::output_frames(894),
0,
"894 samples should still produce 0 frames"
);
}
#[test]
fn test_output_frames_monotonicity() {
let mut prev = ConvStem::output_frames(0);
for n in 1..=100_000 {
let curr = ConvStem::output_frames(n);
assert!(
curr >= prev,
"Monotonicity violated: output_frames({}) = {} < output_frames({}) = {}",
n,
curr,
n - 1,
prev
);
prev = curr;
}
}
#[test]
fn test_output_frames_known_durations() {
assert_eq!(ConvStem::output_frames(16_000), 40);
assert_eq!(ConvStem::output_frames(24_000), 61);
assert_eq!(ConvStem::output_frames(48_000), 123);
assert_eq!(ConvStem::output_frames(160_000), 415);
assert_eq!(ConvStem::output_frames(480_000), 1248);
}
#[test]
fn test_groupnorm_new() {
let gn = GroupNorm::new(4, 16);
assert_eq!(gn.num_groups, 4);
assert_eq!(gn.num_channels, 16);
assert_eq!(gn.weight.len(), 16);
assert_eq!(gn.bias.len(), 16);
assert!((gn.weight[0] - 1.0).abs() < f32::EPSILON);
assert!((gn.bias[0]).abs() < f32::EPSILON);
}
#[test]
fn test_groupnorm_forward_identity() {
let gn = GroupNorm::new(1, 4);
let input = vec![1.0, 2.0, 3.0, 4.0];
let output = gn.forward(&input, 1).expect("forward should succeed");
assert_eq!(output.len(), 4);
let mean: f32 = output.iter().sum::<f32>() / 4.0;
assert!(mean.abs() < 1e-5, "output mean should be ~0, got {mean}");
}
#[test]
fn test_groupnorm_forward_multi_position() {
let gn = GroupNorm::new(2, 4);
let input = vec![
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, ];
let output = gn.forward(&input, 2).expect("forward should succeed");
assert_eq!(output.len(), 8);
assert!(output.iter().all(|x| x.is_finite()));
}
#[test]
fn test_groupnorm_forward_wrong_size() {
let gn = GroupNorm::new(1, 4);
let input = vec![1.0, 2.0, 3.0]; let result = gn.forward(&input, 1);
assert!(result.is_err());
}
}