sinc.rsannotatedsinc.rssource775 lines · 26.1 KB · raw

Based on this windowed sinc interpolation algorithm by Julius O. Smith III: https://ccrma.stanford.edu/~jos/resample/resample.html

4mod performance;
5mod quality;
7use bincode::{Decode, Encode};
8use std::collections::VecDeque;
9use std::marker::PhantomData;
10use std::{array, iter};
11
12const LINEAR_INTERPOLATION_BITS: u32 = 20;
13
14pub trait SincKernel {
15    fn fir() -> &'static [f32];
16
17    fn oversample_factor() -> u32;
18}
19
20#[derive(Debug, Clone, Copy, Default, Encode, Decode)]
21pub struct Quality;
22
23impl SincKernel for Quality {
24    fn fir() -> &'static [f32] {
25        quality::SINC_KERNEL
26    }
27
28    fn oversample_factor() -> u32 {
29        quality::SINC_OVERSAMPLE_FACTOR
30    }
31}
32
33#[derive(Debug, Clone, Copy, Default, Encode, Decode)]
34pub struct Performance;
35
36impl SincKernel for Performance {
37    fn fir() -> &'static [f32] {
38        performance::SINC_KERNEL
39    }
40
41    fn oversample_factor() -> u32 {
42        performance::SINC_OVERSAMPLE_FACTOR
43    }
44}
45
46#[derive(Debug, Clone, Encode, Decode)]
47struct SampleRingBuffer {
48    buffer: Vec<f64>,
49    idx: usize,
50    len: usize,
51}
52
53impl SampleRingBuffer {
54    // Always leave space at the beginning and end so AVX512 loads can read the first and last
55    // samples without going out of bounds
56    const EXTRA_SPACE: usize = 8;
57
58    // For an N-sample buffer, allows roughly 64*N^2 samples before needing to copy back to the
59    // beginning
60    const CAPACITY_MULTIPLIER: usize = 64;
61
62    fn new(required_samples: usize) -> Self {
63        assert_ne!(required_samples, 0);
64
65        let buffer_len = Self::CAPACITY_MULTIPLIER * required_samples;
66
67        Self { buffer: vec![0.0; buffer_len], idx: Self::EXTRA_SPACE, len: 0 }
68    }
69
70    fn ensure_capacity(&mut self, required_samples: usize) {
71        let buffer_len = Self::CAPACITY_MULTIPLIER * required_samples;
72        let additional = buffer_len.saturating_sub(self.buffer.len());
73        if additional != 0 {
74            self.buffer.extend(iter::repeat_n(0.0, additional));
75        }
76    }
77
78    fn push(&mut self, sample: f64) {
79        if self.idx + self.len == self.buffer.len() - Self::EXTRA_SPACE {
80            // Copy from end of buffer to start
81            let left_len = Self::EXTRA_SPACE + self.len;
82            let (left, right) = self.buffer.split_at_mut(left_len);
83
84            let copy_start = self.idx - left_len;
85            left[Self::EXTRA_SPACE..].copy_from_slice(&right[copy_start..copy_start + self.len]);
86
87            self.idx = Self::EXTRA_SPACE;
88        }
89
90        self.buffer[self.idx + self.len] = sample;
91        self.len += 1;
92    }
93
94    fn pop(&mut self) {
95        if self.len != 0 {
96            self.idx += 1;
97            self.len -= 1;
98        }
99    }
100
101    #[cfg(test)]
102    fn as_slice(&self) -> &[f64] {
103        &self.buffer[self.idx..self.idx + self.len]
104    }
105}
106
107#[derive(Debug, Clone, Encode, Decode)]
108pub struct SincResampler<const CHANNELS: usize, Kernel: SincKernel> {
109    input_counter: f64,
110    source_rate: f64,
111    target_rate: f64,
112    ratio: f64,
113    required_samples: usize,
114    input: [SampleRingBuffer; CHANNELS],
115    output: VecDeque<[f64; CHANNELS]>,
116    // Required for this to compile with the generic Kernel type
117    _marker: PhantomData<Kernel>,
118}
119
120impl<const CHANNELS: usize, Kernel: SincKernel> SincResampler<CHANNELS, Kernel> {
121    #[must_use]
122    pub fn new(source_rate: f64, target_rate: f64) -> Self {
123        let ratio = target_rate / source_rate;
124        let required_samples = estimate_required_samples::<Kernel>(ratio);
125
126        let resampler = Self {
127            input_counter: 0.0,
128            source_rate,
129            target_rate,
130            ratio,
131            required_samples,
132            input: array::from_fn(|_| SampleRingBuffer::new(required_samples)),
133            output: VecDeque::with_capacity(48000 / 30),
134            _marker: PhantomData,
135        };
136
137        resampler.log_debug_output();
138
139        resampler
140    }
141
142    #[inline]
143    pub fn collect(&mut self, samples: [f64; CHANNELS]) {
144        for (sample, input) in iter::zip(samples, &mut self.input) {
145            input.push(sample);
146        }
147
148        while self.input[0].len >= self.required_samples {
149            self.generate_output_sample();
150
151            self.input_counter += 1.0 / self.ratio;
152            while self.input_counter >= 1.0 {
153                self.input_counter -= 1.0;
154                for input in &mut self.input {
155                    input.pop();
156                }
157            }
158        }
159    }
160
161    fn generate_output_sample(&mut self) {
162        fn interpolation_idx_float_to_fixed_point(float: f64, oversample_factor: u32) -> u64 {
163            (float * f64::from(1 << LINEAR_INTERPOLATION_BITS) * f64::from(oversample_factor))
164                .round() as u64
165        }
166
167        let fir = Kernel::fir();
168        let oversample_factor = Kernel::oversample_factor();
169
170        let n = self.required_samples / 2;
171
172        // Steps are smaller when downsampling (ratio < 1.0) to lower the low-pass filter's cutoff
173        let scale = if self.ratio < 1.0 { self.ratio } else { 1.0 };
174
175        let step = {
176            let step_float = scale * f64::from(oversample_factor);
177            (step_float * f64::from(1 << LINEAR_INTERPOLATION_BITS)).round() as u64
178        };
179
180        let input_slices: [_; CHANNELS] = array::from_fn(|ch| self.input[ch].buffer.as_slice());
181
182        // Sum the left wing of the input window / right wing of the windowed sinc
183        let interpolation_idx = {
184            let idx_float = scale * self.input_counter.fract();
185            interpolation_idx_float_to_fixed_point(idx_float, oversample_factor)
186        };
187        let l_sum =
188            sum_wing::<true, _>(fir, interpolation_idx, step, input_slices, n + self.input[0].idx);
189
190        // Sum the right wing of the input window / left wing of the windowed sinc
191        let interpolation_idx = {
192            let idx_float = scale * (1.0 - self.input_counter.fract());
193            interpolation_idx_float_to_fixed_point(idx_float, oversample_factor)
194        };
195        let r_sum =
196            sum_wing::<false, _>(fir, interpolation_idx, step, input_slices, n + self.input[0].idx);
197
198        let volume = scale * f64::from(oversample_factor);
199        self.output.push_back(array::from_fn(|ch| volume * (l_sum[ch] + r_sum[ch])));
200    }
201
202    #[inline]
203    #[must_use]
204    pub fn output_buffer_len(&self) -> usize {
205        self.output.len()
206    }
207
208    #[inline]
209    #[must_use]
210    pub fn output_buffer_pop_front(&mut self) -> Option<[f64; CHANNELS]> {
211        self.output.pop_front()
212    }
213
214    pub fn update_source_frequency(&mut self, source_frequency: f64) {
215        self.source_rate = source_frequency;
216        self.handle_frequency_update();
217    }
218
219    pub fn update_output_frequency(&mut self, output_frequency: f64) {
220        self.target_rate = output_frequency;
221        self.handle_frequency_update();
222    }
223
224    fn handle_frequency_update(&mut self) {
225        // TODO adjust input counter?
226        self.ratio = self.target_rate / self.source_rate;
227        self.required_samples = estimate_required_samples::<Kernel>(self.ratio);
228
229        for input in &mut self.input {
230            input.ensure_capacity(self.required_samples);
231        }
232
233        self.log_debug_output();
234    }
235
236    fn log_debug_output(&self) {
237        if !log::log_enabled!(log::Level::Debug) {
238            return;
239        }
240
241        log::debug!("Source frequency: {}", self.source_rate);
242        log::debug!("Target frequency: {}", self.target_rate);
243        log::debug!("Ratio: {}", self.ratio);
244        log::debug!("FIR half-length: {}", Kernel::fir().len());
245        log::debug!("Oversampling factor: {}", Kernel::oversample_factor());
246        log::debug!("Required input samples: {}", self.required_samples);
247    }
248}
249
250fn estimate_required_samples<Kernel: SincKernel>(ratio: f64) -> usize {
251    let fir = Kernel::fir();
252    let oversample_factor = Kernel::oversample_factor();
253
254    let mut step: f64 = oversample_factor.into();
255    if ratio < 1.0 {
256        step *= ratio;
257    }
258
259    let required_half = (fir.len() as f64 / step).ceil() as usize;
260    2 * required_half + 1
261}
262
263fn sum_wing<const REVERSE: bool, const CHANNELS: usize>(
264    fir: &[f32],
265    interpolation_idx: u64,
266    step: u64,
267    input: [&[f64]; CHANNELS],
268    n: usize,
269) -> [f64; CHANNELS] {
270    #[cfg(target_arch = "x86_64")]
271    {
272        use std::sync::LazyLock;
273
274        static AVX512_SUPPORTED: LazyLock<bool> = LazyLock::new(|| {
275            crate::AVX512_ENABLED
276                && is_x86_feature_detected!("avx512f")
277                && is_x86_feature_detected!("avx512vl")
278                && is_x86_feature_detected!("avx512dq")
279        });
280
281        static AVX2_SUPPORTED: LazyLock<bool> = LazyLock::new(|| {
282            crate::AVX2_ENABLED
283                && is_x86_feature_detected!("avx2")
284                && is_x86_feature_detected!("fma")
285        });
286
287        if *AVX512_SUPPORTED {
288            // SAFETY: This CPU supports AVX512 (F + VL + DQ)
289            unsafe {
290                return sum_wing_avx512::<REVERSE, _>(fir, interpolation_idx, step, input, n);
291            }
292        }
293
294        if *AVX2_SUPPORTED {
295            // SAFETY: This CPU supports AVX2 and FMA
296            unsafe {
297                return sum_wing_avx2::<REVERSE, _>(fir, interpolation_idx, step, input, n);
298            }
299        }
300    }
301
302    sum_wing_no_avx::<REVERSE, _>(fir, interpolation_idx, step, input, n)
303}
304
305fn sum_wing_no_avx<const REVERSE: bool, const CHANNELS: usize>(
306    fir: &[f32],
307    mut interpolation_idx: u64,
308    step: u64,
309    input: [&[f64]; CHANNELS],
310    n: usize,
311) -> [f64; CHANNELS] {
312    let mut sum = [0.0; CHANNELS];
313    for i in 0.. {
314        let fir_idx = (interpolation_idx >> LINEAR_INTERPOLATION_BITS) as usize;
315
316        // Check len-1 because last entry is always 0
317        if fir_idx >= fir.len() - 1 {
318            break;
319        }
320
321        // Apply linear interpolation
322        // Coefficient diffs are not cached as the algorithm describes because RAM speed is likely
323        // going to be the bottleneck here, not calculation throughput
324        let linear_factor = (interpolation_idx & ((1 << LINEAR_INTERPOLATION_BITS) - 1)) as f64
325            / f64::from(1 << LINEAR_INTERPOLATION_BITS);
326
327        let coefficient: f64 = fir[fir_idx].into();
328        let next_coeff: f64 = fir[fir_idx + 1].into();
329        let multiplier = coefficient + linear_factor * (next_coeff - coefficient);
330
331        let input_idx = if REVERSE { n - i } else { n + i + 1 };
332
333        for ch in 0..CHANNELS {
334            let in_sample = input[ch][input_idx];
335            sum[ch] += multiplier * in_sample;
336        }
337
338        interpolation_idx += step;
339    }
340
341    sum
342}

SAFETY: Can only be called on a CPU that supports AVX2 and FMA instructions

345#[cfg(target_arch = "x86_64")]
346#[target_feature(enable = "avx2,fma")]
347fn sum_wing_avx2<const REVERSE: bool, const CHANNELS: usize>(
348    fir: &[f32],
349    interpolation_idx: u64,
350    step: u64,
351    input: [&[f64]; CHANNELS],
352    n: usize,
353) -> [f64; CHANNELS] {
354    #[allow(clippy::wildcard_imports)]
355    use std::arch::x86_64::*;
356    use std::hint::cold_path;
357    use std::mem::transmute;
358
359    const LINEAR_FACTOR_MASK: i64 = (1 << LINEAR_INTERPOLATION_BITS) - 1;
360    const LINEAR_FACTOR_MULTIPLIER: f64 = 1.0 / (1 << LINEAR_INTERPOLATION_BITS) as f64;
361
362    // SAFETY: Later code assumes all input slices are the same length and that n < input.len()
363    assert!(n < input[0].len());
364    if CHANNELS > 1 {
365        assert!(input[1..].iter().all(|channel_input| channel_input.len() == input[0].len()));
366    }
367
368    let mut sums = [_mm256_setzero_pd(); CHANNELS];
369
370    let initial_steps = if REVERSE {
371        _mm256_set_epi64x(0, step as i64, (2 * step) as i64, (3 * step) as i64)
372    } else {
373        _mm256_setr_epi64x(0, step as i64, (2 * step) as i64, (3 * step) as i64)
374    };
375
376    let mut interpolation_idxs =
377        _mm256_add_epi64(_mm256_set1_epi64x(interpolation_idx as i64), initial_steps);
378
379    for i in (0..).step_by(4) {
380        let fir_idxs =
381            _mm256_srli_epi64::<{ LINEAR_INTERPOLATION_BITS as i32 }>(interpolation_idxs);
382
383        // SAFETY: Compare to len-1 instead of len because FIR is read using 64-bit loads
384        let in_bounds = _mm256_xor_si256(
385            _mm256_set1_epi32(!0),
386            _mm256_cmpgt_epi64(fir_idxs, _mm256_set1_epi64x((fir.len() - 2) as i64)),
387        );
388        if _mm256_testz_si256(in_bounds, _mm256_set1_epi32(!0)) != 0 {
389            break;
390        }
391
392        let mut linear_factor_numerators =
393            _mm256_and_si256(interpolation_idxs, _mm256_set1_epi64x(LINEAR_FACTOR_MASK));
394
395        // _mm256_cvtepi64_* intrinsics are AVX512-only :(
396        // Shuffle/permute to an i32x4 vector then convert to f64x4
397
398        // 0 x 1 x  2 x 3 x  ->  0 1 0 1  2 3 2 3
399        linear_factor_numerators = _mm256_shuffle_epi32::<0b10_00_10_00>(linear_factor_numerators);
400
401        // 0 1 0 1  2 3 2 3  ->  0 1 2 3  0 1 2 3
402        linear_factor_numerators =
403            _mm256_permute4x64_epi64::<0b10_00_10_00>(linear_factor_numerators);
404
405        let linear_factor_numerators =
406            _mm256_cvtepi32_pd(_mm256_castsi256_si128(linear_factor_numerators));
407
408        let linear_factors =
409            _mm256_mul_pd(linear_factor_numerators, _mm256_set1_pd(LINEAR_FACTOR_MULTIPLIER));
410
411        // SAFETY: Mask is used to prevent out-of-bounds loads
412        // Load as f64s to pull in each coefficient followed by the next coefficient
413        // Pointer cast is fine because gather instructions don't require an aligned pointer
414        #[allow(clippy::cast_ptr_alignment)]
415        let all_coefficients = unsafe {
416            _mm256_mask_i64gather_pd::<4>(
417                _mm256_setzero_pd(),
418                fir.as_ptr().cast::<f64>(),
419                fir_idxs,
420                _mm256_castsi256_pd(in_bounds),
421            )
422        };
423
424        // 0 n0 1 n1  2 n2 3 n3 -> 0 1 n0 n1  2 3 n2 n3
425        let all_coefficients =
426            _mm256_permute_ps::<0b11_01_10_00>(_mm256_castpd_ps(all_coefficients));
427
428        // 0 1 n0 n1  2 3 n2 n3 -> 0 1 2 3  n0 n1 n2 n3
429        let all_coefficients = _mm256_castpd_ps(_mm256_permute4x64_pd::<0b11_01_10_00>(
430            _mm256_castps_pd(all_coefficients),
431        ));
432
433        let coefficients = _mm256_cvtps_pd(_mm256_castps256_ps128(all_coefficients));
434        let next_coefficients = _mm256_cvtps_pd(_mm256_extractf128_ps::<1>(all_coefficients));
435
436        let multipliers = _mm256_fmadd_pd(
437            linear_factors,
438            _mm256_sub_pd(next_coefficients, coefficients),
439            coefficients,
440        );
441
442        let input_idx = if REVERSE {
443            let (idx, overflowed) = n.overflowing_sub(i + 3);
444            if overflowed {
445                // Should never happen, but don't go out of bounds if it does
446                cold_path();
447                panic!("input array index out of bounds; idx={idx}, len={}", input[0].len());
448            }
449            idx
450        } else {
451            let idx = n + i + 1;
452            if idx + 3 >= input[0].len() {
453                // Should never happen, but don't go out of bounds if it does
454                cold_path();
455                panic!("input array index out of bounds; idx={}, len={}", idx + 3, input[0].len());
456            }
457            idx
458        };
459
460        for ch in 0..CHANNELS {
461            // SAFETY: Checked that idx >= 0 and idx + 3 < input.len()
462            let in_samples = unsafe { _mm256_loadu_pd(input[ch].as_ptr().add(input_idx)) };
463            sums[ch] = _mm256_fmadd_pd(multipliers, in_samples, sums[ch]);
464        }
465
466        interpolation_idxs =
467            _mm256_add_epi64(interpolation_idxs, _mm256_set1_epi64x((4 * step) as i64));
468    }
469
470    sums.map(|sum| {
471        let hsum = _mm256_hadd_pd(sum, sum);
472        let components: [f64; 4] = unsafe { transmute(hsum) };
473        components[0] + components[2]
474    })
475}

SAFETY: Can only be called on a CPU that supports AVX512 (F + VL + DQ)

478#[cfg(target_arch = "x86_64")]
479#[target_feature(enable = "avx512f,avx512vl,avx512dq")]
480fn sum_wing_avx512<const REVERSE: bool, const CHANNELS: usize>(
481    fir: &[f32],
482    interpolation_idx: u64,
483    step: u64,
484    input: [&[f64]; CHANNELS],
485    n: usize,
486) -> [f64; CHANNELS] {
487    #[allow(clippy::wildcard_imports)]
488    use std::arch::x86_64::*;
489    use std::hint::cold_path;
490
491    const LINEAR_FACTOR_MASK: i64 = (1 << LINEAR_INTERPOLATION_BITS) - 1;
492    const LINEAR_FACTOR_MULTIPLIER: f64 = 1.0 / (1 << LINEAR_INTERPOLATION_BITS) as f64;
493
494    // SAFETY: Later code assumes all input slices are the same length and that n < input.len()
495    assert!(n < input[0].len());
496    if CHANNELS > 1 {
497        assert!(input[1..].iter().all(|channel_input| channel_input.len() == input[0].len()));
498    }
499
500    let mut sums = [_mm512_setzero_pd(); CHANNELS];
501
502    let initial_step_multipliers = if REVERSE {
503        _mm512_set_epi64(0, 1, 2, 3, 4, 5, 6, 7)
504    } else {
505        _mm512_setr_epi64(0, 1, 2, 3, 4, 5, 6, 7)
506    };
507
508    let mut interpolation_idxs = _mm512_add_epi64(
509        _mm512_set1_epi64(interpolation_idx as i64),
510        _mm512_mullo_epi64(_mm512_set1_epi64(step as i64), initial_step_multipliers),
511    );
512
513    for i in (0..).step_by(8) {
514        let fir_idxs = _mm512_srli_epi64::<{ LINEAR_INTERPOLATION_BITS }>(interpolation_idxs);
515
516        // SAFETY: Compare to len-1 instead of len because FIR is read using 64-bit loads
517        let in_bounds =
518            _mm512_cmplt_epi64_mask(fir_idxs, _mm512_set1_epi64((fir.len() - 1) as i64));
519        if in_bounds == 0 {
520            break;
521        }
522
523        let linear_factor_numerators =
524            _mm512_and_si512(interpolation_idxs, _mm512_set1_epi64(LINEAR_FACTOR_MASK));
525        let linear_factors = _mm512_mul_pd(
526            _mm512_cvtepi64_pd(linear_factor_numerators),
527            _mm512_set1_pd(LINEAR_FACTOR_MULTIPLIER),
528        );
529
530        // SAFETY: Mask is used to prevent out-of-bounds loads
531        // Load as f64s to pull in each coefficient followed by the next coefficient
532        // Pointer cast is fine because gather instructions don't require an aligned pointer
533        #[allow(clippy::cast_ptr_alignment)]
534        let all_coefficients = unsafe {
535            _mm512_mask_i64gather_pd::<4>(
536                _mm512_setzero_pd(),
537                in_bounds,
538                fir_idxs,
539                fir.as_ptr().cast::<f64>(),
540            )
541        };
542
543        // 0 n0 1 n1 2 n2 3 n3 4 n4 5 n5 6 n6 7 n7 -> 0 1 2 3 4 5 6 7 n0 n1 n2 n3 n4 n5 n6 n7
544        let all_coefficients = _mm512_permutexvar_ps(
545            _mm512_setr_epi32(0, 2, 4, 6, 8, 10, 12, 14, 1, 3, 5, 7, 9, 11, 13, 15),
546            _mm512_castpd_ps(all_coefficients),
547        );
548
549        let coefficients = _mm512_cvtps_pd(_mm512_castps512_ps256(all_coefficients));
550        let next_coefficients = _mm512_cvtps_pd(_mm512_extractf32x8_ps::<1>(all_coefficients));
551
552        let multipliers = _mm512_fmadd_pd(
553            linear_factors,
554            _mm512_sub_pd(next_coefficients, coefficients),
555            coefficients,
556        );
557
558        let input_idx = if REVERSE {
559            let (idx, overflowed) = n.overflowing_sub(i + 7);
560            if overflowed {
561                // Should never happen, but don't go out of bounds if it does
562                cold_path();
563                panic!("input array index out of bounds; idx={idx}, len={}", input[0].len());
564            }
565            idx
566        } else {
567            let idx = n + i + 1;
568            if idx + 7 >= input[0].len() {
569                // Should never happen, but don't go out of bounds if it does
570                cold_path();
571                panic!("input array index out of bounds; idx={}, len={}", idx + 7, input[0].len());
572            }
573            idx
574        };
575
576        for ch in 0..CHANNELS {
577            // SAFETY: Checked that idx >= 0 and idx + 7 < input.len()
578            let in_samples = unsafe { _mm512_loadu_pd(input[ch].as_ptr().add(input_idx)) };
579            sums[ch] = _mm512_fmadd_pd(multipliers, in_samples, sums[ch]);
580        }
581
582        interpolation_idxs =
583            _mm512_add_epi64(interpolation_idxs, _mm512_set1_epi64((8 * step) as i64));
584    }
585
586    sums.map(|sum| _mm512_reduce_add_pd(sum))
587}
589pub type QualitySincResampler<const CHANNELS: usize> = SincResampler<CHANNELS, Quality>;
590pub type PerformanceSincResampler<const CHANNELS: usize> = SincResampler<CHANNELS, Performance>;
591
592#[cfg(test)]
593mod tests {
594    use super::*;
595
596    #[test]
597    fn ring_buffer_basic() {
598        let mut buffer = SampleRingBuffer::new(10);
599
600        buffer.push(1.0);
601        buffer.push(2.0);
602        buffer.push(3.0);
603        buffer.push(4.0);
604        buffer.push(5.0);
605
606        assert_eq!(buffer.as_slice(), &[1.0, 2.0, 3.0, 4.0, 5.0]);
607
608        buffer.pop();
609        assert_eq!(buffer.as_slice(), &[2.0, 3.0, 4.0, 5.0]);
610        buffer.pop();
611        assert_eq!(buffer.as_slice(), &[3.0, 4.0, 5.0]);
612        buffer.pop();
613        assert_eq!(buffer.as_slice(), &[4.0, 5.0]);
614        buffer.pop();
615        assert_eq!(buffer.as_slice(), &[5.0]);
616        buffer.pop();
617        assert_eq!(buffer.as_slice(), &[]);
618        buffer.pop();
619        assert_eq!(buffer.as_slice(), &[]);
620    }
621
622    #[test]
623    fn ring_buffer_copy() {
624        let mut buffer = SampleRingBuffer::new(10);
625
626        buffer.buffer.fill_with(|| rand::random_range(-1.0..=1.0));
627
628        let len = 15;
629        let end_start = buffer.buffer.len() - 8 - len;
630
631        buffer.idx = end_start;
632        buffer.len = len;
633
634        // Should trigger a copy from end to start
635        buffer.push(1.0);
636
637        // Validate that the index moved to start
638        assert_eq!(buffer.idx, 8);
639        assert_eq!(buffer.len, len + 1);
640
641        // Validate that samples were copied
642        let mut expected = buffer.buffer[end_start..end_start + len].to_vec();
643        expected.push(1.0);
644        assert_eq!(buffer.as_slice(), expected);
645    }
646
647    // Validates that sum_wing_avx2() and sum_wing_avx512() produce the same results as sum_wing_no_avx()
648    fn sum_wing_test<Kernel: SincKernel>(
649        source_rate: f64,
650        sum_wing_avx_fn: impl Fn(bool, &[f32], u64, u64, [&[f64]; 1], usize) -> [f64; 1],
651    ) {
652        let target_rate = 48000.0;
653        let ratio = target_rate / source_rate;
654        let required_samples = estimate_required_samples::<Kernel>(ratio);
655
656        let scale = if ratio < 1.0 { ratio } else { 1.0 };
657        let step_float = scale
658            * f64::from(Kernel::oversample_factor())
659            * f64::from(1 << LINEAR_INTERPOLATION_BITS);
660        let step = step_float.round() as u64;
661
662        let fir = Kernel::fir();
663
664        let mut buffer = SampleRingBuffer::new(required_samples);
665        buffer.idx = buffer.buffer.len() - SampleRingBuffer::EXTRA_SPACE - required_samples;
666        for _ in 0..2 * required_samples {
667            buffer.push(rand::random());
668        }
669
670        // Use fewer iterations when running in miri because miri is very slow
671        let iterations = cfg_select! {
672            miri => 3,
673            _ => required_samples / 2,
674        };
675
676        for _ in 0..iterations {
677            let n = required_samples / 2 + buffer.idx;
678
679            let interpolation_idx = (rand::random_range(0.0..1.0) * step_float).round() as u64;
680
681            let non_avx_reverse = sum_wing_no_avx::<true, _>(
682                fir,
683                interpolation_idx,
684                step,
685                [buffer.buffer.as_slice()],
686                n,
687            );
688            let non_avx_forward = sum_wing_no_avx::<false, _>(
689                fir,
690                interpolation_idx,
691                step,
692                [buffer.buffer.as_slice()],
693                n,
694            );
695
696            let avx_reverse =
697                sum_wing_avx_fn(true, fir, interpolation_idx, step, [buffer.buffer.as_slice()], n);
698            let avx_forward =
699                sum_wing_avx_fn(false, fir, interpolation_idx, step, [buffer.buffer.as_slice()], n);
700
701            assert!(
702                (non_avx_reverse[0] - avx_reverse[0]).abs() < 1e-9,
703                "{non_avx_reverse:?} == {avx_reverse:?} (source rate {source_rate})"
704            );
705            assert!(
706                (non_avx_forward[0] - avx_forward[0]).abs() < 1e-9,
707                "{non_avx_forward:?} == {avx_forward:?} (source rate {source_rate})",
708            );
709
710            buffer.pop();
711        }
712    }
713
714    const TEST_SOURCE_RATES: &[f64] = &[4000000.0, 55000.0, 48000.0, 20000.0];
715
716    // These tests should ideally be run using miri to verify no out-of-bounds memory reads:
717    //   $ cargo +nightly miri test -p dsp
718    #[cfg(target_arch = "x86_64")]
719    #[test]
720    fn test_sum_wing_avx2() {
721        // SAFETY: Only run test on CPUs that support AVX2 and FMA
722        if !is_x86_feature_detected!("avx2") || !is_x86_feature_detected!("fma") {
723            return;
724        }
725
726        for &source_rate in TEST_SOURCE_RATES {
727            let sum_wing_fn = |reverse,
728                               fir: &[f32],
729                               interpolation_idx,
730                               step,
731                               samples: [&[f64]; 1],
732                               n| {
733                if reverse {
734                    unsafe { sum_wing_avx2::<true, _>(fir, interpolation_idx, step, samples, n) }
735                } else {
736                    unsafe { sum_wing_avx2::<false, _>(fir, interpolation_idx, step, samples, n) }
737                }
738            };
739
740            sum_wing_test::<Performance>(source_rate, sum_wing_fn);
741            sum_wing_test::<Quality>(source_rate, sum_wing_fn);
742        }
743    }
744
745    #[cfg(target_arch = "x86_64")]
746    #[cfg_attr(miri, ignore)] // miri does not support all required AVX512 intrinsics as of 1.99 nightly
747    #[test]
748    fn test_sum_wing_avx512() {
749        // SAFETY: Only run test on CPUs that support AVX512 (F + DQ + VL)
750        if !is_x86_feature_detected!("avx512f")
751            || !is_x86_feature_detected!("avx512dq")
752            || !is_x86_feature_detected!("avx512vl")
753        {
754            return;
755        }
756
757        for &source_rate in TEST_SOURCE_RATES {
758            let sum_wing_fn = |reverse,
759                               fir: &[f32],
760                               interpolation_idx,
761                               step,
762                               samples: [&[f64]; 1],
763                               n| {
764                if reverse {
765                    unsafe { sum_wing_avx512::<true, _>(fir, interpolation_idx, step, samples, n) }
766                } else {
767                    unsafe { sum_wing_avx512::<false, _>(fir, interpolation_idx, step, samples, n) }
768                }
769            };
770
771            sum_wing_test::<Performance>(source_rate, sum_wing_fn);
772            sum_wing_test::<Quality>(source_rate, sum_wing_fn);
773        }
774    }
775}