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

如何创建Keras自定义指标:分类任务中邻类预测可判定为正确

Keras自定义指标:允许相邻类别判定为正确预测

嘿,这个需求在有序分类场景里特别实用——比如情感评分(1-5星)、疾病严重程度分级这类任务,相邻类别其实差异不大,算成"准正确"完全合理。我来给你一步步实现这个自定义指标,代码清晰还能直接用!

核心思路

我们的指标要满足两个判定条件之一就算预测正确:

  • 预测类别和真实类别完全一致(差值为0)
  • 预测类别是真实类别的直接相邻类别(差值的绝对值为1)

在Keras里,自定义指标需要继承tf.keras.metrics.Metric类,并重写三个关键方法:__init__(初始化累计变量)、update_state(每一批次更新统计值)、result(计算最终指标值)。

完整代码实现

import tensorflow as tf

class AdjacentAccuracy(tf.keras.metrics.Metric):
    def __init__(self, name="adjacent_accuracy", **kwargs):
        super().__init__(name=name, **kwargs)
        # 累计正确预测的样本数
        self.correct = self.add_weight(name="correct", initializer="zeros")
        # 累计处理的总样本数
        self.total = self.add_weight(name="total", initializer="zeros")

    def update_state(self, y_true, y_pred, sample_weight=None):
        # 处理真实标签:如果是one-hot编码,先转成整数类别;如果已经是整数则直接用
        if y_true.shape.rank > 1:
            y_true = tf.argmax(y_true, axis=1)
        # 处理预测结果:从概率分布转成整数类别
        y_pred = tf.argmax(y_pred, axis=1)
        
        # 转换为整数类型,避免浮点运算的误差
        y_true = tf.cast(y_true, tf.int32)
        y_pred = tf.cast(y_pred, tf.int32)
        
        # 计算真实标签和预测标签的差值绝对值
        diff = tf.abs(y_true - y_pred)
        # 判定正确:差值为0(完全匹配)或1(相邻类别)
        is_correct = tf.math.logical_or(diff == 0, diff == 1)
        # 转换成浮点型便于累加
        is_correct = tf.cast(is_correct, tf.float32)
        
        # 处理样本权重(可选,如果有需要的话)
        if sample_weight is not None:
            sample_weight = tf.cast(sample_weight, tf.float32)
            is_correct = tf.multiply(is_correct, sample_weight)
        
        # 累计正确数和总样本数
        self.correct.assign_add(tf.reduce_sum(is_correct))
        self.total.assign_add(tf.cast(tf.size(y_true), tf.float32))

    def result(self):
        # 计算准确率:正确数/总样本数
        return tf.divide(self.correct, self.total, name="result")

    def reset_state(self):
        # 重置指标状态(比如每个epoch开始时)
        self.correct.assign(0.0)
        self.total.assign(0.0)

如何使用这个指标

在模型编译的时候,直接把这个自定义指标传进去就行:

# 假设你已经定义好了你的分类模型
model = tf.keras.Sequential([...])

# 编译模型时使用自定义指标
# 如果你的真实标签是稀疏整数(比如直接是0,1,2...),用sparse_categorical_crossentropy
model.compile(
    optimizer="adam",
    loss="sparse_categorical_crossentropy",
    metrics=[AdjacentAccuracy()]
)

# 如果真实标签是one-hot编码,用categorical_crossentropy
# model.compile(
#     optimizer="adam",
#     loss="categorical_crossentropy",
#     metrics=[AdjacentAccuracy()]
# )

注意事项

  • 这个指标默认假设你的类别是连续的整数(比如0,1,2,3,4),如果你的类别是离散不连续的(比如1,3,5),那需要先把类别映射成连续整数再使用。
  • 代码里已经兼容了稀疏标签(整数形式)和one-hot标签两种情况,不用额外修改。
  • 如果需要更灵活的判定(比如允许间隔2个类别也算正确),只需要把diff == 1改成diff <= 2就行,非常容易调整。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 03:32:44