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

类别级精度评估仅输出0类精度问题排查与解决请求

问题解决:类别级指标仅输出0类结果的修复

问题现象

使用tf.keras.preprocessing.image_dataset_from_directory加载数据集(标签模式为int,损失用sparse_categorical_crossentropy)时,自定义的类别精度、召回、精确率指标仅能输出0类的有效结果,其余类别均显示为0.0000e+00。

核心原因

自定义指标函数中错误地使用了K.argmax(y_true, axis=-1)处理标签:

  • 当使用label_mode='int'和sparse_categorical_crossentropy时,输入的y_true是整数型的类别ID张量(shape为(batch_size,)),而非独热编码张量(shape为(batch_size, num_classes))。K.argmax对整数张量的处理逻辑不符合预期,导致后续掩码计算失效,非0类的指标无法正确统计。

修复步骤

  1. 修正标签处理逻辑:移除K.argmax(y_true, axis=-1),直接使用y_true作为真实类别ID。
  2. 补全缺失的指标:当前编译模型时仅添加了0、1类的召回和精确率指标,需补全2、3、4类的对应指标,确保所有类别都能输出结果。
  3. 修复代码语法问题:原代码中部分换行处未添加反斜杠,导致语法错误,需修正。

修复后的完整代码

import tensorflow as tf
from tensorflow.keras import Sequential
from tensorflow.keras.layers import Flatten, Dense
from tensorflow.keras.optimizers import Adam
from tensorflow.keras import backend as K

# 加载数据集
train_ds = tf.keras.preprocessing.image_dataset_from_directory(
    data_dir, labels='inferred', label_mode='int',
    validation_split=0.2,
    subset="training",
    seed=123,
    image_size=(180, 180),
    batch_size=batch_size)

val_ds = tf.keras.preprocessing.image_dataset_from_directory(
    data_dir, labels='inferred', label_mode='int',
    validation_split=0.2,
    subset="validation",
    seed=123,
    image_size=(180, 180),
    batch_size=batch_size)

# 构建模型
resnet_model = Sequential()
pretrained_model = tf.keras.applications.VGG16(
    include_top=False,
    input_shape=(180, 180, 3),
    pooling='avg',
    classes=5,
    weights='imagenet',
    classifier_activation='softmax'
)

for layer in pretrained_model.layers:
    layer.trainable = False

resnet_model.add(pretrained_model)
resnet_model.add(Flatten())
resnet_model.add(Dense(512, activation='relu'))
resnet_model.add(Dense(5, activation='softmax'))

# 自定义类别级指标
def single_class_accuracy(interesting_class_id):
    def acc1(y_true, y_pred):
        class_id_true = y_true  # 直接使用整数标签
        class_id_preds = K.argmax(y_pred, axis=-1)
        accuracy_mask = K.cast(K.equal(class_id_preds, interesting_class_id), 'int32')
        class_acc_tensor = K.cast(K.equal(class_id_true, class_id_preds), 'int32') * accuracy_mask
        class_acc = K.cast(K.sum(class_acc_tensor), 'float32') / K.cast(K.maximum(K.sum(accuracy_mask), 1), 'float32')
        return class_acc
    acc1.__name__ = 'acc_1_{}'.format(interesting_class_id)
    return acc1

def single_class_recall(interesting_class_id):
    def recall(y_true, y_pred):
        class_id_true = y_true  # 直接使用整数标签
        class_id_pred = K.argmax(y_pred, axis=-1)
        recall_mask = K.cast(K.equal(class_id_true, interesting_class_id), 'int32')
        class_recall_tensor = K.cast(K.equal(class_id_true, class_id_pred), 'int32') * recall_mask
        class_recall = K.cast(K.sum(class_recall_tensor), 'float32') / K.cast(K.maximum(K.sum(recall_mask), 1), 'float32')
        return class_recall
    recall.__name__ = 'recall_1_{}'.format(interesting_class_id)
    return recall

def single_class_precision(interesting_class_id):
    def prec(y_true, y_pred):
        class_id_true = y_true  # 直接使用整数标签
        class_id_pred = K.argmax(y_pred, axis=-1)
        precision_mask = K.cast(K.equal(class_id_pred, interesting_class_id), 'int32')
        class_prec_tensor = K.cast(K.equal(class_id_true, class_id_pred), 'int32') * precision_mask
        class_prec = K.cast(K.sum(class_prec_tensor), 'float32') / K.cast(K.maximum(K.sum(precision_mask), 1), 'float32')
        return class_prec
    prec.__name__ = 'prec_1_{}'.format(interesting_class_id)
    return prec

# 编译模型(补全所有类别的指标)
resnet_model.compile(
    optimizer=Adam(lr=0.01),
    loss='sparse_categorical_crossentropy',
    metrics=[
        'accuracy',
        single_class_accuracy(0), single_class_accuracy(1), single_class_accuracy(2), single_class_accuracy(3), single_class_accuracy(4),
        single_class_recall(0), single_class_recall(1), single_class_recall(2), single_class_recall(3), single_class_recall(4),
        single_class_precision(0), single_class_precision(1), single_class_precision(2), single_class_precision(3), single_class_precision(4)
    ]
)

# 训练模型
history = resnet_model.fit(train_ds, validation_data=val_ds, epochs=2)

补充说明

  • 若后续切换为独热编码标签(label_mode='categorical',损失用categorical_crossentropy),则需要恢复K.argmax(y_true, axis=-1)的使用。
  • 训练初期部分类别指标为0可能是模型尚未学习到对应类别的特征,可增加训练轮次或调整模型参数观察变化。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.20 01:15:55