求教:如何在PyTorch代码中集成Sklearn的Precision、Recall、F1 Score指标
集成Precision、Recall、F1 Score的实现方案
手动计算方式(适配你的现有代码)
要在训练循环中加入精确率、召回率和F1值,需在每个epoch里统计真正例(TP)、假正例(FP)、**假负例(FN)**三个核心指标,再通过公式推导最终结果。
步骤1:初始化统计变量
每个epoch开始时,在原有n_correct和n_total基础上,新增三个统计变量:
true_positives = 0 false_positives = 0 false_negatives = 0
步骤2:在batch循环中更新统计量
获取predicted和labels后,保留原准确率统计逻辑,补充TP、FP、FN的计算:
_, predicted = torch.max(outputs, 1) # 原准确率统计逻辑 n_correct += (predicted == labels).sum().item() n_total += labels.shape[0] # 二分类场景下的TP/FP/FN计算(标签为0/1) true_positives += ((predicted == 1) & (labels == 1)).sum().item() false_positives += ((predicted == 1) & (labels == 0)).sum().item() false_negatives += ((predicted == 0) & (labels == 1)).sum().item()
步骤3:计算Precision、Recall、F1 Score
epoch结束后,基于统计值计算指标(处理分母为0的情况避免报错):
accuracy = 100 * n_correct / n_total precision = true_positives / (true_positives + false_positives) if (true_positives + false_positives) > 0 else 0 recall = true_positives / (true_positives + false_negatives) if (true_positives + false_negatives) > 0 else 0 f1_score = 2 * (precision * recall) / (precision + recall) if (precision + recall) > 0 else 0
修改后的完整代码
for epoch in range(num_epochs): # 初始化所有统计变量 n_correct = 0 n_total = 0 true_positives = 0 false_positives = 0 false_negatives = 0 epoch_loss = 0 # 统计整个epoch的平均损失,替代单个batch损失 for i, (words, labels) in enumerate(train_loader): words = words.to(device) labels = labels.to(dtype=torch.long).to(device) # Forward pass outputs = model(words) loss = criterion(outputs, labels) epoch_loss += loss.item() # Backward and optimize optimizer.zero_grad() loss.backward() optimizer.step() # 计算预测结果 _, predicted = torch.max(outputs, 1) # 更新准确率统计 n_correct += (predicted == labels).sum().item() n_total += labels.shape[0] # 更新TP、FP、FN统计(二分类场景) true_positives += ((predicted == 1) & (labels == 1)).sum().item() false_positives += ((predicted == 1) & (labels == 0)).sum().item() false_negatives += ((predicted == 0) & (labels == 1)).sum().item() # 计算epoch级别的指标 accuracy = 100 * n_correct / n_total precision = true_positives / (true_positives + false_positives) if (true_positives + false_positives) > 0 else 0 recall = true_positives / (true_positives + false_negatives) if (true_positives + false_negatives) > 0 else 0 f1_score = 2 * (precision * recall) / (precision + recall) if (precision + recall) > 0 else 0 avg_loss = epoch_loss / len(train_loader) # 保存指标到列表(需提前初始化train_precision、train_recall、train_f1列表) train_losses.append(avg_loss) train_epochs.append(epoch) train_acc.append(accuracy) train_precision.append(precision) train_recall.append(recall) train_f1.append(f1_score) # 打印日志 if (epoch+1) % 10 == 0: print(f'Epoch [{epoch+1}/{num_epochs}], Avg Loss: {avg_loss:.2f}, Acc: {accuracy:.2f}, Precision: {precision:.2f}, Recall: {recall:.2f}, F1: {f1_score:.2f}')
多分类任务适配
如果是多分类任务(标签为0,1,2,...,N-1),可按类别统计后计算宏平均指标:
# 初始化每个类别的统计字典 num_classes = 3 # 替换为你的类别数 class_stats = {cls: {'tp':0, 'fp':0, 'fn':0} for cls in range(num_classes)} # 在batch循环中更新每个类别的统计 for cls in range(num_classes): class_stats[cls]['tp'] += ((predicted == cls) & (labels == cls)).sum().item() class_stats[cls]['fp'] += ((predicted == cls) & (labels != cls)).sum().item() class_stats[cls]['fn'] += ((predicted != cls) & (labels == cls)).sum().item() # 计算宏平均指标 macro_precision = 0 macro_recall = 0 macro_f1 = 0 for cls in class_stats: tp = class_stats[cls]['tp'] fp = class_stats[cls]['fp'] fn = class_stats[cls]['fn'] p = tp/(tp+fp) if (tp+fp)>0 else 0 r = tp/(tp+fn) if (tp+fn)>0 else 0 f = 2*(p*r)/(p+r) if (p+r)>0 else 0 macro_precision += p macro_recall += r macro_f1 += f macro_precision /= num_classes macro_recall /= num_classes macro_f1 /= num_classes
内容的提问来源于stack exchange,提问作者aaron
相关产品推荐
相关产品推荐

