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

