Repository navigation
feat(cuda): add cuFFT-backed DFT - #2080
Conversation
Implement f32 real and complex DFT with arbitrary-axis packing, inverse normalization, onesided output, governed workspace, and a bounded stream-bound cuFFT plan cache. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: d60eb808-7cc6-4abc-b48d-2a6dd3841624
Codecov Report❌ Patch coverage is
Additional details and impacted files@@ Coverage Diff @@
## main #2080 +/- ##
==========================================
+ Coverage 80.27% 80.77% +0.49%
==========================================
Files 422 428 +6
Lines 198030 212725 +14695
Branches 198030 212725 +14695
==========================================
+ Hits 158964 171818 +12854
- Misses 33474 35192 +1718
- Partials 5592 5715 +123
Flags with carried forward coverage won't be shown. Click here to find out more.
🚀 New features to boost your workflow:
|
|
Heads up: this left
Fix is #2100, in your declared direction rather than around it — the kernel already returns No GPU needed to reproduce: it is a source-level contract check, which is why it runs on the no-CUDA lane. |
## Summary - add CUDA `ai.onnx::STFT` v17 using the cuFFT foundation merged in #2080 - preserve the exact CPU/reference and shared shape contracts merged in #2083 - fuse frame extraction, optional windowing, real-to-complex promotion, and batch packing into one NVRTC kernel - execute every `(signal batch, frame)` transform in one cuFFT PlanMany call, then unpack the requested spectrum with one NVRTC kernel ## Claimed surface The CUDA EP claims contiguous f32 signal/output tensors and an optional contiguous f32 window. `frame_step` is a required Int32/Int64 scalar; `frame_length` is an optional Int32/Int64 scalar. Signal shape is `[batch, signal_length, 1|2]`. Real and complex input, full spectrum, and real-input `onesided` (default 1) are supported. Frames are complete, uncentered, and unpadded, exactly matching #2083/ONNX v17. All frames across all signal batches are packed and submitted as one FFT batch. Arbitrary positive frame lengths, including non-power-of-two lengths, are supported by cuFFT. f16/bf16/f64 data, strided signal/window layouts, complex `onesided=1`, malformed scalar ranks/dtypes, zero signal/window extents, and invalid attributes decline at claim time where metadata permits; dynamic zero `frame_step`, mismatched window/frame length, and short signals return typed runtime errors. ## Foundation reuse - Reuses #2080's dynamically loaded cuFFT API, RAII plan ownership, stream binding, external work area, global counters, and bounded 16-entry LRU. DFT and STFT factories receive the same cache instance; no second FFT wrapper/cache exists. - Reuses #2083's already-shared native STFT shape rule and plugin adapter. No shape-rule or adapter census changes were needed. - The plan key includes device, f32/input kind, rank/axis layout, frame length, total FFT batch, and direction. ## Workspace and capture cuFFT auto-allocation remains disabled. Packed frames plus the cuFFT work area come from executor-governed step workspace with overflow-checked offsets and 256-byte alignment. There is no DFT/STFT-side `cudaMalloc`; cuFFT may retain its documented opaque plan/JIT state. CUDA graph capture fails closed with a precise reason: required runtime scalar reads and plan selection are not capture-safe in this implementation. ## GPU validation RTX 4060 Laptop GPU, driver 591.55, pinned CUDA 13.1 wheel environment, serialized: - `cargo test -p onnx-runtime-ep-cuda --features gpu-tests --test stft_gpu -- --test-threads=1` — **9 passed, 0 failed, 0 ignored** - shared cuFFT DFT regression suite — **7 passed, 0 failed, 0 ignored** - CUDA STFT lib/unit selection — **2 passed** - unchanged CPU STFT oracle selection — **13 passed** - shared shape inference STFT selection — **3 passed** - shared plugin shape agreement suite — **4 passed** - source-derived gap guard and covered-op duplicate guard — **1 passed each** - conformance-profile duplicate guard — **1 passed** - `cargo clippy -p onnx-runtime-ep-cuda --all-targets --features gpu-tests -- -D warnings` — passed - `cargo check -p onnx-runtime-ep-cuda-plugin --features cuda` — passed The broad pre-existing `every_covered_op_has_a_conformance_entry` guard remains red only for `PagedAttention`, `Mish`, `Celu`, and `TensorScatter`; the new STFT entry is present and is not in the missing set. ## Mutation evidence Each mutation was applied independently, caught, and reverted: - ignored the window -> `nontrivial_window_is_applied_and_matches_cpu` failed - advanced by `frame_length` instead of `frame_step` -> `overlapping_step_selects_the_middle_frame` failed - omitted the final complete frame -> `final_complete_frame_is_not_dropped` failed - emitted `N/2` instead of `N/2+1` bins -> `onesided_keeps_n_over_two_plus_one_and_matches_full_prefix` failed ## Structural performance/memory evidence For `[batch=2, signal=16, real]`, `frame_length=5`, `frame_step=3`: - frames per signal: **4** - one cuFFT batch: **8 transforms** - explicit custom launches: **2** (fused frame+window pack, output unpack) - cuFFT API executions: **1** (vendor-internal kernel count was not profiled) - governed workspace: **512 bytes** (320 bytes packed complex data, aligned; cuFFT reported zero additional work bytes for this geometry) No latency or throughput claim is made. Remaining optimization opportunities are using the full-spectrum output directly as the in-place FFT buffer, an R2C onesided path, and pre-resolved constant scalars/plans for capture-safe execution. ## Not verified - Linux or H100/H200 execution - CUDA graph capture (explicitly unsupported) - f16/bf16/f64 CUDA STFT - strided CUDA signal/window execution (explicitly declined) - latency/throughput versus CPU, PyTorch, or other GPU FFT stacks - cuFFT vendor-internal kernel launch count Co-authored-by: justinchuby <223556219+Copilot@users.noreply.github.com> Copilot-Session: d60eb808-7cc6-4abc-b48d-2a6dd3841624
…mise (#2099) #2080 added an unconditional `stream().synchronize()` to `dft.rs::run` without listing it in the capture-sync contract, so `CUDA compile (Linux x86_64)` has been red on main and on every branch cut from it. The set difference is exactly one entry. The sync is legitimate: `DftKernel::capture_support` returns `CaptureSupport::unsupported`, and the barrier drains the compute stream before the synchronous default-stream metadata upload so a prior DFT's step-scoped metadata cannot be overwritten. Listed with that justification. The allowlist also carried its own justification in a comment that nothing checked, which made it an unconditional escape hatch: one line silences the contract for a kernel advertising `CaptureSupport::Supported`, and graph capture then breaks with the suite green. `every_allowlisted_file_can_decline_capture` turns that sentence into an assertion, and documents its own limit — it resolves `capture_support` per file, not per kernel, so it is a lower bound. Falsified both directions without a GPU (the contract test is a source scan and runs without `--features cuda`, which is how CI's failure was reproduced here byte-identically): dropping the entry fails the contract test with CI's exact left/right sets; making `dft.rs` advertise `Supported` while still listed fails the new test, while the pre-existing contract test still reports ok. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
#2105) Closes #2104. ## The defect The `changes` job classified a PR with a **two-dot** diff against `github.event.pull_request.base.sha`. That SHA is the base branch tip **at the moment the event fired**, not the PR's merge-base — so the diff also reports every file main gained since the branch diverged, with the sign inverted, since the branch "lacks" them. A pure-markdown PR is then classified as code and takes the full matrix. Found it by checking my own PR #2103, which changes exactly one `.md` file, and seeing all nine lanes queue. | range | files | |---|---| | `git diff --name-only $BASE $HEAD` (two-dot, before) | **9** | | `git diff --name-only $BASE...$HEAD` (three-dot, after) | **1** | | `gh pr view 2103 --json files` — ground truth | **1** | The eight extras are exactly #2091's file list — and `base.sha` **is** #2091's merge commit. (I first wrote #2080 here; see the review response in the comments for how that wrong citation was produced and verified away.) ## It is intermittent, and that is the interesting part My first inference was that the fast path essentially never fires. **That was wrong, and I checked before writing it down.** #2070 (three `.md` files) classified `docs_only=true` and skipped all nine lanes — its branch happened to be level with main when the event fired, so two-dot and three-dot agreed. The trigger is precisely *main gained a commit between the PR's merge-base and the event*. Which means the fast path works whenever you go looking at it on a quiet tree, and silently doesn't on a busy one. Same shape as the stale-base problem: the answer depends on where main was standing, not on the PR. ## Why it is worth fixing even though it is safe The direction is **fail-closed** — more CI, never less — so this is cost and latency, not correctness. It matters because runner capacity is the binding constraint: queue depth has been in the 50s today, and a docs PR taking the full nine-lane matrix displaces work that actually needs it. ## Safety - `push` **stays two-dot** — `before..after` is the push itself, and three-dot would be wrong there. - Failure behaviour is unchanged. If the merge-base is unavailable the diff errors, `|| true` leaves `files` empty, and the existing `elif [ -n "$files" ]` leaves `docs_only=false` → full CI. `fetch-depth: 0` is already set on this job. - The #2077 guard (a `.md` compiled in by `include_str!` is source, not docs) is untouched and is arm C below. ## Battery — 9/9 The step's script was extracted from the YAML and run under the runner's own shell (`bash --noprofile --norc -eo pipefail`), against real SHAs: | arm | case | old | new | |---|---|---|---| | A | docs-only PR, base advanced | `false` ← **the bug** | `true` | | B | genuine code PR | — | `false` | | C | pure `.md` edit to an `include_str!` target (#2077) | — | `false` | | C2 | pure `.md` edit, **not** an embed target — the control | — | `true` | | C3 | mixed `.md` + `.rs` | — | `false` | | D | missing SHAs | — | `false` | | E | unresolvable SHA | — | `false` | | F | unknown event (`schedule`) | — | `false` | | G | push event | `false` | `false` (unchanged) | **C and C2 are the pair that carries the claim.** Without C2, arm C passes on a classifier that answers "code" to literally everything. My first C2 fixture used `docs/execution/CUDA_COVERAGE.md` and the arm failed — that file is *itself* one of the two `.md` embed targets, reached by a relative path. The fixture was wrong, not the code. Worth stating because a control that fails looks exactly like a defect, and I nearly filed it as one. Arm A's `old` column is the positive control for the harness: it proves the battery can produce a FAIL rather than agreeing with everything. ## Expected CI on this PR This PR changes `ci.yml`, so it is code — it must run the **full** matrix, and `docs_only=false` here is the correct answer, not a symptom. #2103 is the one that should flip to docs-only once this lands. Requesting independent Opus review. Co-authored-by: Holden <holden@users.noreply.github.com> Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Summary
ai.onnx::DFTfor f32 real/complex input, forward/inverse, full/onesided output, arbitrary signal axes, and truncating/zero-paddingdft_lengthClaimed surface
The CUDA EP claims contiguous f32 DFT data/output only. Optional
dft_lengthand opset-20axisinputs are Int64 scalars. Rank >= 2, every non-component signal axis, real (last_dim=1) and complex (last_dim=2) input, forward/inverse, full spectrum, and real-inputonesidedare covered. f16, bf16, f64, strided input, and complex-inputonesided=1decline at placement time.Arbitrary ONNX layouts are packed into contiguous
[batch, N]complex f32 storage, transformed with cuFFT C2C, then unpacked. This also centralizes truncation/zero-padding and applies ONNX inverse normalization exactly once.Plans, memory, and capture
Plans are RAII-owned, bound to the EP compute stream, and cached by device, dtype/input kind, rank/axis layout, length, batch, and direction. The LRU holds at most 16 entries; tests force and observe eviction. cuFFT auto-allocation is disabled. Metadata, packed complex data, and the cuFFT work area all come from executor-governed step workspace. cuFFT may retain opaque vendor plan/JIT state, but there is no hidden per-dispatch
cudaMallocfrom the DFT path.CUDA graph capture fails closed: plan selection plus runtime scalar/metadata staging are not capture-safe. The kernel reports that reason through
capture_support().Validation
RTX 4060 Laptop GPU, driver 591.55, serialized, with a registry-rebuilt CUDA 13.1 environment (
nvidia-cufft==12.1.0.78,nvidia-nvjitlink==13.1.115):cargo test -p onnx-runtime-ep-cuda --features gpu-tests --test dft_gpu -- --test-threads=1— 7 passed, 0 failed, 0 ignoredcargo clippy -p onnx-runtime-ep-cuda -p onnx-runtime-ep-cpu --all-targets -- -D warnings— passedcargo check -p onnx-runtime-ep-cuda-plugin --features cuda— passedThe GPU matrix covers explicit N=4 sign convention, inverse normalization, real full-spectrum conjugate symmetry, N/2+1 onesided output, shorter/longer
dft_length, default and non-default batched axes, CPU parity, claim-time declines, plan reuse/eviction, and capture refusal.Mutation checks, each reverted:
1/Nscale -> inverse normalization test failedN/2onesided bins -> N/2+1 test failedThe broader
every_covered_op_has_a_conformance_entryguard remains red on unchanged pre-existing entriesPagedAttention,Mish,Celu, andTensorScatter; the new DFT profile entry is present and is not in that missing set.Dependency and license
PyPI registry metadata was checked before editing.
nvidia-cufft==12.1.0.78andnvidia-nvjitlink==13.1.115both publish Windows x64 and Linux x86_64/aarch64 wheels on the CUDA 13.1 release date. They remain separately installed NVIDIA runtime dependencies governed by NVIDIA's license; this repository does not redistribute or bundle their binaries.STFT follow-up
This PR intentionally stops at the reusable packed-batch FFT/plan foundation. ONNX STFT still needs framing/windowing, optional window input handling, hop/frame-step semantics, output-layout construction, CPU wiring, CUDA wiring, and plan reuse across frames.
Not verified