Skip to main content

tenferro_fft/
lib.rs

1//! FFT extension operations for tenferro.
2//!
3//! This crate is an out-of-tree `ExtensionOp` package with an explicit
4//! [`FftBackend`] capability. [`tenferro_cpu::CpuBackend`] implements the
5//! capability through RustFFT. With the `webgpu` feature,
6//! `tenferro_gpu::webgpu::WebGpuBackend` executes C32 CFFT, F32 one-sided RFFT, and
7//! C32-to-F32 IRFFT through CubeK on its existing WebGPU placement. That first
8//! GPU path supports power-of-two lengths only; unsupported operations and
9//! dtypes return an error and never fall back to CPU or transfer tensor data.
10//! With the `cuda` feature, `tenferro_gpu::cuda::CudaBackend` executes the
11//! supported one-dimensional F32/F64/C32/C64 operations through dynamically
12//! loaded cuFFT without implicit transfers or CPU fallback. The vendor call
13//! synchronizes at the cuFFT FFI boundary; subsequent CUDA postprocessing and
14//! explicit download remain stream-managed. On macOS,
15//! `tenferro_gpu::apple::AppleContext` pairs that Metal backend with a
16//! domain-bound CPU RustFFT backend. Backend choice remains explicit, while
17//! matching managed tensors can be used without an intervening download.
18//! Concrete non-AD execution uses
19//! [`TensorFftExt`] and [`TensorReadFftExt`]. Eager FFTs use
20//! `EagerSessionFftExt` on a borrowed session when `autodiff` is enabled;
21//! consuming in-place transforms retain `EagerTensorFftExt`. Traced graph
22//! construction uses [`TracedTensorFftExt`].
23//!
24//! # Cargo features
25//!
26//! | Feature | Enables |
27//! |---|---|
28//! | `cpu-faer` (default) and the other CPU provider features | Forwarded to `tenferro-cpu`; see its documentation. |
29//! | `autodiff` | The eager surface (`EagerSessionFftExt`, `EagerTensorFftExt`) and AD rules. Adds the `tenferro-ad` dependency. |
30//! | `cuda` | CUDA execution through `tenferro-gpu`. |
31//! | `webgpu` | WebGPU/Metal execution through `tenferro-gpu` (a subset of operations). |
32//! | `rocm` | Placeholder; HIP/ROCm is not implemented. |
33//!
34//! For transforms without AD, use [`TensorFftExt`] / [`TensorReadFftExt`] on
35//! concrete tensors inside a backend session; they need no `autodiff`.
36//!
37//! # Examples
38//!
39//! ```
40//! use num_complex::Complex64;
41//! use tenferro_cpu::CpuBackend;
42//! use tenferro_runtime::{GraphCompiler, Runtime, TracedTensor};
43//! use tenferro_fft::{FftNorm, TracedTensorFftExt};
44//!
45//! let x = TracedTensor::from_vec_col_major(
46//!     vec![4],
47//!     vec![
48//!         Complex64::new(1.0, 0.0),
49//!         Complex64::new(2.0, 0.0),
50//!         Complex64::new(3.0, 0.0),
51//!         Complex64::new(4.0, 0.0),
52//!     ],
53//! )
54//! .unwrap();
55//! let y = x.fft(None, -1, FftNorm::Backward).unwrap();
56//!
57//! let mut compiler = GraphCompiler::new();
58//! let program = compiler.compile(&y).unwrap();
59//! let backend = CpuBackend::new();
60//! let engine_id = tenferro_cpu::runtime_engine_id().unwrap();
61//! let mut builder = Runtime::builder();
62//! builder
63//!     .register_engine(tenferro_cpu::runtime_engine_registration(&backend).unwrap())
64//!     .unwrap();
65//! builder
66//!     .install_extension_module(tenferro_fft::extension_module::<CpuBackend>(engine_id).unwrap())
67//!     .unwrap();
68//! let runtime = builder.build().unwrap();
69//! let out = runtime.run_compiled(&program, &[]).unwrap().pop().unwrap();
70//! assert_eq!(out.shape(), &[4]);
71//! assert_eq!(out.as_slice::<Complex64>().unwrap()[0], Complex64::new(10.0, 0.0));
72//! ```
73//!
74//! ```
75//! # #[cfg(all(feature = "webgpu", target_os = "macos"))]
76//! # {
77//! use num_complex::Complex32;
78//! use tenferro_cpu::CpuBackend;
79//! use tenferro_fft::{FftNorm, TensorFftExt};
80//! use tenferro_gpu::apple::AppleContext;
81//! use tenferro_tensor::{BackendSessionHost, Tensor};
82//!
83//! if let Ok(context) = AppleContext::new() {
84//!     let host = Tensor::from_vec_col_major(
85//!         vec![4],
86//!         vec![Complex32::new(1.0, 0.0); 4],
87//!     ).unwrap();
88//!     let input = context.upload_tensor(&host).unwrap();
89//!     let after_creation = context.transfer_stats();
90//!     let mut cpu = context.cpu_backend().clone();
91//!     let cpu_output = cpu
92//!         .with_backend_session(|session| input.fft(None, 0, FftNorm::Backward, session))?
93//!         .unwrap();
94//!     let mut metal = context.metal_backend().clone();
95//!     let output = metal
96//!         .with_backend_session(|session| input.fft(None, 0, FftNorm::Backward, session))?
97//!         .unwrap();
98//!     metal.synchronize().unwrap();
99//!     assert_eq!(output.shape(), &[4]);
100//!     assert_eq!(cpu_output.shape(), output.shape());
101//!     assert_eq!(context.transfer_stats(), after_creation);
102//! }
103//! # }
104//! # Ok::<(), Box<dyn std::error::Error>>(())
105//! ```
106//!
107//! ```
108//! use num_complex::Complex64;
109//! use tenferro_cpu::CpuBackend;
110//! use tenferro_fft::{FftNorm, TensorFftExt};
111//! use tenferro_tensor::{BackendSessionHost, Tensor};
112//!
113//! let x = Tensor::from_vec_col_major(vec![4], vec![1.0_f64, 2.0, 3.0, 4.0]).unwrap();
114//! let mut backend = CpuBackend::new();
115//! let out = backend
116//!     .with_backend_session(|session| x.fft(None, -1, FftNorm::Backward, session))?
117//!     .unwrap();
118//!
119//! assert_eq!(out.as_slice::<Complex64>().unwrap()[0], Complex64::new(10.0, 0.0));
120//! # Ok::<(), Box<dyn std::error::Error>>(())
121//! ```
122#![cfg_attr(docsrs, feature(doc_cfg))]
123
124use std::any::Any;
125use std::hash::Hasher;
126use std::num::NonZeroUsize;
127use std::sync::Arc;
128
129#[cfg(feature = "autodiff")]
130use tenferro_ad::semantic_extension::{
131    AdValue, ResidualSpec, SemanticAdError, SemanticExtensionRegistryError,
132    SemanticExtensionRuleSet, SemanticLinearTransposeRequest, SemanticLinearTransposeRule,
133    SemanticLinearizeRequest, SemanticLinearizeResult, SemanticLinearizeRule,
134    SemanticPrimalVjpRequest, SemanticPrimalVjpRule,
135};
136use tenferro_cpu::with_cpu_exec_session;
137use tenferro_extension_macros::define_extension_runtime;
138#[cfg(feature = "cuda")]
139use tenferro_gpu::cuda::{with_cuda_exec_session, CudaBackend};
140#[cfg(feature = "webgpu")]
141use tenferro_gpu::webgpu::with_webgpu_exec_session;
142use tenferro_ops::SymDim;
143use tenferro_runtime::extension::{
144    apply, ExtensionCacheStore, ExtensionExecutionContext, ExtensionOp,
145};
146#[cfg(feature = "autodiff")]
147use tenferro_runtime::program::{CoreSemanticOp, ProgramValue, SemanticProgramBuilder};
148use tenferro_runtime::{Error, ErrorPhase, Result, TracedTensor};
149use tenferro_tensor::{
150    BackendSession, CacheStats, DType, ErrorKind, Tensor, TensorBackend, TensorRead,
151    ValidationError,
152};
153
154mod backend;
155mod cache;
156mod cpu;
157#[cfg(feature = "cuda")]
158mod cuda;
159#[cfg(feature = "autodiff")]
160mod eager_ext;
161#[cfg(feature = "autodiff")]
162mod eager_in_place;
163pub mod prelude;
164mod spec;
165#[cfg(feature = "webgpu")]
166mod webgpu;
167
168pub use backend::{FftBackend, FftExecutionCache};
169pub use cache::{
170    fft_plan_cache_selector, FftPlanCache, DEFAULT_FFT_PLAN_CACHE_CAPACITY, FFT_PLAN_CACHE_NAME,
171};
172#[cfg(feature = "autodiff")]
173#[cfg_attr(docsrs, doc(cfg(feature = "autodiff")))]
174pub use eager_ext::{EagerSessionFftExt, EagerTensorFftExt};
175#[cfg(feature = "autodiff")]
176#[cfg_attr(docsrs, doc(cfg(feature = "autodiff")))]
177pub use eager_in_place::EagerFftInPlaceError;
178pub use spec::{FftNorm, FftOperation, FftPlanSpec};
179
180/// Extension family id used by the tenferro FFT extension.
181///
182/// # Examples
183///
184/// ```
185/// assert_eq!(
186///     tenferro_fft::FFT_EXTENSION_FAMILY_ID,
187///     "tenferro-fft.fft.v1"
188/// );
189/// ```
190pub const FFT_EXTENSION_FAMILY_ID: &str = "tenferro-fft.fft.v1";
191
192/// Reusable concrete FFT executor with an explicitly owned backend-neutral cache.
193///
194/// Use this executor for repeated concrete FFT calls that should reuse backend
195/// plans. The immediate [`TensorFftExt`] and [`TensorReadFftExt`] methods stay
196/// one-shot and do not retain hidden process-global, thread-local, or
197/// backend-owned plan state between calls.
198#[derive(Default)]
199pub struct FftExecutor {
200    plans: FftPlanCache,
201}
202
203impl FftExecutor {
204    /// Create an executor from a caller-configured FFT execution cache.
205    pub fn new(plans: FftPlanCache) -> Self {
206        Self { plans }
207    }
208
209    /// Inspect the owned backend-neutral FFT cache.
210    pub const fn plan_cache(&self) -> &FftPlanCache {
211        &self.plans
212    }
213
214    /// Mutably inspect or configure the owned backend-neutral FFT cache.
215    pub fn plan_cache_mut(&mut self) -> &mut FftPlanCache {
216        &mut self.plans
217    }
218
219    /// Snapshot aggregate statistics for every backend cache namespace.
220    pub fn cache_stats(&self) -> CacheStats {
221        self.plans.stats()
222    }
223
224    /// Remove every retained backend plan or workspace from this executor.
225    pub fn clear_cache(&mut self) {
226        self.plans.clear();
227    }
228
229    /// Execute a complex or full-spectrum real FFT while reusing owned plans.
230    ///
231    /// # Errors
232    ///
233    /// Returns [`tenferro_tensor::Error::Validation`] with `AxisOutOfBounds` or
234    /// `InvalidArgument` for invalid `axis`/`n`,
235    /// [`tenferro_tensor::Error::Extension`] with [`ErrorKind::Unsupported`]
236    /// for unsupported dtypes, a typed capability error when the session does
237    /// not expose an FFT execution capability, or a typed backend source for
238    /// execution.
239    pub fn fft(
240        &mut self,
241        input: &Tensor,
242        n: Option<usize>,
243        axis: isize,
244        norm: FftNorm,
245        session: &mut dyn BackendSession,
246    ) -> tenferro_tensor::Result<Tensor> {
247        self.execute(
248            input,
249            concrete_fft_operation("FftExecutor::fft", input.dtype())?,
250            "FftExecutor::fft",
251            n,
252            axis,
253            norm,
254            session,
255        )
256    }
257
258    /// Execute an inverse complex FFT while reusing owned plans.
259    ///
260    /// # Errors
261    ///
262    /// Returns [`tenferro_tensor::Error::Validation`] with `AxisOutOfBounds` or
263    /// `InvalidArgument` for invalid `axis`/`n`,
264    /// [`tenferro_tensor::Error::Extension`] with [`ErrorKind::Unsupported`]
265    /// for a non-complex input, a typed capability error when the session does
266    /// not expose an FFT execution capability, or a typed backend source for
267    /// execution.
268    pub fn ifft(
269        &mut self,
270        input: &Tensor,
271        n: Option<usize>,
272        axis: isize,
273        norm: FftNorm,
274        session: &mut dyn BackendSession,
275    ) -> tenferro_tensor::Result<Tensor> {
276        self.execute(
277            input,
278            concrete_ifft_operation("FftExecutor::ifft", input.dtype())?,
279            "FftExecutor::ifft",
280            n,
281            axis,
282            norm,
283            session,
284        )
285    }
286
287    /// Execute a real FFT while reusing owned plans.
288    ///
289    /// # Errors
290    ///
291    /// Returns [`tenferro_tensor::Error::Validation`] with `AxisOutOfBounds` or
292    /// `InvalidArgument` for invalid `axis`/`n`,
293    /// [`tenferro_tensor::Error::Extension`] with [`ErrorKind::Unsupported`]
294    /// for a non-real input, a typed capability error when the session does not
295    /// expose an FFT execution capability, or a typed backend source for
296    /// execution.
297    pub fn rfft(
298        &mut self,
299        input: &Tensor,
300        n: Option<usize>,
301        axis: isize,
302        norm: FftNorm,
303        session: &mut dyn BackendSession,
304    ) -> tenferro_tensor::Result<Tensor> {
305        self.execute(
306            input,
307            concrete_rfft_operation("FftExecutor::rfft", input.dtype())?,
308            "FftExecutor::rfft",
309            n,
310            axis,
311            norm,
312            session,
313        )
314    }
315
316    /// Execute an inverse real FFT while reusing owned plans.
317    ///
318    /// # Errors
319    ///
320    /// Returns [`tenferro_tensor::Error::Validation`] with `AxisOutOfBounds`,
321    /// `InvalidArgument`, or spectrum-length details,
322    /// [`tenferro_tensor::Error::Extension`] with [`ErrorKind::Unsupported`]
323    /// for a non-complex input, a typed capability error when the session does
324    /// not expose an FFT execution capability, or a typed backend source for
325    /// execution.
326    pub fn irfft(
327        &mut self,
328        input: &Tensor,
329        n: Option<usize>,
330        axis: isize,
331        norm: FftNorm,
332        session: &mut dyn BackendSession,
333    ) -> tenferro_tensor::Result<Tensor> {
334        self.execute(
335            input,
336            concrete_irfft_operation("FftExecutor::irfft", input.dtype())?,
337            "FftExecutor::irfft",
338            n,
339            axis,
340            norm,
341            session,
342        )
343    }
344
345    #[allow(clippy::too_many_arguments)]
346    fn execute(
347        &mut self,
348        input: &Tensor,
349        operation: FftOperation,
350        op_name: &'static str,
351        n: Option<usize>,
352        axis: isize,
353        norm: FftNorm,
354        session: &mut dyn BackendSession,
355    ) -> tenferro_tensor::Result<Tensor> {
356        let spec = concrete_fft_spec(
357            op_name,
358            operation,
359            input.dtype(),
360            input.shape(),
361            n,
362            axis,
363            norm,
364        )?;
365        // The executor calls the concrete backend directly (no internal
366        // session entry); the built-in dispatch only bridges the borrowed
367        // session to its FFT execution capability.
368        with_fft_exec_session(session, op_name, |backend| {
369            backend.execute_fft(
370                input,
371                &spec,
372                FftExecutionCache::caller_owned(&mut self.plans),
373            )
374        })
375    }
376}
377
378/// FFT extension methods for [`TracedTensor`].
379pub trait TracedTensorFftExt {
380    /// Build a traced complex or full-spectrum real FFT.
381    ///
382    /// # Errors
383    ///
384    /// Returns `Error::Validation` with `AxisOutOfBounds` or
385    /// `InvalidArgument` for invalid `axis`/`n`, or `Error::Extension` with
386    /// `ErrorKind::Unsupported` for integer, boolean, or otherwise unsupported
387    /// dtypes.
388    ///
389    /// # Deferred errors
390    ///
391    /// Symbolic axis extents and extension execution failures are checked at
392    /// compile or execution time after concrete inputs are bound.
393    fn fft(&self, n: Option<usize>, axis: isize, norm: FftNorm) -> Result<TracedTensor>;
394
395    /// Build a traced inverse complex FFT.
396    ///
397    /// # Errors
398    ///
399    /// Returns `Error::Validation` with `AxisOutOfBounds` or
400    /// `InvalidArgument` for invalid `axis`/`n`, or `Error::Extension` with
401    /// `ErrorKind::Unsupported` when the input is not `C32`/`C64`.
402    ///
403    /// # Deferred errors
404    ///
405    /// Symbolic shape and extension execution failures may be deferred to
406    /// compile or execution.
407    fn ifft(&self, n: Option<usize>, axis: isize, norm: FftNorm) -> Result<TracedTensor>;
408
409    /// Build a traced one-sided real FFT.
410    ///
411    /// # Errors
412    ///
413    /// Returns `Error::Validation` with `AxisOutOfBounds` or
414    /// `InvalidArgument` for invalid `axis`/`n`, or `Error::Extension` with
415    /// `ErrorKind::Unsupported` when the input is not `F32`/`F64`.
416    ///
417    /// # Deferred errors
418    ///
419    /// Symbolic shape and extension execution failures may be deferred to
420    /// compile or execution.
421    fn rfft(&self, n: Option<usize>, axis: isize, norm: FftNorm) -> Result<TracedTensor>;
422
423    /// Build a traced inverse one-sided real FFT.
424    ///
425    /// # Errors
426    ///
427    /// Returns `Error::Validation` with `AxisOutOfBounds` or
428    /// `InvalidArgument` for invalid `axis`/`n` or spectrum length, or
429    /// `Error::Extension` with `ErrorKind::Unsupported` for non-complex input.
430    ///
431    /// # Deferred errors
432    ///
433    /// Symbolic spectrum lengths and extension execution failures may be
434    /// deferred to compile or execution.
435    fn irfft(&self, n: Option<usize>, axis: isize, norm: FftNorm) -> Result<TracedTensor>;
436}
437
438impl TracedTensorFftExt for TracedTensor {
439    fn fft(&self, n: Option<usize>, axis: isize, norm: FftNorm) -> Result<TracedTensor> {
440        fft(self, n, axis, norm)
441    }
442
443    fn ifft(&self, n: Option<usize>, axis: isize, norm: FftNorm) -> Result<TracedTensor> {
444        ifft(self, n, axis, norm)
445    }
446
447    fn rfft(&self, n: Option<usize>, axis: isize, norm: FftNorm) -> Result<TracedTensor> {
448        rfft(self, n, axis, norm)
449    }
450
451    fn irfft(&self, n: Option<usize>, axis: isize, norm: FftNorm) -> Result<TracedTensor> {
452        irfft(self, n, axis, norm)
453    }
454}
455
456/// Backend-explicit FFT methods for concrete [`Tensor`] values.
457///
458/// This is the non-AD immediate execution surface. It uses unsuffixed method
459/// names because the receiver is an owned compact tensor value. Use
460/// [`TensorReadFftExt`] when the input is a borrowed view or other
461/// [`TensorRead`] value.
462///
463/// Direct calls intentionally use a call-local one-shot FFT plan cache. Use
464/// [`FftExecutor`] for repeated concrete calls with stable transform lengths,
465/// or traced/runtime execution when the runtime should own the extension cache.
466///
467/// # Examples
468///
469/// ```
470/// use num_complex::Complex64;
471/// use tenferro_cpu::CpuBackend;
472/// use tenferro_fft::{FftNorm, TensorFftExt};
473/// use tenferro_tensor::{BackendSessionHost, Tensor};
474///
475/// let input = Tensor::from_vec_col_major(vec![4], vec![1.0_f64, 2.0, 3.0, 4.0])?;
476/// let mut backend = CpuBackend::new();
477///
478/// let spectrum = backend
479///     .with_backend_session(|session| input.fft(None, -1, FftNorm::Backward, session))??;
480/// assert_eq!(spectrum.shape(), &[4]);
481/// assert_eq!(spectrum.as_slice::<Complex64>()?[0], Complex64::new(10.0, 0.0));
482/// # Ok::<(), tenferro_tensor::Error>(())
483/// ```
484pub trait TensorFftExt {
485    /// Execute a one-dimensional FFT along `axis`.
486    ///
487    /// # Errors
488    ///
489    /// Returns `Error::Validation` with `AxisOutOfBounds` or `InvalidArgument`
490    /// for `axis`/`n`, `Error::Extension` with `ErrorKind::Unsupported` for an
491    /// integer or boolean input, a typed capability error when the session
492    /// does not expose an FFT execution capability, or a typed backend source
493    /// for execution.
494    fn fft(
495        &self,
496        n: Option<usize>,
497        axis: isize,
498        norm: FftNorm,
499        session: &mut dyn BackendSession,
500    ) -> tenferro_tensor::Result<Tensor>;
501
502    /// Execute a one-dimensional inverse FFT along `axis`.
503    ///
504    /// # Errors
505    ///
506    /// Returns `Error::Validation` with `AxisOutOfBounds` or `InvalidArgument`
507    /// for `axis`/`n`, `Error::Extension` with `ErrorKind::Unsupported` for a
508    /// non-complex input, a typed capability error when the session does not
509    /// expose an FFT execution capability, or a typed backend source for
510    /// execution.
511    fn ifft(
512        &self,
513        n: Option<usize>,
514        axis: isize,
515        norm: FftNorm,
516        session: &mut dyn BackendSession,
517    ) -> tenferro_tensor::Result<Tensor>;
518
519    /// Execute a one-dimensional real FFT along `axis`.
520    ///
521    /// # Errors
522    ///
523    /// Returns `Error::Validation` with `AxisOutOfBounds` or `InvalidArgument`
524    /// for `axis`/`n`, `Error::Extension` with `ErrorKind::Unsupported` for a
525    /// non-`F32`/`F64` input, a typed capability error when the session does
526    /// not expose an FFT execution capability, or a typed backend source for
527    /// execution.
528    fn rfft(
529        &self,
530        n: Option<usize>,
531        axis: isize,
532        norm: FftNorm,
533        session: &mut dyn BackendSession,
534    ) -> tenferro_tensor::Result<Tensor>;
535
536    /// Execute a one-dimensional inverse real FFT along `axis`.
537    ///
538    /// # Errors
539    ///
540    /// Returns `Error::Validation` with `AxisOutOfBounds`, `InvalidArgument`,
541    /// or spectrum-length details, `Error::Extension` with
542    /// `ErrorKind::Unsupported` for a non-complex input, a typed capability
543    /// error when the session does not expose an FFT execution capability, or
544    /// a typed backend source for execution.
545    fn irfft(
546        &self,
547        n: Option<usize>,
548        axis: isize,
549        norm: FftNorm,
550        session: &mut dyn BackendSession,
551    ) -> tenferro_tensor::Result<Tensor>;
552}
553
554impl TensorFftExt for Tensor {
555    fn fft(
556        &self,
557        n: Option<usize>,
558        axis: isize,
559        norm: FftNorm,
560        session: &mut dyn BackendSession,
561    ) -> tenferro_tensor::Result<Tensor> {
562        let spec = concrete_fft_spec(
563            "TensorFftExt::fft",
564            concrete_fft_operation("TensorFftExt::fft", self.dtype())?,
565            self.dtype(),
566            self.shape(),
567            n,
568            axis,
569            norm,
570        )?;
571        with_fft_exec_session(session, "TensorFftExt::fft", |backend| {
572            execute_concrete_fft_op(self, &spec, backend)
573        })
574    }
575
576    fn ifft(
577        &self,
578        n: Option<usize>,
579        axis: isize,
580        norm: FftNorm,
581        session: &mut dyn BackendSession,
582    ) -> tenferro_tensor::Result<Tensor> {
583        let spec = concrete_fft_spec(
584            "TensorFftExt::ifft",
585            concrete_ifft_operation("TensorFftExt::ifft", self.dtype())?,
586            self.dtype(),
587            self.shape(),
588            n,
589            axis,
590            norm,
591        )?;
592        with_fft_exec_session(session, "TensorFftExt::ifft", |backend| {
593            execute_concrete_fft_op(self, &spec, backend)
594        })
595    }
596
597    fn rfft(
598        &self,
599        n: Option<usize>,
600        axis: isize,
601        norm: FftNorm,
602        session: &mut dyn BackendSession,
603    ) -> tenferro_tensor::Result<Tensor> {
604        let spec = concrete_fft_spec(
605            "TensorFftExt::rfft",
606            concrete_rfft_operation("TensorFftExt::rfft", self.dtype())?,
607            self.dtype(),
608            self.shape(),
609            n,
610            axis,
611            norm,
612        )?;
613        with_fft_exec_session(session, "TensorFftExt::rfft", |backend| {
614            execute_concrete_fft_op(self, &spec, backend)
615        })
616    }
617
618    fn irfft(
619        &self,
620        n: Option<usize>,
621        axis: isize,
622        norm: FftNorm,
623        session: &mut dyn BackendSession,
624    ) -> tenferro_tensor::Result<Tensor> {
625        let spec = concrete_fft_spec(
626            "TensorFftExt::irfft",
627            concrete_irfft_operation("TensorFftExt::irfft", self.dtype())?,
628            self.dtype(),
629            self.shape(),
630            n,
631            axis,
632            norm,
633        )?;
634        with_fft_exec_session(session, "TensorFftExt::irfft", |backend| {
635            execute_concrete_fft_op(self, &spec, backend)
636        })
637    }
638}
639
640/// Backend-explicit FFT methods for read-only tensor inputs.
641///
642/// The `_read` suffix follows the repository convention for APIs that
643/// explicitly accept [`TensorRead`] values such as borrowed views.
644///
645/// Direct read calls intentionally materialize through a call-local one-shot
646/// FFT plan cache. Use [`FftExecutor`] on compact owned tensors when repeated
647/// concrete calls should retain backend plans across calls.
648///
649/// # Examples
650///
651/// ```
652/// use num_complex::Complex64;
653/// use tenferro_cpu::CpuBackend;
654/// use tenferro_fft::{FftNorm, TensorReadFftExt};
655/// use tenferro_tensor::{BackendSessionHost, TensorRead, TensorView};
656///
657/// let shape = [4usize];
658/// let data = [1.0_f64, 2.0, 3.0, 4.0];
659/// let input = TensorRead::from_view(TensorView::f64(&shape, &data)?);
660/// let mut backend = CpuBackend::new();
661///
662/// let spectrum = backend
663///     .with_backend_session(|session| input.fft_read(None, -1, FftNorm::Backward, session))??;
664/// assert_eq!(spectrum.as_slice::<Complex64>()?[0], Complex64::new(10.0, 0.0));
665/// # Ok::<(), tenferro_tensor::Error>(())
666/// ```
667pub trait TensorReadFftExt {
668    /// Execute a one-dimensional FFT along `axis`.
669    ///
670    /// # Errors
671    ///
672    /// Returns `Error::Validation` with `AxisOutOfBounds` or `InvalidArgument`
673    /// for `axis`/`n`, `Error::Extension` with `ErrorKind::Unsupported` for an
674    /// integer or boolean input, a typed capability error when the session
675    /// does not expose an FFT execution capability, or a typed backend source
676    /// for materialization or execution.
677    fn fft_read(
678        &self,
679        n: Option<usize>,
680        axis: isize,
681        norm: FftNorm,
682        session: &mut dyn BackendSession,
683    ) -> tenferro_tensor::Result<Tensor>;
684
685    /// Execute a one-dimensional inverse FFT along `axis`.
686    ///
687    /// # Errors
688    ///
689    /// Returns `Error::Validation` with `AxisOutOfBounds` or `InvalidArgument`
690    /// for `axis`/`n`, `Error::Extension` with `ErrorKind::Unsupported` for a
691    /// non-complex input, a typed capability error when the session does not
692    /// expose an FFT execution capability, or a typed backend source for
693    /// materialization.
694    fn ifft_read(
695        &self,
696        n: Option<usize>,
697        axis: isize,
698        norm: FftNorm,
699        session: &mut dyn BackendSession,
700    ) -> tenferro_tensor::Result<Tensor>;
701
702    /// Execute a one-dimensional real FFT along `axis`.
703    ///
704    /// # Errors
705    ///
706    /// Returns `Error::Validation` with `AxisOutOfBounds` or `InvalidArgument`
707    /// for `axis`/`n`, `Error::Extension` with `ErrorKind::Unsupported` for a
708    /// non-`F32`/`F64` input, a typed capability error when the session does
709    /// not expose an FFT execution capability, or a typed backend source for
710    /// materialization.
711    fn rfft_read(
712        &self,
713        n: Option<usize>,
714        axis: isize,
715        norm: FftNorm,
716        session: &mut dyn BackendSession,
717    ) -> tenferro_tensor::Result<Tensor>;
718
719    /// Execute a one-dimensional inverse real FFT along `axis`.
720    ///
721    /// # Errors
722    ///
723    /// Returns `Error::Validation` with `AxisOutOfBounds`, `InvalidArgument`,
724    /// or spectrum-length details, `Error::Extension` with
725    /// `ErrorKind::Unsupported` for a non-complex input, a typed capability
726    /// error when the session does not expose an FFT execution capability, or
727    /// a typed backend source for materialization.
728    fn irfft_read(
729        &self,
730        n: Option<usize>,
731        axis: isize,
732        norm: FftNorm,
733        session: &mut dyn BackendSession,
734    ) -> tenferro_tensor::Result<Tensor>;
735}
736
737impl TensorReadFftExt for TensorRead<'_> {
738    fn fft_read(
739        &self,
740        n: Option<usize>,
741        axis: isize,
742        norm: FftNorm,
743        session: &mut dyn BackendSession,
744    ) -> tenferro_tensor::Result<Tensor> {
745        with_fft_exec_session(session, "TensorReadFftExt::fft_read", |backend| {
746            execute_concrete_fft_read_op(
747                self,
748                concrete_fft_operation("TensorReadFftExt::fft_read", self.dtype())?,
749                "TensorReadFftExt::fft_read",
750                n,
751                axis,
752                norm,
753                backend,
754            )
755        })
756    }
757
758    fn ifft_read(
759        &self,
760        n: Option<usize>,
761        axis: isize,
762        norm: FftNorm,
763        session: &mut dyn BackendSession,
764    ) -> tenferro_tensor::Result<Tensor> {
765        with_fft_exec_session(session, "TensorReadFftExt::ifft_read", |backend| {
766            execute_concrete_fft_read_op(
767                self,
768                concrete_ifft_operation("TensorReadFftExt::ifft_read", self.dtype())?,
769                "TensorReadFftExt::ifft_read",
770                n,
771                axis,
772                norm,
773                backend,
774            )
775        })
776    }
777
778    fn rfft_read(
779        &self,
780        n: Option<usize>,
781        axis: isize,
782        norm: FftNorm,
783        session: &mut dyn BackendSession,
784    ) -> tenferro_tensor::Result<Tensor> {
785        with_fft_exec_session(session, "TensorReadFftExt::rfft_read", |backend| {
786            execute_concrete_fft_read_op(
787                self,
788                concrete_rfft_operation("TensorReadFftExt::rfft_read", self.dtype())?,
789                "TensorReadFftExt::rfft_read",
790                n,
791                axis,
792                norm,
793                backend,
794            )
795        })
796    }
797
798    fn irfft_read(
799        &self,
800        n: Option<usize>,
801        axis: isize,
802        norm: FftNorm,
803        session: &mut dyn BackendSession,
804    ) -> tenferro_tensor::Result<Tensor> {
805        with_fft_exec_session(session, "TensorReadFftExt::irfft_read", |backend| {
806            execute_concrete_fft_read_op(
807                self,
808                concrete_irfft_operation("TensorReadFftExt::irfft_read", self.dtype())?,
809                "TensorReadFftExt::irfft_read",
810                n,
811                axis,
812                norm,
813                backend,
814            )
815        })
816    }
817}
818
819#[derive(Debug, thiserror::Error)]
820enum FftError {
821    #[error("{op} does not support dtype {dtype:?}; expected {expected}")]
822    UnsupportedDType {
823        op: &'static str,
824        dtype: DType,
825        expected: &'static str,
826    },
827}
828
829#[derive(Clone, Debug, PartialEq)]
830struct FftOp {
831    operation: FftOperation,
832    axis: usize,
833    n: Option<usize>,
834    norm: FftNorm,
835}
836
837impl FftOp {
838    fn new(operation: FftOperation, axis: usize, n: Option<usize>, norm: FftNorm) -> Self {
839        Self {
840            operation,
841            axis,
842            n,
843            norm,
844        }
845    }
846
847    #[cfg(feature = "autodiff")]
848    fn c2c_adjoint(&self) -> Option<Self> {
849        match self.operation {
850            FftOperation::C2cForward => Some(Self {
851                operation: FftOperation::C2cInverse,
852                axis: self.axis,
853                n: self.n,
854                norm: self.norm.c2c_adjoint(),
855            }),
856            FftOperation::C2cInverse => Some(Self {
857                operation: FftOperation::C2cForward,
858                axis: self.axis,
859                n: self.n,
860                norm: self.norm.c2c_adjoint(),
861            }),
862            FftOperation::R2cFull | FftOperation::R2cOnesided | FftOperation::C2r => None,
863        }
864    }
865}
866
867impl ExtensionOp for FftOp {
868    fn family_id(&self) -> &'static str {
869        FFT_EXTENSION_FAMILY_ID
870    }
871
872    fn payload_hash(&self, hasher: &mut dyn Hasher) {
873        let operation = match self.operation {
874            FftOperation::C2cForward => 0,
875            FftOperation::C2cInverse => 1,
876            FftOperation::R2cOnesided => 2,
877            FftOperation::R2cFull => 3,
878            FftOperation::C2r => 4,
879        };
880        hasher.write_u8(operation);
881        hasher.write_usize(self.axis);
882        match self.n {
883            Some(n) => {
884                hasher.write_u8(1);
885                hasher.write_usize(n);
886            }
887            None => hasher.write_u8(0),
888        }
889        let norm = match self.norm {
890            FftNorm::Backward => 0,
891            FftNorm::Forward => 1,
892            FftNorm::Ortho => 2,
893        };
894        hasher.write_u8(norm);
895    }
896
897    fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
898        other
899            .as_any()
900            .downcast_ref::<FftOp>()
901            .is_some_and(|that| self == that)
902    }
903
904    fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
905        Arc::new(self.clone())
906    }
907
908    fn as_any(&self) -> &dyn Any {
909        self
910    }
911
912    fn input_count(&self) -> usize {
913        1
914    }
915
916    fn output_count(&self) -> usize {
917        1
918    }
919
920    fn semantic_effects(&self) -> tenferro_ops::ext_op::ExtensionEffectDeclaration<'_> {
921        tenferro_ops::ext_op::ExtensionEffectDeclaration::Declared(&[])
922    }
923
924    fn semantic_aliases(&self) -> tenferro_ops::ext_op::ExtensionAliasDeclaration<'_> {
925        tenferro_ops::ext_op::ExtensionAliasDeclaration::AllFresh
926    }
927
928    fn infer_output_meta(
929        &self,
930        ctx: &mut tenferro_ops::ExtensionShapeContext<'_>,
931    ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
932        let input_dtype = ctx.input_dtype(0)?;
933        let input_shape = ctx.input_shape(0)?;
934        if self.axis >= input_shape.len() {
935            return Err(tenferro_tensor::Error::axis_out_of_bounds(
936                "tenferro-fft",
937                self.axis,
938                input_shape.len(),
939            ));
940        }
941
942        let mut out_shape = input_shape.to_vec();
943        let output_dtype = match self.operation {
944            FftOperation::C2cForward | FftOperation::C2cInverse => {
945                if !matches!(input_dtype, DType::C32 | DType::C64) {
946                    return Err(tensor_unsupported_dtype(
947                        "tenferro-fft",
948                        input_dtype,
949                        "C32 or C64",
950                    ));
951                }
952                input_dtype
953            }
954            FftOperation::R2cFull | FftOperation::R2cOnesided => {
955                let len = transform_len_dim(self.n, &input_shape[self.axis]);
956                out_shape[self.axis] = if self.operation.is_onesided() {
957                    len / 2usize + 1usize
958                } else {
959                    len
960                };
961                match input_dtype {
962                    DType::F32 => DType::C32,
963                    DType::F64 => DType::C64,
964                    _ => {
965                        return Err(tensor_unsupported_dtype(
966                            "tenferro-fft",
967                            input_dtype,
968                            "F32 or F64",
969                        ));
970                    }
971                }
972            }
973            FftOperation::C2r => {
974                out_shape[self.axis] = output_dim_c2r(&input_shape[self.axis], self.n)?;
975                match input_dtype {
976                    DType::C32 => DType::F32,
977                    DType::C64 => DType::F64,
978                    _ => {
979                        return Err(tensor_unsupported_dtype(
980                            "tenferro-fft",
981                            input_dtype,
982                            "C32 or C64",
983                        ));
984                    }
985                }
986            }
987        };
988
989        if self.operation.is_c2c() {
990            out_shape[self.axis] = transform_len_dim(self.n, &input_shape[self.axis]);
991        }
992
993        Ok(vec![(output_dtype, out_shape)])
994    }
995}
996
997/// Run a concrete FFT body against the built-in FFT execution sessions
998/// carried by `session` (CPU/CUDA/WebGPU), returning a typed capability error
999/// when the session does not expose an FFT execution capability.
1000///
1001/// This is the built-in dispatch shared by the concrete FFT surface; callers
1002/// never downcast themselves (issue #1680 Phase 3). Third-party
1003/// [`FftBackend`] implementations remain supported through the SPI trait, but
1004/// the concrete op path is built-in-session only.
1005fn with_fft_exec_session<X>(
1006    session: &mut dyn BackendSession,
1007    op: &'static str,
1008    f: impl FnOnce(&mut dyn FftBackend) -> tenferro_tensor::Result<X>,
1009) -> tenferro_tensor::Result<X> {
1010    // The capability branches are mutually exclusive, so `f` runs exactly
1011    // once. Probe the marker first, then re-extract the same exec session and
1012    // run the concrete body on it (FnOnce cannot be captured by several
1013    // branch closures).
1014    if with_cpu_exec_session(session, |_| ()).is_some() {
1015        return with_cpu_exec_session(session, |exec| f(exec as &mut dyn FftBackend))
1016            .expect("marker probe matched a CPU execution session");
1017    }
1018    #[cfg(feature = "cuda")]
1019    if with_cuda_exec_session(session, |_| ()).is_some() {
1020        return with_cuda_exec_session(session, |exec| f(exec as &mut dyn FftBackend))
1021            .expect("marker probe matched a CUDA execution session");
1022    }
1023    #[cfg(feature = "webgpu")]
1024    if with_webgpu_exec_session(session, |_| ()).is_some() {
1025        return with_webgpu_exec_session(session, |exec| f(exec as &mut dyn FftBackend))
1026            .expect("marker probe matched a WebGPU execution session");
1027    }
1028    Err(tenferro_tensor::Error::unsupported(
1029        op,
1030        "selected backend session does not expose an FFT execution capability",
1031    ))
1032}
1033
1034fn execute_concrete_fft_op(
1035    input: &Tensor,
1036    spec: &FftPlanSpec,
1037    backend: &mut dyn FftBackend,
1038) -> tenferro_tensor::Result<Tensor> {
1039    let mut plans = FftPlanCache::with_capacity(NonZeroUsize::MIN);
1040    backend.execute_fft(input, spec, FftExecutionCache::caller_owned(&mut plans))
1041}
1042
1043#[allow(clippy::too_many_arguments)]
1044fn execute_concrete_fft_read_op(
1045    input: &TensorRead<'_>,
1046    operation: FftOperation,
1047    op_name: &'static str,
1048    n: Option<usize>,
1049    axis: isize,
1050    norm: FftNorm,
1051    backend: &mut dyn FftBackend,
1052) -> tenferro_tensor::Result<Tensor> {
1053    let spec = concrete_fft_spec(
1054        op_name,
1055        operation,
1056        input.dtype(),
1057        input.shape(),
1058        n,
1059        axis,
1060        norm,
1061    )?;
1062    let mut plans = FftPlanCache::with_capacity(NonZeroUsize::MIN);
1063    backend.execute_fft_read(
1064        input.clone(),
1065        &spec,
1066        FftExecutionCache::caller_owned(&mut plans),
1067    )
1068}
1069
1070#[allow(clippy::too_many_arguments)]
1071fn concrete_fft_spec(
1072    op: &'static str,
1073    operation: FftOperation,
1074    input_dtype: DType,
1075    input_shape: &[usize],
1076    n: Option<usize>,
1077    axis: isize,
1078    norm: FftNorm,
1079) -> tenferro_tensor::Result<FftPlanSpec> {
1080    validate_concrete_n(op, n)?;
1081    let axis = normalize_concrete_axis(op, axis, input_shape.len())?;
1082    validated_fft_plan_spec(op, operation, input_dtype, input_shape, n, axis, norm)
1083}
1084
1085#[allow(clippy::too_many_arguments)]
1086fn validated_fft_plan_spec(
1087    op: &'static str,
1088    operation: FftOperation,
1089    input_dtype: DType,
1090    input_shape: &[usize],
1091    n: Option<usize>,
1092    axis: usize,
1093    norm: FftNorm,
1094) -> tenferro_tensor::Result<FftPlanSpec> {
1095    validate_concrete_n(op, n)?;
1096    validate_operation_dtype(op, operation, input_dtype)?;
1097    validate_axis(op, input_shape, axis)?;
1098    validate_concrete_transform_len(op, input_shape, n, axis)?;
1099    if operation == FftOperation::C2r {
1100        output_shape_c2r(input_shape, axis, n)?;
1101    }
1102    Ok(FftPlanSpec::new(
1103        operation,
1104        axis,
1105        n,
1106        norm,
1107        input_dtype,
1108        input_shape.to_vec(),
1109    ))
1110}
1111
1112fn concrete_fft_operation(op: &'static str, dtype: DType) -> tenferro_tensor::Result<FftOperation> {
1113    match dtype {
1114        DType::C32 | DType::C64 => Ok(FftOperation::C2cForward),
1115        DType::F32 | DType::F64 => Ok(FftOperation::R2cFull),
1116        DType::I32 | DType::I64 | DType::Bool | DType::External(_) => {
1117            Err(tensor_unsupported_dtype(op, dtype, "F32, F64, C32, or C64"))
1118        }
1119    }
1120}
1121
1122fn concrete_ifft_operation(
1123    op: &'static str,
1124    dtype: DType,
1125) -> tenferro_tensor::Result<FftOperation> {
1126    match dtype {
1127        DType::C32 | DType::C64 => Ok(FftOperation::C2cInverse),
1128        DType::F32 | DType::F64 | DType::I32 | DType::I64 | DType::Bool | DType::External(_) => {
1129            Err(tensor_unsupported_dtype(op, dtype, "C32 or C64"))
1130        }
1131    }
1132}
1133
1134fn concrete_rfft_operation(
1135    op: &'static str,
1136    dtype: DType,
1137) -> tenferro_tensor::Result<FftOperation> {
1138    match dtype {
1139        DType::F32 | DType::F64 => Ok(FftOperation::R2cOnesided),
1140        DType::C32 | DType::C64 | DType::I32 | DType::I64 | DType::Bool | DType::External(_) => {
1141            Err(tensor_unsupported_dtype(op, dtype, "F32 or F64"))
1142        }
1143    }
1144}
1145
1146fn concrete_irfft_operation(
1147    op: &'static str,
1148    dtype: DType,
1149) -> tenferro_tensor::Result<FftOperation> {
1150    match dtype {
1151        DType::C32 | DType::C64 => Ok(FftOperation::C2r),
1152        DType::F32 | DType::F64 | DType::I32 | DType::I64 | DType::Bool | DType::External(_) => {
1153            Err(tensor_unsupported_dtype(op, dtype, "C32 or C64"))
1154        }
1155    }
1156}
1157
1158fn validate_operation_dtype(
1159    op: &'static str,
1160    operation: FftOperation,
1161    dtype: DType,
1162) -> tenferro_tensor::Result<()> {
1163    let supported = match operation {
1164        FftOperation::C2cForward | FftOperation::C2cInverse | FftOperation::C2r => {
1165            matches!(dtype, DType::C32 | DType::C64)
1166        }
1167        FftOperation::R2cFull | FftOperation::R2cOnesided => {
1168            matches!(dtype, DType::F32 | DType::F64)
1169        }
1170    };
1171    if supported {
1172        Ok(())
1173    } else {
1174        Err(tensor_unsupported_dtype(
1175            op,
1176            dtype,
1177            expected_dtype_description(operation),
1178        ))
1179    }
1180}
1181
1182fn validate_concrete_n(op: &'static str, n: Option<usize>) -> tenferro_tensor::Result<()> {
1183    if n == Some(0) {
1184        return Err(tenferro_tensor::Error::invalid_argument(
1185            op,
1186            "n",
1187            "transform length must be positive",
1188        ));
1189    }
1190    Ok(())
1191}
1192
1193fn validate_concrete_transform_len(
1194    op: &'static str,
1195    input_shape: &[usize],
1196    n: Option<usize>,
1197    axis: usize,
1198) -> tenferro_tensor::Result<()> {
1199    if n.is_none() && input_shape.get(axis).copied() == Some(0) {
1200        return Err(tenferro_tensor::Error::invalid_argument(
1201            op,
1202            "n",
1203            "transform length must be positive",
1204        ));
1205    }
1206    Ok(())
1207}
1208
1209fn normalize_concrete_axis(
1210    op: &'static str,
1211    axis: isize,
1212    rank: usize,
1213) -> tenferro_tensor::Result<usize> {
1214    if rank == 0 {
1215        return Err(tenferro_tensor::Error::invalid_argument(
1216            op,
1217            "rank",
1218            "FFT requires rank >= 1",
1219        ));
1220    }
1221    let normalized = if axis >= 0 {
1222        axis as usize
1223    } else {
1224        rank.checked_sub(axis.unsigned_abs()).ok_or_else(|| {
1225            tenferro_tensor::Error::axis_out_of_bounds(op, axis.unsigned_abs(), rank)
1226        })?
1227    };
1228    if normalized >= rank {
1229        return Err(tenferro_tensor::Error::axis_out_of_bounds(
1230            op, normalized, rank,
1231        ));
1232    }
1233    Ok(normalized)
1234}
1235
1236fn tensor_unsupported_dtype(
1237    op: &'static str,
1238    dtype: DType,
1239    expected: &'static str,
1240) -> tenferro_tensor::Error {
1241    tenferro_tensor::Error::extension(
1242        op,
1243        FFT_EXTENSION_FAMILY_ID,
1244        ErrorKind::Unsupported,
1245        FftError::UnsupportedDType {
1246            op,
1247            dtype,
1248            expected,
1249        },
1250    )
1251}
1252
1253#[cfg(feature = "autodiff")]
1254#[derive(Debug)]
1255struct FftAdRule;
1256
1257#[cfg(feature = "autodiff")]
1258impl SemanticLinearizeRule for FftAdRule {
1259    fn family_id(&self) -> &'static str {
1260        FFT_EXTENSION_FAMILY_ID
1261    }
1262
1263    fn linearize(
1264        &self,
1265        request: SemanticLinearizeRequest<'_>,
1266        builder: &mut SemanticProgramBuilder,
1267    ) -> std::result::Result<SemanticLinearizeResult, SemanticAdError> {
1268        let fft_op = semantic_fft_payload(request.op(), SemanticAdRuleKind::Linearize)?;
1269        if !fft_op.operation.is_c2c() {
1270            return Err(semantic_fft_unsupported(
1271                fft_op.operation,
1272                SemanticAdRuleKind::Linearize,
1273            ));
1274        }
1275        let tangent = match request.tangent_inputs()[0] {
1276            AdValue::Absent => AdValue::Absent,
1277            AdValue::Value(tangent) => {
1278                AdValue::Value(builder.add_extension(Arc::new(fft_op.clone()), &[tangent])?[0])
1279            }
1280        };
1281        Ok(SemanticLinearizeResult::new([tangent], []))
1282    }
1283}
1284
1285#[cfg(feature = "autodiff")]
1286impl SemanticLinearTransposeRule for FftAdRule {
1287    fn family_id(&self) -> &'static str {
1288        FFT_EXTENSION_FAMILY_ID
1289    }
1290
1291    fn residual_mask(&self) -> ResidualSpec {
1292        // The variable-length adjoint path reads primal input 0 as a tensor
1293        // (ShapeOf + PadToMatch against the original input).
1294        ResidualSpec::input(0)
1295    }
1296
1297    fn linear_transpose(
1298        &self,
1299        request: SemanticLinearTransposeRequest<'_>,
1300        builder: &mut SemanticProgramBuilder,
1301    ) -> std::result::Result<Box<[AdValue]>, SemanticAdError> {
1302        Ok([semantic_fft_adjoint(
1303            request.op(),
1304            request.cotangent_outputs()[0],
1305            request.active_inputs()[0],
1306            request.primal_input_value(0)?,
1307            request.residual_mask(),
1308            builder,
1309        )?]
1310        .into())
1311    }
1312}
1313
1314#[cfg(feature = "autodiff")]
1315impl SemanticPrimalVjpRule for FftAdRule {
1316    fn family_id(&self) -> &'static str {
1317        FFT_EXTENSION_FAMILY_ID
1318    }
1319
1320    fn residual_mask(&self) -> ResidualSpec {
1321        ResidualSpec::input(0)
1322    }
1323
1324    fn primal_vjp(
1325        &self,
1326        request: SemanticPrimalVjpRequest<'_>,
1327        builder: &mut SemanticProgramBuilder,
1328    ) -> std::result::Result<Box<[AdValue]>, SemanticAdError> {
1329        Ok([semantic_fft_adjoint(
1330            request.op(),
1331            request.cotangent_outputs()[0],
1332            request.active_inputs()[0],
1333            request.primal_input_value(0)?,
1334            request.residual_mask(),
1335            builder,
1336        )?]
1337        .into())
1338    }
1339}
1340
1341#[cfg(feature = "autodiff")]
1342#[derive(Clone, Copy)]
1343enum SemanticAdRuleKind {
1344    Linearize,
1345    Transpose,
1346}
1347
1348#[cfg(feature = "autodiff")]
1349fn semantic_fft_payload(
1350    op: &dyn ExtensionOp,
1351    role: SemanticAdRuleKind,
1352) -> std::result::Result<&FftOp, SemanticAdError> {
1353    op.as_any().downcast_ref::<FftOp>().ok_or_else(|| {
1354        semantic_fft_unsupported_family(
1355            FFT_EXTENSION_FAMILY_ID,
1356            role,
1357            "FFT semantic AD received an incompatible extension payload",
1358        )
1359    })
1360}
1361
1362#[cfg(feature = "autodiff")]
1363fn semantic_fft_adjoint(
1364    op: &dyn ExtensionOp,
1365    cotangent: AdValue,
1366    active: bool,
1367    primal_input: ProgramValue,
1368    residual_mask: ResidualSpec,
1369    builder: &mut SemanticProgramBuilder,
1370) -> std::result::Result<AdValue, SemanticAdError> {
1371    if !active {
1372        return Ok(AdValue::Absent);
1373    }
1374    let AdValue::Value(cotangent) = cotangent else {
1375        return Ok(AdValue::Absent);
1376    };
1377    let fft_op = semantic_fft_payload(op, SemanticAdRuleKind::Transpose)?;
1378    if !fft_op.operation.is_c2c() {
1379        return Err(semantic_fft_unsupported(
1380            fft_op.operation,
1381            SemanticAdRuleKind::Transpose,
1382        ));
1383    }
1384    let adjoint_op = fft_op
1385        .c2c_adjoint()
1386        .ok_or_else(|| semantic_fft_unsupported(fft_op.operation, SemanticAdRuleKind::Transpose))?;
1387    let adjoint = builder.add_extension(Arc::new(adjoint_op), &[cotangent])?[0];
1388    restore_semantic_c2c_adjoint_input_length(builder, adjoint, primal_input, residual_mask, fft_op)
1389        .map(AdValue::Value)
1390}
1391
1392#[cfg(feature = "autodiff")]
1393fn restore_semantic_c2c_adjoint_input_length(
1394    builder: &mut SemanticProgramBuilder,
1395    adjoint: ProgramValue,
1396    primal_input: ProgramValue,
1397    residual_mask: ResidualSpec,
1398    fft_op: &FftOp,
1399) -> std::result::Result<ProgramValue, SemanticAdError> {
1400    let Some(transform_len) = fft_op.n else {
1401        return Ok(adjoint);
1402    };
1403    debug_assert!(
1404        residual_mask.declares_input(0),
1405        "fft transpose read primal input 0 as a tensor operand but the residual mask does not \
1406         declare it; declare it in the fft rule's residual mask"
1407    );
1408    let input_len = builder
1409        .value_metadata(primal_input)?
1410        .shape()
1411        .get(fft_op.axis)
1412        .and_then(|extent| extent.as_exact())
1413        .and_then(|dim| match dim {
1414            tenferro_ops::dim_expr::DimExpr::Const(value) => Some(*value),
1415            _ => None,
1416        });
1417    if input_len == Some(transform_len) {
1418        return Ok(adjoint);
1419    }
1420
1421    let size = builder.add_op(
1422        CoreSemanticOp::ShapeOf { axis: fft_op.axis },
1423        &[primal_input],
1424    )?[0];
1425    let truncated = builder.add_op(
1426        CoreSemanticOp::DynamicTruncate { axis: fft_op.axis },
1427        &[adjoint, size],
1428    )?[0];
1429    Ok(builder.add_op(
1430        CoreSemanticOp::PadToMatch { axis: fft_op.axis },
1431        &[truncated, primal_input],
1432    )?[0])
1433}
1434
1435#[cfg(feature = "autodiff")]
1436fn semantic_fft_unsupported(operation: FftOperation, role: SemanticAdRuleKind) -> SemanticAdError {
1437    semantic_fft_unsupported_family(
1438        fft_ad_family_id(operation),
1439        role,
1440        "FFT operation has no semantic AD rule",
1441    )
1442}
1443
1444#[cfg(feature = "autodiff")]
1445fn semantic_fft_unsupported_family(
1446    family_id: &'static str,
1447    role: SemanticAdRuleKind,
1448    message: impl Into<String>,
1449) -> SemanticAdError {
1450    SemanticAdError::Unsupported {
1451        family_id,
1452        role: match role {
1453            SemanticAdRuleKind::Linearize => {
1454                tenferro_ad::semantic_extension::SemanticAdRuleRole::Linearize
1455            }
1456            SemanticAdRuleKind::Transpose => {
1457                tenferro_ad::semantic_extension::SemanticAdRuleRole::LinearTranspose
1458            }
1459        },
1460        message: message.into(),
1461    }
1462}
1463
1464/// Return the semantic-program FFT extension AD rule set.
1465#[cfg(feature = "autodiff")]
1466///
1467/// # Errors
1468///
1469/// Returns [`SemanticExtensionRegistryError::MalformedFamilyId`] if the FFT
1470/// family identifier is invalid, or
1471/// [`SemanticExtensionRegistryError::DuplicateRule`] if a rule for the family
1472/// and role is already registered.
1473pub fn semantic_ad_rules(
1474) -> std::result::Result<SemanticExtensionRuleSet, SemanticExtensionRegistryError> {
1475    SemanticExtensionRuleSet::new()
1476        .with_linearize(Arc::new(FftAdRule))?
1477        .with_linear_transpose(Arc::new(FftAdRule))?
1478        .with_primal_vjp(Arc::new(FftAdRule))
1479}
1480
1481pub(crate) fn execute_fft_extension_reads_session(
1482    op: &FftOp,
1483    inputs: &[TensorRead<'_>],
1484    ctx: &mut ExtensionExecutionContext<'_, dyn BackendSession + '_>,
1485) -> tenferro_tensor::Result<Vec<Tensor>> {
1486    let (session, caches) = ctx.parts_mut();
1487    execute_fft_extension_reads_on_session(op, inputs, session, caches)
1488}
1489
1490fn execute_fft_extension_reads_for_capability<B: FftBackend + ?Sized>(
1491    op: &FftOp,
1492    inputs: &[TensorRead<'_>],
1493    session: &mut B,
1494    caches: &mut ExtensionCacheStore,
1495) -> tenferro_tensor::Result<Vec<Tensor>> {
1496    if inputs.len() != 1 {
1497        return Err(tenferro_tensor::Error::invalid_argument(
1498            "tenferro-fft",
1499            "inputs",
1500            format!("expected 1 input, got {}", inputs.len()),
1501        ));
1502    }
1503    let input = &inputs[0];
1504    session.validate_fft_read_input(fft_op_name(op.operation), input)?;
1505    let spec = validated_fft_plan_spec(
1506        fft_op_name(op.operation),
1507        op.operation,
1508        input.dtype(),
1509        input.shape(),
1510        op.n,
1511        op.axis,
1512        op.norm,
1513    )?;
1514    let output = session.execute_fft_read(
1515        input.clone(),
1516        &spec,
1517        FftExecutionCache::runtime_owned(caches),
1518    )?;
1519    Ok(vec![output])
1520}
1521
1522fn execute_fft_extension_reads_on_session(
1523    op: &FftOp,
1524    inputs: &[TensorRead<'_>],
1525    session: &mut dyn BackendSession,
1526    caches: &mut ExtensionCacheStore,
1527) -> tenferro_tensor::Result<Vec<Tensor>> {
1528    if let Some(result) = with_cpu_exec_session(session, |session| {
1529        execute_fft_extension_reads_for_capability(op, inputs, session, caches)
1530    }) {
1531        return result;
1532    }
1533    #[cfg(feature = "cuda")]
1534    if let Some(result) = with_cuda_exec_session(session, |session| {
1535        execute_fft_extension_reads_for_capability(op, inputs, session, caches)
1536    }) {
1537        return result;
1538    }
1539    #[cfg(feature = "webgpu")]
1540    if let Some(result) = with_webgpu_exec_session(session, |session| {
1541        execute_fft_extension_reads_for_capability(op, inputs, session, caches)
1542    }) {
1543        return result;
1544    }
1545    Err(tenferro_tensor::Error::unsupported(
1546        fft_op_name(op.operation),
1547        "selected backend session does not expose an FFT execution capability",
1548    ))
1549}
1550
1551define_extension_runtime! {
1552    runtime = FftRuntime,
1553    family_id = FFT_EXTENSION_FAMILY_ID,
1554    op_type = FftOp,
1555    execute_in_session = execute_fft_extension_reads_in_session,
1556    session_supported = fft_session_supported,
1557    backend_bound = TensorBackend,
1558}
1559
1560/// Adapter from the scheduler/`apply_eager` borrowed-session shape to the
1561/// existing FFT session executor. Reuses the same forward kernel already shared
1562/// by the owner and eager paths; do not reimplement it here.
1563fn execute_fft_extension_reads_in_session(
1564    op: &FftOp,
1565    session: &mut dyn BackendSession,
1566    caches: &mut ExtensionCacheStore,
1567    inputs: &[TensorRead<'_>],
1568) -> tenferro_tensor::Result<Vec<Tensor>> {
1569    let mut ctx = ExtensionExecutionContext::new(session, caches);
1570    execute_fft_extension_reads_session(op, inputs, &mut ctx)
1571}
1572
1573fn fft_session_supported<B: tenferro_tensor::TensorBackend + 'static>(_op: &FftOp) -> bool {
1574    // The session executor routes CPU/CUDA/WebGPU through their FftBackend exec
1575    // sessions; keep scheduler-session admission consistent with the backends
1576    // that `execute_fft_extension_reads_on_session` actually handles.
1577    let type_id = std::any::TypeId::of::<B>();
1578    type_id == std::any::TypeId::of::<tenferro_cpu::CpuBackend>() || {
1579        #[cfg(feature = "cuda")]
1580        {
1581            type_id == std::any::TypeId::of::<CudaBackend>()
1582        }
1583        #[cfg(not(feature = "cuda"))]
1584        {
1585            false
1586        }
1587    }
1588}
1589
1590/// Build a one-dimensional FFT along `axis`.
1591///
1592/// Complex inputs use a complex-to-complex transform. Real inputs use a
1593/// real-to-complex transform that returns the full complex spectrum.
1594///
1595/// # Examples
1596///
1597/// ```
1598/// use num_complex::Complex64;
1599/// use tenferro_cpu::CpuBackend;
1600/// use tenferro_runtime::{GraphCompiler, Runtime, TracedTensor};
1601/// use tenferro_fft::{FftNorm, TracedTensorFftExt};
1602///
1603/// let x = TracedTensor::from_vec_col_major(vec![2], vec![Complex64::new(1.0, 0.0), Complex64::new(2.0, 0.0)]).unwrap();
1604/// let y = x.fft(None, -1, FftNorm::Backward).unwrap();
1605///
1606/// let mut compiler = GraphCompiler::new();
1607/// let program = compiler.compile(&y).unwrap();
1608/// let backend = CpuBackend::new();
1609/// let engine_id = tenferro_cpu::runtime_engine_id().unwrap();
1610/// let mut builder = Runtime::builder();
1611/// builder
1612///     .register_engine(tenferro_cpu::runtime_engine_registration(&backend).unwrap())
1613///     .unwrap();
1614/// builder
1615///     .install_extension_module(tenferro_fft::extension_module::<CpuBackend>(engine_id).unwrap())
1616///     .unwrap();
1617/// let runtime = builder.build().unwrap();
1618/// let out = runtime.run_compiled(&program, &[]).unwrap().pop().unwrap();
1619/// assert_eq!(out.as_slice::<Complex64>().unwrap()[0], Complex64::new(3.0, 0.0));
1620/// ```
1621fn fft(input: &TracedTensor, n: Option<usize>, axis: isize, norm: FftNorm) -> Result<TracedTensor> {
1622    let operation = runtime_forward_fft_operation(input.dtype)?;
1623    apply_unary_fft("fft", input, operation, n, axis, norm)
1624}
1625
1626/// Build a one-dimensional inverse FFT along `axis`.
1627///
1628/// # Examples
1629///
1630/// ```
1631/// use num_complex::Complex64;
1632/// use tenferro_cpu::CpuBackend;
1633/// use tenferro_runtime::{GraphCompiler, Runtime, TracedTensor};
1634/// use tenferro_fft::{FftNorm, TracedTensorFftExt};
1635///
1636/// let spectrum = TracedTensor::from_vec_col_major(vec![2], vec![Complex64::new(3.0, 0.0), Complex64::new(-1.0, 0.0)]).unwrap();
1637/// let y = spectrum.ifft(None, -1, FftNorm::Backward).unwrap();
1638///
1639/// let mut compiler = GraphCompiler::new();
1640/// let program = compiler.compile(&y).unwrap();
1641/// let backend = CpuBackend::new();
1642/// let engine_id = tenferro_cpu::runtime_engine_id().unwrap();
1643/// let mut builder = Runtime::builder();
1644/// builder
1645///     .register_engine(tenferro_cpu::runtime_engine_registration(&backend).unwrap())
1646///     .unwrap();
1647/// builder
1648///     .install_extension_module(tenferro_fft::extension_module::<CpuBackend>(engine_id).unwrap())
1649///     .unwrap();
1650/// let runtime = builder.build().unwrap();
1651/// let out = runtime.run_compiled(&program, &[]).unwrap().pop().unwrap();
1652/// assert_eq!(out.as_slice::<Complex64>().unwrap()[0], Complex64::new(1.0, 0.0));
1653/// ```
1654fn ifft(
1655    input: &TracedTensor,
1656    n: Option<usize>,
1657    axis: isize,
1658    norm: FftNorm,
1659) -> Result<TracedTensor> {
1660    require_runtime_dtype("ifft", input.dtype, &[DType::C32, DType::C64], "C32 or C64")?;
1661    apply_unary_fft("ifft", input, FftOperation::C2cInverse, n, axis, norm)
1662}
1663
1664/// Build a one-dimensional real FFT along `axis`.
1665///
1666/// The output keeps only the Hermitian one-sided spectrum with axis length
1667/// `n / 2 + 1`.
1668///
1669/// # Examples
1670///
1671/// ```
1672/// use num_complex::Complex64;
1673/// use tenferro_cpu::CpuBackend;
1674/// use tenferro_runtime::{GraphCompiler, Runtime, TracedTensor};
1675/// use tenferro_fft::{FftNorm, TracedTensorFftExt};
1676///
1677/// let x = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
1678/// let y = x.rfft(None, -1, FftNorm::Backward).unwrap();
1679///
1680/// let mut compiler = GraphCompiler::new();
1681/// let program = compiler.compile(&y).unwrap();
1682/// let backend = CpuBackend::new();
1683/// let engine_id = tenferro_cpu::runtime_engine_id().unwrap();
1684/// let mut builder = Runtime::builder();
1685/// builder
1686///     .register_engine(tenferro_cpu::runtime_engine_registration(&backend).unwrap())
1687///     .unwrap();
1688/// builder
1689///     .install_extension_module(tenferro_fft::extension_module::<CpuBackend>(engine_id).unwrap())
1690///     .unwrap();
1691/// let runtime = builder.build().unwrap();
1692/// let out = runtime.run_compiled(&program, &[]).unwrap().pop().unwrap();
1693/// assert_eq!(out.shape(), &[2]);
1694/// assert_eq!(out.as_slice::<Complex64>().unwrap()[0], Complex64::new(3.0, 0.0));
1695/// ```
1696fn rfft(
1697    input: &TracedTensor,
1698    n: Option<usize>,
1699    axis: isize,
1700    norm: FftNorm,
1701) -> Result<TracedTensor> {
1702    require_runtime_dtype("rfft", input.dtype, &[DType::F32, DType::F64], "F32 or F64")?;
1703    apply_unary_fft("rfft", input, FftOperation::R2cOnesided, n, axis, norm)
1704}
1705
1706/// Build a one-dimensional inverse real FFT along `axis`.
1707///
1708/// If `n` is `None`, the output length is inferred as twice one less than the
1709/// input spectrum length.
1710///
1711/// # Examples
1712///
1713/// ```
1714/// use num_complex::Complex64;
1715/// use tenferro_cpu::CpuBackend;
1716/// use tenferro_runtime::{GraphCompiler, Runtime, TracedTensor};
1717/// use tenferro_fft::{FftNorm, TracedTensorFftExt};
1718///
1719/// let spectrum = TracedTensor::from_vec_col_major(
1720///     vec![2],
1721///     vec![Complex64::new(3.0, 0.0), Complex64::new(-1.0, 0.0)],
1722/// )
1723/// .unwrap();
1724/// let y = spectrum.irfft(Some(2), -1, FftNorm::Backward).unwrap();
1725///
1726/// let mut compiler = GraphCompiler::new();
1727/// let program = compiler.compile(&y).unwrap();
1728/// let backend = CpuBackend::new();
1729/// let engine_id = tenferro_cpu::runtime_engine_id().unwrap();
1730/// let mut builder = Runtime::builder();
1731/// builder
1732///     .register_engine(tenferro_cpu::runtime_engine_registration(&backend).unwrap())
1733///     .unwrap();
1734/// builder
1735///     .install_extension_module(tenferro_fft::extension_module::<CpuBackend>(engine_id).unwrap())
1736///     .unwrap();
1737/// let runtime = builder.build().unwrap();
1738/// let out = runtime.run_compiled(&program, &[]).unwrap().pop().unwrap();
1739/// assert_eq!(out.as_slice::<f64>().unwrap(), &[1.0, 2.0]);
1740/// ```
1741fn irfft(
1742    input: &TracedTensor,
1743    n: Option<usize>,
1744    axis: isize,
1745    norm: FftNorm,
1746) -> Result<TracedTensor> {
1747    require_runtime_dtype(
1748        "irfft",
1749        input.dtype,
1750        &[DType::C32, DType::C64],
1751        "C32 or C64",
1752    )?;
1753    apply_unary_fft("irfft", input, FftOperation::C2r, n, axis, norm)
1754}
1755
1756fn apply_unary_fft(
1757    op_name: &'static str,
1758    input: &TracedTensor,
1759    operation: FftOperation,
1760    n: Option<usize>,
1761    axis: isize,
1762    norm: FftNorm,
1763) -> Result<TracedTensor> {
1764    let concrete_shape = input.try_concrete_shape();
1765    let op = Arc::new(prepare_runtime_fft_op(
1766        op_name,
1767        operation,
1768        input.rank,
1769        concrete_shape.as_deref(),
1770        n,
1771        axis,
1772        norm,
1773    )?);
1774    let mut outputs = apply(op, &[input])?;
1775    outputs
1776        .pop()
1777        .ok_or_else(|| Error::Internal("FFT extension declares exactly one output".into()))
1778}
1779
1780fn normalize_axis(op: &'static str, axis: isize, rank: usize) -> Result<usize> {
1781    if rank == 0 {
1782        return Err(runtime_invalid_argument(
1783            op,
1784            "rank",
1785            "FFT requires rank >= 1",
1786        ));
1787    }
1788    let normalized = if axis >= 0 {
1789        axis as usize
1790    } else {
1791        rank.checked_sub(axis.unsigned_abs())
1792            .ok_or_else(|| runtime_axis_out_of_bounds(op, axis.unsigned_abs(), rank))?
1793    };
1794    if normalized >= rank {
1795        return Err(runtime_axis_out_of_bounds(op, normalized, rank));
1796    }
1797    Ok(normalized)
1798}
1799
1800fn validate_n(op: &'static str, n: Option<usize>) -> Result<()> {
1801    if n == Some(0) {
1802        return Err(runtime_invalid_argument(
1803            op,
1804            "n",
1805            "transform length must be positive",
1806        ));
1807    }
1808    Ok(())
1809}
1810
1811fn prepare_runtime_fft_op(
1812    op: &'static str,
1813    operation: FftOperation,
1814    rank: usize,
1815    concrete_shape: Option<&[usize]>,
1816    n: Option<usize>,
1817    axis: isize,
1818    norm: FftNorm,
1819) -> Result<FftOp> {
1820    validate_n(op, n)?;
1821    let axis = normalize_axis(op, axis, rank)?;
1822    if n.is_none() && concrete_shape.and_then(|shape| shape.get(axis).copied()) == Some(0) {
1823        return Err(runtime_invalid_argument(
1824            op,
1825            "n",
1826            "transform length must be positive",
1827        ));
1828    }
1829    if operation == FftOperation::C2r {
1830        if let Some(shape) = concrete_shape {
1831            output_shape_c2r(shape, axis, n)?;
1832        }
1833    }
1834    Ok(FftOp::new(operation, axis, n, norm))
1835}
1836
1837fn runtime_forward_fft_operation(dtype: DType) -> Result<FftOperation> {
1838    match dtype {
1839        DType::C32 | DType::C64 => Ok(FftOperation::C2cForward),
1840        DType::F32 | DType::F64 => Ok(FftOperation::R2cFull),
1841        DType::I32 | DType::I64 | DType::Bool | DType::External(_) => Err(
1842            runtime_unsupported_dtype("fft", dtype, "F32, F64, C32, or C64"),
1843        ),
1844    }
1845}
1846
1847fn require_runtime_dtype(
1848    op: &'static str,
1849    dtype: DType,
1850    supported: &[DType],
1851    expected: &'static str,
1852) -> Result<()> {
1853    if supported.contains(&dtype) {
1854        Ok(())
1855    } else {
1856        Err(runtime_unsupported_dtype(op, dtype, expected))
1857    }
1858}
1859
1860fn runtime_invalid_argument(
1861    op: &'static str,
1862    argument: &'static str,
1863    message: impl Into<String>,
1864) -> Error {
1865    Error::validation(
1866        op,
1867        ErrorPhase::GraphBuild,
1868        ValidationError::InvalidArgument {
1869            argument,
1870            message: message.into(),
1871        },
1872    )
1873}
1874
1875fn runtime_axis_out_of_bounds(op: &'static str, axis: usize, rank: usize) -> Error {
1876    Error::validation(
1877        op,
1878        ErrorPhase::GraphBuild,
1879        ValidationError::AxisOutOfBounds { axis, rank },
1880    )
1881}
1882
1883fn runtime_unsupported_dtype(op: &'static str, dtype: DType, expected: &'static str) -> Error {
1884    Error::extension(
1885        op,
1886        ErrorPhase::GraphBuild,
1887        FFT_EXTENSION_FAMILY_ID,
1888        ErrorKind::Unsupported,
1889        FftError::UnsupportedDType {
1890            op,
1891            dtype,
1892            expected,
1893        },
1894    )
1895}
1896
1897fn transform_len_dim(n: Option<usize>, input_dim: &SymDim) -> SymDim {
1898    n.map(SymDim::from).unwrap_or_else(|| input_dim.clone())
1899}
1900
1901fn expected_dtype_description(operation: FftOperation) -> &'static str {
1902    match operation {
1903        FftOperation::C2cForward | FftOperation::C2cInverse | FftOperation::C2r => "C32 or C64",
1904        FftOperation::R2cFull | FftOperation::R2cOnesided => "F32 or F64",
1905    }
1906}
1907
1908fn fft_op_name(operation: FftOperation) -> &'static str {
1909    match operation {
1910        FftOperation::C2cForward => "fft",
1911        FftOperation::C2cInverse => "ifft",
1912        FftOperation::R2cFull | FftOperation::R2cOnesided => "rfft",
1913        FftOperation::C2r => "irfft",
1914    }
1915}
1916
1917#[cfg(feature = "autodiff")]
1918fn fft_ad_family_id(operation: FftOperation) -> &'static str {
1919    match operation {
1920        FftOperation::C2cForward | FftOperation::C2cInverse => FFT_EXTENSION_FAMILY_ID,
1921        FftOperation::R2cFull | FftOperation::R2cOnesided => "tenferro-fft.rfft.v1",
1922        FftOperation::C2r => "tenferro-fft.irfft.v1",
1923    }
1924}
1925
1926fn output_shape_c2c(
1927    shape: &[usize],
1928    axis: usize,
1929    n: Option<usize>,
1930) -> tenferro_tensor::Result<Vec<usize>> {
1931    let len = transform_len(shape, axis, n)?;
1932    let mut out_shape = shape.to_vec();
1933    out_shape[axis] = len;
1934    Ok(out_shape)
1935}
1936
1937fn output_shape_r2c(
1938    shape: &[usize],
1939    axis: usize,
1940    n: Option<usize>,
1941    onesided: bool,
1942) -> tenferro_tensor::Result<Vec<usize>> {
1943    let len = transform_len(shape, axis, n)?;
1944    let mut out_shape = shape.to_vec();
1945    out_shape[axis] = if onesided { len / 2 + 1 } else { len };
1946    Ok(out_shape)
1947}
1948
1949fn output_shape_c2r(
1950    shape: &[usize],
1951    axis: usize,
1952    n: Option<usize>,
1953) -> tenferro_tensor::Result<Vec<usize>> {
1954    validate_axis("irfft", shape, axis)?;
1955    let input_len = shape[axis];
1956    let len = match n {
1957        Some(len) => len,
1958        None => default_c2r_output_len(input_len)?,
1959    };
1960    if len == 0 {
1961        return Err(tenferro_tensor::Error::invalid_argument(
1962            "irfft",
1963            "output length",
1964            "must be positive",
1965        ));
1966    }
1967    validate_c2r_spectrum_len(input_len, len)?;
1968    let mut out_shape = shape.to_vec();
1969    out_shape[axis] = len;
1970    Ok(out_shape)
1971}
1972
1973fn output_dim_c2r(input_dim: &SymDim, n: Option<usize>) -> tenferro_tensor::Result<SymDim> {
1974    match (input_dim.constant_value(), n) {
1975        (Some(input_len), Some(output_len)) => {
1976            if output_len == 0 {
1977                return Err(tenferro_tensor::Error::invalid_argument(
1978                    "irfft",
1979                    "output length",
1980                    "must be positive",
1981                ));
1982            }
1983            validate_c2r_spectrum_len(input_len, output_len)?;
1984            Ok(SymDim::from(output_len))
1985        }
1986        (Some(input_len), None) => Ok(SymDim::from(default_c2r_output_len(input_len)?)),
1987        (None, Some(output_len)) => {
1988            if output_len == 0 {
1989                return Err(tenferro_tensor::Error::invalid_argument(
1990                    "irfft",
1991                    "output length",
1992                    "must be positive",
1993                ));
1994            }
1995            Ok(SymDim::from(output_len))
1996        }
1997        (None, None) => Ok((input_dim.clone() - 1usize) * 2usize),
1998    }
1999}
2000
2001fn default_c2r_output_len(input_len: usize) -> tenferro_tensor::Result<usize> {
2002    if input_len == 0 {
2003        return Err(tenferro_tensor::Error::invalid_argument(
2004            "irfft",
2005            "input spectrum axis length",
2006            "must be positive",
2007        ));
2008    }
2009    input_len
2010        .checked_sub(1)
2011        .and_then(|len| len.checked_mul(2))
2012        .ok_or_else(|| {
2013            tenferro_tensor::Error::invalid_argument(
2014                "irfft",
2015                "default output length",
2016                "overflows usize",
2017            )
2018        })
2019}
2020
2021fn validate_c2r_spectrum_len(
2022    input_len: usize,
2023    output_len: usize,
2024) -> tenferro_tensor::Result<usize> {
2025    let expected = output_len / 2 + 1;
2026    if input_len != expected {
2027        return Err(tenferro_tensor::Error::invalid_argument(
2028            "irfft",
2029            "spectrum",
2030            format!(
2031                "one-sided spectrum axis length mismatch: expected {expected} for output length {output_len}, got {input_len}"
2032            ),
2033        ));
2034    }
2035    Ok(expected)
2036}
2037
2038fn transform_len(shape: &[usize], axis: usize, n: Option<usize>) -> tenferro_tensor::Result<usize> {
2039    validate_axis("fft", shape, axis)?;
2040    let len = n.unwrap_or(shape[axis]);
2041    if len == 0 {
2042        return Err(tenferro_tensor::Error::invalid_argument(
2043            "fft",
2044            "transform length",
2045            "must be positive",
2046        ));
2047    }
2048    Ok(len)
2049}
2050
2051fn validate_axis(op: &'static str, shape: &[usize], axis: usize) -> tenferro_tensor::Result<()> {
2052    if axis >= shape.len() {
2053        return Err(tenferro_tensor::Error::axis_out_of_bounds(
2054            op,
2055            axis,
2056            shape.len(),
2057        ));
2058    }
2059    Ok(())
2060}
2061
2062#[cfg(test)]
2063mod concrete_tests;
2064
2065#[cfg(test)]
2066mod tests {
2067    use super::*;
2068
2069    #[test]
2070    fn fft_infer_output_meta_rejects_invalid_trait_inputs_without_panicking() {
2071        let op = FftOp::new(FftOperation::R2cOnesided, 0, None, FftNorm::Backward);
2072        let shape = [SymDim::from(4usize)];
2073
2074        assert!(
2075            tenferro_ops::ext_op::invoke_extension_shape_inference(&op, &[], &[&shape]).is_err()
2076        );
2077        assert!(
2078            tenferro_ops::ext_op::invoke_extension_shape_inference(&op, &[DType::F64], &[])
2079                .is_err()
2080        );
2081        assert!(tenferro_ops::ext_op::invoke_extension_shape_inference(
2082            &op,
2083            &[DType::I64],
2084            &[&shape]
2085        )
2086        .is_err());
2087
2088        let bad_axis = FftOp::new(FftOperation::C2cForward, 2, None, FftNorm::Backward);
2089        assert!(tenferro_ops::ext_op::invoke_extension_shape_inference(
2090            &bad_axis,
2091            &[DType::C64],
2092            &[&shape]
2093        )
2094        .is_err());
2095    }
2096
2097    #[test]
2098    fn checked_shape_product_rejects_overflow_before_allocation() {
2099        let err = cpu::checked_shape_product("fft", "output", &[usize::MAX, 2])
2100            .expect_err("overflowing output shape should be rejected");
2101
2102        assert!(err.to_string().contains("overflows usize"), "{err}");
2103    }
2104
2105    #[test]
2106    fn irfft_default_output_length_rejects_overflow() {
2107        let err = output_shape_c2r(&[usize::MAX], 0, None)
2108            .expect_err("default irfft output length should reject overflow");
2109
2110        assert!(err.to_string().contains("overflows usize"), "{err}");
2111    }
2112
2113    #[test]
2114    fn normalize_axis_handles_large_rank_without_isize_cast_wrap() {
2115        assert_eq!(normalize_axis("fft", 0, usize::MAX).unwrap(), 0);
2116        assert_eq!(
2117            normalize_axis("fft", -1, usize::MAX).unwrap(),
2118            usize::MAX - 1
2119        );
2120        assert!(normalize_axis("fft", isize::MIN, 3).is_err());
2121    }
2122
2123    #[test]
2124    fn axis_lane_layout_rejects_stride_overflow() {
2125        let err = cpu::LaneLayout::new(&[usize::MAX, 2], 1, 2)
2126            .expect_err("lane layout should reject stride overflow");
2127
2128        assert!(err.to_string().contains("overflows usize"), "{err}");
2129    }
2130
2131    #[cfg(feature = "autodiff")]
2132    #[test]
2133    fn fft_semantic_rules_emit_extension_first_jvp_and_length_restoring_transpose() {
2134        use tenferro_ops::dim_expr::DimExpr;
2135        use tenferro_runtime::program::{ProgramInputSpec, SemanticOpRef, SemanticProgramBuilder};
2136
2137        let fft_op = FftOp::new(FftOperation::C2cForward, 0, Some(2), FftNorm::Backward);
2138        let mut source = SemanticProgramBuilder::new();
2139        let source_input = source
2140            .input(ProgramInputSpec::new(DType::C64, [DimExpr::Const(4)]))
2141            .unwrap();
2142        let source_output = source
2143            .add_extension(Arc::new(fft_op), &[source_input])
2144            .unwrap()[0];
2145        let source = source.finish(&[source_output]).unwrap();
2146        let operation = source.program.operations().next().unwrap();
2147
2148        let rules = semantic_ad_rules().unwrap();
2149        let mut destination = SemanticProgramBuilder::new();
2150        let primal = destination
2151            .input(ProgramInputSpec::new(DType::C64, [DimExpr::Const(4)]))
2152            .unwrap();
2153        let tangent = destination
2154            .input(ProgramInputSpec::new(DType::C64, [DimExpr::Const(4)]))
2155            .unwrap();
2156        let primal_output = destination
2157            .add_extension(
2158                Arc::new(FftOp::new(
2159                    FftOperation::C2cForward,
2160                    0,
2161                    Some(2),
2162                    FftNorm::Backward,
2163                )),
2164                &[primal],
2165            )
2166            .unwrap()[0];
2167        let linearized = rules
2168            .linearize_operation(
2169                operation,
2170                &[primal],
2171                &[primal_output],
2172                &[AdValue::Value(tangent)],
2173                &[true],
2174                &mut destination,
2175            )
2176            .unwrap();
2177        let AdValue::Value(tangent_output) = linearized.tangent_outputs()[0] else {
2178            panic!("FFT tangent must be active");
2179        };
2180        let cotangent_inputs = rules
2181            .linear_transpose_operation(
2182                operation,
2183                &[primal],
2184                &[primal_output],
2185                &[AdValue::Value(tangent_output)],
2186                &[true],
2187                linearized.residuals(),
2188                &mut destination,
2189            )
2190            .unwrap();
2191        let AdValue::Value(cotangent_input) = cotangent_inputs[0] else {
2192            panic!("FFT cotangent must be active");
2193        };
2194        let frozen = destination
2195            .finish(&[tangent_output, cotangent_input])
2196            .unwrap();
2197        let operations: Vec<_> = frozen.program.operations().collect();
2198        assert!(
2199            operations
2200                .iter()
2201                .filter(|operation| matches!(operation.op(), SemanticOpRef::Extension(_)))
2202                .count()
2203                >= 3
2204        );
2205        assert!(operations.iter().any(|operation| matches!(
2206            operation.op(),
2207            SemanticOpRef::Core(CoreSemanticOp::DynamicTruncate { axis: 0 })
2208        )));
2209        assert!(operations.iter().any(|operation| matches!(
2210            operation.op(),
2211            SemanticOpRef::Core(CoreSemanticOp::PadToMatch { axis: 0 })
2212        )));
2213    }
2214
2215    #[cfg(feature = "autodiff")]
2216    #[test]
2217    fn fft_semantic_rules_run_through_whole_program_jvp_and_vjp() {
2218        use tenferro_ad::AdContext;
2219        use tenferro_ops::dim_expr::DimExpr;
2220        use tenferro_runtime::program::{ProgramInputSpec, SemanticOpRef, SemanticProgramBuilder};
2221
2222        let mut builder = SemanticProgramBuilder::new();
2223        let input = builder
2224            .input(ProgramInputSpec::new(DType::C64, [DimExpr::Const(4)]))
2225            .unwrap();
2226        let output = builder
2227            .add_extension(
2228                Arc::new(FftOp::new(
2229                    FftOperation::C2cForward,
2230                    0,
2231                    Some(2),
2232                    FftNorm::Backward,
2233                )),
2234                &[input],
2235            )
2236            .unwrap()[0];
2237        let source = builder.finish(&[output]).unwrap();
2238        let ad = AdContext::builder()
2239            .with_semantic_extension_rules(semantic_ad_rules().unwrap())
2240            .unwrap()
2241            .build()
2242            .unwrap();
2243
2244        let jvp = ad.jvp_program(&source, &[true]).unwrap();
2245        assert_eq!(jvp.derivative_input_indices(), &[Some(1)]);
2246        assert!(matches!(
2247            jvp.frozen().program.operations().last().unwrap().op(),
2248            SemanticOpRef::Extension(op) if op.family_id() == FFT_EXTENSION_FAMILY_ID
2249        ));
2250
2251        let vjp = ad.vjp_program(&source, &[true], &[true]).unwrap();
2252        assert_eq!(vjp.derivative_output_indices(), &[Some(0)]);
2253        assert!(vjp.frozen().program.operations().any(|operation| matches!(
2254            operation.op(),
2255            SemanticOpRef::Core(CoreSemanticOp::PadToMatch { axis: 0 })
2256                | SemanticOpRef::Core(CoreSemanticOp::DynamicTruncate { axis: 0 })
2257        )));
2258    }
2259}