如何使用torch.fft.fft2输出与旧版PyTorch的torch.fft一致的结果
问题1:旧版torch.fft的存储逻辑
PyTorch 1.1.0版本还没有原生复数张量类型,因此所有复数结果都要通过实张量拆分存储。torch.fft会把FFT计算得到的复数结果的实部、虚部沿张量的最后一个维度拼接,比如输入是形状为[256, 256]的2D实矩阵,调用torch.fft计算2D FFT时,输出形状为[256, 256, 2],最后一个维度的第0位存对应位置复数的实部,第1位存虚部。你看到的文档描述“形状与输入一致”是针对复数输入的场景:如果输入本身就是最后一个维度为2的实张量(用来表示复数输入),输出的形状就和输入完全相同。
你贴出的旧版输出全为实数值,是因为你只截取了输出张量最后一个维度第0位(实部)的部分数据。
问题2:用torch.fft.fft2得到和旧版一致的结果
只需要把新版torch.fft.fft2输出的复数张量拆分为实部、虚部,再沿最后一个维度拼接即可,参考代码如下:
import torch import numpy as np # 生成测试输入 input = torch.from_numpy(np.random.rand(256,256)) # 新版2D FFT计算 new_fft = torch.fft.fft2(input) # 拼接实部虚部,得到和旧版torch.fft完全一致的结果 old_style_fft = torch.stack([new_fft.real, new_fft.imag], dim=-1)
如果你只需要旧版输出的实部,直接取new_fft.real即可。
内容的提问来源于stack exchange,提问作者youban
相关产品推荐
相关产品推荐

