use crate::dsp::mirror_buf::MirroredBuffer;
pub const WAVENET_MAX_NUM_FRAMES: usize = 64;
pub const LAYER_ARRAY_BUFFER_PADDING: usize = 24;
pub const MAX_KERNEL: usize = 16;
pub struct WavenetProcessContext<'a> {
pub condition: &'a [f32],
pub head_input: &'a mut [f32],
pub output: &'a mut [f32],
pub layer_buffer: &'a [f32],
pub buffer_start: usize,
pub num_frames: usize,
pub block: &'a mut [f32],
pub is_first_layer: bool,
pub seed: Option<&'a [f32]>,
}
#[repr(align(64))]
#[derive(Clone)]
pub struct WaveNetLayerState {
pub layer_buffer: MirroredBuffer<f32>,
pub buffer_start: usize,
pub receptive_field_size: usize,
}
impl WaveNetLayerState {
pub fn new(
channels: usize,
receptive_field_size: usize,
alloc_num: usize,
) -> std::io::Result<Self> {
let min_buffer_frames =
receptive_field_size + (LAYER_ARRAY_BUFFER_PADDING + 1) * WAVENET_MAX_NUM_FRAMES;
let buffer = MirroredBuffer::<f32>::new_aligned(min_buffer_frames * channels, channels)?;
let actual_buffer_frames = buffer.size() / channels;
let jitter = (alloc_num % LAYER_ARRAY_BUFFER_PADDING) + 1;
let start = actual_buffer_frames * 2 - (WAVENET_MAX_NUM_FRAMES * jitter);
if start < receptive_field_size {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!(
"buffer_start ({}) is smaller than receptive_field_size ({}); \
increase LAYER_ARRAY_BUFFER_PADDING or reduce RF",
start, receptive_field_size
),
));
}
Ok(Self {
layer_buffer: buffer,
buffer_start: start,
receptive_field_size,
})
}
pub fn try_clone(&self) -> std::io::Result<Self> {
Ok(Self {
layer_buffer: self.layer_buffer.try_clone()?,
buffer_start: self.buffer_start,
receptive_field_size: self.receptive_field_size,
})
}
pub fn advance_frames(&mut self, num_frames: usize, channels: usize) {
self.buffer_start += num_frames;
let buffer_frames = self.layer_buffer.size() / channels;
if self.buffer_start + WAVENET_MAX_NUM_FRAMES > buffer_frames * 2 {
self.buffer_start -= buffer_frames;
}
}
}