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

TensorFlow中损失函数内可训练变量k1/k2/k3不更新,如何加入训练列表?

如何让自定义损失中的权重变量成为可训练变量?

你的问题核心是:自定义损失函数中用到的k1、k2、k3虽然声明为tf.Variable,但未被纳入模型的可训练变量列表,优化器无法对其计算梯度并更新。以下是两种实用解决方案:


方案一:用Keras Loss类封装(推荐)

将可训练权重和损失逻辑封装到继承自tf.keras.losses.Loss的类中,Keras会自动将类内的可训练变量注册到模型的可训练集合中。

代码实现:

import tensorflow as tf

# 假设cosine_dist和custom___error是你已定义的函数
def cosine_dist(a, b):
    # 你的余弦距离实现
    pass

def custom___error(y_true, y_pred):
    # 你的自定义误差实现
    pass

class CustomWeightedLoss(tf.keras.losses.Loss):
    def __init__(self, max_bits, name="custom_weighted_loss"):
        super().__init__(name=name)
        self.MAX_BITS = max_bits
        # 初始化可训练权重(trainable=True为默认,显式声明更清晰)
        self.k1 = tf.Variable(0.5, dtype=tf.float32, trainable=True)
        self.k2 = tf.Variable(0.5, dtype=tf.float32, trainable=True)
        self.k3 = tf.Variable(0.5, dtype=tf.float32, trainable=True)

    def call(self, y_true, y_pred):
        # 基础交叉熵损失
        base_loss = tf.reduce_mean(tf.nn.sigmoid_cross_entropy_with_logits(labels=y_true, logits=y_pred))
        # 余弦距离项
        cosine0 = cosine_dist(y_true[:, :self.MAX_BITS], y_pred[:, :self.MAX_BITS])
        # 自定义误差项
        error = custom___error(y_true, y_pred)
        # 加权总损失
        total_loss = base_loss * self.k1 + cosine0 * self.k2 + self.k3 * error
        return total_loss

# 初始化损失实例
loss_fn = CustomWeightedLoss(max_bits=MAX_BITS)
# 编译模型
model.compile(loss=loss_fn, optimizer=optimizer, metrics=['accuracy', custom_error, custom___error])

方案二:手动添加变量到模型可训练列表

如果不想改动损失函数的结构,可以手动将k1、k2、k3加入模型的trainable_variables集合,让优化器能识别到它们。

代码实现:

# 定义可训练变量(确保trainable=True)
k1 = tf.Variable(0.5, dtype=tf.float32, trainable=True)
k2 = tf.Variable(0.5, dtype=tf.float32, trainable=True)
k3 = tf.Variable(0.5, dtype=tf.float32, trainable=True)

# 手动将变量添加到模型的可训练变量列表(必须在compile之前执行)
model.trainable_variables.extend([k1, k2, k3])

# 保持原有的损失函数和编译逻辑
def custom_loss(y_true, y_pred, k1, k2, k3):
    loss = tf.reduce_mean(tf.nn.sigmoid_cross_entropy_with_logits(labels=y_true, logits=y_pred))
    cosine0 = cosine_dist(y_true[:MAX_BITS], y_pred[:MAX_BITS])
    error = custom___error(y_true, y_pred)
    loss = loss * k1 + cosine0 * k2 + k3 * error
    return loss

model.compile(loss=lambda y_true, y_pred: custom_loss(y_true, y_pred, k1, k2, k3), 
              optimizer=optimizer, 
              metrics=['accuracy', custom_error, custom___error])

注意事项:

  • 无论哪种方案,都要确保变量初始化时设置trainable=True(默认值为True,但显式声明更不易出错)
  • 手动添加变量的操作必须在model.compile()之前执行,否则优化器不会识别新添加的变量
  • 方案一更符合Keras的设计规范,便于后续维护和扩展,优先推荐

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 16:05:52