如何让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
相关产品推荐
相关产品推荐

