use crate::error::{WhisperError, WhisperResult};
#[derive(Debug, Clone)]
pub struct Conv1dConfig {
pub channels: usize,
pub kernel_size: usize,
pub causal: bool,
pub bias: bool,
}
impl Conv1dConfig {
#[must_use]
pub fn lfm2_2_6b() -> Self {
Self {
channels: 2048,
kernel_size: 4,
causal: true,
bias: false,
}
}
#[must_use]
pub const fn cache_len(&self) -> usize {
if self.causal {
self.kernel_size - 1
} else {
0
}
}
pub fn validate(&self) -> WhisperResult<()> {
if self.channels == 0 {
return Err(WhisperError::Model("channels must be > 0".into()));
}
if self.kernel_size == 0 {
return Err(WhisperError::Model("kernel_size must be > 0".into()));
}
Ok(())
}
}
#[derive(Debug)]
pub struct Conv1d {
pub config: Conv1dConfig,
pub weight: Vec<f32>,
pub bias: Option<Vec<f32>>,
pub depthwise: bool,
}
impl Conv1d {
pub fn new_depthwise(config: Conv1dConfig) -> WhisperResult<Self> {
config.validate()?;
let weight_size = config.channels * config.kernel_size;
let bias = if config.bias {
Some(vec![0.0; config.channels])
} else {
None
};
Ok(Self {
config,
weight: vec![0.0; weight_size],
bias,
depthwise: true,
})
}
pub fn new_standard(config: Conv1dConfig) -> WhisperResult<Self> {
config.validate()?;
let weight_size = config.channels * config.channels * config.kernel_size;
let bias = if config.bias {
Some(vec![0.0; config.channels])
} else {
None
};
Ok(Self {
config,
weight: vec![0.0; weight_size],
bias,
depthwise: false,
})
}
pub fn forward(
&self,
input: &[f32],
seq_len: usize,
cache: Option<&[f32]>,
) -> WhisperResult<Vec<f32>> {
let c = self.config.channels;
if input.len() != seq_len * c {
return Err(WhisperError::Model(format!(
"input length {} != seq_len * channels ({})",
input.len(),
seq_len * c
)));
}
if self.depthwise {
self.forward_depthwise(input, seq_len, cache)
} else {
self.forward_standard(input, seq_len, cache)
}
}
fn forward_depthwise(
&self,
input: &[f32],
seq_len: usize,
cache: Option<&[f32]>,
) -> WhisperResult<Vec<f32>> {
let c = self.config.channels;
let k = self.config.kernel_size;
let cache_len = self.config.cache_len();
let padded = if self.config.causal {
let mut p = vec![0.0f32; (cache_len + seq_len) * c];
if let Some(cache_data) = cache {
if cache_data.len() >= cache_len * c {
p[..cache_len * c].copy_from_slice(&cache_data[..cache_len * c]);
}
}
p[cache_len * c..].copy_from_slice(input);
p
} else {
input.to_vec()
};
let padded_len = if self.config.causal {
cache_len + seq_len
} else {
seq_len
};
let mut output = vec![0.0f32; seq_len * c];
for ch in 0..c {
for t in 0..seq_len {
let out_idx = t * c + ch;
let mut sum = 0.0f32;
for ki in 0..k {
let t_in = t + ki;
if t_in < padded_len {
let in_idx = t_in * c + ch;
let w_idx = ch * k + ki;
sum += padded[in_idx] * self.weight[w_idx];
}
}
if let Some(ref b) = self.bias {
sum += b[ch];
}
output[out_idx] = sum;
}
}
Ok(output)
}
fn forward_standard(
&self,
input: &[f32],
seq_len: usize,
cache: Option<&[f32]>,
) -> WhisperResult<Vec<f32>> {
let c = self.config.channels;
let k = self.config.kernel_size;
let cache_len = self.config.cache_len();
let padded = if self.config.causal {
let mut p = vec![0.0f32; (cache_len + seq_len) * c];
if let Some(cache_data) = cache {
if cache_data.len() >= cache_len * c {
p[..cache_len * c].copy_from_slice(&cache_data[..cache_len * c]);
}
}
p[cache_len * c..].copy_from_slice(input);
p
} else {
input.to_vec()
};
let padded_len = if self.config.causal {
cache_len + seq_len
} else {
seq_len
};
let mut output = vec![0.0f32; seq_len * c];
for out_ch in 0..c {
for t in 0..seq_len {
let out_idx = t * c + out_ch;
let mut sum = 0.0f32;
for in_ch in 0..c {
for ki in 0..k {
let t_in = t + ki;
if t_in < padded_len {
let in_idx = t_in * c + in_ch;
let w_idx = (out_ch * c + in_ch) * k + ki;
sum += padded[in_idx] * self.weight[w_idx];
}
}
}
if let Some(ref b) = self.bias {
sum += b[out_ch];
}
output[out_idx] = sum;
}
}
Ok(output)
}
#[must_use]
pub fn get_new_cache(&self, input: &[f32], seq_len: usize) -> Vec<f32> {
let c = self.config.channels;
let cache_len = self.config.cache_len();
if seq_len <= cache_len {
let mut cache = vec![0.0f32; cache_len * c];
let start = (cache_len - seq_len) * c;
cache[start..].copy_from_slice(input);
cache
} else {
let start = (seq_len - cache_len) * c;
input[start..].to_vec()
}
}
#[must_use]
pub fn num_params(&self) -> usize {
let weight_params = self.weight.len();
let bias_params = self.bias.as_ref().map_or(0, Vec::len);
weight_params + bias_params
}
#[must_use]
pub fn memory_bytes(&self) -> usize {
self.num_params() * std::mem::size_of::<f32>()
}
}
#[derive(Debug, Clone)]
pub struct ConvCache {
pub data: Vec<f32>,
pub cache_len: usize,
pub channels: usize,
}
impl ConvCache {
#[must_use]
pub fn new(cache_len: usize, channels: usize) -> Self {
Self {
data: vec![0.0; cache_len * channels],
cache_len,
channels,
}
}
pub fn update(&mut self, input: &[f32], seq_len: usize) {
if seq_len >= self.cache_len {
let start = (seq_len - self.cache_len) * self.channels;
self.data.copy_from_slice(&input[start..]);
} else {
let shift = seq_len * self.channels;
let keep = self.data.len() - shift;
self.data.copy_within(shift.., 0);
self.data[keep..].copy_from_slice(input);
}
}
pub fn reset(&mut self) {
self.data.fill(0.0);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_conv1d_config_lfm2() {
let config = Conv1dConfig::lfm2_2_6b();
assert_eq!(config.channels, 2048);
assert_eq!(config.kernel_size, 4);
assert!(config.causal);
assert!(!config.bias);
assert_eq!(config.cache_len(), 3);
assert!(config.validate().is_ok());
}
#[test]
fn test_conv1d_depthwise_new() {
let config = Conv1dConfig {
channels: 16,
kernel_size: 3,
causal: true,
bias: false,
};
let conv = Conv1d::new_depthwise(config).expect("should create conv");
assert!(conv.depthwise);
assert_eq!(conv.weight.len(), 16 * 3);
assert!(conv.bias.is_none());
}
#[test]
fn test_conv1d_standard_new() {
let config = Conv1dConfig {
channels: 8,
kernel_size: 3,
causal: true,
bias: true,
};
let conv = Conv1d::new_standard(config).expect("should create conv");
assert!(!conv.depthwise);
assert_eq!(conv.weight.len(), 8 * 8 * 3);
assert!(conv.bias.is_some());
assert_eq!(conv.bias.as_ref().map(|b| b.len()), Some(8));
}
#[test]
fn test_conv1d_forward_shape() {
let config = Conv1dConfig {
channels: 4,
kernel_size: 3,
causal: true,
bias: false,
};
let conv = Conv1d::new_depthwise(config).expect("should create conv");
let seq_len = 5;
let input = vec![1.0f32; seq_len * 4];
let output = conv
.forward(&input, seq_len, None)
.expect("forward should succeed");
assert_eq!(output.len(), seq_len * 4);
}
#[test]
fn test_conv1d_depthwise_forward() {
let config = Conv1dConfig {
channels: 2,
kernel_size: 2,
causal: true,
bias: false,
};
let mut conv = Conv1d::new_depthwise(config).expect("should create conv");
conv.weight = vec![1.0, 0.0, 0.0, 1.0];
let seq_len = 3;
let input = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let output = conv
.forward(&input, seq_len, None)
.expect("forward should succeed");
assert_eq!(output.len(), 6);
}
#[test]
fn test_conv1d_with_cache() {
let config = Conv1dConfig {
channels: 2,
kernel_size: 3,
causal: true,
bias: false,
};
let mut conv = Conv1d::new_depthwise(config).expect("should create conv");
for w in &mut conv.weight {
*w = 0.5;
}
let input1 = vec![1.0, 1.0, 2.0, 2.0, 3.0, 3.0]; let _output1 = conv
.forward(&input1, 3, None)
.expect("first forward should succeed");
let cache = conv.get_new_cache(&input1, 3);
assert_eq!(cache.len(), 2 * 2);
let input2 = vec![4.0, 4.0, 5.0, 5.0]; let output2 = conv
.forward(&input2, 2, Some(&cache))
.expect("second forward should succeed");
assert_eq!(output2.len(), 4);
}
#[test]
fn test_conv_cache() {
let mut cache = ConvCache::new(3, 2);
assert_eq!(cache.data, vec![0.0; 6]);
let input1 = vec![1.0, 1.0, 2.0, 2.0, 3.0, 3.0, 4.0, 4.0]; cache.update(&input1, 4);
assert_eq!(cache.data, vec![2.0, 2.0, 3.0, 3.0, 4.0, 4.0]);
let input2 = vec![5.0, 5.0, 6.0, 6.0]; cache.update(&input2, 2);
assert_eq!(cache.data, vec![4.0, 4.0, 5.0, 5.0, 6.0, 6.0]);
}
#[test]
fn test_conv_cache_reset() {
let mut cache = ConvCache::new(2, 4);
let input = vec![1.0; 8];
cache.update(&input, 2);
cache.reset();
assert!(cache.data.iter().all(|&x| x == 0.0));
}
#[test]
fn test_conv1d_num_params() {
let config = Conv1dConfig {
channels: 16,
kernel_size: 4,
causal: true,
bias: true,
};
let depthwise = Conv1d::new_depthwise(config.clone()).expect("should create");
assert_eq!(depthwise.num_params(), 16 * 4 + 16);
let standard = Conv1d::new_standard(config).expect("should create");
assert_eq!(standard.num_params(), 16 * 16 * 4 + 16);
}
}