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

PyTorch结合Wandb实现验证集各类别Top1准确率绘图

分类别Top1准确率Wandb可视化实现方案

你的现有验证逻辑已经逐类计算了对应Top1准确率,只需要小幅调整代码即可完成分指标上报和可视化,具体操作如下:


1. 修改验证循环逻辑,收集全部分类别准确率

你当前的validate函数存在两个逻辑问题:一是没有保存每个类别的单独准确率,二是最终返回值是最后一个验证集的Top1准确率,而非计算好的整体平均。修改后的验证函数如下:

def validate(val_loaders, model, criterion, args):
    overall_top1 = 0
    # 新增字典存储每个类别的Top1/Top5准确率
    class_acc = {}
    for nuisance, val_loader in val_loaders:
        batch_time = AverageMeter('Time', ':6.3f', Summary.NONE)
        losses = AverageMeter('Loss', ':.4e', Summary.NONE)
        top1 = AverageMeter('Acc@1', ':6.2f', Summary.AVERAGE)
        top5 = AverageMeter('Acc@5', ':6.2f', Summary.AVERAGE)
        progress = ProgressMeter(
            len(val_loader),
            [batch_time, losses, top1, top5],
            prefix=f'Test {nuisance}: ')

        # 切换到评估模式
        model.eval()

        with torch.no_grad():
            end = time.time()
            for i, (images, target) in enumerate(val_loader):
                if args.gpu is not None:
                    images = images.cuda(args.gpu, non_blocking=True)
                if torch.cuda.is_available():
                    target = target.cuda(args.gpu, non_blocking=True)

                # 计算模型输出与损失
                output = model(images)
                loss = criterion(output, target)

                # 计算准确率并记录
                acc1, acc5 = accuracy(output, target, topk=(1, 5))
                losses.update(loss.item(), images.size(0))
                top1.update(acc1[0], images.size(0))
                top5.update(acc5[0], images.size(0))

                # 记录耗时
                batch_time.update(time.time() - end)
                end = time.time()

                if i % args.print_freq == 0:
                    progress.display(i)

            progress.display_summary()
        
        # 保存当前类别的准确率
        class_acc[nuisance] = {
            "top1": top1.avg.item(),
            "top5": top5.avg.item()
        }
        overall_top1 += top1.avg
    overall_top1 /= len(val_loaders)
    # 返回整体准确率 + 分品类准确率字典
    return overall_top1.item(), class_acc

2. 上报指标到Wandb,自动生成可视化图表

在训练循环中调用验证函数后,把整体准确率和分品类准确率一起传入wandb.log即可:

  • 如果给分品类指标加统一前缀,Wandb会自动把同前缀的指标合并到同一张分组折线图里,方便跨类别对比
  • 也可以根据需要自定义柱状图等其他类型的统计图表

示例上报代码:

# 训练循环中调用验证
overall_acc, class_acc = validate(val_loaders, model, criterion, args)

# 构造要上报的指标字典
log_dict = {
    "val/top1/overall": overall_acc,
    "epoch": current_epoch # 传入当前epoch数,作为图表的x轴
}
# 逐类别加入指标
for cls_name, acc in class_acc.items():
    log_dict[f"val/top1/{cls_name}"] = acc["top1"]
    log_dict[f"val/top5/{cls_name}"] = acc["top5"]

# 一次性上报所有指标
wandb.log(log_dict)

如果需要单独展示每个epoch下各类别准确率的柱状对比图,可以额外加代码生成柱状图上报:

# 生成各类别Top1准确率柱状图
bar_chart = wandb.plot.bar(
    wandb.Table(data=[[k, v["top1"]] for k,v in class_acc.items()], columns=["category", "top1_acc"]),
    "category",
    "top1_acc",
    title="Per-category Top1 Accuracy"
)
log_dict["val/top1/per_category_bar"] = bar_chart
wandb.log(log_dict)

上报完成后可以直接在Wandb面板里调整图表样式、分组规则,不需要额外修改代码。


附原有验证加载器定义,逻辑本身不需要调整:

val_nuisances = ['shape', 'pose', 'texture', 'context', 'weather']
val_loaders = []
for nuisance in val_nuisances:
    val_loaders.append((nuisance, torch.utils.data.DataLoader(
        datasets.ImageFolder(os.path.join(valdir, nuisance), transforms.Compose([
            transforms.Resize(256),
            transforms.CenterCrop(224),
            transforms.ToTensor(),
            normalize,
        ])),
        batch_size=args.batch_size, shuffle=False,
        num_workers=args.workers, pin_memory=True,
    )))

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.02 09:18:27