PyTorch中基于振幅与相位的逆FFT实现问题
正确的FFT逆变换实现方法
你之前的错误在于只构建了复数的虚部,FFT的复数结果是实部 + 虚部×i的形式,必须完整构建这个复数张量才能进行逆FFT。
步骤说明
- 从修改后的振幅(
modified_amp)和相位(modified_pha)重新构建复数FFT结果:
实部 = 振幅 × cos(相位)
虚部 = 振幅 × sin(相位) - 用
torch.complex()把实部和虚部组合成复数张量 - 调用
torch.fft.ifft()执行逆变换 - 最后取结果的实部(因为原始输入是实数信号,逆FFT后可能存在微小的数值虚部,需要剔除)
完整代码示例
import torch def calculate_fft(x): fft_im = torch.fft.fft(x.clone()) # bx3xhxw,复数张量 fft_amp = torch.sqrt(fft_im.real**2 + fft_im.imag**2) fft_pha = torch.atan2(fft_im.imag, fft_im.real) return fft_amp, fft_pha # 假设已经得到修改后的振幅和相位 modified_amp, modified_pha = ... # 这里替换为你修改后的结果 # 重新构建复数FFT张量 reconstructed_fft = torch.complex(modified_amp * torch.cos(modified_pha), modified_amp * torch.sin(modified_pha)) # 执行逆FFT并取实部 reconstructed_x = torch.fft.ifft(reconstructed_fft).real
补充说明
如果你的输入是图像类的多维张量(比如代码里的bx3xhxw),更适合用2D FFT处理空间维度(h和w),对应的函数是torch.fft.fft2和torch.fft.ifft2,示例如下:
# 2D FFT版本的提取函数 def calculate_fft2d(x): fft_im = torch.fft.fft2(x.clone()) # bx3xhxw,2D复数FFT结果 fft_amp = torch.sqrt(fft_im.real**2 + fft_im.imag**2) fft_pha = torch.atan2(fft_im.imag, fft_im.real) return fft_amp, fft_pha # 对应的逆变换 reconstructed_fft2d = torch.complex(modified_amp * torch.cos(modified_pha), modified_amp * torch.sin(modified_pha)) reconstructed_x2d = torch.fft.ifft2(reconstructed_fft2d).real
内容的提问来源于stack exchange,提问作者FeiiYin
相关产品推荐
相关产品推荐

