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

如何为CSRNet添加Precision、Recall、F1指标计算与可视化功能

CSRNet添加Precision、Recall、F1指标计算与可视化方案

CSRNet输出为人群密度图,Precision/Recall/F1属于分类指标,需要先设置密度阈值,将预测和真实密度图转为二值图(密度≥阈值判定为「存在人头/正样本」,否则为负样本),再计算混淆矩阵得到对应指标,修改步骤如下:


第一步 修改train.py代码

1. 新增全局变量

在train.py开头的全局变量定义区域添加以下代码:

# 新增验证指标存储列表,用于后续可视化
Val_Precision_list = []
Val_Recall_list = []
Val_F1_list = []
# 密度阈值,可根据数据集实际情况调整,常规取值范围0.01~0.1
DENSITY_THRESHOLD = 0.05

2. 替换validate函数,添加指标计算逻辑

def validate(val_list, model, criterion):
    print ('begin test')
    test_loader = torch.utils.data.DataLoader(
    dataset.listDataset(val_list,
                   shuffle=False,
                   transform=transforms.Compose([
                       transforms.ToTensor(),transforms.Normalize(mean=[0.485, 0.456, 0.406],
                                     std=[0.229, 0.224, 0.225]),
                   ]),  train=False),
    batch_size=args.batch_size)    
    
    model.eval()
      
    mae = 0
    # 混淆矩阵参数初始化
    total_tp = 0
    total_fp = 0
    total_fn = 0

    for i,(img, target) in enumerate(test_loader):
        img = img.cuda()
        img = Variable(img)
        with torch.no_grad():
            output = model(img)
        
        # 原有MAE计算逻辑保留
        mae += abs(output.data.sum()-target.sum().type(torch.FloatTensor).cuda())

        # 新增二值化&混淆矩阵计算
        pred_np = output.detach().cpu().numpy().squeeze()
        target_np = target.numpy().squeeze()

        # 转为二值图
        pred_binary = (pred_np >= DENSITY_THRESHOLD).astype(np.int32)
        target_binary = (target_np >= DENSITY_THRESHOLD).astype(np.int32)

        # 统计TP、FP、FN
        tp = np.sum(np.logical_and(pred_binary == 1, target_binary == 1))
        fp = np.sum(np.logical_and(pred_binary == 1, target_binary == 0))
        fn = np.sum(np.logical_and(pred_binary == 0, target_binary == 1))

        total_tp += tp
        total_fp += fp
        total_fn += fn
        
    mae = mae/len(test_loader)
    # 计算三大指标,加极小值1e-8防止除零错误
    precision = total_tp / (total_tp + total_fp + 1e-8)
    recall = total_tp / (total_tp + total_fn + 1e-8)
    f1 = 2 * precision * recall / (precision + recall + 1e-8)

    # 存入列表用于后续可视化
    Val_Precision_list.append(precision)
    Val_Recall_list.append(recall)
    Val_F1_list.append(f1)

    print(' * MAE {mae:.3f} | Precision {pre:.3f} | Recall {rec:.3f} | F1 {f1:.3f}'
              .format(mae=mae, pre=precision, rec=recall, f1=f1))

    return mae

3. 替换main函数末尾的可视化代码

x1 = range(0, args.epochs)
    y1 = Train_loss_list
    y2 = Val_Precision_list
    y3 = Val_Recall_list
    y4 = Val_F1_list

    plt.subplots_adjust(left=None, bottom=None, right=None, top=None,
                       wspace=None, hspace=0.8)
    # 第一幅:训练损失曲线
    plt.subplot(4, 1, 1)
    plt.plot(x1, y1, label="Train Loss", color='blue')
    plt.title('Train Loss vs. Epochs')
    plt.xlabel('Epochs')
    plt.ylabel('Loss')
    plt.legend()

    # 第二幅:Precision曲线
    plt.subplot(4, 1, 2)
    plt.plot(x1, y2, label="Val Precision", color='green')
    plt.title('Validation Precision vs. Epochs')
    plt.xlabel('Epochs')
    plt.ylabel('Precision')
    plt.legend()

    # 第三幅:Recall曲线
    plt.subplot(4, 1, 3)
    plt.plot(x1, y3, label="Val Recall", color='orange')
    plt.title('Validation Recall vs. Epochs')
    plt.xlabel('Epochs')
    plt.ylabel('Recall')
    plt.legend()

    # 第四幅:F1曲线
    plt.subplot(4, 1, 4)
    plt.plot(x1, y4, label="Val F1", color='red')
    plt.title('Validation F1 vs. Epochs')
    plt.xlabel('Epochs')
    plt.ylabel('F1')
    plt.legend()

    plt.savefig("C:/Users/Gigabyte pc/Desktop/COUNTING/CSRNet-pytorch-master/PartA2_all_metrics_curve.jpg")

第二步 可选修改val.py(单图测试输出指标)

如果需要单张图片测试时也输出对应指标,可以在val.py读取完真实密度图后添加以下代码:

# 单图二值化指标计算
pred_density = output.detach().cpu().numpy().squeeze()
gt_density = temp_1
threshold = 0.05
pred_bin = (pred_density >= threshold).astype(int)
gt_bin = (gt_density >= threshold).astype(int)
tp = np.logical_and(pred_bin ==1, gt_bin==1).sum()
fp = np.logical_and(pred_bin ==1, gt_bin==0).sum()
fn = np.logical_and(pred_bin ==0, gt_bin==1).sum()
p = tp/(tp+fp+1e-8)
r = tp/(tp+fn+1e-8)
f1 = 2*p*r/(p+r+1e-8)
print(f"单图指标:Precision {p:.3f} | Recall {r:.3f} | F1 {f1:.3f}")

注意事项

  • 密度阈值DENSITY_THRESHOLD需要根据你所用的数据集标注情况调整,建议先抽样几张图统计真实密度图的人头区域密度值,再选定最合适的阈值
  • 若需要更精准的人头级指标而非像素级,可以先对密度图做局部极大值检测提取人头坐标,再和真实人头坐标做IoU匹配后计算Precision/Recall/F1,效果会比像素级更符合实际检测需求

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 07:24:03