为何自行实现的N维Fast Fourier Transform与NumPy结果不一致?
N维FFT实现与NumPy结果不一致的问题
我自己实现了N维FFT代码,但运行结果和NumPy的FFT函数对不上。已经验证过一维FFT功能正常,问题应该出在N维的实现里。
我的N维FFT代码
import numpy as np import math as m def nffourier(f, direct): dim = f.ndim N = f.shape G = np.zeros(f.shape, dtype=complex) G = f for k in range(dim): for i in range(N[k]): aux = G[(slice(None),) * (k) + (i,)] trans = ffourier(aux, direct) G[(slice(None),) * (k) + (i,)] = trans return G
已验证正常的一维FFT代码
def ffourier(f, direct): N = len(f) M = int(m.log(N)/m.log(2)) G = [] order = [] for i in range(N): order.append(int(bin(i)[2:])) digitos = len(str(order[0])) # 修正原代码中未定义的aux问题 for i in range(N): contenido_aux = str(int(order[i])) aux_len = len(contenido_aux) if aux_len < digitos: contenido_aux = '0' * (digitos - aux_len) + contenido_aux G.append(contenido_aux) for i in range(len(G)): G[i] = G[i][::-1] # 位反转重排 G = [int(g, 2) for g in G] G = [f[g] for g in G] signo = -1 if direct else 1 kmax = 1 for alfa in range(1, M+1): w1 = np.exp(signo * 1j * 2 * m.pi / (2**alfa)) kmax *= 2 W = 1 for k in range(kmax//2): for s in range(0, N, kmax): T0 = G[s + k] T1 = G[s + k + kmax//2] * W G[s + k] = T0 + T1 G[s + k + kmax//2] = T0 - T1 W *= w1 # 归一化 cte = 1 / m.sqrt(N) G = [val * cte for val in G] return np.array(G) # 改为返回numpy数组,方便后续处理
问题分析与修复方案
1. 数组引用问题(核心错误)
原N维代码中,G = f是直接让G指向原数组的引用,而非创建副本。这意味着在处理第一个维度时修改G,会直接改变原数组f,导致后续维度的变换基于已经被修改过的数据,结果完全错误。
修复:将G = f改为G = f.copy(),确保操作的是原数组的副本:
G = f.copy()
2. 维度遍历的效率与正确性优化
原代码中对每个维度k,遍历该维度的每个索引i并单独处理切片,逻辑正确但冗余。可以用NumPy的np.apply_along_axis简化,同时避免手动循环的潜在错误:
for k in range(dim): G = np.apply_along_axis(lambda x: ffourier(x, direct), axis=k, arr=G)
这样会自动沿着指定维度对所有元素做一维FFT,逻辑更清晰。
3. 归一化与NumPy的一致性问题
NumPy的FFT默认不做归一化:正变换np.fft.fft输出原始幅度,逆变换np.fft.ifft会除以总长度N。而你的一维FFT每次都会乘以1/sqrt(N),导致N维变换后总归一化系数是1/sqrt(N1*N2*...*Nd),和NumPy的行为不一致。
如果要和NumPy对齐,修改一维FFT的归一化部分:
if direct: # 正变换无归一化,对齐NumPy cte = 1.0 else: # 逆变换除以N,对齐NumPy cte = 1.0 / N G = [val * cte for val in G]
4. 一维FFT中的潜在bug
原一维代码中digitos = len(aux)这一行的aux未定义,会导致运行报错,已在修正后的代码中改为digitos = len(str(order[0])),确保获取正确的二进制位数。
内容的提问来源于stack exchange,提问作者elenafer
相关产品推荐
相关产品推荐

