TensorFlow中STFT与逆STFT信号复原异常问题排查
我在TensorFlow和PyTorch中分别实现了STFT与逆STFT,PyTorch版本可成功复原原始信号,但TensorFlow版本无法正常完成信号复原。
TensorFlow实现代码
import tensorflow as tf import math def spectro( signal: tf.Tensor, n_fft: int = 4096, hop_length: int = None ) -> tf.Tensor: if hop_length is None: hop_length = n_fft // 4 assert hop_length == n_fft // 4, "hop_length should be n_fft // 4" frames = int(math.ceil(signal.shape[-1] / hop_length)) pad_left = (hop_length // 2 * 3) pad_right = pad_left + frames * hop_length - signal.shape[-1] pad_second = (n_fft // 2) # shape (batch, channel, signal) -> (batch, channel, signal + pad_left + pad_right) padded_signal = STFTUtils.pad_1d(signal, (pad_left, pad_right), "REFLECT") padded_signal = STFTUtils.pad_1d(padded_signal, (pad_second, pad_second), "REFLECT") window_fn = tf.signal.hann_window *other, length = padded_signal.shape padded_signal = tf.reshape(padded_signal,[-1,length]) stfts = tf.signal.stft( padded_signal, frame_length=n_fft, frame_step=hop_length, fft_length=n_fft, window_fn=window_fn, ) # shape (batch, channel, fft_bins, fft_length // 2 + 1) _,frames_sec, frequency = stfts.shape stfts = stfts[..., : , :-1] stfts = stfts[..., 2 : 2 + frames, : ] return stfts # shape (batch, channel, fft_bins - 4, fft_length // 2) def inverse_spectro(spectrograms, hop_length=None, signal_length=None) -> tf.Tensor: spectrograms = tf.pad(spectrograms, [[0, 0], [0, 0], [2,2], [0,1]]) *other_dims, num_freqs = spectrograms.shape n_fft = 2 * num_freqs - 2 if hop_length is None: hop_length = n_fft // 4 window_fn = tf.signal.hann_window signal = tf.signal.inverse_stft( spectrograms, frame_length=n_fft, frame_step=hop_length, fft_length=n_fft, window_fn=window_fn ) # shape (batch, channel, signal_length + padding) padding = (hop_length // 2 * 3) + (n_fft // 2) # Amount of padding that was applied signal = signal[..., padding : signal_length + padding] return signal # shape (batch, channel, signal_length)
PyTorch实现代码
import torch as th def spectro(x, n_fft=512, hop_length=None, pad=0): *other, length = x.shape x = x.reshape(-1, length) is_mps = x.device.type == 'mps' if is_mps: x = x.cpu() #print("before stft", x, x.shape, torch.sum(torch.square(x))) z = th.stft(x, n_fft * (1 + pad), hop_length or n_fft // 4, window=th.hann_window(n_fft).to(x), win_length=n_fft, normalized=False, center=True, return_complex=True, pad_mode='reflect') _, freqs, frame = z.shape return z.view(*other, freqs, frame) def ispectro(z, hop_length=None, length=None, pad=0): *other, freqs, frames = z.shape n_fft = 2 * freqs - 2 z = z.view(-1, freqs, frames) win_length = n_fft // (1 + pad) is_mps = z.device.type == 'mps' if is_mps: z = z.cpu() x = th.istft(z, n_fft, hop_length, window=th.hann_window(win_length).to(z.real), win_length=win_length, normalized=False, length=length, center=True) _, length = x.shape return x.view(*other, length)
问题分析与修复建议
1. 手动裁剪与补全逻辑不匹配
TensorFlow版spectro中对STFT结果做了两次不必要的裁剪:
- 频率维度:
stfts = stfts[..., : , :-1]丢弃最后一个频率分量 - 时间维度:
stfts = stfts[..., 2 : 2 + frames, : ]丢弃前2帧,仅保留手动计算的frames个帧
逆变换时用tf.pad补全的是零值(默认CONSTANT模式),和正变换的反射填充完全不一致,直接导致边缘信息丢失,无法准确复原。
修复:移除正变换中的手动裁剪逻辑,让tf.signal.stft自动生成完整结果;若必须裁剪,逆变换需补全与正变换填充方式一致的内容(而非零值)。
2. 自定义填充与TF内置STFT填充冲突
tf.signal.stft默认会自动对信号做中心填充(pad_end=True),保证能完整分割所有帧。但你手动做了两次反射填充:
padded_signal = STFTUtils.pad_1d(signal, (pad_left, pad_right), "REFLECT") padded_signal = STFTUtils.pad_1d(padded_signal, (pad_second, pad_second), "REFLECT")
这导致信号被过度填充,逆变换时计算的padding值完全不匹配实际填充量,最终裁剪出的原始信号位置错误。
修复:移除手动填充逻辑,改用tf.signal.stft默认的自动填充;若要手动控制填充,需设置pad_end=False避免重复填充。
3. 帧数量计算错误
你手动计算的frames = int(math.ceil(signal.shape[-1] / hop_length))和tf.signal.stft实际生成的帧数量不一致。STFT帧数量的正确计算公式为:
frames = ceil((signal_length - frame_length) / frame_step) + 1
手动计算忽略了帧长度的影响,导致裁剪后的STFT帧数量错误,逆变换无法还原正确长度的信号。
修复:直接使用tf.shape(stfts)[-2]获取实际帧数量,不要手动计算。
4. 逆变换信号裁剪逻辑错误
inverse_spectro中计算的padding = (hop_length // 2 * 3) + (n_fft // 2)未考虑STFT自动填充的部分,导致裁剪起始和结束位置完全错误。此外,tf.signal.inverse_stft返回的信号长度由输入STFT的帧数量、帧步长和帧长度决定,不能用手动计算的signal_length + padding来裁剪。
修复:若使用STFT默认填充,可通过tf.signal.inverse_stft_window_fn生成窗口,结合frame_length和frame_step计算需要裁剪的填充量;或直接基于逆变换结果调整到目标长度。
5. 窗口函数归一化差异
PyTorch的istft会自动处理窗口重叠相加的归一化,而TensorFlow的tf.signal.inverse_stft需要显式指定forward_window_fn匹配正变换窗口,否则会出现幅度偏差。当前仅传入window_fn,未保证逆变换窗口与正变换完全一致,也未满足重叠相加的条件(窗口重叠部分之和为常数)。
修复:逆变换中使用tf.signal.inverse_stft_window_fn生成对应逆窗口:
forward_window = tf.signal.hann_window(n_fft) inverse_window = tf.signal.inverse_stft_window_fn(hop_length, forward_window) signal = tf.signal.inverse_stft( spectrograms, frame_length=n_fft, frame_step=hop_length, fft_length=n_fft, window_fn=lambda x: inverse_window )
内容的提问来源于stack exchange,提问作者Bisnu Sarkar

