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

如何在Keras中每个epoch显示各类别准确率?

嘿,这个需求超实用!其实完全不用修改内置的callback,咱们自己写个自定义回调(或者在训练循环里加段小逻辑)就能轻松搞定。下面分两种最常用的深度学习框架,给你具体讲怎么实现:

Keras/TensorFlow 实现方式:自定义Callback

Keras的回调机制很灵活,咱们可以继承keras.callbacks.Callback写一个专属的回调类,在每个epoch结束时自动计算并打印每个类别的准确率。

具体代码示例

import tensorflow as tf
from tensorflow.keras import layers, models
import numpy as np

# 加载并预处理MNIST数据
(x_train, y_train), (x_test, y_test) = tf.keras.datasets.mnist.load_data()
x_train = x_train / 255.0
x_test = x_test / 255.0

# 构建一个简单的全连接模型
model = models.Sequential([
    layers.Flatten(input_shape=(28, 28)),
    layers.Dense(128, activation='relu'),
    layers.Dense(10, activation='softmax')
])

# 注意这里不用指定metrics='accuracy',因为咱们要自己算类别准确率
model.compile(optimizer='adam', loss='sparse_categorical_crossentropy')

# 自定义回调类
class ClasswiseAccuracyCallback(tf.keras.callbacks.Callback):
    def __init__(self, validation_data):
        super().__init__()
        self.x_val, self.y_val = validation_data  # 传入要评估的数据集
    
    def on_epoch_end(self, epoch, logs=None):
        # 获取模型对验证集的预测结果(转成类别索引)
        y_pred = np.argmax(self.model.predict(self.x_val, verbose=0), axis=1)
        class_acc_results = []
        
        # 遍历0-9每个数字类别
        for cls in range(10):
            # 筛选出当前类别的所有样本
            cls_samples_mask = (self.y_val == cls)
            # 计算该类别的正确预测数和总样本数
            correct_count = np.sum(y_pred[cls_samples_mask] == self.y_val[cls_samples_mask])
            total_count = np.sum(cls_samples_mask)
            # 计算准确率(避免除以0的情况)
            cls_acc = correct_count / total_count if total_count > 0 else 0.0
            class_acc_results.append(f"类别 {cls}: {cls_acc:.4f}")
        
        # 打印结果
        print(f"\nEpoch {epoch+1} 各类别准确率:")
        print("\n".join(class_acc_results))

# 初始化回调,传入验证集(如果要算训练集准确率,换成x_train和y_train就行)
class_acc_callback = ClasswiseAccuracyCallback(validation_data=(x_test, y_test))

# 开始训练,把自定义回调加入callbacks列表
model.fit(x_train, y_train, epochs=5, batch_size=32, 
          validation_data=(x_test, y_test),
          callbacks=[class_acc_callback])

逻辑说明

这个回调会在每个epoch结束后,自动用你指定的数据集(这里是验证集)计算每个类别的准确率。核心就是针对每个类别单独统计正确数和总样本数,再计算比值。


PyTorch 实现方式:训练循环中添加统计逻辑

PyTorch没有内置的Callback系统,但咱们可以在每个epoch的训练结束后,手动调用一个统计函数来计算类别准确率,同样很简单。

具体代码示例

import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader
from torchvision import datasets, transforms

# 数据预处理
transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.1307,), (0.3081,))  # MNIST的均值和标准差
])

# 加载MNIST数据集
train_dataset = datasets.MNIST(root='./data', train=True, download=True, transform=transform)
test_dataset = datasets.MNIST(root='./data', train=False, download=True, transform=transform)

train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)
test_loader = DataLoader(test_dataset, batch_size=32, shuffle=False)

# 定义简单的全连接模型
class MNISTNet(nn.Module):
    def __init__(self):
        super(MNISTNet, self).__init__()
        self.flatten = nn.Flatten()
        self.fc1 = nn.Linear(28*28, 128)
        self.fc2 = nn.Linear(128, 10)
    
    def forward(self, x):
        x = self.flatten(x)
        x = torch.relu(self.fc1(x))
        x = self.fc2(x)
        return x

# 初始化模型、损失函数和优化器
model = MNISTNet()
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)

# 定义计算类别准确率的函数
def get_classwise_accuracy(model, dataloader, device):
    model.eval()  # 切换到评估模式
    class_correct = [0] * 10  # 每个类别的正确预测数
    class_total = [0] * 10    # 每个类别的总样本数
    
    with torch.no_grad():  # 关闭梯度计算,节省内存
        for images, labels in dataloader:
            images, labels = images.to(device), labels.to(device)
            outputs = model(images)
            _, predicted = torch.max(outputs, 1)  # 获取预测的类别索引
            
            # 逐个样本统计
            for label, pred in zip(labels, predicted):
                class_total[label] += 1
                class_correct[label] += (pred == label).item()
    
    # 生成每个类别的准确率结果
    class_acc_list = []
    for cls in range(10):
        acc = class_correct[cls] / class_total[cls] if class_total[cls] > 0 else 0.0
        class_acc_list.append(f"类别 {cls}: {acc:.4f}")
    return class_acc_list

# 选择设备(GPU优先)
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model.to(device)

# 开始训练循环
epochs = 5
for epoch in range(epochs):
    model.train()  # 切换到训练模式
    running_loss = 0.0
    
    # 训练批次循环
    for inputs, labels in train_loader:
        inputs, labels = inputs.to(device), labels.to(device)
        
        optimizer.zero_grad()
        outputs = model(inputs)
        loss = criterion(outputs, labels)
        loss.backward()
        optimizer.step()
        
        running_loss += loss.item()
    
    # 每个epoch结束后计算并打印结果
    print(f"\nEpoch {epoch+1} 训练损失: {running_loss/len(train_loader):.4f}")
    # 计算验证集的类别准确率
    test_class_accs = get_classwise_accuracy(model, test_loader, device)
    print("验证集各类别准确率:")
    print("\n".join(test_class_accs))

逻辑说明

在每个epoch的训练完成后,调用get_classwise_accuracy函数,遍历数据集统计每个类别的正确数和总样本数,最后计算并打印每个类别的准确率。如果想查看训练集的类别准确率,只需要把test_loader换成train_loader即可。


总的来说,不管用哪个框架,核心逻辑都是在每个epoch结束后,针对每个类别单独统计正确预测数和该类别的总样本数,再计算二者的比值。完全不需要修改内置的callback,自己写点小代码就能实现需求~

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.29 08:40:51