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
相关产品推荐
相关产品推荐

