Scipy与FFTW的fftshift输出不一致,如何让C++结果匹配Python?
让C++ xtensor的fftshift输出与Scipy完全一致的解决方案
问题描述
我用C++的xtensor生成meshgrid后执行fftshift操作,代码如下:
const std::size_t size = 5; auto ar = xt::meshgrid(xt::arange<double>(0, size), xt::arange<double>(0, size)); int translate = (size + 1) / 2; xt::xarray<double> x = std::get<0>(ar) - translate; xt::xarray<double> y = std::get<1>(ar) - translate; xt::xarray<double> xy_ = xt::stack(xt::xtuple(x, y)); auto p = xt::fftw::fftshift(xy_); std::cout << p << std::endl;
输出结果:
{{{-3., -2., -1., 0., 1.}, {-3., -2., -1., 0., 1.}, {-3., -2., -1., 0., 1.}, {-3., -2., -1., 0., 1.}, {-3., -2., -1., 0., 1.}}, {{-1., -1., -1., -1., -1.}, { 0., 0., 0., 0., 0.}, { 1., 1., 1., 1., 1.}, {-3., -3., -3., -3., -3.}, {-2., -2., -2., -2., -2.}}}
而Python中用相同逻辑执行fftshift的代码:
import numpy as np from scipy.fftpack import fftshift size = 5 mat = np.mgrid[:size, :size] - int((size + 1)/2) fftshifted_mat = fftshift(mat) print(fftshifted_mat)
输出结果:
[[[ 0 1 -3 -2 -1] [ 0 1 -3 -2 -1] [ 0 1 -3 -2 -1] [ 0 1 -3 -2 -1] [ 0 1 -3 -2 -1]] [[ 0 0 0 0 0] [ 1 1 1 1 1] [-3 -3 -3 -3 -3] [-2 -2 -2 -2 -2] [-1 -1 -1 -1 -1]]]
需要让C++中FFTW的fftshift输出与Scipy完全一致,但尝试过xt::roll、xt::transpose+xt::swap、手动循环移位均未成功。
尝试过程
- 最初用
xt::roll循环移位的代码:
for (std::size_t axis = 0; axis < xy_.shape().size(); ++axis){ std::size_t dim_size = xy_.shape()[axis]; std::size_t shift = (dim_size - 1) / 2; xy_ = xt::roll(xy_, shift, axis); }
仅当size=5或125时能得到与Scipy一致的结果,其他尺寸不生效。
- 后来参考实现了手动roll版本,能复现结果但速度较慢:
template <typename T> void fftshift_roll(xt::xarray<T>& array) { std::size_t ndims = array.dimension(); std::vector<std::ptrdiff_t> shift_indices(ndims); for (std::size_t i = 0; i < ndims; ++i) { std::ptrdiff_t shift = static_cast<std::ptrdiff_t>(array.shape(i)) / 2; shift_indices[i] = shift; } for (std::size_t i = 0; i < ndims; ++i) { auto rolled = xt::roll(array, shift_indices[i], i); array = xt::view(rolled, xt::all(), xt::all()); } }
正确实现方案
核心差异: Scipy的fftshift逻辑是对每个维度n,移动floor(n/2)个位置;而xt::fftw::fftshift遵循FFTW的默认逻辑,奇数维度移动(n+1)/2,这是两者结果不一致的根本原因。
要实现和Scipy完全一致的fftshift,需要严格按照Scipy的移位逻辑来实现,以下是优化后的高效版本,避免了多次拷贝带来的性能损耗:
template <typename T> xt::xarray<T> scipy_fftshift(const xt::xarray<T>& array) { std::size_t ndims = array.dimension(); std::vector<std::ptrdiff_t> shifts(ndims); for (std::size_t i = 0; i < ndims; ++i) { shifts[i] = static_cast<std::ptrdiff_t>(array.shape(i)) / 2; } // 一次性对所有维度执行roll操作,减少中间拷贝 return xt::roll(array, shifts); }
使用示例
替换原代码中的xt::fftw::fftshift调用即可:
// 原代码中的xy_生成后 auto p = scipy_fftshift(xy_); std::cout << p << std::endl;
该版本通过一次性传递所有维度的移位参数给xt::roll,既保证了逻辑与Scipy完全对齐,又避免了循环中多次赋值的性能开销,任意尺寸下都能得到和Scipy一致的结果。
内容的提问来源于stack exchange,提问作者Jarwin
相关产品推荐
相关产品推荐

