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
相关产品推荐
相关产品推荐

