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

自定义损失函数训练音频模型遇无语音文件时梯度缺失问题

问题描述

我在用自定义损失函数微调一个预训练的音频模型。输入带语音的文件时,梯度计算正常;但输入无语音的文件时,就会报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。

  • 针对性修复

    1. 保留所有分支的梯度依赖
      不要直接返回常数,而是返回一个与模型输出相关但数值极小的项,确保梯度链不中断。例如:

      • 谐波误差的else分支修改:
        lambda: tf.reduce_mean(model_frame_h, keepdims=True) * 1e-10
        
        这样既保证输出形状和process_frame一致,又保留了和模型变量的依赖关系,梯度可以正常回传,同时极小的系数不会影响损失值的有效性。
      • 共振峰误差的else分支修改:
        lambda: tf.tile(tf.reduce_mean(model_frame_f, keepdims=True)*1e-10, [8])
        
        确保输出形状与calculate_formant_error返回的(8,)一致。
    2. 简化条件判断,避免不必要的控制流
      原代码用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的判断逻辑。

    3. 确保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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.20 08:57:15