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

TensorFlow验证指标不显示:与训练指标重名问题排查

问题:自定义test_step后无法显示val_my_metric

使用model.compile(metrics=[MyMetric])配置模型后,每个epoch结束仅能看到loss、val_loss和my_metric,但看不到val_my_metric。

调试TensorFlow代码后发现:trainer.py中的fit()逻辑正常,但在CallbackList.on_epoch_end中,通过python_utils.pythonify_logs(logs)扁平化字典时,compile_metrics和val_compile_metrics的键被存储为同一名称,导致验证阶段的指标被训练阶段的指标覆盖。

注:未修改Trainer逻辑,仅按官方文档子类化了train和test_step,代码如下:

def update_metrics(self, loss, f1, f2):
    # From Cholet's Keras docs:
    for metric in self.metrics:
        if metric.name == "loss":
            metric.update_state(loss)
        else:
            metric.update_state(f1, f2)

def test_step(self, data):
    """
    
    Args:
        data (tf.Tensor): Input batch of shape [B, 2, H, W, C]
        
    Returns:
        dict: Dictionary containing the loss
    """
    f1, f2, p1, p2, z1, z2 = self(data, training=False)
    loss = self.compute_loss(p1, p2, z1, z2)
    
    self.update_metrics(loss, f1, f2)

    # Return metrics
    return {m.name: m.result() for m in self.metrics}

原因分析

代码存在两个核心问题:

  1. test_step返回的指标名称与训练阶段完全一致,导致Keras合并日志时,验证阶段的指标被训练阶段的同名指标覆盖。
  2. 使用self.metrics更新指标时,训练和验证阶段共用同一批指标实例,没有区分训练/验证专用的指标集合,导致Keras无法自动生成val_前缀的指标键。

解决方案

方案1:手动给验证指标添加val_前缀

修改test_step的返回逻辑,给除loss外的指标手动加上val_前缀,避免键冲突:

def test_step(self, data):
    f1, f2, p1, p2, z1, z2 = self(data, training=False)
    loss = self.compute_loss(p1, p2, z1, z2)
    
    self.update_metrics(loss, f1, f2)

    # 生成带val_前缀的验证指标字典
    val_metrics = {}
    for m in self.metrics:
        if m.name == "loss":
            val_metrics["loss"] = m.result()
        else:
            val_metrics[f"val_{m.name}"] = m.result()
    return val_metrics

方案2:使用Keras内置的compiled_metrics区分训练/验证指标

Keras编译时会自动为训练和验证阶段创建独立的指标实例,存储在self.compiled_metrics中。修改代码使用这个集合来更新和返回指标:

def update_metrics(self, loss, f1, f2):
    # 更新损失(Keras会自动处理val_loss)
    self.compiled_loss.update_state(loss)
    # 更新编译时指定的验证指标
    for metric in self.compiled_metrics:
        metric.update_state(f1, f2)

def test_step(self, data):
    f1, f2, p1, p2, z1, z2 = self(data, training=False)
    loss = self.compute_loss(p1, p2, z1, z2)
    
    self.update_metrics(loss, f1, f2)

    # 返回损失和编译后的验证指标(Keras会自动添加val_前缀)
    return {"loss": loss} | {m.name: m.result() for m in self.compiled_metrics}

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 23:29:58