You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

TensorFlow中STFT与逆STFT信号复原异常问题排查

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.21 21:55:55