自定义损失函数训练音频模型遇无语音文件时梯度缺失问题
我在用自定义损失函数微调一个预训练的音频模型。输入带语音的文件时,梯度计算正常;但输入无语音的文件时,就会报No gradient provided for any variable错误。能返回损失值,但所有变量的梯度都是None,而且和损失值大小无关。已经确认模型变量确实参与了输入文件的降噪变换,所以不是变量没参与计算的问题。怀疑问题出在共振峰误差、谐波误差不满足计算条件时的处理逻辑上,相关代码如下:
# Update the code snippet to calculate conditions for harmonic and formant errors is_voiced_h = tf.cond( tf.logical_and( tf.logical_and(i < num_frames - 1, tf.equal(speechs[i], 1)), tf.logical_and(tf.equal(speechs[i], 1), tf.equal(speechs[i + 1], 1) )), lambda: True, # Condition met lambda: False) # Out of bounds or condition not met is_voiced_f = tf.cond(tf.logical_and( tf.logical_and(i < num_frames - 1, tf.equal(speechs[i], 1)), tf.logical_or(tf.equal(speechs[i + 1], 1), tf.equal(speechs[i + 2], 1)) ), lambda: True, # Condition met lambda: False) # Out of bounds or condition not met is_unvoiced_f = tf.cond(tf.logical_and(tf.logical_and(i< num_frames - 2, tf.equal(speechs[i], 3)), tf.logical_or(tf.equal(speechs[i + 1], 3), tf.equal(speechs[i + 2], 3))), lambda: True, # Condition met lambda: False) # Out of bounds or condition not met if i < (num_frames-2)//2: not_flagged_f = tf.cond(tf.logical_and(i < num_frames - 2, # Check out of bounds tf.logical_or(tf.equal(flags[i], False), tf.logical_or(tf.equal(flags[i + 1], False), tf.equal(flags[i + 2], False)))), lambda: True, # Condition met lambda: False) # out of bounds or condition not met else: not_flagged_f = True pitch = process_audio_tf(ref_frame_h, sample_rate=sample_rate) pitch = pitch[0] epsilon = 1e-5 harmonic_error = tf.cond(is_voiced_h, lambda: process_frame(ref_frame_h, model_frame_h, pitch, sample_rate), lambda: tf.zeros((1,), dtype=tf.float32)) + epsilon formant_error = tf.cond(tf.logical_and(tf.logical_or(is_voiced_f, is_unvoiced_f), not_flagged_f), lambda: calculate_formant_error(ref_frame_f, model_frame_f, sample_rate, formants[i:i + 1,...][0, 0], formants[i:i + 1,...][0, 1], n_fft), lambda: tf.zeros((8,), dtype=tf.float32)) + epsilon #tf.print('harmonic_error from body', harmonic_error) #tf.print('formant_error from body', formant_error) harmonic_errors = harmonic_errors.write(i, harmonic_error) formant_errors = formant_errors.write(i, formant_error) return i + 1, harmonic_errors, formant_errors _, harmonic_errors, formant_errors = tf.while_loop(lambda i, *_: i < num_frames, body, [0, harmonic_errors, formant_errors]) harmonic_errors = harmonic_errors.stack() formant_errors = formant_errors.stack()
核心问题:梯度链断裂
当输入无语音文件时,is_voiced_h、is_voiced_f、is_unvoiced_f等条件全为False,此时harmonic_error和formant_error都进入返回tf.zeros的分支。这个分支的输出是纯常数,完全不依赖模型的输出张量(model_frame_h/model_frame_f),导致模型变量到损失值的梯度链直接断裂,最终所有变量梯度为None。针对性修复
保留所有分支的梯度依赖
不要直接返回常数,而是返回一个与模型输出相关但数值极小的项,确保梯度链不中断。例如:- 谐波误差的else分支修改:
这样既保证输出形状和lambda: tf.reduce_mean(model_frame_h, keepdims=True) * 1e-10process_frame一致,又保留了和模型变量的依赖关系,梯度可以正常回传,同时极小的系数不会影响损失值的有效性。 - 共振峰误差的else分支修改:
确保输出形状与lambda: tf.tile(tf.reduce_mean(model_frame_f, keepdims=True)*1e-10, [8])calculate_formant_error返回的(8,)一致。
- 谐波误差的else分支修改:
简化条件判断,避免不必要的控制流
原代码用tf.cond返回布尔值是冗余操作,直接用张量逻辑运算生成布尔张量即可,减少控制流带来的梯度追踪复杂度:is_voiced_h = tf.logical_and( tf.logical_and(i < num_frames - 1, tf.equal(speechs[i], 1)), tf.equal(speechs[i + 1], 1) )同理修改
is_voiced_f、is_unvoiced_f和not_flagged_f的判断逻辑。确保tf.while_loop的梯度兼容性
显式设置back_prop=True(默认开启,但显式指定更稳妥),同时保证循环内所有张量的形状固定,避免动态形状导致梯度无法追踪:_, harmonic_errors, formant_errors = tf.while_loop( lambda i, *_: i < num_frames, body, [0, harmonic_errors, formant_errors], back_prop=True )
内容的提问来源于stack exchange,提问作者paras Joshi

