use crate::common::diagnostics::NamErrorCode;
use crate::math::common::AlignedVec;
use crate::math::common::Avx2Math;
use crate::math::common::Avx512Math;
use crate::math::common::Avx512VnniBf16Math;
use crate::math::common::dispatch::InstructionSet;
use crate::math::common::dispatch::SimdMathConfig;
use crate::math::common::traits::SimdMath;
use crate::math::dsp::fft::RfftPlanner;
use log::info;
pub struct ConvEngine {
fft_size: usize,
n_bins: usize,
partition_size: usize,
num_partitions: usize,
h_fdl_re: AlignedVec<f32>,
h_fdl_im: AlignedVec<f32>,
fdl_re: AlignedVec<f32>,
fdl_im: AlignedVec<f32>,
fdl_idx: usize,
input_buf: AlignedVec<f32>,
rfft: RfftPlanner<f32>,
fft_buf_re: AlignedVec<f32>,
fft_buf_im: AlignedVec<f32>,
acc_re: AlignedVec<f32>,
acc_im: AlignedVec<f32>,
output_buf: AlignedVec<f32>,
output_start: usize,
isa: InstructionSet,
}
impl ConvEngine {
pub fn new(ir: &[f32], partition_size: usize) -> Result<Self, NamErrorCode> {
assert!(partition_size > 0, "partition_size must be positive");
let fft_size = (2 * partition_size).next_power_of_two();
let n_bins = fft_size / 2 + 1;
let output_start = fft_size - partition_size;
let num_partitions = if ir.is_empty() {
0
} else {
ir.len().div_ceil(partition_size)
};
let mut rfft = RfftPlanner::<f32>::new(fft_size);
let h_fdl_part_len = num_partitions * n_bins;
let mut h_fdl_re = AlignedVec::new(h_fdl_part_len, 0.0_f32)?;
let mut h_fdl_im = AlignedVec::new(h_fdl_part_len, 0.0_f32)?;
let mut ir_buf = vec![0.0f32; fft_size];
let mut tmp_re = vec![0.0f32; n_bins];
let mut tmp_im = vec![0.0f32; n_bins];
for p in 0..num_partitions {
let ir_start = p * partition_size;
let ir_end = (ir_start + partition_size).min(ir.len());
ir_buf.fill(0.0);
for (i, &sample) in ir[ir_start..ir_end].iter().enumerate() {
ir_buf[i] = sample;
}
rfft.process_forward(&ir_buf, &mut tmp_re, &mut tmp_im);
let base = p * n_bins;
for k in 0..n_bins {
h_fdl_re[base + k] = tmp_re[k];
h_fdl_im[base + k] = tmp_im[k];
}
}
let fdl_part_len = num_partitions * n_bins;
let fdl_re = AlignedVec::new(fdl_part_len, 0.0_f32)?;
let fdl_im = AlignedVec::new(fdl_part_len, 0.0_f32)?;
let input_buf = AlignedVec::new(fft_size, 0.0_f32)?;
let fft_buf_re = AlignedVec::new(n_bins, 0.0_f32)?;
let fft_buf_im = AlignedVec::new(n_bins, 0.0_f32)?;
let acc_re = AlignedVec::new(n_bins, 0.0_f32)?;
let acc_im = AlignedVec::new(n_bins, 0.0_f32)?;
let output_buf = AlignedVec::new(fft_size, 0.0_f32)?;
let isa = SimdMathConfig::current().instruction_set;
if num_partitions == 0 {
info!(
"[Conv] Engine built: passthrough (empty IR), partition={}, fft={}",
partition_size, fft_size
);
} else {
info!(
"[Conv] Engine built: {} IR samples, partition={}, fft={}, {} partitions, isa={:?}",
ir.len(),
partition_size,
fft_size,
num_partitions,
isa
);
}
Ok(Self {
fft_size,
n_bins,
partition_size,
num_partitions,
h_fdl_re,
h_fdl_im,
fdl_re,
fdl_im,
fdl_idx: 0,
input_buf,
rfft,
fft_buf_re,
fft_buf_im,
acc_re,
acc_im,
output_buf,
output_start,
isa,
})
}
#[inline(always)]
pub fn partition_size(&self) -> usize {
self.partition_size
}
#[inline(always)]
pub fn fft_size(&self) -> usize {
self.fft_size
}
#[inline(always)]
pub fn num_partitions(&self) -> usize {
self.num_partitions
}
#[inline(always)]
pub fn latency_samples(&self) -> usize {
self.partition_size
}
#[inline(always)]
pub fn is_passthrough(&self) -> bool {
self.num_partitions == 0
}
#[inline]
pub fn process(&mut self, input: &[f32], output: &mut [f32]) {
debug_assert_eq!(input.len(), self.partition_size);
debug_assert_eq!(output.len(), self.partition_size);
if self.num_partitions == 0 {
unsafe {
core::ptr::copy_nonoverlapping(
input.as_ptr(),
output.as_mut_ptr(),
self.partition_size,
);
}
return;
}
let in_len = self.fft_size;
let out_start = self.output_start;
unsafe {
core::ptr::copy(
self.input_buf.as_ptr().add(self.partition_size),
self.input_buf.as_mut_ptr(),
in_len - self.partition_size,
);
}
unsafe {
core::ptr::copy_nonoverlapping(
input.as_ptr(),
self.input_buf.as_mut_ptr().add(out_start),
self.partition_size,
);
}
self.rfft
.process_forward(&self.input_buf, &mut self.fft_buf_re, &mut self.fft_buf_im);
let fdl_base = self.fdl_idx * self.n_bins;
self.fdl_re[fdl_base..fdl_base + self.n_bins]
.copy_from_slice(&self.fft_buf_re[..self.n_bins]);
self.fdl_im[fdl_base..fdl_base + self.n_bins]
.copy_from_slice(&self.fft_buf_im[..self.n_bins]);
let p_count = self.num_partitions;
let n_bins = self.n_bins;
if p_count == 1 {
let fdl_start = self.fdl_idx * self.n_bins;
unsafe {
match self.isa {
InstructionSet::Avx512VnniBf16 => Avx512VnniBf16Math::complex_mac_overwrite(
&self.h_fdl_re[..n_bins],
&self.h_fdl_im[..n_bins],
&self.fdl_re[fdl_start..fdl_start + n_bins],
&self.fdl_im[fdl_start..fdl_start + n_bins],
&mut self.acc_re[..n_bins],
&mut self.acc_im[..n_bins],
),
InstructionSet::Avx512 => Avx512Math::complex_mac_overwrite(
&self.h_fdl_re[..n_bins],
&self.h_fdl_im[..n_bins],
&self.fdl_re[fdl_start..fdl_start + n_bins],
&self.fdl_im[fdl_start..fdl_start + n_bins],
&mut self.acc_re[..n_bins],
&mut self.acc_im[..n_bins],
),
InstructionSet::Avx2 => Avx2Math::complex_mac_overwrite(
&self.h_fdl_re[..n_bins],
&self.h_fdl_im[..n_bins],
&self.fdl_re[fdl_start..fdl_start + n_bins],
&self.fdl_im[fdl_start..fdl_start + n_bins],
&mut self.acc_re[..n_bins],
&mut self.acc_im[..n_bins],
),
}
}
} else {
self.acc_re[..n_bins].fill(0.0);
self.acc_im[..n_bins].fill(0.0);
for p in 0..p_count {
let fdl_p = (self.fdl_idx + p_count - p) % p_count;
let fdl_start = fdl_p * self.n_bins;
let h_start = p * self.n_bins;
unsafe {
match self.isa {
InstructionSet::Avx512VnniBf16 => {
Avx512VnniBf16Math::complex_mac_accumulate(
&self.h_fdl_re[h_start..h_start + n_bins],
&self.h_fdl_im[h_start..h_start + n_bins],
&self.fdl_re[fdl_start..fdl_start + n_bins],
&self.fdl_im[fdl_start..fdl_start + n_bins],
&mut self.acc_re[..n_bins],
&mut self.acc_im[..n_bins],
)
}
InstructionSet::Avx512 => Avx512Math::complex_mac_accumulate(
&self.h_fdl_re[h_start..h_start + n_bins],
&self.h_fdl_im[h_start..h_start + n_bins],
&self.fdl_re[fdl_start..fdl_start + n_bins],
&self.fdl_im[fdl_start..fdl_start + n_bins],
&mut self.acc_re[..n_bins],
&mut self.acc_im[..n_bins],
),
InstructionSet::Avx2 => Avx2Math::complex_mac_accumulate(
&self.h_fdl_re[h_start..h_start + n_bins],
&self.h_fdl_im[h_start..h_start + n_bins],
&self.fdl_re[fdl_start..fdl_start + n_bins],
&self.fdl_im[fdl_start..fdl_start + n_bins],
&mut self.acc_re[..n_bins],
&mut self.acc_im[..n_bins],
),
}
}
}
}
self.rfft
.process_inverse(&mut self.acc_re, &mut self.acc_im, &mut self.output_buf);
unsafe {
core::ptr::copy_nonoverlapping(
self.output_buf.as_ptr().add(out_start),
output.as_mut_ptr(),
self.partition_size,
);
}
self.fdl_idx += 1;
if self.fdl_idx >= p_count {
self.fdl_idx = 0;
}
}
}
#[cfg(test)]
#[path = "conv_test.rs"]
mod conv_test;