训练跑起来只是第一步,跑完什么都看不懂才是最折磨人的。我真正被 W&B 这个工具打动,是在一个周一下午:手头 12 组不同超参的实验同时跑,我一边盯着终端翻日志,一边手动往 Excel 里抄 loss 值,抄到第 7 组的时候忽然发现,有一组实验的 batch_size 我忘记记了。那一刻我意识到,训练过程可视化真正要解决的不是“画一条曲线”,而是“让每次实验都有一条完整、可追溯、可对比的生命周期记录”。
这篇是“Pytorch 可视化”系列的第五篇。前面几篇聊过怎么用 matplotlib 画 Loss 曲线、怎么把中间特征图可视化,这一篇我决定集中写一个工具——WandB(Weights & Biases),它是我目前最推荐的 Pytorch 训练过程可视化方案,没有之一。它能把训练中的 loss、准确率、学习率、梯度分布、显存占用、验证集指标全部统一记录到一个可视化管理页面里,自动按超参数分组对比,还能做超参搜索、模型文件版本管理、团队共享报告。这篇文章我会从安装登录、四个核心 API、改造完整训练循环,一路写到 Sweeps 自动调参、Artifacts 版本管理、离线同步,以及我在实际使用中踩过的那些报错坑。
不管你是刚搭好 Pytorch 环境、正准备认真调模型的初学者,还是已经用 TensorBoard 很久、觉得实验对比和团队分享不方便的进阶玩家,这篇都适用。下面所有内容都是我自己用下来的经验,你可以直接当作一份能抄作业的攻略。
1. 为什么训练可视化我最后选了 WandB
1.1 TensorBoard 够用,但不够好用
TensorBoard 我前后用了大概两年。作为 Pytorch 官方出品,它确实轻量、免安装,用 SummaryWriter 写几个 scalars 就能看曲线,单机训练完全没有问题。但当我开始同时跑一组对比实验时,它的短板就特别明显了:
第一是实验对比太割裂。TensorBoard 确实支持同时勾选多个 run 查看,但日志文件一多,页面加载变慢,run 的名字一长就分不清谁是谁。而且它把所有本地事件文件都压在你自己电脑上,换一台机器训练就看不到历史,除非你再把日志拷过去。
第二是超参数没法跟指标自动绑定。TensorBoard 的 hparams 插件虽然存在,但我使用下来的体验是配置麻烦,而且每次都要手动指定哪些超参要记录、要展示,一段时间不维护,新实验就忘了加。可我记得最重要的恰恰是当时用了什么优化器、什么 batch_size、多少步衰减的学习率。
第三是团队协作几乎为零。同事想看你某个实验的曲线,要么你把截图发过去,要么给他开端口,要么把 events 文件打包发给他。在团队里做模型实验,这效率实在太低了。
说白了,TensorBoard 是一个“单机版画线工具”,不是实验管理平台。当你只有一两个实验、自己调试用,它完全够;但当你把训练当成一场持续进行的“实验管理”时,工具就会变成瓶颈。
1.2 WandB 到底解决了什么
WandB 的定位和 TensorBoard 不同。它把“记录训练指标”这件事直接做成了服务:你在训练脚本里调用 log,数据会传到 wandb 项目页面;打开浏览器,你能在一个页面里看到这个项目下所有 run 的曲线,支持按超参数值给曲线着色,点开任意一个 run 就能看到完整配置、输出日志、训练指标、梯度直方图、甚至模型权重分布。
对我来说最实用的功能有三个:
- 实验自动归档:每次 wandb.init 都是一个新 run,训练时用了什么参数、什么环境、什么代码版本,都会自动归档到页面,不用自己再维护一份实验记录表。
- 超参与指标联动:曲线图旁边可以直接看 config 表,某个 run 用的 lr、batch_size、优化器参数一目了然;对比 0.001 和 0.0001 两组实验时,直接勾选两个 run,颜色和名称都区分好了。
- 链路完整:训练日志、指标曲线、模型权重文件、数据集版本、生成报告全在同一个工作区里,复现实验时不用在微信里翻聊天记录找模型文件放在哪。
还有一个很现实的原因:WandB 有免费额度,个人和初创团队用起来完全够。数据默认存云端,你在任何一台能开浏览器的机器上都能看。它也支持离线模式,后面我会单独讲。
1.3 选型对比:TensorBoard、Mlflow、WandB 与手写日志
我最初自建过一套“轮子”——直接在训练脚本里把 loss 打印到日志,再写个 matplotlib 把日志解析画图。听起来可控,但过了一周连自己都不想看,因为要写大量解析代码,而且不同的实验脚本日志格式还经常不统一。
选 W&B 前我也大致对比了一圈:
| 方案 | 指标可视化 | 超参记录 | 实验对比 | 协作分享 | 上手成本 | 备注 |
|---|---|---|---|---|---|---|
| TensorBoard | 好 | 一般 | 一般 | 几乎为零 | 低 | 单机调试首选 |
| Mlflow | 一般 | 好 | 好 | 一般 | 中 | 更偏向模型注册与工程化 |
| WandB | 很好 | 很好 | 很好 | 很好 | 低 | 一体化实验管理,本系列第五篇主角 |
| 手写日志+matplotlib | 自控 | 自己写 | 自己写 | 自己写 | 高 | 适合完全内网隔离的极端场景 |
结论是:单机调试、小团队做算法验证、教学演示、入门 Pytorch,WandB 的综合体验最均衡。而且它和 Pytorch 的集成非常自然,训练循环里每 step 调一次 log 就行,原来的代码结构基本不用动。这也是为什么我在这个系列里,把训练过程可视化这一篇给了 WandB 而不是继续写 matplotlib。
2. 三分钟跑通 WandB 基础配置
2.1 安装与登录,先让第一个 run 跑起来
安装很简单,直接用 pip:
pip install wandb装完之后在终端执行:
wandb login首次登录会让你去官网注册账号,拿到一个 API key,粘贴回来就完成授权。key 会写入用户目录下的~/.netrc,之后在当前机器上运行时不需要重复登录。
如果你是在脚本里自动化跑,更推荐用环境变量,避免交互式粘贴:
import os os.environ["WANDB_API_KEY"] = "你的-api-key" os.environ["WANDB_PROJECT"] = "pytorch-visualization-demo" os.environ["WANDB_ENTITY"] = "你的用户名或团队名"把 key 写环境变量还有个好处:多人共用的训练服务器上,每位同事可以有自己的 key,不会因为.netrc被其他人覆盖而互相干扰。注意,绝对不要把 API key 提交到公开代码仓库,这是最容易踩的安全坑。
提示:如果运行环境没有外网,或者你暂时不想把数据传到云端,可以先执行
wandb offline切换成本地离线模式,记录的数据会写在本机缓存目录,之后再统一同步。这个用法在后面第 4.3 节详细讲。
2.2 记住这四个 API 就够入门
WandB 的接口很多,但 90% 的日常使用只需要下面四个:
wandb.init(),每次训练开始前调用,创建一条 run 记录。project 参数指定项目名,config 参数把超参打包传进去,后续所有指标都会挂在这个 run 名下。
wandb.config,init 时传入的超参集合。它既是一个可读的对象,也可以当作字典操作。训练脚本里直接读wandb.config.lr,整个实验的超参数就能自动归档到页面上。
wandb.log(),最核心的接口。每次调用传入一个字典,比如wandb.log({"loss": 0.32, "acc": 0.91}),页面就会把这条记录追加到对应 run 的曲线里。支持一次传多个 key,间隔一定 step 调用一次即可,不用每个 batch 都调。
wandb.watch(),用来监控模型参数和梯度。传入 model 之后,它会自动在设定的频率下记录权重直方图、梯度直方图,方便你观察有没有梯度消失、爆炸,或者某些层长时间不更新。
最小示例大概长这样:
import wandb wandb.init(project="my-demo", config={"lr": 1e-3, "batch_size": 64}) # 训练循环里调用 for step in range(100): loss = compute_loss(step) wandb.log({"loss": loss}) wandb.watch(model, log="all", log_freq=100) wandb.finish()2.3 理解 run 的生命周期:id、tags、notes 与断点续跑
每次wandb.init()都会生成一个唯一的 run,即使脚本崩溃或者手动退出,这个 run 依然会保留在项目页面里。很多新手困惑“我重新跑了一下,怎么又多了一个 run”——这是 WandB 的默认设计:每次启动都是一次新的训练记录。
如果想在同一个 run 上续跑,需要手动指定 id:
run = wandb.init( project="my-demo", id="your-previous-run-id", resume="must", # 或 "allow" )resume="must"表示必须续跑旧 run,找不到就直接报错,适合训练中断后恢复的场景。resume="allow"则表示能找到就续,找不到就新建。
除了 id,我强烈建议用tags、notes、group三个字段维护实验组织:
wandb.init( project="resnet-cifar", name="resnet18-bs128-lr1e-3", group="resnet18-batchsize-compare", job_type="train", tags=["resnet18", "cifar10", "baseline"], notes="第一次正式跑resnet18,作为后续实验基线", )group特别适合把同维度对比实验归到一组;tags方便筛选;notes可以写备注,比在团队群里发一句“这个实验是 xxx 跑的好模型”靠谱多了。
2.4 用 config 管理超参数,别再把参数写在变量名里
我见过很多项目,超参数散落在脚本各处,跑完后根本分不清“这个 run 用的 lr 是 1e-3 还是 1e-4”。WandB 的 config 就是来解决这个问题的:
config = wandb.config config.lr = 1e-3 config.batch_size = 64 config.epochs = 30之后在训练代码里直接读取config.lr,页面会自动把这一组配置归档到 run 详情页。你回看任何一个 run,点开就知道当时用了什么配置,完全不用靠记忆。
注意:不要在训练过程中随意修改 config 的值。尤其不要把动态学习率写进 config,比如
config.lr = scheduler.get_last_lr()。学习率这类逐 step 变化的值,应该用wandb.log({"lr": current_lr})当作普通指标记录,config 里只保留初始静态超参。
3. 实战:改造一个 Pytorch 训练循环
3.1 完整示例:给 MNIST 的 CNN 训练接上 WandB
这里我给出一个可以直接复制的完整示例。任务很简单,MNIST 手写数字分类,用一个两层卷积的 CNN,训练 5 个 epoch。重点是展示 WandB 在真实训练循环里应该插在哪些位置。
import torch from torch import nn, optim from torch.utils.data import DataLoader from torchvision import datasets, transforms import wandb device = "cuda" if torch.cuda.is_available() else "cpu" # 1. 初始化 run,传入超参 run = wandb.init( project="pytorch-visualization-demo", name="cnn-mnist-baseline", config={ "lr": 1e-3, "batch_size": 64, "epochs": 5, "optimizer": "adam", }, tags=["mnist", "baseline"], notes="用于验证wandb接入流程的第一个run", ) config = wandb.config # 2. 数据准备 transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)), ]) train_set = datasets.MNIST(root="./data", train=True, download=True, transform=transform) val_set = datasets.MNIST(root="./data", train=False, download=True, transform=transform) train_loader = DataLoader(train_set, batch_size=config.batch_size, shuffle=True, num_workers=2) val_loader = DataLoader(val_set, batch_size=256, shuffle=False, num_workers=2) # 3. 模型与优化器 class CNN(nn.Module): def __init__(self): super().__init__() self.features = nn.Sequential( nn.Conv2d(1, 32, 3, padding=1), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(32, 64, 3, padding=1), nn.ReLU(), nn.MaxPool2d(2), ) self.classifier = nn.Sequential( nn.Flatten(), nn.Linear(64 * 7 * 7, 256), nn.ReLU(), nn.Linear(256, 10), ) def forward(self, x): return self.classifier(self.features(x)) model = CNN().to(device) criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=config.lr) # 4. 监控模型参数与梯度 wandb.watch(model, criterion=criterion, log="all", log_freq=50) # 5. 验证函数 def evaluate(model, loader): model.eval() correct, total = 0, 0 with torch.no_grad(): for x, y in loader: x, y = x.to(device), y.to(device) out = model(x) pred = out.argmax(dim=1) total += y.size(0) correct += (pred == y).sum().item() return correct / total # 6. 训练循环 for epoch in range(config.epochs): model.train() train_loss, train_correct, train_total = 0.0, 0, 0 for step, (x, y) in enumerate(train_loader): x, y = x.to(device), y.to(device) optimizer.zero_grad() out = model(x) loss = criterion(out, y) loss.backward() optimizer.step() train_loss += loss.item() * x.size(0) pred = out.argmax(dim=1) train_correct += (pred == y).sum().item() train_total += y.size(0) # 每 100 个 batch 记录一次训练侧指标 if step % 100 == 0: wandb.log({ "batch_loss": loss.item(), "batch_acc": (pred == y).float().mean().item(), "epoch": epoch, }) # epoch 结束后跑一次验证集,并记录验证准确率 val_acc = evaluate(model, val_loader) avg_train_loss = train_loss / train_total train_acc = train_correct / train_total wandb.log({ "epoch": epoch, "train_loss": avg_train_loss, "train_acc": train_acc, "val_acc": val_acc, }) print(f"epoch {epoch} | train_loss {avg_train_loss:.4f} | " f"train_acc {train_acc:.4f} | val_acc {val_acc:.4f}") # 7. 保存模型并记录为 artifact torch.save(model.state_dict(), "cnn_mnist_baseline.pt") artifact = wandb.Artifact("cnn-mnist", type="model") artifact.add_file("cnn_mnist_baseline.pt") run.log_artifact(artifact) run.finish()这段代码结构很常规,你只需要把模型结构换成自己的网络,数据加载换成自己的数据,WandB 部分完全不需要动。
3.2 拆解每个记录点的意义
我把刚才示例里的记录点逐个说明一下,避免你照抄完不知道为什么要写这些。
第一个记录点是训练循环里的:
wandb.log({ "batch_loss": loss.item(), "batch_acc": (pred == y).float().mean().item(), "epoch": epoch, })它记录的是当前 batch 的 loss 和准确率。为什么取名 batch_loss 而不是 loss?因为一个 epoch 结束后我还会记录一个 train_loss,两者语义不同:一个是实时噪声大的 batch 值,一个是经过全量平均的 epoch 值。如果你只用一个 key,后面再记录另一个值,曲线会被覆盖或混杂,区分开更清晰。
第二个记录点在 epoch 结束之后:
wandb.log({ "epoch": epoch, "train_loss": avg_train_loss, "train_acc": train_acc, "val_acc": val_acc, })这里三个指标之间的横轴必须统一。我的习惯是显式传epoch作为时间轴的一部分。如果不传step,WandB 默认按照 log 的调用次数递增 x 轴。问题在于训练循环里每 100 个 batch 调一次 log,epoch 结束又调一次,这两类曲线的横轴语义会不一致:batch_loss 是 0, 100, 200... 的步数,val_acc 是 0, 1, 2... 的轮数。统一在 log 里带上epoch这个 key,页面会自动把 val 相关曲线放到 epoch 维度上,方便和 train 指标对看。
第三个是wandb.watch(model, criterion=criterion, log="all", log_freq=50)。watch 会把模型参数的权重分布、梯度分布记录成直方图。训练跑到一半如果你发现 loss 不掉,去看直方图通常能直接发现问题:比如某一层梯度全变零,或者权重数值异常大。但注意不要log="all"的同时把log_freq设太密,否则序列化开销会拖慢训练,后面第 3.4 节会展开。
3.3 记录 Loss、准确率之外的指标
真实项目很少只看 loss 和 acc。下面三个自定义指标是最常见也最容易被问到的:
学习率曲线,尤其是使用 StepLR、CosineAnnealingLR 这类动态调度器时,学习率变化直接决定 loss 走势。每次 scheduler 更新后顺手记录:
wandb.log({"lr": scheduler.get_last_lr()[0]})梯度范数,判断梯度是否消失或爆炸。PyTorch 里可以先对梯度做 clip,再把 clip 前的总范数记录:
total_norm = 0.0 for p in model.parameters(): if p.grad is not None: param_norm = p.grad.data.norm(2) total_norm += param_norm.item() ** 2 total_norm = total_norm ** 0.5 wandb.log({"grad_norm": total_norm})F1、IoU 这类验证集指标,在完成一个 epoch 的验证后计算。比如分割任务可以这样:
val_iou = compute_iou(pred_masks, target_masks) wandb.log({"val_iou": val_iou})还有一个实用技巧:WandB 不只支持数值曲线,还支持把本地的图片、音频、表格直接记录到页面。例如:
wandb.log({"val_examples": [wandb.Image(img), wandb.Image(mask)]})训练中随手存一批验证样本的预测结果,比截图传群再写“这是第 20 个 epoch 的效果”高效得多。
3.4 日志频率与性能:别让记录拖慢训练
WandB 默认是边训练边把数据异步上传到服务端,不会因为网络问题阻塞训练,但高频调用wandb.log和wandb.watch依然会产生 CPU 序列化和 IO 开销。我实测下来的经验是:
| 记录内容 | 推荐频率 | 原因 |
|---|---|---|
| batch 级 loss / acc | 每 50~200 个 batch | 画曲线密度足够,开销很小 |
| epoch 级指标 | 每个 epoch 一次 | 天然低频,无压力 |
| 学习率 | 每次 scheduler 更新后 | 本来就是要看变化趋势 |
| 梯度直方图 / 权重直方图 | 每 100~500 个 batch | 高频直方图序列化开销大 |
| 验证集图片 | 每个 epoch 存 8~16 张 | 每 epoch 一次即可,避免占空间 |
另一个性能大坑是:不要在验证循环内部调wandb.log。如果你写for x, y in val_loader: wandb.log({"loss_per_batch": ...}),验证集几百个 batch 就会产生几百条记录,页面曲线又乱又卡,还白白拖慢验证速度。验证集指标应该等整个验证循环跑完,汇总成单个数值再记录。
4. 进阶能力:Sweeps、Artifacts 与协作
4.1 Sweeps:一键自动搜索超参数
当你不再手动一组组跑实验,而是想让 WandB 自动帮你搜索超参数时,用 Sweeps。它内置了随机搜索、网格搜索和贝叶斯搜索三种算法。我目前使用最多的是贝叶斯搜索,因为它会根据历史试点的结果来指导下一组参数,同样的试点数下效果往往更好。
用法分两步。第一步写一个 sweep 配置文件,比如sweep.yaml:
program: train.py method: bayes metric: name: val_acc goal: maximize parameters: lr: distribution: log_uniform min: 0.0001 max: 0.01 batch_size: values: [32, 64, 128] hidden_dim: distribution: int_uniform min: 128 max: 512第二步在终端启动:
wandb sweep sweep.yaml启动后终端会输出一个 sweep_id,然后启动一个或多个 agent:
wandb agent <sweep_id>每个 agent 会不断从参数空间里采样生成新的 run,自动调用train.py。你可以在同一台机器上起多个 agent 并行,也可以在多台机器上各起一个,最终搜索结果都会汇总到同一个 sweep 页面里。
这里有一个细节:学习率这类跨度极大、影响非线性的超参,我建议用log_uniform而不是uniform。因为 0.0001 和 0.001 之间的差距与 0.1 和 0.2 之间的差距,对训练的影响完全不同,在 log 空间采样更合理。这也是我在实际搜索中踩过坑后换过来的经验。
4.2 Artifacts:模型和数据集进入版本管理
Artifacts 是 WandB 给我惊喜最多的一块。它把“模型文件”和“数据集”也纳入版本管理。训练完一个模型,不再只是存一个.pt文件,还会自动记录是哪个 run 产生的、当时的超参是什么、父 artifact 是哪个数据版本。
记录模型 artifact 的代码其实就在刚才第 3.1 节里:
artifact = wandb.Artifact("cnn-mnist", type="model") artifact.add_file("cnn_mnist_baseline.pt") run.log_artifact(artifact)重新运行训练时,可以读取某个历史版本:
artifact = run.use_artifact("cnn-mnist:latest") model_path = artifact.file()数据集同样可以管理。比如一份经过预处理的训练数据,第一次处理完写入数据集 artifact,之后每次实验都从同一个数据版本读取。这样如果团队里有人改过数据处理逻辑,你能明确知道某个 run 用的到底是哪一版数据,避免“这个结果怎么复现不出来”的争论。
注意:不要在训练循环里把 checkpoints 每轮都 log 一次。只保存最优模型或最后几个 epoch 的模型就够了,否则 artifact 体积会迅速膨胀,页面也会变得难以维护。
4.3 离线环境下的 WandB:先本地记录再同步
很多训练环境并没有外网,或者你对数据上传有顾虑。这时候 WandB 的本地模式就很重要。
在终端或脚本里设置:
wandb offline或者用环境变量:
os.environ["WANDB_MODE"] = "offline"离线模式下,wandb.log、wandb.watch照常工作,所有记录会写入本地的.wandb缓存目录,不会做网络上传。等这台机器恢复网络之后,再执行:
wandb sync同步命令会把缓存里所有 run 的记录上传到云端项目。我自己的习惯是:在实验室几台没有公网的机器上,全部开启离线模式,跑完实验后用一台有网络的机器统一wandb sync,这样既保留了 WandB 的完整功能,又不影响训练环境的网络限制。
如果你是公司内网环境,不想把数据传到外部服务,WandB 也有私有化部署方案,但门槛相对高,我建议小团队先在云端免费额度上跑通,再评估是否需要私有化。
4.4 用 Report 把实验结果变成共享页面
训练跑完了,曲线都在 WandB 里了,但总不能把网页链接甩给同事让他自己一个个 run 点开看。WandB 的 Report 功能就是把一堆曲线、表格、文字说明组合成一份可分享的试验报告。
操作上很简单:在项目页面的 Runs 列表里勾选几组要展示的 run,然后点击“Create report”,把损失曲线、验证集准确率、准确率对比表这些面板拖进报告里,再补两段文字说明结论,生成一个固定链接。之后每次训练完把新 run 加进报告,链接不用变,团队里所有人都能看到最新的对比结果。
这个功能在我写周报、和算法组同步进展时帮了很大忙。以前要把截图一张张贴上去,现在直接丢一个链接,同事自己点开看曲线,还能交互式缩放。
5. 常见问题与排查技巧实录
5.1 高频报错速查表
我整理了一张速查表,都是我在训练中实际遇到并排查过的问题:
| 现象 | 最常见原因 | 解决办法 |
|---|---|---|
| 初始化时卡在网络请求 | 当前环境无法访问 WandB 公网服务 | 切换 offline 模式,或设置 WANDB_MODE=offline |
| 提示 API key 未找到 | 没有执行 wandb login,或环境变量未设置 | 执行wandb login,或设置 WANDB_API_KEY |
| 曲线画了一部分就停住 | 训练崩溃,run 被标记为 failed | 修复脚本后使用 resume="allow" 续跑 |
| 曲线横轴乱了,train 和 val 对不上 | log 中 step 语义不统一 | 显式在 log 里带上 epoch 或 global_step |
| 多卡训练出现重复 run | 每个进程都调用了 wandb.init | 只在 rank 0 进程 init 和 log |
| watch 之后显存明显变大 | log="all" 且频率过高 | 改为 log="gradients",log_freq 调大 |
| 页面数据不更新,本地控制台也无输出 | 网络上传线程异常 | 检查缓存目录,必要时重启脚本并 resume |
5.2 初始化卡住与登录失效
最让新手头疼的就是wandb.init()卡住不动。通常是当前环境访问不了 WandB 云端服务。排查思路很简单:先看终端有没有出现Network error之类的字样,如果没有,再尝试临时切到 offline 模式:
wandb offline python train.py如果能正常跑且生成了本地缓存,说明确实卡在网络同步。你可以在有外网的机器上用wandb sync同步结果,不影响训练数据本身。
登录失效也比较常见,尤其多人共用一台服务器时。.netrc被覆盖之后,再次运行就会出现 key 找不到的报错。我习惯直接用环境变量:
export WANDB_API_KEY="你的key"写入~/.bashrc或者项目启动脚本里,比依赖.netrc要稳定得多。
5.3 曲线没显示或数据丢失
如果你调用了wandb.log但页面上看不到曲线,大概率是横轴步数语义出了问题。比如你在验证循环里也调用了 log,和训练循环的 step 混在一起,曲线刷新时前后点顺序错乱,页面会把数据覆盖或者画出很奇怪的多段折线。
解决办法是统一所有 log 的步数语义。我通常在脚本里维护一个全局的global_step:
global_step += 1 wandb.log({"train_loss": loss}, step=global_step)这样无论训练循环、验证循环还是学习率更新,都基于同一个全局计数器,曲线一定整齐。
另一个常见问题是数据量太大导致图表加载慢。几十万条 batch 日志点会让浏览器渲染很吃力。这时候解决思路不是删数据,而是只保留关键频率的记录点,或者用页面上自带的平滑系数看图。
5.4 分布式训练与多进程避坑
现在不少项目用 DDP 多卡训练。WandB 在这块的坑非常经典:如果你在每个进程里都执行wandb.init(),项目页面就会出现 N 个一模一样的 run,曲线全被淹没,还很难清理。
正确做法是只在主进程里记录。代码里加一个判断:
import torch.distributed as dist if dist.is_initialized(): is_main = dist.get_rank() == 0 else: is_main = True run = None if is_main: run = wandb.init(project="...") # 训练循环里 if is_main: wandb.log({"loss": loss.item()})对于DataParallel这种单进程多卡模式,不需要做这个处理,整个进程只有一个 run。
数据加载方面,DDP 场景记得用DistributedSampler保证每个进程拿到不同子集,WandB 记录的是主进程负责的那部分数据指标。如果采样器写错导致多个进程看到的数据完全相同,最后汇总出的指标可能重复放大某种异常,排查起来特别费劲。
6. 最后说点个人经验
我从 TensorBoard 切成 WandB 之后,最直观的感受不是“曲线更好看了”,而是“实验记录这件事终于自动了”。我现在写任何 Pytorch 脚本,不管是个小消融实验还是完整训练,都会顺手把 WandB 接上:init 一下、config 传参数、log 几个指标、watch 一下模型,五分钟的事。但省下来的时间远超五分钟——我再也不用翻聊天记录确认“那个 0.993 的模型到底是哪次跑出来的”,也不用在周报里一个个贴损失图。
个人建议的接入顺序是:先把第 3 节的最小改造跑通,确认基础流程没问题;再用 config 和 tags 规范项目组织;接着上 Sweeps 做超参搜索;最后才是 Artifacts 和 Report 这类团队协作功能。一步到位反而容易因为细节不熟而气馁。
这个系列如果在“Pytorch 可视化”上继续往深走,下一步我会想写训练完模型转换部署方向的内容,有读者之前在后台问过 Pytorch 转 ONNX 的细节。不管你是做图像、NLP 还是强化学习——强化学习里每个 episode 的 reward 曲线用 WandB 看同样很合适——先把 WandB 接进训练流程,后面所有实验效率都会上一个台阶。