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::traits::SimdMath;
impl super::LinearFftState {
#[cold]
pub fn reset(&mut self) {
self.fdl_re.fill(0.0f32);
self.fdl_im.fill(0.0f32);
self.input_buf.fill(0.0f32);
self.fft_re.fill(0.0f32);
self.fft_im.fill(0.0f32);
self.acc_re.fill(0.0f32);
self.acc_im.fill(0.0f32);
self.output_buf.fill(0.0f32);
self.tail_output_buf.fill(0.0f32);
self.fdl_write_idx = self.num_partitions.saturating_sub(1);
self.sample_counter = 0;
}
pub fn process_tail_block(&mut self, input_window: &[f32]) {
let p = self.p;
let block_size = 2 * p;
let num_bins = self.num_bins;
let num_partitions = self.num_partitions;
debug_assert_eq!(input_window.len(), block_size);
if num_partitions == 0 {
return;
}
self.input_buf[..block_size].copy_from_slice(input_window);
self.rfft
.process_forward(&self.input_buf, &mut self.fft_re, &mut self.fft_im);
self.acc_re[..num_bins].fill(0.0);
self.acc_im[..num_bins].fill(0.0);
unsafe {
match self.isa {
InstructionSet::Avx512VnniBf16 => Avx512VnniBf16Math::complex_mac_accumulate(
&self.h_fdl_re[..num_bins],
&self.h_fdl_im[..num_bins],
&self.fft_re[..num_bins],
&self.fft_im[..num_bins],
&mut self.acc_re[..num_bins],
&mut self.acc_im[..num_bins],
),
InstructionSet::Avx512 => Avx512Math::complex_mac_accumulate(
&self.h_fdl_re[..num_bins],
&self.h_fdl_im[..num_bins],
&self.fft_re[..num_bins],
&self.fft_im[..num_bins],
&mut self.acc_re[..num_bins],
&mut self.acc_im[..num_bins],
),
InstructionSet::Avx2 => Avx2Math::complex_mac_accumulate(
&self.h_fdl_re[..num_bins],
&self.h_fdl_im[..num_bins],
&self.fft_re[..num_bins],
&self.fft_im[..num_bins],
&mut self.acc_re[..num_bins],
&mut self.acc_im[..num_bins],
),
}
}
for k in 1..num_partitions {
let input_idx = (self.fdl_write_idx + num_partitions - k) % num_partitions;
let fdl_start = input_idx * num_bins;
let h_start = k * num_bins;
unsafe {
match self.isa {
InstructionSet::Avx512VnniBf16 => Avx512VnniBf16Math::complex_mac_accumulate(
&self.h_fdl_re[h_start..h_start + num_bins],
&self.h_fdl_im[h_start..h_start + num_bins],
&self.fdl_re[fdl_start..fdl_start + num_bins],
&self.fdl_im[fdl_start..fdl_start + num_bins],
&mut self.acc_re[..num_bins],
&mut self.acc_im[..num_bins],
),
InstructionSet::Avx512 => Avx512Math::complex_mac_accumulate(
&self.h_fdl_re[h_start..h_start + num_bins],
&self.h_fdl_im[h_start..h_start + num_bins],
&self.fdl_re[fdl_start..fdl_start + num_bins],
&self.fdl_im[fdl_start..fdl_start + num_bins],
&mut self.acc_re[..num_bins],
&mut self.acc_im[..num_bins],
),
InstructionSet::Avx2 => Avx2Math::complex_mac_accumulate(
&self.h_fdl_re[h_start..h_start + num_bins],
&self.h_fdl_im[h_start..h_start + num_bins],
&self.fdl_re[fdl_start..fdl_start + num_bins],
&self.fdl_im[fdl_start..fdl_start + num_bins],
&mut self.acc_re[..num_bins],
&mut self.acc_im[..num_bins],
),
}
}
}
{
let fdl_base = self.fdl_write_idx * num_bins;
self.fdl_re[fdl_base..fdl_base + num_bins].copy_from_slice(&self.fft_re);
self.fdl_im[fdl_base..fdl_base + num_bins].copy_from_slice(&self.fft_im);
}
self.fdl_write_idx += 1;
if self.fdl_write_idx >= num_partitions {
self.fdl_write_idx = 0;
}
self.rfft
.process_inverse(&mut self.acc_re, &mut self.acc_im, &mut self.output_buf);
self.tail_output_buf[..p].copy_from_slice(&self.output_buf[p..block_size]);
}
}