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

Keras自定义损失函数:如何在每个epoch结束后跟踪各子损失值

解决自定义多损失在Epoch结束后跟踪的问题

你的问题核心在于:self.losses里的各个损失张量是和模型的输入占位符绑定的计算图节点,直接用K.eval()去求值时,没有提供对应的输入数据,所以才会触发占位符未喂值的错误。下面给你几个更优雅的解决方案,不用额外跑预测也能跟踪这些子损失:

方案一:利用Keras的指标系统直接记录子损失

你可以在自定义损失函数里,把每个子损失作为模型的指标来维护,Keras会自动帮你在每个batch和epoch结束时累计平均数值,不需要额外传递self.losses给回调。

修改后的代码示例:

def train_test(self):
    def custom_loss(y_true, y_pred):
        # 先计算各个子损失
        rmse_loss = ...  # 你的RMSE计算逻辑
        ce_loss = ...    # 你的交叉熵计算逻辑
        other_loss = ... # 你的OtherLoss计算逻辑
        total_loss = rmse_loss + ce_loss + other_loss
        
        # 关键:更新各个指标的累计值
        self.rmse_metric.update_state(rmse_loss)
        self.ce_metric.update_state(ce_loss)
        self.other_metric.update_state(other_loss)
        
        return total_loss
    
    # 初始化用于累计子损失的指标
    self.rmse_metric = tf.keras.metrics.Mean(name='RMSE')
    self.ce_metric = tf.keras.metrics.Mean(name='CrossEntropy')
    self.other_metric = tf.keras.metrics.Mean(name='OtherLoss')
    
    # 模型构建部分保持不变
    logits = keras.layers.Dense(365, activation=keras.activations.softmax)(concat)
    self.model = keras.Model(inputs=[...], outputs=logits)
    self.model.compile(optimizer=keras.optimizers.Adam(0.001), loss=custom_loss)
    
    # 自定义回调读取指标值
    class TrackLossCallback(keras.callbacks.Callback):
        def __init__(self, rmse_metric, ce_metric, other_metric):
            self.rmse_metric = rmse_metric
            self.ce_metric = ce_metric
            self.other_metric = other_metric
            
        def on_epoch_end(self, epoch, logs={}):
            # 直接获取指标的epoch平均数值
            print(f"Epoch {epoch+1} - RMSE: {self.rmse_metric.result().numpy():.4f}")
            print(f"Epoch {epoch+1} - CrossEntropy: {self.ce_metric.result().numpy():.4f}")
            print(f"Epoch {epoch+1} - OtherLoss: {self.other_metric.result().numpy():.4f}")
            # 重置指标,避免下一个epoch累计上一轮的数据
            self.rmse_metric.reset_states()
            self.ce_metric.reset_states()
            self.other_metric.reset_states()
    
    self.history = self.model.fit_generator(
        generator=self.train_data,
        steps_per_epoch=train_data_size//FLAGS.batch_size,
        epochs=5,
        callbacks=[TrackLossCallback(self.rmse_metric, self.ce_metric, self.other_metric)])

方案二:拆分损失函数,在compile时指定多个损失

如果你愿意把每个子损失封装成独立函数,Keras支持在编译时同时指定多个损失/指标,这样每个子损失会自动被跟踪并记录到history中,回调里直接从logs字典读取即可:

def train_test(self):
    # 定义独立的子损失函数
    def rmse_loss(y_true, y_pred):
        return ...  # 你的RMSE计算逻辑
    
    def ce_loss(y_true, y_pred):
        return ...  # 你的交叉熵计算逻辑
    
    def other_loss(y_true, y_pred):
        return ...  # 你的OtherLoss计算逻辑
    
    # 总损失函数
    def total_loss(y_true, y_pred):
        return rmse_loss(y_true, y_pred) + ce_loss(y_true, y_pred) + other_loss(y_true, y_pred)
    
    # 模型构建部分不变
    logits = keras.layers.Dense(365, activation=keras.activations.softmax)(concat)
    self.model = keras.Model(inputs=[...], outputs=logits)
    
    # 编译时指定总损失和需要跟踪的子损失指标
    self.model.compile(
        optimizer=keras.optimizers.Adam(0.001),
        loss=total_loss,
        metrics=[rmse_loss, ce_loss, other_loss]
    )
    
    # 回调直接读取logs中的指标值
    class TrackLossCallback(keras.callbacks.Callback):
        def on_epoch_end(self, epoch, logs={}):
            print(f"Epoch {epoch+1} - RMSE: {logs['rmse_loss']:.4f}")
            print(f"Epoch {epoch+1} - CrossEntropy: {logs['ce_loss']:.4f}")
            print(f"Epoch {epoch+1} - OtherLoss: {logs['other_loss']:.4f}")
    
    self.history = self.model.fit_generator(
        generator=self.train_data,
        steps_per_epoch=train_data_size//FLAGS.batch_size,
        epochs=5,
        callbacks=[TrackLossCallback()])

原代码报错的原因解释

你原代码里的self.losses[key]是计算图中的张量,它依赖于模型的输入占位符(比如报错里的input_3)。当你在回调里调用K.eval()时,并没有给这个张量提供对应的输入数据,TensorFlow找不到占位符的数值,所以抛出了InvalidArgumentError。上面的两个方案要么通过指标系统自动累计数值,要么让Keras自动处理损失的计算和记录,都避开了直接求值未绑定数据的张量问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.12 04:14:38