基于FFT实现卷积的ftConv函数异常?求原因及修复方案
基于FFT的卷积函数
ftConv输出错误的原因及修复方案 问题背景
我正在编写基于卷积定理、使用快速傅里叶变换(FFT)实现卷积的代码,包含4个核心函数:
Conv:常规卷积函数;fft2d:计算图像的二维FFT;ifft2d:计算逆傅里叶变换(ifft2d<int>(fft2d<int>(image))可还原原图像);ftConv:基于FFT的卷积函数,将图像或核零填充至两者最大尺寸,分别计算FFT后逐元素相乘,再执行ifft2d。
但ftConv的结果与Conv完全不同,测试代码及输出如下:
测试代码
int main() { std::vector<std::vector<int>> image = {{1, 2, 0, 4, 5}, {6, 0, 8, 0, 10}, {0, 12, 13, 14, 0}, {16, 0, 18, 0, 20}, {21, 22, 0, 24, 25}}; std::vector<std::vector<int>> kernel = {{-1, 0, 1}, {-2, 0, 2}, {-1, 0, 1}}; auto result1 = Conv(image, kernel); auto result2 = ftConv(image, kernel); for (const auto& row : result1) { for (auto pixel : row) std::cout << (double) pixel << " "; std::cout << std::endl; } std::cout << std::endl; for (const auto& row : result2) { for (auto pixel : row) std::cout << (double) pixel << " "; std::cout << std::endl; } return 0; }
输出结果
0 0 0 0 0 0 16 4 -4 0 0 30 4 -22 0 0 -4 4 16 0 0 0 0 0 0 -15 -15 31 4 -4 -57 -7 29 41 -6 -37 2 19 21 -6 3 4 11 -15 -4 21 5 6 -29 -4
错误原因分析
混淆卷积与互相关定义
你的Conv函数实际实现的是互相关(直接用核与图像窗口相乘求和),而FFT卷积定理要求的是标准卷积——需要先将核翻转180度再与图像做互相关。直接用原核参与FFT相乘,会导致结果与常规互相关完全不匹配。零填充尺寸错误
为了用FFT实现与常规卷积一致的线性卷积,填充后的尺寸应为图像尺寸 + 核尺寸 - 1,而非取两者的最大尺寸。取最大尺寸会触发循环卷积的混叠效应,结果完全偏离预期。核的填充位置与方向错误
当前代码将核居中填充,但FFT卷积要求核填充到左上角(对应循环卷积的对齐逻辑),且未翻转核,导致结果相位偏移、位置错位。ifft2d索引反转错误ifft2d最后转换实数部分的循环中,行列索引写反,导致图像被转置,进一步加剧结果错误。
修复方案
1. 修正ifft2d的索引错误
将ifft2d末尾的实数转换循环改为:
for (size_t i = 0; i < imgHeight; ++i) for (size_t j = 0; j < imgWidth; ++j) result[i][j] = static_cast<T>(std::real(imgITransformed[i][j]));
确保结果的行列维度与输入一致,避免图像转置。
2. 调整零填充尺寸为线性卷积要求的大小
在ftConv中修改填充尺寸计算:
size_t resWidth = imgWidth + kWidth - 1; size_t resHeight = imgHeight + kHeight - 1;
消除循环卷积的混叠效应,得到正确的线性卷积结果。
3. 对核进行180度翻转
在填充核之前,先翻转核:
// 翻转核(旋转180度) std::vector<std::vector<U>> flippedKernel(kHeight, std::vector<U>(kWidth)); for (size_t i = 0; i < kHeight; ++i) for (size_t j = 0; j < kWidth; ++j) flippedKernel[i][j] = kernel[kHeight - 1 - i][kWidth - 1 - j];
让FFT卷积的逻辑匹配常规互相关的结果。
4. 修正图像与核的填充位置
将图像和翻转后的核填充到左上角,确保卷积对齐:
// 图像填充到左上角 for (size_t i = 0; i < imgHeight; ++i) for (size_t j = 0; j < imgWidth; ++j) imgPadded[i][j] = image[i][j]; // 翻转后的核填充到左上角 for (size_t i = 0; i < kHeight; ++i) for (size_t j = 0; j < kWidth; ++j) kPadded[i][j] = flippedKernel[i][j];
5. 提取与原Conv匹配的有效区域(可选)
如果需要和原Conv的输出尺寸一致,可从线性卷积结果中提取有效区域:
int kCenter = kHeight / 2; std::vector<std::vector<T>> validResult(imgHeight, std::vector<T>(imgWidth, 0)); for (int i = kCenter; i < imgHeight - kCenter; ++i) for (int j = kCenter; j < imgWidth - kCenter; ++j) validResult[i][j] = static_cast<T>(std::round(result[i][j]));
用std::round处理FFT计算带来的浮点精度误差。
修复后的完整ftConv函数
template <class T, class U> std::vector<std::vector<T>> ftConv(const std::vector<std::vector<T>>& image, const std::vector<std::vector<U>>& kernel) { size_t imgHeight = image.size(), imgWidth = image[0].size(); size_t kHeight = kernel.size(), kWidth = kernel[0].size(); // 修正填充尺寸为线性卷积所需大小 size_t resWidth = imgWidth + kWidth - 1; size_t resHeight = imgHeight + kHeight - 1; std::vector<std::vector<T>> imgPadded(resHeight, std::vector<T>(resWidth, 0)); std::vector<std::vector<T>> kPadded(resHeight, std::vector<T>(resWidth, 0)); // 图像填充到左上角 for (size_t i = 0; i < imgHeight; ++i) for (size_t j = 0; j < imgWidth; ++j) imgPadded[i][j] = image[i][j]; // 翻转核(旋转180度) std::vector<std::vector<U>> flippedKernel(kHeight, std::vector<U>(kWidth)); for (size_t i = 0; i < kHeight; ++i) for (size_t j = 0; j < kWidth; ++j) flippedKernel[i][j] = kernel[kHeight - 1 - i][kWidth - 1 - j]; // 翻转后的核填充到左上角 for (size_t i = 0; i < kHeight; ++i) for (size_t j = 0; j < kWidth; ++j) kPadded[i][j] = flippedKernel[i][j]; freqMat imgTransformed = fft2d<T>(imgPadded); freqMat kTransformed = fft2d<U>(kPadded); freqMat resultTransformed(resHeight, freqVec(resWidth)); for (size_t i = 0; i < resHeight; ++i) for (size_t j = 0; j < resWidth; ++j) resultTransformed[i][j] = imgTransformed[i][j] * kTransformed[i][j]; std::vector<std::vector<T>> result = ifft2d<T>(resultTransformed); // 提取与原Conv对应的有效区域 int kCenter = kHeight / 2; std::vector<std::vector<T>> validResult(imgHeight, std::vector<T>(imgWidth, 0)); for (int i = kCenter; i < imgHeight - kCenter; ++i) for (int j = kCenter; j < imgWidth - kCenter; ++j) validResult[i][j] = static_cast<T>(std::round(result[i][j])); return validResult; }
内容的提问来源于stack exchange,提问作者Michael Shkarubski
相关产品推荐
相关产品推荐

