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

Keras 3.4中多输出模型损失权重随Epoch动态调整难题

在Keras 3.4中动态调整多输出模型的损失权重

Keras 3.x确实收紧了loss_weights的类型限制,不能再用张量变量直接传入。要实现按Epoch动态调整损失权重,最靠谱的方式是自定义总损失函数,把权重用tf.Variable单独维护,在回调里更新变量值,让损失函数自动使用最新权重计算总损失。

具体实现步骤

1. 定义可更新的权重变量

用tf.Variable保存各损失分量的权重,这部分变量会在训练过程中被动态修改:

import tensorflow as tf
from tensorflow.keras import callbacks, losses, metrics, optimizers

# 初始化权重变量
loss_weight1 = tf.Variable(1.0, dtype=tf.float32)
loss_weight2 = tf.Variable(0.0, dtype=tf.float32)

2. 自定义总损失函数

这个函数接收模型的多个输出和对应标签,分别计算每个输出的损失后,用当前权重加权求和:

def custom_total_loss(y_true, y_pred):
    # 适配字典形式的输出与标签
    loss1 = losses.BinaryCrossentropy()(y_true['loss1'], y_pred['loss1'])
    loss2 = losses.BinaryCrossentropy()(y_true['loss2'], y_pred['loss2'])
    # 加权求和得到总损失
    return loss_weight1 * loss1 + loss_weight2 * loss2

如果你的模型输出是列表形式(而非字典),可以改用索引访问,比如y_true[0]、y_pred[0]。

3. 修改回调类

直接在回调里用assign方法更新tf.Variable的值:

class WeightAdjuster(callbacks.Callback):
    def __init__(self, w1, w2):
        self.w1 = w1
        self.w2 = w2
    
    def on_epoch_end(self, epoch, logs=None):
        # 按你的逻辑更新权重,示例为w1加1、w2减1
        self.w1.assign(self.w1 + 1)
        self.w2.assign(self.w2 - 1)
        # 可选:打印当前权重用于调试
        print(f"\nEpoch {epoch+1}: 更新后权重 - loss1: {self.w1.numpy()}, loss2: {self.w2.numpy()}")

4. 编译并训练模型

编译时不需要指定loss_weights,把自定义损失函数传入loss参数,同时传入回调:

# 假设你已经定义好多输出模型nn
optimizer = optimizers.Adam()
metrics_dict = {'loss1': 'accuracy', 'loss2': 'accuracy'}

# 编译模型
nn.compile(
    loss=custom_total_loss,
    optimizer=optimizer,
    metrics=metrics_dict
)

# 初始化回调
weight_callback = WeightAdjuster(loss_weight1, loss_weight2)

# 开始训练
nn.fit(
    training_generator,
    epochs=epochs,
    validation_data=validation_generator,
    steps_per_epoch=None,
    callbacks=[weight_callback]
)

为什么旧方法失效?

Keras 3.x对loss_weights做了严格的类型校验,仅接受Python float/int。编译时这些值会被转换成静态张量,后续修改传入的字典或外部变量不会影响模型内部的计算逻辑——因为模型已经缓存了初始权重值,不会再读取外部字典的变化。而自定义损失函数直接依赖tf.Variable,每次计算损失时都会取变量的当前值,因此能实现动态调整。

注意事项

  • 如果权重更新逻辑需要结合验证集指标,可以在on_epoch_end中通过logs参数获取当前epoch的训练/验证指标,再调整权重。
  • 确保权重变量的数据类型与损失计算的 dtype 一致,避免类型错误。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.18 02:50:17