news 2026/8/29 23:20:41

【Bug已解决】best way of tqdm for data loader 解决方案

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
【Bug已解决】best way of tqdm for data loader 解决方案

【Bug已解决】best way of tqdm for data loader 解决方案

问题描述

在深度学习训练中,使用DataLoader加载数据时,进度条是监控训练进度的重要工具。tqdm是 Python 中最流行的进度条库,但将其与 PyTorch 的DataLoader正确集成时,开发者经常遇到各种问题。

典型的问题场景包括:

  1. 进度条在多进程 DataLoader 下不显示或显示混乱
  2. 进度条的总数(total)不正确,显示为?/?
  3. 每个 epoch 结束后进度条不关闭,导致输出堆积
  4. 在 Jupyter Notebook 中进度条显示异常
  5. 多 GPU 训练时进度条重复显示
  6. 进度条信息不够丰富,缺少 loss、学习率等关键指标
  7. num_workers > 0时 tqdm 报错或卡死

这些问题的核心在于理解tqdm的工作机制以及与 PyTorchDataLoader多进程数据加载的交互方式。

错误复现

场景一:进度条总数不正确

from tqdm import tqdm from torch.utils.data import DataLoader dataloader = DataLoader(dataset, batch_size=32, shuffle=True) # 错误:没有指定 total for batch in tqdm(dataloader): # 处理 batch pass # 输出:it/s 而不是 it/s [00:32<00:01, 3.12it/s] # 没有进度百分比和预计剩余时间

场景二:多进程下进度条混乱

dataloader = DataLoader(dataset, batch_size=32, num_workers=4) for batch in tqdm(dataloader): pass # 在多进程下,进度条可能闪烁、重复或完全不显示 # 因为子进程的输出干扰了主进程的进度条

场景三:进度条不关闭

for epoch in range(10): for batch in tqdm(dataloader, desc=f'Epoch {epoch}'): # 训练代码 pass # 没有 close 进度条,输出堆积 # 终端中堆积了大量未关闭的进度条

场景四:Jupyter Notebook 显示异常

# 在 Jupyter 中使用标准 tqdm for batch in tqdm(dataloader): pass # 输出多行文本,而不是单行动态更新的进度条

场景五:缺少训练指标

for batch in tqdm(dataloader): loss = model(batch) # 进度条只显示进度,不显示 loss # 需要另外 print loss,导致输出混乱

根因分析

1.tqdm与迭代器的关系

tqdm包装一个可迭代对象,通过计算已迭代次数和总长度的比值来显示进度。对于DataLoader,如果len()方法可用,tqdm可以自动推断总长度;但在某些情况下(如使用IterableDataset),len()不可用,需要手动指定total

2. 多进程数据加载的影响

num_workers > 0时,DataLoader使用子进程加载数据。子进程的stdout输出可能干扰主进程的tqdm进度条刷新。此外,子进程中的tqdm实例可能与主进程冲突。

3.tqdm的刷新机制

tqdm通过\r(回车符)实现单行更新。在终端中这工作良好,但在 Jupyter Notebook 或日志文件中,\r可能不被正确处理,导致多行输出。

4. 进度条生命周期

每个tqdm实例都应该在使用后关闭,以释放资源和清理输出。使用with语句或显式调用close()可以确保正确清理。

解决方案

方案一:基本用法(推荐)

from tqdm import tqdm from torch.utils.data import DataLoader dataloader = DataLoader(dataset, batch_size=32, shuffle=True) # 使用 with 语句确保正确关闭 with tqdm(dataloader, desc='Training', total=len(dataloader)) as pbar: for batch in pbar: # 训练代码 loss = train_step(batch) # 更新进度条信息 pbar.set_postfix({'loss': f'{loss:.4f}'})

方案二:Jupyter Notebook 专用

from tqdm.notebook import tqdm as tqdm_notebook # 在 Jupyter 中使用 notebook 版本 for batch in tqdm_notebook(dataloader, desc='Training'): # 训练代码 pass

方案三:封装训练进度管理器

from tqdm import tqdm from typing import Optional, Dict import torch from torch.utils.data import DataLoader class TrainingProgress: """训练进度管理器""" def __init__(self, total: Optional[int] = None, desc: str = 'Training', use_notebook: bool = False): self.total = total self.desc = desc self.use_notebook = use_notebook self.pbar = None def __enter__(self): tqdm_cls = tqdm_notebook if self.use_notebook else tqdm self.pbar = tqdm_cls(total=self.total, desc=self.desc) return self def __exit__(self, exc_type, exc_val, exc_tb): if self.pbar: self.pbar.close() def update(self, n: int = 1): if self.pbar: self.pbar.update(n) def set_postfix(self, info: Dict): if self.pbar: self.pbar.set_postfix(info) def set_description(self, desc: str): if self.pbar: self.pbar.set_description(desc)

方案四:集成训练指标的进度条

class MetricTracker: """跟踪和显示训练指标""" def __init__(self): self.metrics = {} self.counts = {} def update(self, metrics: Dict[str, float]): for key, value in metrics.items(): if key not in self.metrics: self.metrics[key] = 0.0 self.counts[key] = 0 self.metrics[key] += value self.counts[key] += 1 def get_averages(self) -> Dict[str, float]: return {k: v / self.counts[k] for k, v in self.metrics.items()} def reset(self): self.metrics = {} self.counts = {}

完整修复代码

""" 完整的 tqdm 与 DataLoader 集成方案 涵盖:基本进度条、训练指标显示、多进程支持、Jupyter兼容、断点续训 """ import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader, TensorDataset from tqdm import tqdm from typing import Optional, Dict, List, Callable import time import sys import os # ============================================ # 进度条管理器 # ============================================ class ProgressBarManager: """全面的进度条管理器""" def __init__(self, use_notebook: bool = False, file=None, ncols: Optional[int] = None): """ Args: use_notebook: 是否在 Jupyter Notebook 中使用 file: 输出文件(默认 stderr) ncols: 进度条宽度 """ self.use_notebook = use_notebook self.file = file or sys.stderr self.ncols = ncols self._tqdm_cls = self._get_tqdm_class() def _get_tqdm_class(self): """获取合适的 tqdm 类""" try: if self.use_notebook or 'IPython' in sys.modules: from tqdm.notebook import tqdm as nb_tqdm return nb_tqdm except ImportError: pass return tqdm def create_pbar(self, dataloader: DataLoader, desc: str = '', total: Optional[int] = None, **kwargs) -> tqdm: """创建进度条 Args: dataloader: PyTorch DataLoader desc: 描述文字 total: 总批次数(None 则自动推断) **kwargs: 传递给 tqdm 的额外参数 """ if total is None: try: total = len(dataloader) except TypeError: total = None # IterableDataset 可能没有 len defaults = { 'desc': desc, 'total': total, 'file': self.file, 'ncols': self.ncols, 'leave': True, # 完成后保留进度条 'dynamic_ncols': True, # 动态调整宽度 'mininterval': 0.1, # 最小更新间隔 'ascii': False, # 使用 Unicode 字符 } defaults.update(kwargs) return self._tqdm_cls(dataloader, **defaults) def wrap_dataloader(self, dataloader: DataLoader, desc: str = '', metrics_fn: Optional[Callable] = None, **kwargs): """包装 DataLoader,返回带进度条的迭代器 Args: dataloader: PyTorch DataLoader desc: 描述文字 metrics_fn: 接收 batch 数据,返回指标字典的函数 **kwargs: 传递给 tqdm 的额外参数 """ pbar = self.create_pbar(dataloader, desc=desc, **kwargs) for batch in pbar: if metrics_fn is not None: metrics = metrics_fn(batch) pbar.set_postfix(metrics) yield batch pbar.close() # ============================================ # 训练指标跟踪器 # ============================================ class MetricTracker: """跟踪训练指标""" def __init__(self, metrics_names: List[str] = None): self.metrics_names = metrics_names or [] self.reset() def reset(self): self.values = {name: 0.0 for name in self.metrics_names} self.counts = {name: 0 for name in self.metrics_names} self.history = {name: [] for name in self.metrics_names} def update(self, name: str, value: float, n: int = 1): if name not in self.values: self.values[name] = 0.0 self.counts[name] = 0 ![配图](https://i-blog.csdnimg.cn/img_convert/231e2c718b4e3564784a03e6429765a3.png) self.history[name] = [] self.values[name] += value * n self.counts[name] += n def get_average(self, name: str) -> float: if name not in self.counts or self.counts[name] == 0: return 0.0 return self.values[name] / self.counts[name] def get_all_averages(self) -> Dict[str, float]: return {name: self.get_average(name) for name in self.values} def record_epoch(self): for name in self.values: self.history[name].append(self.get_average(name)) def get_postfix_dict(self) -> Dict[str, str]: return {name: f'{self.get_average(name):.4f}' for name in self.values} # ============================================ # 完整训练器 # ============================================ class TrainerWithProgress: """带进度条的训练器""" def __init__(self, model, optimizer, criterion, device='cpu', use_notebook=False, log_file=None): self.model = model self.optimizer = optimizer self.criterion = criterion self.device = torch.device(device) self.model.to(self.device) self.pbar_manager = ProgressBarManager( use_notebook=use_notebook, file=log_file ) self.metric_tracker = MetricTracker(['loss', 'accuracy']) self.epoch_history = [] def train_epoch(self, dataloader, epoch, log_interval=10): """训练一个 epoch""" self.model.train() self.metric_tracker.reset() desc = f'Epoch {epoch:3d}' with self.pbar_manager.create_pbar( dataloader, desc=desc, unit='batch' ) as pbar: for batch_idx, (data, target) in enumerate(pbar): data, target = data.to(self.device), target.to(self.device) self.optimizer.zero_grad() output = self.model(data) loss = self.criterion(output, target) loss.backward() self.optimizer.step() # 计算指标 with torch.no_grad(): pred = output.argmax(dim=1) correct = pred.eq(target).sum().item() accuracy = 100. * correct / target.size(0) self.metric_tracker.update('loss', loss.item()) self.metric_tracker.update('accuracy', accuracy) # 更新进度条 postfix = self.metric_tracker.get_postfix_dict() postfix['lr'] = f'{self.optimizer.param_groups[0]["lr"]:.2e}' pbar.set_postfix(postfix) # 记录 epoch 结果 epoch_metrics = self.metric_tracker.get_all_averages() self.epoch_history.append(epoch_metrics) return epoch_metrics def validate(self, dataloader, epoch): """验证""" self.model.eval() self.metric_tracker.reset() desc = f'Valid {epoch:3d}' with self.pbar_manager.create_pbar( dataloader, desc=desc, unit='batch', colour='green' ) as pbar: with torch.no_grad(): for data, target in pbar: data, target = data.to(self.device), target.to(self.device) output = self.model(data) loss = self.criterion(output, target) pred = output.argmax(dim=1) correct = pred.eq(target).sum().item() accuracy = 100. * correct / target.size(0) self.metric_tracker.update('loss', loss.item()) self.metric_tracker.update('accuracy', accuracy) pbar.set_postfix(self.metric_tracker.get_postfix_dict()) return self.metric_tracker.get_all_averages() def fit(self, train_loader, val_loader, num_epochs): """完整训练""" print(f"训练开始: {num_epochs} epochs") print(f"设备: {self.device}") print(f"训练批次: {len(train_loader)}") print(f"验证批次: {len(val_loader)}") print("=" * 70) for epoch in range(1, num_epochs + 1): train_metrics = self.train_epoch(train_loader, epoch) val_metrics = self.validate(val_loader, epoch) print(f" -> Train Loss: {train_metrics['loss']:.4f}, " f"Train Acc: {train_metrics['accuracy']:.2f}%") print(f" -> Val Loss: {val_metrics['loss']:.4f}, " f"Val Acc: {val_metrics['accuracy']:.2f}%") print("-" * 70) print("训练完成!") # ============================================ # 使用示例 # ============================================ def demo_basic_tqdm(): """基本 tqdm 用法""" print("=" * 60) print("示例 1: 基本 tqdm 用法") print("=" * 60) # 创建数据 dataset = TensorDataset( torch.randn(100, 10), torch.randint(0, 3, (100,)) ) dataloader = DataLoader(dataset, batch_size=16, shuffle=True) # 基本用法 print("\n1. 基本进度条:") for batch in tqdm(dataloader, desc='Basic', unit='batch'): time.sleep(0.01) # 带 postfix print("\n2. 带 postfix 的进度条:") pbar = tqdm(dataloader, desc='WithPostfix', unit='batch') for batch_idx, (data, target) in enumerate(pbar): loss = torch.rand(1).item() pbar.set_postfix({'loss': f'{loss:.4f}', 'batch': batch_idx}) time.sleep(0.01) pbar.close() # 使用 with 语句 print("\n3. 使用 with 语句:") with tqdm(dataloader, desc='WithContext', unit='batch') as pbar: for batch in pbar: pbar.set_postfix({'status': 'training'}) time.sleep(0.01) print() def demo_progress_manager(): """进度条管理器用法""" print("=" * 60) print("示例 2: 进度条管理器") print("=" * 60) dataset = TensorDataset( torch.randn(200, 10), torch.randint(0, 3, (200,)) ) dataloader = DataLoader(dataset, batch_size=32, shuffle=True) manager = ProgressBarManager() # 使用 wrap_dataloader print("\n带指标的进度条:") def compute_metrics(batch): data, target = batch return {'batch_size': data.size(0)} for batch in manager.wrap_dataloader( dataloader, desc='Managed', metrics_fn=compute_metrics ): time.sleep(0.01) print() def demo_metric_tracker(): """指标跟踪器用法""" print("=" * 60) print("示例 3: 指标跟踪器") print("=" * 60) tracker = MetricTracker(['loss', 'accuracy']) # 模拟训练 for i in range(10): loss = 1.0 / (i + 1) acc = 50 + i * 5 tracker.update('loss', loss) tracker.update('accuracy', acc) print(f"平均 Loss: {tracker.get_average('loss'):.4f}") print(f"平均 Accuracy: {tracker.get_average('accuracy'):.2f}%") print(f"所有指标: {tracker.get_all_averages()}") print(f"Postfix: {tracker.get_postfix_dict()}") print() def demo_full_training(): """完整训练示例""" print("=" * 60) print("示例 4: 完整训练流程") print("=" * 60) # 创建数据 torch.manual_seed(42) X = torch.randn(500, 10) y = (X @ torch.randn(10, 3)).argmax(dim=1) dataset = TensorDataset(X, y) train_size = 400 val_size = 100 train_ds, val_ds = torch.utils.data.random_split(dataset, [train_size, val_size]) train_loader = DataLoader(train_ds, batch_size=32, shuffle=True) val_loader = DataLoader(val_ds, batch_size=32) # 创建模型 model = nn.Sequential( nn.Linear(10, 64), nn.ReLU(), nn.Dropout(0.2), nn.Linear(64, 32), nn.ReLU(), nn.Linear(32, 3), ) optimizer = optim.Adam(model.parameters(), lr=0.001) criterion = nn.CrossEntropyLoss() # 创建训练器 trainer = TrainerWithProgress( model=model, optimizer=optimizer, criterion=criterion, ) # 训练 trainer.fit(train_loader, val_loader, num_epochs=5) print() def demo_multiprocess_safe(): """多进程安全的进度条""" print("=" * 60) print("示例 5: 多进程 DataLoader 的进度条") print("=" * 60) dataset = TensorDataset( torch.randn(100, 10), torch.randint(0, 3, (100,)) ) # num_workers > 0 时的正确用法 dataloader = DataLoader( dataset, batch_size=16, shuffle=True, num_workers=2, # 多进程加载 persistent_workers=True, # 避免重复创建进程 ) # tqdm 在主进程中使用,不受子进程影响 with tqdm(dataloader, desc='MultiWorker', unit='batch') as pbar: for batch in pbar: time.sleep(0.02) pbar.set_postfix({'workers': 2}) print() def demo_custom_format(): """自定义进度条格式""" print("=" * 60) print("示例 6: 自定义进度条格式") print("=" * 60) dataset = TensorDataset( torch.randn(100, 10), torch.randint(0, 3, (100,)) ) dataloader = DataLoader(dataset, batch_size=16, shuffle=True) # 自定义格式 bar_format = '{desc}: {percentage:3.0f}%|{bar}| {n_fmt}/{total_fmt} [{elapsed}<{remaining}, {rate_fmt}]' with tqdm(dataloader, desc='Custom', unit='batch', bar_format=bar_format, colour='blue') as pbar: for batch in pbar: time.sleep(0.01) print() def demo_nested_progress(): """嵌套进度条(epoch + batch)""" print("=" * 60) print("示例 7: 嵌套进度条") print("=" * 60) dataset = TensorDataset( torch.randn(100, 10), torch.randint(0, 3, (100,)) ) dataloader = DataLoader(dataset, batch_size=16, shuffle=True) num_epochs = 3 # 外层 epoch 进度条 epoch_pbar = tqdm(range(num_epochs), desc='Overall', position=0) for epoch in epoch_pbar: # 内层 batch 进度条 batch_pbar = tqdm(dataloader, desc=f'Epoch {epoch}', position=1, leave=False) for batch in batch_pbar: time.sleep(0.01) batch_pbar.set_postfix({'loss': f'{torch.rand(1).item():.4f}'}) batch_pbar.close() epoch_pbar.set_postfix({'status': f'epoch {epoch} done'}) epoch_pbar.close() print() if __name__ == '__main__': demo_basic_tqdm() demo_progress_manager() demo_metric_tracker() demo_full_training() demo_multiprocess_safe() demo_custom_format() demo_nested_progress() print("=" * 60) print("所有示例执行完毕!") print("=" * 60)

常见陷阱与注意事项

1. 始终指定total

# 推荐:显式指定 total for batch in tqdm(dataloader, total=len(dataloader)): pass # 对于 IterableDataset,手动指定 for batch in tqdm(dataloader, total=estimated_batches): pass

2. 使用with语句或close()

# 推荐:使用 with with tqdm(dataloader) as pbar: for batch in pbar: pass # 或者显式 close pbar = tqdm(dataloader) for batch in pbar: pass pbar.close() # 必须关闭!

3. Jupyter Notebook 中使用tqdm.notebook

# 在 Jupyter 中使用 from tqdm.notebook import tqdm # 而不是 from tqdm import tqdm # 这会在 Jupyter 中显示异常

4. 多进程下避免子进程中的 tqdm

# 错误:在 collate_fn 或 Dataset.__getitem__ 中使用 tqdm def collate_fn(batch): # 不要在这里用 tqdm! return default_collate(batch) # 正确:只在主进程的迭代中使用 for batch in tqdm(dataloader): pass

5.set_postfix的性能

# 频繁更新 postfix 可能影响性能 for batch in tqdm(dataloader): pbar.set_postfix({'loss': loss.item()}) # 每个 batch 都更新 # 可以降低更新频率 for batch_idx, batch in enumerate(tqdm(dataloader)): if batch_idx % 10 == 0: pbar.set_postfix({'loss': loss.item()})

6. 日志文件中的进度条

# 输出到文件时禁用进度条刷新 with open('train.log', 'w') as f: for batch in tqdm(dataloader, file=f, mininterval=1.0): pass

7. 分布式训练中的进度条

# 只在 rank 0 显示进度条 if dist.get_rank() == 0: pbar = tqdm(dataloader) else: pbar = dataloader # 其他 rank 不使用 tqdm for batch in pbar: pass

总结

在 PyTorch DataLoader 中使用 tqdm 进度条,关键要点如下:

  1. 始终指定total:使用total=len(dataloader)确保进度条显示正确的百分比和预计时间。

  2. 使用with语句:确保进度条正确关闭,避免输出堆积。

  3. Jupyter 中使用tqdm.notebook:在 Notebook 环境中使用专用版本,避免显示异常。

  4. 只在主进程使用 tqdm:多进程 DataLoader 中,子进程不应使用 tqdm,避免输出混乱。

  5. 使用set_postfix显示指标:将 loss、accuracy 等指标集成到进度条中,避免额外的 print 输出。

  6. 合理设置更新频率mininterval参数控制刷新频率,避免过于频繁的更新影响性能。

  7. 分布式训练中限制 rank 0:多 GPU 训练时只在主进程显示进度条。

  8. 封装进度管理器:将 tqdm 逻辑封装成可复用的类,提高代码可维护性。

通过正确使用 tqdm,可以显著提升训练过程的可观测性,实时监控训练进度和指标变化,快速发现训练中的问题。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/8/29 23:18:45

LIS25BA骨传导拾音实战:TDM接口与振动信号处理全记录

去年调一块骨传导方案的板子&#xff0c;被风噪折磨了一个星期。TWS耳机里那两颗MEMS麦克风&#xff0c;在户外风一吹&#xff0c;采集到的全是呼啸声&#xff0c;语音完全提不出来。后来把骨传导拾音从麦克风换成了LIS25BA——一颗带TDM接口的低噪声、高带宽3轴数字输出加速度…

作者头像 李华
网站建设 2026/8/29 23:18:03

用友2016校招前端笔试题:JavaScript闭包、盒模型等基础考点解密

很多人问我当年用友2016校招web前端笔试题考了什么。2016年正好是前端圈一个很有意思的节点&#xff1a;jQuery还统治着大量存量项目&#xff0c;AngularJS热潮稍退&#xff0c;React开始快速占领新项目&#xff0c;Vue也进入了不少团队的技术选型名单。在这股框架热潮里&#…

作者头像 李华
网站建设 2026/8/29 23:17:47

四足机器人步态规划与运动控制核心解析

简介&#xff1a;在足式机器人的运动控制研究中&#xff0c;步态是决定其动态性能与稳定性的核心概念。步态规划通过设定各腿的运动相位与支撑顺序&#xff0c;使四足机器人能够在不同地形上实现协调移动。理解步态原理有助于优化机器人的能量效率与负载能力&#xff0c;技术价…

作者头像 李华
网站建设 2026/8/29 23:14:17

Matlab方程求解实战:从线性代数到微分方程的核心工具与避坑指南

1. 项目概述&#xff1a;为什么方程求解是Matlab的基石如果你用过Matlab&#xff0c;哪怕只是画过一张简单的正弦波图&#xff0c;你大概率也已经在后台调用了它的方程求解能力。方程求解&#xff0c;这个听起来有点“数学课”味道的词&#xff0c;其实是Matlab这座大厦最核心的…

作者头像 李华
网站建设 2026/8/29 23:13:13

HCI_HARDWARE_ERROR_EVENT 与 ISR 延迟误差:蓝牙控制器异常排查实录

HCI_HARDWARE_ERROR_EVENT 与 ISR 延迟误差&#xff1a;一次完整的蓝牙控制器异常排查实录最近在调试一款基于低功耗蓝牙芯片的物联网模组时&#xff0c;遇到了一个非常棘手的稳定性问题。设备在长时间运行后&#xff0c;会随机出现连接断开&#xff0c;并且在调试日志中频繁看…

作者头像 李华
网站建设 2026/8/29 23:10:46

游戏服务端日志分析与数据库工具安全使用指南

简介&#xff1a;游戏服务端日志&#xff08;如ItemLog.BIN&#xff09;和配置脚本&#xff08;如Player.lua&#xff09;是运维与调试的关键数据载体&#xff0c;其解析依赖于对二进制日志结构和Lua逻辑的底层理解&#xff1b;MDBQuery.exe等数据库查询工具虽能高效读取Access…

作者头像 李华