PyTorch训练中tqdm进度条遭意外打印干扰的问题排查
问题分析与解决方案
为什么会重复打印版本信息?
你的print(f"torch version:...")语句处于全局作用域,且未用if __name__ == "__main__":包裹。如果主脚本被engine.py(或其他模块)导入,或是代码逻辑触发了全局代码的重复执行(比如循环中意外触发模块重新加载),这些全局代码就会反复运行。结合你的场景,最可能的原因是engine.train()执行时,主脚本被当作模块重复导入,导致每次tqdm循环迭代时,全局的print语句都会被执行一遍。
主函数为何会关联到这些全局打印语句?
这里并非主函数“访问”打印语句,而是全局代码的执行时机问题:main()函数定义在全局作用域中,当脚本被直接运行或导入时,全局作用域的代码(包括print语句)会先于函数执行。如果脚本被重复导入,全局代码会再次执行,print语句就会重复输出,而main()调用train()的过程刚好触发了这个重复导入逻辑。
如何避免重复打印?
最规范的解决方法是把所有仅在脚本直接运行时需要执行的代码,全部包裹在if __name__ == "__main__":块中,确保这些代码只在脚本作为主程序运行时执行,被其他模块导入时不会触发。修改后的代码示例如下:
import torch import torchvision from torchvision import transforms import data_setup # 定义主函数(仅在被调用时执行) def main(): # 版本信息仅在主函数执行时打印一次 print(f"torch version: {torch.__version__}") print(f"torchvision version: {torchvision.__version__}") load_data() manual_transforms = transforms.Compose([]) train_dataloader, test_dataloader, class_names = data_setup.create_dataloaders() results = engine.train(model=model, train_dataloader=train_dataloader, test_dataloader=test_dataloader, optimizer=optimizer, loss_fn=loss_fn, epochs=5, device=device) # 脚本直接运行时才执行入口逻辑 if __name__ == "__main__": # 初始化训练所需的变量(model、optimizer等) device = "cuda" if torch.cuda.is_available() else "cpu" # 假设model、optimizer、loss_fn已提前定义/导入 main()
核心逻辑是:if __name__ == "__main__":是Python的程序入口判断,只有脚本被直接运行时,__name__才会等于"__main__",块内代码才会执行;当脚本被其他模块导入时,__name__会变为模块名,块内代码不会运行,从根源上避免了全局代码被重复执行。
内容的提问来源于stack exchange,提问作者Jose Ramon
相关产品推荐
相关产品推荐

