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

求教:如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.20 03:01:21