PyTorch训练分类器教程报错求助:BrokenPipeError问题解决
解决PyTorch DataLoader引发的BrokenPipeError和RuntimeError
我之前在Windows上跑PyTorch的CIFAR10教程时,也碰到过完全一样的报错!这个问题的核心原因是Windows系统的多进程机制和PyTorch DataLoader的交互问题,报错信息里其实已经给出了明确的解决提示——必须把你的主程序逻辑放到if __name__ == '__main__':代码块中。
具体解决步骤:
- 把所有涉及
DataLoader初始化、迭代(比如你代码里的dataiter = iter(trainloader))以及图像展示的代码,全部包裹在if __name__ == '__main__':代码块内。 - 如果不需要把脚本打包成可执行文件,
freeze_support()这一行可以省略。
为什么要这么做?
Windows系统创建子进程用的是spawn模式,这种模式会重新导入你的主脚本文件。如果没有if __name__ == '__main__':的判断,子进程会重复执行全局范围内的所有代码(包括重新创建DataLoader),进而引发进程启动冲突,最终导致BrokenPipeError和RuntimeError。而Linux/macOS用的是fork模式,不会有这个问题,所以官方教程里可能没特意说明,但Windows环境下必须遵守这个规范。
修改后的完整示例代码
import torch import torchvision import torchvision.transforms as transforms import matplotlib.pyplot as plt import numpy as np # 定义图像展示函数(这个可以放在全局) def imshow(img): img = img / 2 + 0.5 # 反归一化 npimg = img.numpy() plt.imshow(np.transpose(npimg, (1, 2, 0))) plt.show() # 别忘了加这一行,不然图像可能不显示 # 所有主逻辑放到这里! if __name__ == '__main__': # 数据预处理和DataLoader初始化 transform = transforms.Compose( [transforms.ToTensor(), transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))]) trainset = torchvision.datasets.CIFAR10(root='./data', train=True, download=True, transform=transform) trainloader = torch.utils.data.DataLoader(trainset, batch_size=4, shuffle=True, num_workers=2) classes = ('plane', 'car', 'bird', 'cat', 'deer', 'dog', 'frog', 'horse', 'ship', 'truck') # 获取随机训练图像并展示 dataiter = iter(trainloader) images, labels = dataiter.next() imshow(torchvision.utils.make_grid(images)) print(' '.join('%5s' % classes[labels[j]] for j in range(4)))
额外提示:
- 如果你设置了
num_workers>0(多进程加载数据),这个规范是必须的;如果把num_workers改成0(单进程),可能暂时不会报错,但不推荐这么做,因为会影响数据加载效率。 - 确保
plt.show()被调用,不然图像可能只会在后台处理而不会弹出窗口。
内容的提问来源于stack exchange,提问作者Clém Grt
相关产品推荐
相关产品推荐

