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

TensorFlow多GPU训练时损失函数出现NaN的问题排查

多GPU训练STFT损失出现NaN的排查与解决

问题背景

使用双3090Ti通过tf.distribute.MirroredStrategy搭建多GPU训练环境,损失函数包含波形损失+多分辨率STFT损失。仅用波形损失时多GPU训练正常;加入STFT损失后,单GPU训练(两块GPU分别测试)正常,但多GPU训练时,tf.signal.stft计算后出现NaN(输入数据无NaN)。

可能原因及对应解决方法

1. 数值稳定性不足

多GPU环境下,设备间张量计算的精度差异或极小值,会在STFT损失的对数、除法操作中被放大为NaN。

解决方法:

  • 强化对数计算的数值保护:在LogSTFTMagnitudeLoss中,给谱图幅值添加更大的epsilon,避免取对数时遇到0:
    class LogSTFTMagnitudeLoss(tf.keras.layers.Layer):
        def call(self, x_mag, y_mag):
            # 增大epsilon防止log(0)
            safe_x = tf.maximum(x_mag, 1e-6)
            safe_y = tf.maximum(y_mag, 1e-6)
            loss = tf.reduce_mean(tf.abs(tf.math.log(safe_y) - tf.math.log(safe_x)))
            return loss
    
  • 避免分母为0:在SpectralConvergenceLoss的除法操作中,给分母添加epsilon:
    class SpectralConvergenceLoss(tf.keras.layers.Layer):
        def call(self, x_mag, y_mag):        
            numerator = tf.norm(y_mag - x_mag, ord=2)
            # 添加epsilon避免除以0
            denominator = tf.maximum(tf.norm(y_mag, ord=2), 1e-6)
            loss = numerator / denominator
            return loss
    

2. 分布式策略下Layer初始化异常

MultiResolutionSTFTLoss中创建的STFTLoss实例未在strategy.scope()内完成初始化,导致设备间张量分布错乱。

解决方法:

将自定义Loss Layer的初始化延迟到build阶段(此时已处于策略作用域内):

class custom_loss(tf.keras.losses.Loss):
    def __init__(self, BATCH_SIZE=-1, extra=0.0, **kwargs):
        super().__init__(**kwargs)
        self.BATCH_SIZE = BATCH_SIZE
        self.mrstft_loss = None

    def build(self, input_shape):
        # 在build阶段初始化STFT损失,确保处于分布式策略作用域
        self.mrstft_loss = MultiResolutionSTFTLoss(batch_size=self.BATCH_SIZE)
        super().build(input_shape)

    def call(self, y_true, y_pred):
        if self.mrstft_loss is None:
            self.mrstft_loss = MultiResolutionSTFTLoss(batch_size=self.BATCH_SIZE)
        # 原call方法的其余逻辑不变

3. 分布式数据集维度不匹配

多GPU训练时,每个replica处理的是BATCH_SIZE_PER_REPLICA的数据,若STFT计算的padding、帧分割逻辑未适配分布式维度,会产生无效数值。

解决方法:

  • 检查分布式数据集输出的y_true、y_pred形状,确保与单GPU环境下的维度一致(批量维度应为BATCH_SIZE_PER_REPLICA)。
  • 优化stft函数的窗口生成逻辑,确保窗口与帧长度匹配:
    def stft(x, fft_size, hop_size, win_length, window):
        window_fn = tf.signal.hann_window
        pad_amount = fft_size // 2
        x = tf.pad(x, [[0, 0], [pad_amount, pad_amount]], mode='REFLECT')
        # 明确生成与win_length匹配的窗口
        window = window_fn(win_length, dtype=x.dtype)
        x_stft = tf.signal.stft(
            x, 
            fft_length=fft_size, 
            frame_step=hop_size, 
            frame_length=win_length, 
            window_fn=lambda win_len, dtype: window
        )
        real = tf.math.real(x_stft)
        imag = tf.math.imag(x_stft)
        magnitude = tf.sqrt(tf.maximum(real**2 + imag**2, 1e-6))
        return magnitude
    

4. 损失聚合方式错误

多GPU环境下直接用tf.reduce_mean计算损失,未考虑全局批量大小,会导致设备间聚合异常产生NaN。

解决方法:

使用tf.nn.compute_average_loss正确聚合多GPU损失:

def call(self, y_true, y_pred):
    # 原逻辑计算各项损失
    wavefrom_loss = tf.reduce_mean(tf.abs(y_true - y_pred))
    sc_loss, mag_loss = self.mrstft_loss(y_true, y_pred)
    
    # 用全局批量大小计算平均损失
    total_loss = wavefrom_loss + sc_loss + mag_loss
    return tf.nn.compute_average_loss(total_loss, global_batch_size=self.BATCH_SIZE)

额外调试建议

  • 在stft函数中添加数值检查,定位NaN出现的具体环节:
    def stft(x, fft_size, hop_size, win_length, window):
        # 计算x_stft后
        real = tf.math.real(x_stft)
        imag = tf.math.imag(x_stft)
        tf.debugging.check_numerics(real, "Real part has NaN/Inf")
        tf.debugging.check_numerics(imag, "Imag part has NaN/Inf")
        # 其余逻辑不变
    
  • 打印多GPU环境下每个replica的输入数据形状和部分数值,确认设备间数据一致性。

内容的提问来源于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.30 16:14:52