1#![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
180pub const FFT_EXTENSION_FAMILY_ID: &str = "tenferro-fft.fft.v1";
191
192#[derive(Default)]
199pub struct FftExecutor {
200 plans: FftPlanCache,
201}
202
203impl FftExecutor {
204 pub fn new(plans: FftPlanCache) -> Self {
206 Self { plans }
207 }
208
209 pub const fn plan_cache(&self) -> &FftPlanCache {
211 &self.plans
212 }
213
214 pub fn plan_cache_mut(&mut self) -> &mut FftPlanCache {
216 &mut self.plans
217 }
218
219 pub fn cache_stats(&self) -> CacheStats {
221 self.plans.stats()
222 }
223
224 pub fn clear_cache(&mut self) {
226 self.plans.clear();
227 }
228
229 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 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 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 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 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
378pub trait TracedTensorFftExt {
380 fn fft(&self, n: Option<usize>, axis: isize, norm: FftNorm) -> Result<TracedTensor>;
394
395 fn ifft(&self, n: Option<usize>, axis: isize, norm: FftNorm) -> Result<TracedTensor>;
408
409 fn rfft(&self, n: Option<usize>, axis: isize, norm: FftNorm) -> Result<TracedTensor>;
422
423 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
456pub trait TensorFftExt {
485 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 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 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 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
640pub trait TensorReadFftExt {
668 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 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 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 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
997fn 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 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 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#[cfg(feature = "autodiff")]
1466pub 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
1560fn 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 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
1590fn 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
1626fn 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
1664fn 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
1706fn 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}