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

Keras自定义无监督损失函数报错:No gradients provided for any variable

解决Keras自定义无监督损失函数的梯度追踪问题

核心问题分析

你的代码触发ValueError: No gradients provided for any variable的原因在于几个破坏梯度传播链的操作:

  • 手动切断梯度链:使用K.variable(log_prob)将张量列表转换为独立变量,这会彻底断开与之前y_pred - mu等运算节点的关联,导致梯度无法回溯到模型参数。
  • 张量转numpy破坏梯度:循环中mu.astype(np.float32)将张量转为numpy数组(若mu_est是张量),或直接使用numpy常数(若mu_est是numpy数组),都会切断梯度传播路径。
  • 静态形状获取隐患:n_samples, _ = y_pred.shape获取静态形状,在动态batch场景下会失效,还可能引发维度不兼容问题。

修正后的代码

import tensorflow as tf
from tensorflow.keras import backend as K
import numpy as np

def custom_loss(mu_est, pi_k_est, sigma2, d):
    def loss(y_true, y_pred):
        # 动态重塑,兼容任意batch大小
        y_pred = tf.reshape(y_pred, [-1, d])
        n_samples = tf.shape(y_pred)[0]
        
        log_prob_list = []
        # 遍历mu_est的索引而非直接枚举元素,避免张量转numpy
        for k in range(tf.shape(mu_est)[0]):
            mu = tf.cast(mu_est[k], tf.float32)
            y = y_pred - mu
            log_prob = K.sum(K.square(y), axis=1)
            log_prob_list.append(log_prob)
        
        # 用tf.stack堆叠张量列表,完整保留梯度链
        log_prob = tf.stack(log_prob_list, axis=1)
        log_prob = K.cast(log_prob, tf.float64)
        
        # 统一所有运算的数值类型,避免隐性转换问题
        pi_k_est = tf.cast(pi_k_est, tf.float64)
        sigma2 = tf.cast(sigma2, tf.float64)
        d_float = tf.cast(d, tf.float64)
        
        weighted_log_prob = K.log(pi_k_est) - 0.5 * (d_float * K.log(2 * np.pi * sigma2) + log_prob / sigma2)
        temp = K.logsumexp(weighted_log_prob, axis=1)
        return - K.sum(temp) / tf.cast(n_samples, tf.float64)
    return loss

关键修改说明

  1. 用tf.stack()替代K.variable():tf.stack会将张量列表堆叠为一个新张量,完全保留原计算路径的梯度信息,不会切断传播链。
  2. 全程使用张量操作:所有类型转换用tf.cast,避免将张量转为numpy数组;若mu_est、pi_k_est是需要优化的参数,必须定义为tf.Variable类型,而非numpy数组。
  3. 动态形状获取:用tf.shape(y_pred)[0]动态获取样本数量,兼容训练时的动态batch大小。
  4. 统一数值类型:将所有参与运算的张量统一为float64,避免类型不匹配导致的隐性转换,同时保证梯度计算精度。

额外注意事项

  • 无监督任务中y_true仅为占位符,训练时可传入与y_pred形状匹配的任意张量(如全零张量)。
  • 若mu_est、pi_k_est是模型的可训练参数,需确保它们被正确注册到模型的可训练变量集合中,否则梯度仍无法追踪到这些参数。

内容的提问来源于stack exchange,提问作者James Yu

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 04:17:14