PyTorch中2D卷积高效伪逆的实现方法问询
高效实现2D卷积算子的伪逆(PyTorch + FFT方案)
核心思路
卷积作为线性算子,其伪逆可通过傅里叶变换在频域高效计算——时域卷积等价于频域逐元素乘积,伪逆操作只需对卷积核的傅里叶变换做Moore-Penrose伪逆处理,再逆变换回时域即可,完全规避了直接构造卷积矩阵的低效问题。
原理说明
- 卷积定理:输入
x与卷积核K的时域卷积y = K * x,对应频域中Y = F(K) ⊙ F(x)(F代表傅里叶变换,⊙代表逐元素乘积)。 - 卷积算子
A的Moore-Penrose伪逆A⁺,在频域对应F(A)的伪逆,计算公式为:
其中F(A⁺) = conj(F(A)) / (conj(F(A)) ⊙ F(A) + ε)conj是复共轭操作,ε为小正则化项,用于避免频域零值导致的除零错误。 - 伪逆结果
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
相关产品推荐
相关产品推荐

