use super::super::head::A2HeadConv;
use super::super::layer::A2Layer;
use super::super::params::{A2_DILATIONS, A2_KERNEL_SIZES, A2_NUM_LAYERS};
use super::a2_receptive_field;
use crate::dsp::mirror_buf::MirroredBuffer;
use crate::math::common::AlignedVec;
use crate::models::wavenet::common::WAVENET_MAX_NUM_FRAMES;
use serde_json::Value;
use std::sync::Arc;
pub mod prewarm;
pub mod process;
pub struct WaveNetA2<const CH: usize> {
pub layers: Vec<A2Layer>,
pub rechannel_w_f32: AlignedVec<f32>,
pub head_conv: Option<A2HeadConv>,
pub head_accum: AlignedVec<f32>,
pub head_write_pos: usize,
pub head_ring_mask: usize,
pub layer_buffers: Vec<MirroredBuffer<f32>>,
pub layer_ring_sizes: Vec<usize>,
pub layer_lookbacks: Vec<usize>,
pub layer_buffer_starts: Vec<usize>,
pub layer_in: AlignedVec<f32>,
pub receptive_field_size: usize,
pub max_buffer_size: usize,
pub layer_raw: Option<Value>,
pub z_scratch: AlignedVec<f32>,
pub rt_status: Option<Arc<crate::common::spsc::RtStatusFlags>>,
pub prewarm_on_reset: bool,
}
impl<const CH: usize> WaveNetA2<CH> {
pub fn new() -> anyhow::Result<Self> {
let rf = a2_receptive_field();
let max_buf = WAVENET_MAX_NUM_FRAMES;
let head_ring_size = (rf + max_buf + 1).next_power_of_two();
let head_ring_mask = head_ring_size - 1;
let mut layer_buffers = Vec::with_capacity(A2_NUM_LAYERS);
let mut layer_ring_sizes = Vec::with_capacity(A2_NUM_LAYERS);
let mut layer_lookbacks = Vec::with_capacity(A2_NUM_LAYERS);
let mut layer_buffer_starts = Vec::with_capacity(A2_NUM_LAYERS);
for i in 0..A2_NUM_LAYERS {
let max_lookback = (A2_KERNEL_SIZES[i] - 1) * A2_DILATIONS[i];
let cap = max_lookback + max_buf + 1;
let mb = MirroredBuffer::<f32>::new(cap * CH)?;
let ring_size = mb.size();
layer_buffers.push(mb);
layer_ring_sizes.push(ring_size);
layer_lookbacks.push(max_lookback * CH);
layer_buffer_starts.push(ring_size);
}
Ok(Self {
layers: Vec::with_capacity(A2_NUM_LAYERS),
rechannel_w_f32: AlignedVec::new(CH, 0.0f32)
.expect("allocation should succeed for test-sized buffers"),
head_conv: None,
head_accum: AlignedVec::new(head_ring_size * CH, 0.0f32)
.expect("allocation should succeed for test-sized buffers"),
head_write_pos: rf,
head_ring_mask,
layer_buffers,
layer_ring_sizes,
layer_lookbacks,
layer_buffer_starts,
layer_in: AlignedVec::new(CH * max_buf, 0.0f32)
.expect("allocation should succeed for test-sized buffers"),
receptive_field_size: rf,
max_buffer_size: max_buf,
layer_raw: None,
z_scratch: AlignedVec::new(CH, 0.0f32)
.expect("allocation should succeed for test-sized buffers"),
rt_status: None,
prewarm_on_reset: true,
})
}
#[inline(always)]
pub fn channels(&self) -> usize {
CH
}
pub fn inject_rt_status(&mut self, rt_status: Arc<crate::common::spsc::RtStatusFlags>) {
self.rt_status = Some(rt_status);
}
#[inline(always)]
pub fn receptive_field(&self) -> usize {
self.receptive_field_size
}
pub fn set_max_buffer_size(&mut self, max_buf: usize) -> anyhow::Result<()> {
if max_buf < self.max_buffer_size {
return Ok(());
}
if max_buf == self.max_buffer_size {
let rf = self.receptive_field_size;
let ha_len = self.head_accum.len();
self.head_accum[..ha_len].fill(0.0);
self.head_write_pos = rf;
for buf in self.layer_buffers.iter_mut() {
let len = buf.size();
buf[..len].fill(0.0);
}
self.layer_buffer_starts
.copy_from_slice(&self.layer_ring_sizes);
let li_len = self.layer_in.len();
self.layer_in[..li_len].fill(0.0);
return Ok(());
}
self.max_buffer_size = max_buf;
let rf = self.receptive_field_size;
self.layer_buffers.clear();
self.layer_ring_sizes.clear();
self.layer_lookbacks.clear();
self.layer_buffer_starts.clear();
for i in 0..A2_NUM_LAYERS {
let max_lookback = (A2_KERNEL_SIZES[i] - 1) * A2_DILATIONS[i];
let cap = max_lookback + max_buf + 1;
let mb = MirroredBuffer::<f32>::new(cap * CH)?;
let ring_size = mb.size();
self.layer_buffers.push(mb);
self.layer_ring_sizes.push(ring_size);
self.layer_lookbacks.push(max_lookback * CH);
self.layer_buffer_starts.push(ring_size);
}
self.layer_in = AlignedVec::new(CH * max_buf, 0.0f32)
.expect("allocation should succeed for test-sized buffers");
let head_ring_size = (rf + max_buf + 1).next_power_of_two();
self.head_ring_mask = head_ring_size - 1;
self.head_accum = AlignedVec::new(head_ring_size * CH, 0.0f32)
.expect("allocation should succeed for test-sized buffers");
self.head_write_pos = rf;
Ok(())
}
pub fn reset(&mut self, _sample_rate: u32, max_buffer_size: usize) -> anyhow::Result<()> {
self.set_max_buffer_size(max_buffer_size)?;
if self.prewarm_on_reset {
self.prewarm();
}
Ok(())
}
#[inline(always)]
pub fn has_weights(&self) -> bool {
!self.layers.is_empty()
}
pub fn set_layer_raw(&mut self, raw: Option<Value>) {
self.layer_raw = raw;
}
}