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

Keras多输出自定义模型的损失与梯度计算问题咨询

自定义多输出适配模型代码

class CustomModel(keras.Model):
    def __init__(self, *args, **kwargs):
        super().__init__(*args, **kwargs)
        self.loss_tracker = keras.metrics.Mean(name="loss")
        self.mae_metric = keras.metrics.MeanAbsoluteError(name="mae")
        self.loss_functions=None

    def compile(self, *args, loss=None, **kwargs):
        self.loss_functions = loss  # Store the provided loss functions
        super().compile(*args, **kwargs)

    def train_step(self, data):
        x, y = data

        with tf.GradientTape() as tape:
            # Forward pass
            y_pred = self(x, training=True)  

            # Compute losses for each output
            loss2 = keras.losses.kullback_leibler_divergence(y_pred[2], y_pred[0])
            loss3 = keras.losses.kullback_leibler_divergence(y_pred[2], y_pred[1])

            # Aggregate total loss
            total_loss =loss2+loss3

        # Compute gradients
        trainable_vars = self.trainable_variables
        gradients = tape.gradient([loss2,loss3], trainable_vars)

        # Update weights
        self.optimizer.apply_gradients(zip(gradients, trainable_vars))

        # Update metrics
        self.loss_tracker.update_state(total_loss)
        self.mae_metric.update_state(y[0], y_pred[0])  # Update metrics for output1
        return {"loss": self.loss_tracker.result(), "mae": self.mae_metric.result()}

    @property
    def metrics(self):
        return [self.loss_tracker, self.mae_metric]

模型定义

inputs = keras.Input(shape=(32,))
output1 = keras.layers.Dense(1)(inputs)
hidden1 = keras.layers.Dense(10)(inputs)
output2 = keras.layers.Dense(1)(hidden1)
hidden2 = keras.layers.Dense(10)(inputs)
output3 = keras.layers.Dense(1)(hidden2)

model = CustomModel(inputs, [output1,output2, output3])

模型编译代码

理想情况下需在train_step中覆盖各输出的损失函数,以下是模型编译代码:

model.compile(optimizer="adam",loss=["mse","kullback_leibler_divergence",None])
x = np.random.random((1000, 32))
y = np.random.random((1000, 1))
model.fit(x,y)

技术问题解答

1. 为output3传入None是否能实现无损失约束?

不能。因为你完全重写了train_step方法,编译时传入的loss参数只是被存在self.loss_functions变量里,但当前train_step的损失计算是硬编码的loss2和loss3,完全没调用这个存储的损失列表。所以不管你给output3传什么,只要train_step里没涉及output3的损失计算,它就不会有损失约束;但依赖编译时的None来实现,在当前自定义逻辑里是无效的。

2. 编译时指定的output1的mse损失会被train_step里的自定义损失覆盖吗?

不止是覆盖,是完全被忽略了。你的train_step里根本没调用self.loss_functions中的mse损失,而是自己计算了KL散度相关的损失,编译时指定的mse完全没参与到损失计算和梯度更新流程中。相当于编译时的loss设置在这个自定义模型里是摆设,除非你在train_step里主动调用self.loss_functions里对应的损失函数。

3. tape.gradient([loss2,loss3], trainable_vars)的工作机制是什么?是否loss2仅作用于output1、loss3仅作用于output2?

  • 工作机制:当传入损失列表给tape.gradient时,TensorFlow会分别计算每个损失相对于所有可训练变量的梯度,最终返回一个和可训练变量数量一致的列表,每个元素是对应变量的所有损失梯度的总和。简单说就是先算loss2对每个变量的梯度,再算loss3对每个变量的梯度,把两者相加得到该变量的最终梯度。
  • 不是。loss2是y_pred[2]和y_pred[0]的KL散度,这个损失会关联所有影响这两个输出的可训练变量——包括output1的Dense层参数、output3的Dense层参数,以及它们上游的共享层参数;同理loss3关联的是output2和output3的相关参数。所以两个损失都会影响output3的参数,并非各自只作用于output1和output2。

另外,你尝试传入[loss2,loss3,None]报错是因为tape.gradient不接受None作为损失元素,它要求所有传入的损失都是有效的张量,None会直接触发错误。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 17:17:50