use super::NamModel;
use super::linear_fft::LinearFftState;
use super::sealed;
use crate::common::diagnostics::NamErrorCode;
use crate::dsp::mirror_buf::MirroredBuffer;
use crate::loader::nam_json::LinearImplementation;
use crate::math::common::AlignedVec;
use log::warn;
#[derive(Debug)]
pub enum LinearMode {
Direct,
Fft(Box<LinearFftState>),
}
pub struct LinearModel {
pub weights: AlignedVec<f32>,
pub bias: f32,
pub history: MirroredBuffer<f32>,
pub write_pos: usize,
pub receptive_field: usize,
double_limit: usize,
pub prewarm_on_reset: bool,
pub implementation: LinearImplementation,
pub mode: LinearMode,
}
const FFT_AUTO_THRESHOLD: usize = 256;
const fn largest_power_of_two_le(n: usize) -> usize {
if n == 0 {
return 0;
}
let mut v = n;
let mut r = 1;
while v > 1 {
r <<= 1;
v >>= 1;
}
r
}
fn select_partition_size(receptive_field: usize) -> usize {
let max_p = receptive_field / 2;
largest_power_of_two_le(max_p.max(1))
}
impl LinearModel {
pub fn new(
weights: Vec<f32>,
bias: f32,
implementation: LinearImplementation,
) -> std::io::Result<Self> {
let receptive_field = weights.len();
let mode = Self::resolve_mode(implementation, receptive_field, &weights)
.map_err(|e| std::io::Error::new(std::io::ErrorKind::OutOfMemory, format!("{e}")))?;
let mut aligned = AlignedVec::from_vec(weights)
.map_err(|e| std::io::Error::new(std::io::ErrorKind::OutOfMemory, format!("{e}")))?;
aligned.reverse();
let history = MirroredBuffer::<f32>::new(receptive_field)?;
let limit = history.size();
let double_limit = limit.checked_mul(2).ok_or_else(|| {
std::io::Error::new(std::io::ErrorKind::InvalidInput, "Limit overflow")
})?;
Ok(Self {
weights: aligned,
bias,
history,
write_pos: limit,
receptive_field,
double_limit,
prewarm_on_reset: true,
implementation,
mode,
})
}
fn resolve_mode(
implementation: LinearImplementation,
receptive_field: usize,
weights: &[f32],
) -> Result<LinearMode, NamErrorCode> {
match implementation {
LinearImplementation::Direct => Ok(LinearMode::Direct),
LinearImplementation::Auto => {
if receptive_field >= FFT_AUTO_THRESHOLD {
let p = select_partition_size(receptive_field);
if p < receptive_field {
return Ok(LinearMode::Fft(Box::new(LinearFftState::new(p, weights)?)));
}
}
Ok(LinearMode::Direct)
}
LinearImplementation::Fft => {
if receptive_field < FFT_AUTO_THRESHOLD {
warn!(
"[Linear] Fft requested but receptive_field={receptive_field} < {FFT_AUTO_THRESHOLD} \
— falling back to Direct"
);
return Ok(LinearMode::Direct);
}
let p = select_partition_size(receptive_field);
Ok(LinearMode::Fft(Box::new(LinearFftState::new(p, weights)?)))
}
}
}
}
mod process;
impl sealed::Sealed for LinearModel {}
impl NamModel for LinearModel {
#[inline(always)]
fn process(&mut self, input: &[f32], output: &mut [f32]) {
unsafe { self.process(input, output) };
}
#[cold]
fn prewarm(&mut self, num_samples: usize) {
self.prewarm(num_samples);
}
fn reset(&mut self, sample_rate: u32, max_buffer_size: usize) -> anyhow::Result<()> {
if self.prewarm_on_reset {
self.reset(sample_rate, max_buffer_size);
}
Ok(())
}
fn prewarm_samples(&self) -> usize {
0
}
fn prewarm_on_reset(&self) -> bool {
self.prewarm_on_reset
}
fn set_prewarm_on_reset(&mut self, val: bool) {
self.prewarm_on_reset = val;
}
}
#[cfg(test)]
#[path = "linear_test.rs"]
mod tests;