use crate::math::common::{AlignedVec, SimdMath};
use crate::models::wavenet::PostStackHead;
use crate::models::wavenet::common::WAVENET_MAX_NUM_FRAMES;
use super::block::ConvNetBlock;
#[repr(align(64))]
pub struct ConvNetModel {
pub blocks: Vec<ConvNetBlock>,
pub head_scale: f32,
pub receptive_field_size: usize,
pub post_stack_head: Option<PostStackHead>,
pub head_output_scratch: AlignedVec<f32>,
pub(crate) scratch_a: AlignedVec<f32>,
pub(crate) scratch_b: AlignedVec<f32>,
pub prewarm_on_reset: bool,
pub linear_head: Option<LinearHead>,
}
#[derive(Clone)]
#[repr(align(64))]
pub struct LinearHead {
pub weight: AlignedVec<f32>,
pub bias: AlignedVec<f32>,
pub in_ch: usize,
pub out_ch: usize,
}
impl ConvNetModel {
pub fn in_channels(&self) -> usize {
self.blocks.first().map(|b| b.conv.in_ch).unwrap_or(1)
}
pub fn out_channels(&self) -> usize {
if let Some(ref head) = self.post_stack_head {
head.out_channels()
} else if let Some(ref linear) = self.linear_head {
linear.out_ch
} else {
self.blocks.last().map(|b| b.conv.out_ch).unwrap_or(1)
}
}
pub fn process(&mut self, input: &[f32], output: &mut [f32]) {
unsafe { crate::math::common::dispatch_simd!(self, process_internal, input, output) };
}
#[inline(always)]
unsafe fn process_internal<M: SimdMath>(&mut self, input: &[f32], output: &mut [f32]) {
let total_frames = input.len();
if total_frames == 0 || self.blocks.is_empty() {
output[..total_frames].fill(0.0);
return;
}
let out_ch = self.out_channels();
let mut pos = 0;
while pos < total_frames {
let num_frames = (total_frames - pos).min(WAVENET_MAX_NUM_FRAMES);
let in_slice = &input[pos..pos + num_frames];
let num_blocks = self.blocks.len();
let blocks_ptr = self.blocks.as_mut_ptr();
let first_out_ch = unsafe { (*blocks_ptr).conv.out_ch };
let dst_a = &mut self.scratch_a[..num_frames * first_out_ch];
unsafe {
(*blocks_ptr).process_block_internal::<M>(in_slice, dst_a, num_frames);
}
let mut src_is_a = true;
for i in 1..num_blocks {
let curr = unsafe { &mut *blocks_ptr.add(i) };
let curr_out_ch = curr.conv.out_ch;
if src_is_a {
let src = &self.scratch_a
[..num_frames * unsafe { (*blocks_ptr.add(i - 1)).conv.out_ch }];
let dst = &mut self.scratch_b[..num_frames * curr_out_ch];
unsafe {
curr.process_block_internal::<M>(src, dst, num_frames);
}
} else {
let src = &self.scratch_b
[..num_frames * unsafe { (*blocks_ptr.add(i - 1)).conv.out_ch }];
let dst = &mut self.scratch_a[..num_frames * curr_out_ch];
unsafe {
curr.process_block_internal::<M>(src, dst, num_frames);
}
}
src_is_a = !src_is_a;
}
let last_result_in_a = (num_blocks - 1).is_multiple_of(2);
let last_out_ch = unsafe { (*blocks_ptr.add(num_blocks - 1)).conv.out_ch };
let last_slice = if last_result_in_a {
&self.scratch_a[..num_frames * last_out_ch]
} else {
&self.scratch_b[..num_frames * last_out_ch]
};
if let Some(ref mut head_proc) = self.post_stack_head {
let head_out_ch = head_proc.out_channels();
let head_scratch = &mut self.head_output_scratch[..num_frames * head_out_ch];
unsafe {
head_proc.process_block(last_slice, head_scratch, num_frames);
}
let out_start = pos * out_ch;
let out_slice = &mut output[out_start..out_start + num_frames * out_ch];
out_slice.copy_from_slice(head_scratch);
unsafe {
M::apply_gain(out_slice, self.head_scale);
}
} else if let Some(ref linear) = self.linear_head {
let lh_out_ch = linear.out_ch;
let out_start = pos * out_ch;
let out_slice = &mut output[out_start..out_start + num_frames * lh_out_ch];
out_slice.fill(linear.bias[0]);
for f in 0..num_frames {
let src = &last_slice[f * linear.in_ch..(f + 1) * linear.in_ch];
let dst = &mut out_slice[f * lh_out_ch..(f + 1) * lh_out_ch];
for (o, dst_val) in dst.iter_mut().enumerate().take(lh_out_ch) {
let mut acc = linear.bias[o];
let row_start = o * linear.in_ch;
for (i, &src_val) in src.iter().enumerate().take(linear.in_ch) {
acc += src_val * linear.weight[row_start + i];
}
*dst_val = acc;
}
}
unsafe {
M::apply_gain(out_slice, self.head_scale);
}
} else {
let out_start = pos * out_ch;
let out_slice = &mut output[out_start..out_start + num_frames * out_ch];
out_slice.copy_from_slice(last_slice);
unsafe {
M::apply_gain(out_slice, self.head_scale);
}
}
pos += num_frames;
}
}
#[cold]
pub fn prewarm(&mut self) {
unsafe {
crate::math::common::dispatch_simd!(self, prewarm_internal);
}
}
#[inline(always)]
#[cold]
unsafe fn prewarm_internal<M: SimdMath>(&mut self) {
let num_blocks = self.blocks.len();
if num_blocks == 0 {
return;
}
let blocks_ptr = self.blocks.as_mut_ptr();
for i in 0..num_blocks {
unsafe {
(*blocks_ptr.add(i)).prewarm_internal::<M>();
}
}
if let Some(ref mut head_proc) = self.post_stack_head {
head_proc.prewarm();
}
}
}
#[cfg(test)]
#[path = "convnet_model_test.rs"]
mod tests;