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

