ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

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

【Bug已解决】best way of tqdm for data loader 解决方案 【Bug已解决】best way of tqdm for data loader 解决方案问题描述在深度学习训练中使用DataLoader加载数据时进度条是监控训练进度的重要工具。tqdm是 Python 中最流行的进度条库但将其与 PyTorch 的DataLoader正确集成时开发者经常遇到各种问题。典型的问题场景包括进度条在多进程 DataLoader 下不显示或显示混乱进度条的总数total不正确显示为?/?每个 epoch 结束后进度条不关闭导致输出堆积在 Jupyter Notebook 中进度条显示异常多 GPU 训练时进度条重复显示进度条信息不够丰富缺少 loss、学习率等关键指标num_workers 0时 tqdm 报错或卡死这些问题的核心在于理解tqdm的工作机制以及与 PyTorchDataLoader多进程数据加载的交互方式。错误复现场景一进度条总数不正确from tqdm import tqdm from torch.utils.data import DataLoader dataloader DataLoader(dataset, batch_size32, shuffleTrue) # 错误没有指定 total for batch in tqdm(dataloader): # 处理 batch pass # 输出it/s 而不是 it/s [00:3200:01, 3.12it/s] # 没有进度百分比和预计剩余时间场景二多进程下进度条混乱dataloader DataLoader(dataset, batch_size32, num_workers4) for batch in tqdm(dataloader): pass # 在多进程下进度条可能闪烁、重复或完全不显示 # 因为子进程的输出干扰了主进程的进度条场景三进度条不关闭for epoch in range(10): for batch in tqdm(dataloader, descfEpoch {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可以自动推断总长度但在某些情况下如使用IterableDatasetlen()不可用需要手动指定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_size32, shuffleTrue) # 使用 with 语句确保正确关闭 with tqdm(dataloader, descTraining, totallen(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, descTraining): # 训练代码 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(totalself.total, descself.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, fileNone, 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, descdesc, **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, devicecpu, use_notebookFalse, log_fileNone): self.model model self.optimizer optimizer self.criterion criterion self.device torch.device(device) self.model.to(self.device) self.pbar_manager ProgressBarManager( use_notebookuse_notebook, filelog_file ) self.metric_tracker MetricTracker([loss, accuracy]) self.epoch_history [] def train_epoch(self, dataloader, epoch, log_interval10): 训练一个 epoch self.model.train() self.metric_tracker.reset() desc fEpoch {epoch:3d} with self.pbar_manager.create_pbar( dataloader, descdesc, unitbatch ) 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(dim1) 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 fValid {epoch:3d} with self.pbar_manager.create_pbar( dataloader, descdesc, unitbatch, colourgreen ) 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(dim1) 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}, fTrain Acc: {train_metrics[accuracy]:.2f}%) print(f - Val Loss: {val_metrics[loss]:.4f}, fVal 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_size16, shuffleTrue) # 基本用法 print(\n1. 基本进度条:) for batch in tqdm(dataloader, descBasic, unitbatch): time.sleep(0.01) # 带 postfix print(\n2. 带 postfix 的进度条:) pbar tqdm(dataloader, descWithPostfix, unitbatch) 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, descWithContext, unitbatch) 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_size32, shuffleTrue) 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, descManaged, metrics_fncompute_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(fPostfix: {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(dim1) 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_size32, shuffleTrue) val_loader DataLoader(val_ds, batch_size32) # 创建模型 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(), lr0.001) criterion nn.CrossEntropyLoss() # 创建训练器 trainer TrainerWithProgress( modelmodel, optimizeroptimizer, criterioncriterion, ) # 训练 trainer.fit(train_loader, val_loader, num_epochs5) 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_size16, shuffleTrue, num_workers2, # 多进程加载 persistent_workersTrue, # 避免重复创建进程 ) # tqdm 在主进程中使用不受子进程影响 with tqdm(dataloader, descMultiWorker, unitbatch) 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_size16, shuffleTrue) # 自定义格式 bar_format {desc}: {percentage:3.0f}%|{bar}| {n_fmt}/{total_fmt} [{elapsed}{remaining}, {rate_fmt}] with tqdm(dataloader, descCustom, unitbatch, bar_formatbar_format, colourblue) 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_size16, shuffleTrue) num_epochs 3 # 外层 epoch 进度条 epoch_pbar tqdm(range(num_epochs), descOverall, position0) for epoch in epoch_pbar: # 内层 batch 进度条 batch_pbar tqdm(dataloader, descfEpoch {epoch}, position1, leaveFalse) 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: fepoch {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, totallen(dataloader)): pass # 对于 IterableDataset手动指定 for batch in tqdm(dataloader, totalestimated_batches): pass2. 使用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): pass5.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, filef, mininterval1.0): pass7. 分布式训练中的进度条# 只在 rank 0 显示进度条 if dist.get_rank() 0: pbar tqdm(dataloader) else: pbar dataloader # 其他 rank 不使用 tqdm for batch in pbar: pass总结在 PyTorch DataLoader 中使用 tqdm 进度条关键要点如下始终指定total使用totallen(dataloader)确保进度条显示正确的百分比和预计时间。使用with语句确保进度条正确关闭避免输出堆积。Jupyter 中使用tqdm.notebook在 Notebook 环境中使用专用版本避免显示异常。只在主进程使用 tqdm多进程 DataLoader 中子进程不应使用 tqdm避免输出混乱。使用set_postfix显示指标将 loss、accuracy 等指标集成到进度条中避免额外的 print 输出。合理设置更新频率mininterval参数控制刷新频率避免过于频繁的更新影响性能。分布式训练中限制 rank 0多 GPU 训练时只在主进程显示进度条。封装进度管理器将 tqdm 逻辑封装成可复用的类提高代码可维护性。通过正确使用 tqdm可以显著提升训练过程的可观测性实时监控训练进度和指标变化快速发现训练中的问题。
返回列表