Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 23 additions & 0 deletions crates/onnx-runtime-ep-cpu/src/kernels/dft.rs
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,29 @@ pub static DFT_VDSP_TEST_HITS: AtomicU64 = AtomicU64::new(0);
/// Dispatch counter for the radix-2 FFT fallback path.
pub static DFT_FFT_TEST_HITS: AtomicU64 = AtomicU64::new(0);

/// Power-of-two transforms that took *a* fast path — whichever one this target
/// has.
///
/// `DFT_FFT_TEST_HITS` alone does not answer that question. On Apple targets
/// `DftPlan::new` builds a vDSP setup for every power-of-two `n >= 4`, and
/// `transform` returns from that branch before the radix-2 one, so a caller
/// asserting on the radix-2 counter there is asserting that the platform's own
/// fast path was *not* used — false by construction, and nothing to do with the
/// property it meant to check.
///
/// What every target shares is that a power-of-two transform must not fall back
/// to `naive_dft_into`. That is what this sums, so a caller can assert the
/// property instead of one platform's route to it.
pub fn fast_path_hits() -> u64 {
let radix2 = DFT_FFT_TEST_HITS.load(Ordering::Relaxed);
#[cfg(any(target_os = "macos", target_os = "ios"))]
{
return radix2 + DFT_VDSP_TEST_HITS.load(Ordering::Relaxed);
}
#[cfg(not(any(target_os = "macos", target_os = "ios")))]
radix2
}

pub struct DftFactory;

impl KernelFactory for DftFactory {
Expand Down
10 changes: 5 additions & 5 deletions crates/onnx-runtime-ep-cpu/src/kernels/stft.rs
Original file line number Diff line number Diff line change
Expand Up @@ -263,11 +263,10 @@ fn positive_scalar(name: &str, input: &TensorView<'_>) -> Result<usize> {
#[cfg(test)]
mod tests {
use super::*;
use crate::kernels::dft::DFT_FFT_TEST_HITS;
use crate::kernels::dft::fast_path_hits;
use crate::kernels::testutil::Owned;
use onnx_runtime_ep_api::TensorView;
use onnx_runtime_ir::{Attribute, NodeId};
use std::sync::atomic::Ordering;

fn node(onesided: i64) -> Node {
let mut node = Node::new(NodeId(0), "STFT", vec![], vec![]);
Expand Down Expand Up @@ -353,16 +352,17 @@ mod tests {
let signal = Owned::f32(&[1, 8, 1], &values);
let step = Owned::i64(&[], &[2]);
let length = Owned::i64(&[], &[4]);
let before = DFT_FFT_TEST_HITS.load(Ordering::Relaxed);
let before = fast_path_hits();
let output = execute(&signal, &step, None, Some(&length), 0, &[1, 3, 4, 2]).unwrap();
let after = DFT_FFT_TEST_HITS.load(Ordering::Relaxed);
let after = fast_path_hits();

let input: Vec<f64> = values.iter().map(|&value| value as f64).collect();
assert_close(&output.to_f32(), &reference(&input, 1, 2, 4, None, false));
assert_eq!(output.shape[1], 3, "the last eligible frame must be kept");
assert!(
after >= before + 3,
"each power-of-two frame must use the radix-2 FFT path"
"each power-of-two frame must take a fast path rather than the naive DFT \
(radix-2 everywhere, vDSP on Apple targets); before={before} after={after}"
);
// The middle frame starts at sample 2. A non-overlapping increment
// would instead transform samples 4..8 and fail this comparison.
Expand Down
Loading