正确做法是用 tqdm 包裹迭代器并显式传入 total=len(train_loader);多卡时仅 rank 0 使用 tqdm;set_postfix 中须传入 .item() 后的标量;避免干扰 logging/wandb。

训练循环里加 tqdm,别直接套 tqdm(train_loader)
直接对 DataLoader 对象套 tqdm 看似简单,但会导致 epoch 进度条卡在 0%,尤其当 num_workers > 0 时——因为 tqdm 无法感知子进程里的实际 batch 数。正确做法是用 tqdm 包裹迭代器本身,并显式传入 total=len(train_loader)。
实操建议:
- 用
from tqdm import tqdm,别用from tqdm.auto import tqdm(后者在 Jupyter 和终端行为不一致,容易在 CI 或脚本中意外静默) - 初始化时写成
pbar = tqdm(train_loader, total=len(train_loader), desc=f"Epoch {epoch}"),total必须显式传,不能依赖自动推断 - 如果
train_loader是无限迭代器(如搭配IterableDataset),必须手动控制break条件,并用pbar.update(1)更新进度
动态打印 loss 和 metric,别在 pbar.set_postfix() 里塞未计算的变量
pbar.set_postfix() 是更新右端状态栏的唯一安全方式,但它只接受标量或字符串;若传入 loss.item() 前没调用 .item(),会引发 RuntimeError: Can't call numpy() on Tensor that requires grad;若传入未 detach 的 tensor,还可能拖慢训练甚至 OOM。
常见错误现象:
-
pbar.set_postfix({"loss": loss})→ 报错或内存泄漏 -
pbar.set_postfix({"acc": acc}),但acc是torch.Tensor且未 .item() → 进度条卡顿、GPU 显存缓慢上涨
正确写法:
loss_val = loss.item()
acc_val = acc.item() if isinstance(acc, torch.Tensor) else acc
pbar.set_postfix({"loss": f"{loss_val:.4f}", "acc": f"{acc_val:.3f}"})
多卡 DDP 训练下 tqdm 只在 rank 0 打印
在 torch.nn.parallel.DistributedDataParallel 场景中,所有 rank 都执行同一段训练循环,若每个 rank 都启一个 tqdm,终端会刷出 N 条重叠进度条,且 set_postfix 的值还是各自 local batch 的指标,不具备全局代表性。
实操建议:
- 只在
rank == 0时初始化和更新tqdm,其余 rank 直接用原生train_loader - 全局指标(如 epoch 平均 loss)需通过
torch.distributed.all_reduce()汇总后再在 rank 0 打印,不能直接用单卡值 - 示例判断逻辑:
if rank == 0: pbar = tqdm(...),循环内if rank == 0: pbar.set_postfix(...); pbar.update(1)
避免 tqdm 干扰 logging 或 wandb 日志上报
tqdm 默认刷新 stdout,而 logging 和 wandb.log() 也常往 stdout/stderr 写内容,容易导致进度条被日志冲乱、换行错位,甚至 wandb 的 step 计数错乱。
关键处理点:
- 初始化
tqdm时加参数file=sys.stdout(显式绑定,避免被 logging 重定向干扰) - 禁用
tqdm自动刷新:设leave=True, dynamic_ncols=True, smoothing=0.1,减少跳变 - wandb 上报指标时,确保
step基于真实 batch index(而非pbar.n),因为pbar.n在异常中断后可能不准
最易被忽略的是:tqdm 的内部计数器 pbar.n 和你手动维护的 global_step 不一定同步——尤其在 resume 训练或使用梯度累积时,务必用独立变量控制日志节奏。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











