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

PyTorch中2D卷积高效伪逆的实现方法问询

高效实现2D卷积算子的伪逆(PyTorch + FFT方案)

核心思路

卷积作为线性算子,其伪逆可通过傅里叶变换在频域高效计算——时域卷积等价于频域逐元素乘积,伪逆操作只需对卷积核的傅里叶变换做Moore-Penrose伪逆处理,再逆变换回时域即可,完全规避了直接构造卷积矩阵的低效问题。

原理说明

  1. 卷积定理:输入x与卷积核K的时域卷积y = K * x,对应频域中Y = F(K) ⊙ F(x)(F代表傅里叶变换,⊙代表逐元素乘积)。
  2. 卷积算子A的Moore-Penrose伪逆A⁺,在频域对应F(A)的伪逆,计算公式为:
    F(A⁺) = conj(F(A)) / (conj(F(A)) ⊙ F(A) + ε)
    
    其中conj是复共轭操作,ε为小正则化项,用于避免频域零值导致的除零错误。
  3. 伪逆结果x_hat = A⁺y,可通过x_hat = F⁻¹(F(A⁺) ⊙ F(y))得到(F⁻¹代表逆傅里叶变换)。

代码实现

import torch
import torch.nn.functional as F

def conv_pinv(y, K, padding=1, eps=1e-8):
    b, c_out, h, w = y.shape
    c_in, _, k_h, k_w = K.shape
    
    # 1. 将卷积核填充到与输入特征图相同尺寸,计算傅里叶变换
    K_padded = F.pad(K, (0, w - k_w, 0, h - k_h))
    K_fft = torch.fft.fft2(K_padded, dim=(2, 3))
    
    # 2. 计算输入y的傅里叶变换
    y_fft = torch.fft.fft2(y, dim=(2, 3))
    
    # 3. 频域计算伪逆:转置通道后取共轭,再除以(自身模平方+正则项)
    K_fft_conj = torch.conj(K_fft.transpose(0, 1))
    denominator = torch.abs(K_fft) ** 2 + eps
    pinv_fft = K_fft_conj / denominator.unsqueeze(0)
    
    # 4. 频域乘积后做逆傅里叶变换,提取实部(输入为实数,虚部可忽略)
    x_hat_fft = torch.einsum('bchw,cchw->bchw', y_fft, pinv_fft)
    x_hat = torch.fft.ifft2(x_hat_fft, dim=(2, 3)).real
    
    return x_hat

# 测试验证
c = 32
K = torch.randn(c, c, 3, 3)
x = torch.randn(1, c, 128, 128)
y = F.conv2d(x, K, padding=1)

x_hat = conv_pinv(y, K, padding=1)

# 验证伪逆性质:A(A⁺y) ≈ y
y_recon = F.conv2d(x_hat, K, padding=1)
print(f"重构误差均值: {torch.mean(torch.abs(y_recon - y)):.6f}")

关键注意事项

  • 尺寸匹配:卷积核必须填充到与输入特征图相同尺寸,确保傅里叶变换后维度一致。
  • 正则化项:eps需根据场景调整,过小易引发数值不稳定,过大则会增加重构误差。
  • 通道处理:卷积核通道维度为[c_out, c_in, k_h, k_w],计算共轭前需转置通道,对应卷积算子的伴随算子逻辑。
  • 实部提取:由于输入和卷积核均为实数,逆傅里叶变换后的虚部接近0,取实部即可得到有效结果。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 11:15:28