pub mod common;
pub mod conv1d;
pub mod conv1d_dual;
pub mod conv1d_dyn;
pub mod conv1d_dyn_dual;
pub mod conv_input;
pub mod dense;
pub mod dense_dyn;
pub mod layer;
pub mod layer_array;
pub mod layer_array_dyn;
pub mod layer_dyn;
pub mod model;
pub mod model_dyn;
pub mod post_stack_head;
use super::NamModel;
use super::sealed;
impl<const CH: usize, const K: usize, const HEAD: usize> sealed::Sealed
for model::WaveNetModel<CH, K, HEAD>
{
}
impl sealed::Sealed for model_dyn::WaveNetModelDyn {}
impl<const CH: usize, const K: usize, const HEAD: usize> NamModel
for model::WaveNetModel<CH, K, HEAD>
{
fn process(&mut self, input: &[f32], output: &mut [f32]) {
self.process(input, output);
}
fn prewarm(&mut self, _num_samples: usize) {
self.prewarm();
}
fn prewarm_samples(&self) -> usize {
self.array1.receptive_field_size + self.array2.receptive_field_size
}
fn prewarm_on_reset(&self) -> bool {
self.prewarm_on_reset
}
fn set_prewarm_on_reset(&mut self, val: bool) {
self.prewarm_on_reset = val;
}
}
impl NamModel for model_dyn::WaveNetModelDyn {
fn process(&mut self, input: &[f32], output: &mut [f32]) {
self.process(input, output);
}
fn prewarm(&mut self, _num_samples: usize) {
self.prewarm();
}
fn prewarm_samples(&self) -> usize {
let mut rf: usize = self.arrays.iter().map(|a| a.receptive_field_size).sum();
if let Some(ref cond_dsp) = self.condition_dsp {
rf += cond_dsp.prewarm_samples();
}
if let Some(ref head_proc) = self.post_stack_head {
rf += head_proc.receptive_field() - 1;
}
rf
}
fn set_max_buffer_size(&mut self, max_buf: usize) -> anyhow::Result<()> {
if let Some(ref mut cond_dsp) = self.condition_dsp {
cond_dsp.set_max_buffer_size(max_buf)?;
}
Ok(())
}
fn prewarm_on_reset(&self) -> bool {
self.prewarm_on_reset
}
fn set_prewarm_on_reset(&mut self, val: bool) {
self.prewarm_on_reset = val;
if let Some(ref mut cond_dsp) = self.condition_dsp {
cond_dsp.set_prewarm_on_reset(val);
}
}
}
pub use common::{
LAYER_ARRAY_BUFFER_PADDING, MAX_KERNEL, WAVENET_MAX_NUM_FRAMES, WaveNetLayerState,
WavenetProcessContext,
};
pub use conv1d::Conv1d;
pub use conv1d_dyn::Conv1dDyn;
pub use dense::DenseLayer;
pub use dense_dyn::DenseLayerDyn;
pub use layer::WaveNetLayer;
pub use layer_array::WaveNetLayerArray;
pub use layer_array_dyn::WaveNetLayerArrayDyn;
pub use layer_dyn::WaveNetLayerDyn;
pub use model::WaveNetModel;
pub use model_dyn::WaveNetModelDyn;
pub use post_stack_head::PostStackHead;
#[cfg(test)]
mod wavenet_test;