PyTorch实战:用傅里叶变换给图像做‘体检’,分离振幅与相位(附完整代码)

当我们观察一张照片时,看到的只是像素的排列组合。但就像医生通过X光片能看到骨骼结构一样,傅里叶变换能让我们看到图像的"内部构造"——那些隐藏在像素背后的频域信息。本文将带你用PyTorch的torch.fft模块,像医生解读体检报告一样,拆解图像的振幅谱和相位谱,理解它们各自对图像重建的独特贡献。

1. 频域思维:图像的另一种表达方式

想象你正在听一首交响乐。时域表示就是声音波形随时间的变化,而频域表示则告诉你每个音符(频率成分)的强度和出现时间。图像处理也是如此——空间域是我们熟悉的像素网格,频域则揭示了图像中各种"纹理频率"的分布情况。

傅里叶变换的核心价值在于:

  • 振幅谱:记录每个频率成分的"音量"大小
  • 相位谱:记录每个频率成分的"演奏时机"
  • 频域可逆性:就像乐谱能还原为音乐,频域数据能完美重建原始图像
import torch
import matplotlib.pyplot as plt

def load_image_to_tensor(image_path, size=(256,256)):
    from PIL import Image
    transform = torchvision.transforms.Compose([
        transforms.Resize(size),
        transforms.ToTensor(),
    ])
    return transform(Image.open(image_path)).unsqueeze(0)  # 添加batch维度

2. PyTorch频域诊断工具箱

PyTorch的torch.fft模块提供了完整的频域操作接口。与NumPy的实现相比,它天然支持GPU加速,并能与自动微分系统无缝集成。

2.1 核心操作三步曲

  1. 频域转换torch.fft.fft2执行二维傅里叶变换
  2. 频谱中心化torch.fft.fftshift将零频移到频谱中心
  3. 成分分离torch.abs获取振幅,torch.angle获取相位
def image_spectrum_analysis(img_tensor):
    # 转换为灰度(单通道)
    if img_tensor.shape[1] == 3:
        img_tensor = 0.299 * img_tensor[:,0] + 0.587 * img_tensor[:,1] + 0.114 * img_tensor[:,2]
        img_tensor = img_tensor.unsqueeze(1)
    
    # 执行FFT
    fft = torch.fft.fft2(img_tensor)
    fft_shifted = torch.fft.fftshift(fft)
    
    # 获取振幅和相位
    magnitude = torch.abs(fft_shifted)
    phase = torch.angle(fft_shifted)
    
    return magnitude, phase

2.2 可视化诊断报告

为了直观理解频域信息,我们需要对原始频谱数据进行可视化处理:

处理步骤 目的 数学操作
对数变换 压缩动态范围 20 * log(1 + magnitude)
归一化 适配显示范围 (spectrum - min) / (max - min)
伪彩色 增强视觉区分 应用jetviridis色图
def visualize_spectrum(magnitude, phase):
    # 振幅谱处理
    log_magnitude = 20 * torch.log1p(magnitude)
    norm_mag = (log_magnitude - log_magnitude.min()) / (log_magnitude.max() - log_magnitude.min())
    
    # 相位谱处理
    norm_phase = (phase + torch.pi) / (2 * torch.pi)  # 将[-π, π]映射到[0,1]
    
    # 创建画布
    plt.figure(figsize=(15,5))
    
    # 绘制振幅谱
    plt.subplot(1,3,1)
    plt.imshow(norm_mag[0,0].cpu().numpy(), cmap='jet')
    plt.title('Amplitude Spectrum'), plt.axis('off')
    
    # 绘制相位谱
    plt.subplot(1,3,2)
    plt.imshow(norm_phase[0,0].cpu().numpy(), cmap='hsv')
    plt.title('Phase Spectrum'), plt.axis('off')
    
    # 绘制原始图像
    plt.subplot(1,3,3)
    plt.imshow(img_tensor[0,0].cpu().numpy(), cmap='gray')
    plt.title('Original Image'), plt.axis('off')
    
    plt.tight_layout()
    plt.show()

3. 成分分离实验:谁决定了图像的本质?

通过控制变量实验,我们可以直观感受振幅和相位各自的作用:

3.1 仅保留相位信息

实验方法:将振幅谱替换为常数,仅用原始相位谱重建图像

def reconstruct_from_phase(phase):
    constant = torch.mean(torch.abs(fft_shifted))
    reconstructed = constant * torch.exp(1j * phase)
    img_recon = torch.abs(torch.fft.ifft2(torch.fft.ifftshift(reconstructed)))
    return img_recon

现象观察

  • 重建图像保留了原始图像的结构轮廓
  • 细节纹理变得模糊不清
  • 证明相位信息主导图像的结构识别

3.2 仅保留振幅信息

实验方法:将相位谱替换为随机噪声,仅用原始振幅谱重建图像

def reconstruct_from_magnitude(magnitude):
    random_phase = 2 * torch.pi * torch.rand_like(magnitude) - torch.pi
    reconstructed = magnitude * torch.exp(1j * random_phase)
    img_recon = torch.abs(torch.fft.ifft2(torch.fft.ifftshift(reconstructed)))
    return img_recon

现象观察

  • 重建图像呈现类似"星云"的纹理模式
  • 完全丢失原始图像的结构信息
  • 证明振幅信息决定图像的视觉风格

4. 实战应用:频域编辑技巧

理解了振幅和相位的作用后,我们可以进行有针对性的频域编辑:

4.1 频域混合术

将图像A的振幅谱与图像B的相位谱结合,会产生有趣的视觉效果:

def frequency_domain_mixing(img1, img2):
    # 获取图像1的振幅谱
    mag1, _ = image_spectrum_analysis(img1)
    
    # 获取图像2的相位谱
    _, phase2 = image_spectrum_analysis(img2)
    
    # 混合重建
    mixed = mag1 * torch.exp(1j * phase2)
    recon = torch.abs(torch.fft.ifft2(torch.fft.ifftshift(mixed)))
    
    return recon

4.2 自适应频域滤波

基于振幅谱的统计特性实现智能滤波:

def adaptive_filter(img_tensor, keep_ratio=0.2):
    magnitude, phase = image_spectrum_analysis(img_tensor)
    
    # 计算振幅阈值
    sorted_mag = torch.sort(magnitude.flatten()).values
    threshold = sorted_mag[int((1-keep_ratio)*len(sorted_mag))]
    
    # 创建掩膜
    mask = (magnitude > threshold).float()
    
    # 应用滤波
    filtered = mask * magnitude * torch.exp(1j * phase)
    recon = torch.abs(torch.fft.ifft2(torch.fft.ifftshift(filtered)))
    
    return recon

5. 完整工作流示例

下面是一个端到端的图像频域分析流程,包含从加载到可视化的所有步骤:

import torch
import torchvision.transforms as transforms
from PIL import Image
import matplotlib.pyplot as plt

# 1. 图像加载与预处理
image_path = "your_image.jpg"
img_tensor = load_image_to_tensor(image_path)

# 2. 频域分析
magnitude, phase = image_spectrum_analysis(img_tensor)

# 3. 可视化诊断
visualize_spectrum(magnitude, phase)

# 4. 成分分离重建
img_phase_only = reconstruct_from_phase(phase)
img_mag_only = reconstruct_from_magnitude(magnitude)

# 5. 显示对比结果
plt.figure(figsize=(12,4))
plt.subplot(1,3,1), plt.imshow(img_tensor[0,0].cpu().numpy(), cmap='gray')
plt.title('Original'), plt.axis('off')
plt.subplot(1,3,2), plt.imshow(img_phase_only[0,0].cpu().numpy(), cmap='gray')
plt.title('Phase Only'), plt.axis('off')
plt.subplot(1,3,3), plt.imshow(img_mag_only[0,0].cpu().numpy(), cmap='gray')
plt.title('Magnitude Only'), plt.axis('off')
plt.show()

在实际项目中,这种频域分析方法可以帮助我们:

  • 检测图像中的周期性噪声(如摩尔纹)
  • 设计更智能的图像增强算法
  • 开发基于频域特征的图像检索系统
  • 优化神经网络中的频域注意力机制
Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐