use super::conv::ConvEngine;
use crate::common::diagnostics::NamErrorCode;
use crate::math::common::AlignedVec;
pub struct CabSimAdapter {
engine: Box<ConvEngine>,
partition: usize,
input_buf: AlignedVec<f32>,
output_buf: AlignedVec<f32>,
output_scratch: AlignedVec<f32>,
input_count: usize,
output_read: usize,
output_write: usize,
}
impl CabSimAdapter {
pub fn new(engine: Box<ConvEngine>) -> Result<Self, NamErrorCode> {
let partition = engine.partition_size();
Ok(Self {
engine,
partition,
input_buf: AlignedVec::new(2 * partition, 0.0_f32)?,
output_buf: AlignedVec::new(2 * partition, 0.0_f32)?,
output_scratch: AlignedVec::new(partition, 0.0_f32)?,
input_count: 0,
output_read: 0,
output_write: 0,
})
}
#[inline(always)]
pub fn partition_size(&self) -> usize {
self.partition
}
#[inline(always)]
pub fn latency_samples(&self) -> usize {
self.partition
}
#[inline(always)]
pub fn is_passthrough(&self) -> bool {
self.engine.is_passthrough()
}
#[inline(always)]
pub fn num_partitions(&self) -> usize {
self.engine.num_partitions()
}
#[inline(always)]
pub fn engine(&self) -> &ConvEngine {
&self.engine
}
#[inline(always)]
pub fn engine_mut(&mut self) -> &mut ConvEngine {
&mut self.engine
}
#[inline(always)]
pub fn needs_flush(&self) -> bool {
self.input_count > 0 || self.output_read < self.output_write
}
#[inline(always)]
pub fn tail_samples(&self) -> usize {
if self.engine.is_passthrough() {
return 0;
}
self.engine.num_partitions().saturating_mul(self.partition) + self.partition }
pub fn process_variable(&mut self, input: &[f32], output: &mut [f32]) {
let sub_n = input.len();
assert!(sub_n <= self.partition, "sub-block exceeds partition_size");
assert_eq!(output.len(), sub_n);
if self.engine.is_passthrough() {
output.copy_from_slice(input);
return;
}
if sub_n > 0 {
self.input_buf[self.input_count..self.input_count + sub_n].copy_from_slice(input);
self.input_count += sub_n;
}
while self.input_count >= self.partition {
self.engine.process(
&self.input_buf[..self.partition],
&mut self.output_scratch[..self.partition],
);
if self.output_read > 0 {
let remaining = self.output_write - self.output_read;
if remaining > 0 {
self.output_buf
.copy_within(self.output_read..self.output_write, 0);
}
self.output_write = remaining;
self.output_read = 0;
}
self.output_buf[self.output_write..self.output_write + self.partition]
.copy_from_slice(&self.output_scratch[..self.partition]);
self.output_write += self.partition;
let remaining = self.input_count - self.partition;
if remaining > 0 {
self.input_buf
.copy_within(self.partition..self.input_count, 0);
}
self.input_count = remaining;
}
let available = self.output_write - self.output_read;
let n = sub_n.min(available);
if n > 0 {
output[..n].copy_from_slice(&self.output_buf[self.output_read..self.output_read + n]);
self.output_read += n;
}
output[n..].fill(0.0);
if self.output_read >= self.output_write {
self.output_read = 0;
self.output_write = 0;
}
}
}
#[cfg(test)]
#[path = "adapter_test.rs"]
mod adapter_test;