diff --git a/onnxscript/function_libs/torch_lib/ops/fft.py b/onnxscript/function_libs/torch_lib/ops/fft.py index 30d086a5d5..d426d3cd50 100644 --- a/onnxscript/function_libs/torch_lib/ops/fft.py +++ b/onnxscript/function_libs/torch_lib/ops/fft.py @@ -172,17 +172,33 @@ def aten__fft_r2c( transformed = op.Unsqueeze(self, axes=[-1]) for idx, dimension in enumerate(reversed(dim)): + dft_length = op.Squeeze( + op.Shape(transformed, start=dimension, end=dimension + 1), + axes=[0], + ) transformed = _fftn_onnx_normalization( transformed, normalization, - op.Shape(transformed, start=dimension, end=dimension + 1), + dft_length, inverse=False, ) if idx > 0: - transformed = op.DFT(transformed, axis=dimension, inverse=False, onesided=False) + transformed = op.DFT( + transformed, + dft_length=dft_length, + axis=dimension, + inverse=False, + onesided=False, + ) else: # Torch computes one-sided FFT on the last dimension only. - transformed = op.DFT(transformed, axis=dimension, inverse=False, onesided=onesided) + transformed = op.DFT( + transformed, + dft_length=dft_length, + axis=dimension, + inverse=False, + onesided=onesided, + ) if unsqueeze_first_dim: transformed = op.Squeeze(transformed, axes=[0]) diff --git a/tests/function_libs/torch_lib/e2e_ops_tests.py b/tests/function_libs/torch_lib/e2e_ops_tests.py index b07f57ce20..e5b040b001 100644 --- a/tests/function_libs/torch_lib/e2e_ops_tests.py +++ b/tests/function_libs/torch_lib/e2e_ops_tests.py @@ -387,6 +387,23 @@ def forward(self, x): ) _testing.assert_onnx_program(onnx_program) + def test_rfft_produces_correct_dft_length(self): + class RFFTModel(torch.nn.Module): + def forward(self, x): + x = torch.fft.rfft(x, n=512, dim=-1) + return x.real**2 + x.imag**2 + + x = torch.randn(4, 512, dtype=torch.float32) + onnx_program = torch.onnx.export( + RFFTModel(), + (x,), + opset_version=20, + dynamo=True, + optimize=False, + ) + + self.assertEqual(onnx_program.model.graph.outputs[0].shape, [4, 512 // 2 + 1]) + def test_avg_pool(self): class Model(torch.nn.Module): def forward(self, x2d, x3d, x4d, x5d):