You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何让PyTorch中旧版torch.fft与新版torch.fft.fft兼容?

如何让新版torch.fft.fft输出与PyTorch v1.7.1及旧版torch.fft语义一致的结果?

问题说明

PyTorch v1.7.1及更早版本的torch.fft接受**[实部, 虚部]格式的实数张量作为输入,而新版torch.fft.fft要求输入的最深维度为复数格式**,两者输入维度存在差异。

我在处理以下对应2D图像场景的3D输入时遇到了问题:

a = torch.tensor([[[1.0, 2.0], [3.0, 4.0], [5.0, 6.0], [7.0, 8.0]],
                  [[2.0, 3.0], [4.0, 5.0], [6.0, 7.0], [8.0, 9.0]],
                  [[3.0, 4.0], [5.0, 6.0], [7.0, 8.0], [9.0,10.0]]])
# 旧版调用方式:print(torch.fft(a, signal_ndim=2, normalized=False))

已实现的1D数据场景兼容示例

针对对应1D数据场景的2D输入,我已完成兼容适配,代码及输出如下:

旧版(PyTorch v1.7.1)代码

b = torch.tensor([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]])
print(b.shape)
print(torch.fft(b, signal_ndim=1, normalized=False))

输出:

torch.Size([3, 2])
tensor([[ 9.0000, 12.0000],
        [-4.7321, -1.2679],
        [-1.2679, -4.7321]])

新版兼容代码

import torch.fft
b = torch.tensor([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]])
# 将[实部,虚部]转换为复数张量
b = b[:,0] + 1j * b[:,1]
# 保持维度匹配
b = torch.unsqueeze(b, 1)
print(b)
print(b.shape)
print(torch.fft.fft(b, dim=0))

输出:

tensor([[1.+2.j],
        [3.+4.j],
        [5.+6.j]])
torch.Size([3, 1])
tensor([[ 9.0000+12.0000j],
        [-4.7321-1.2679j],
        [-1.2679-4.7321j]])

注:在PyTorch v1.7.1中可同时使用新旧API,但因模块名冲突不能同时导入。

寻求帮助

恳请提供3D输入(对应2D图像场景)下的适配方案,让新版torch.fft.fft输出与旧版torch.fft(a, signal_ndim=2, normalized=False)语义完全一致的结果。

内容的提问来源于stack exchange,提问作者modu01

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.29 06:27:03