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

Keras自定义容忍±2偏差的Accuracy指标报错求解

温度分类预测自定义±2容差准确率实现

场景说明

  • 构建面向特定日期、时刻的温度概率分布预测模型,温度取值共设46个档位,所有温度四舍五入取整后作为独立分类类别处理
  • 训练阶段使用3年时长的时间序列温度数据集,提取hour_of_day、dayofweek、month三类时间周期特征
  • 评估规则调整:不再要求预测值与真实值完全相等,二者差值在±2范围内即判定为预测正确,例如真实值为5、预测值为7时,差值为2,计为预测匹配

原有代码问题

最初编写的自定义指标代码如下:

def accuracy(y_true, y_pred):
  y_true_class = K.argmax(y_true, axis=-1)
  y_pred_class = K.argmax(y_pred, axis=-1)
  if K.any(K.abs(y_true_class-y_pred_class) <= 2) :
    matches = K.cast(K.equal(y_true_class, y_pred_class), 'int32')
    accuracy = K.mean(matches)

  return accuracy

模型训练时触发报错,核心提示为:

ValueError: 'accuracy' must also be initialized in the else branch

代码存在三处核心错误:

  • 用Python原生if分支处理Keras张量运算:Keras构建计算图阶段无法通过Python原生条件判断张量的布尔值,所有逻辑必须通过张量运算实现
  • 变量accuracy仅在if分支内定义,不满足计算图要求的所有分支变量初始化规则
  • 匹配逻辑不符合需求:即使进入if分支,内部仍通过K.equal判断类别完全相等,完全没有实现±2容差的判定规则

正确实现代码

不需要任何条件分支,直接通过张量运算完成全样本的容差判定即可:

import tensorflow.keras.backend as K

def tolerance_accuracy(y_true, y_pred):
    # 将one-hot格式的标签、预测结果转为类别索引
    y_true_cls = K.argmax(y_true, axis=-1)
    y_pred_cls = K.argmax(y_pred, axis=-1)
    # 计算每个样本预测类别与真实类别的绝对差值
    cls_diff = K.abs(y_true_cls - y_pred_cls)
    # 差值≤2记为预测正确,转换为0/1浮点张量
    correct_mask = K.cast(cls_diff <= 2, 'float32')
    # 计算正确样本占比即为准确率
    return K.mean(correct_mask)

使用方式

模型编译阶段将该自定义函数传入metrics参数即可:

model.compile(
    optimizer='adam',
    loss='categorical_crossentropy',
    metrics=[tolerance_accuracy]
)

该实现完全基于张量运算,无原生Python条件分支,不会触发计算图相关报错,且逻辑完全匹配±2容差的判定要求:差值为0(完全匹配)、1、2的样本均会被计为预测正确,差值≥3的样本计为错误。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 18:39:28