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

TensorFlow v2.16中tf.keras.metrics.mean_absolute_error调用失效求助

解决TensorFlow v2.16中MAE计算报错的问题

问题分析

升级到TensorFlow v2.16后,tf.keras.metrics.mean_absolute_error这类便捷函数的调用逻辑发生了变化——这类函数原本是为训练流程中累加式评估设计的,新版本对输入处理方式更严格,直接传入张量会触发兼容错误,尤其在TF Keras独立封装为keras._tf_keras的场景下更容易出现此类问题。

解决方案

以下两种方案均可解决问题,按需选择:

方案1:使用TensorFlow底层运算直接计算(推荐)

绕过Keras metrics封装,用基础张量运算计算指标,彻底避免版本兼容问题,逻辑更直观:

def evaluate_preds(y_true, y_pred):
    # 统一转换为float32类型
    y_true = tf.cast(y_true, dtype=tf.float32)
    y_pred = tf.cast(y_pred, dtype=tf.float32)

    # 直接用张量运算计算各项指标
    mae = tf.math.reduce_mean(tf.abs(y_true - y_pred))
    mse = tf.math.reduce_mean(tf.square(y_true - y_pred))
    rmse = tf.sqrt(mse)
    # MAPE计算加入分母防零处理
    mape = tf.math.reduce_mean(tf.abs((y_true - y_pred) / tf.where(y_true == 0, tf.ones_like(y_true), y_true))) * 100
    mase = mean_absolute_scaled_error(y_true, y_pred)
    
    return {"mae": mae.numpy(),
            "mse": mse.numpy(),
            "rmse": rmse.numpy(),
            "mape": mape.numpy(),
            "mase": mase.numpy()}

方案2:遵循Keras metrics标准使用流程

若坚持使用Keras metrics,需按照「实例化类→更新状态→获取结果」的标准流程调用,而非直接使用便捷函数:

def evaluate_preds(y_true, y_pred):
    y_true = tf.cast(y_true, dtype=tf.float32)
    y_pred = tf.cast(y_pred, dtype=tf.float32)

    # 实例化metrics类并完成计算流程
    mae_metric = tf.keras.metrics.MeanAbsoluteError()
    mae_metric.update_state(y_true, y_pred)
    mae = mae_metric.result()

    mse_metric = tf.keras.metrics.MeanSquaredError()
    mse_metric.update_state(y_true, y_pred)
    mse = mse_metric.result()
    
    rmse = tf.sqrt(mse)

    mape_metric = tf.keras.metrics.MeanAbsolutePercentageError()
    mape_metric.update_state(y_true, y_pred)
    mape = mape_metric.result()

    mase = mean_absolute_scaled_error(y_true, y_pred)
    
    return {"mae": mae.numpy(),
            "mse": mse.numpy(),
            "rmse": rmse.numpy(),
            "mape": mape.numpy(),
            "mase": mase.numpy()}

额外检查项

  • 确认y_true和y_pred的形状完全匹配,包括batch维度和特征维度,新版本对形状不匹配的容忍度更低;
  • Apple Silicon环境下,需确保tensorflow-metal版本与tensorflow-macos版本严格对应,可通过conda list核对版本号。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 18:53:10