PyTorch设置num_workers>0时触发RuntimeError问题求助
解决PyTorch Lightning数据加载多进程RuntimeError问题
问题根源分析
这个错误本质是Windows系统下多进程启动的机制限制(默认用spawn而非fork),加上你的代码中大部分初始化逻辑都在全局作用域——当num_workers>0时,子进程会重新执行整个脚本的全局代码,从而引发重复初始化的冲突。你之前加了if __name__ == '__main__'但没生效,大概率是没把所有关键初始化逻辑都包裹到这个判断块里。
正确的代码重构方案
你需要把所有涉及数据集、数据加载器、模型实例化、训练启动的代码,甚至绘图代码都放到if __name__ == '__main__'保护块中,只把导入、函数定义和类定义留在全局作用域。修改后的代码如下:
import torchvision from torchvision import transforms import torchmetrics import pytorch_lightning as pl from pytorch_lightning.callbacks import ModelCheckpoint from pytorch_lightning.loggers import TensorBoardLogger from tqdm import tqdm import numpy as np import matplotlib.pyplot as plt def load_file(path): return np.load(path).astype(np.float32) # 数据变换定义可留在全局作用域 train_transforms = transforms.Compose([ transforms.ToTensor(), transforms.Normalize(0.49, 0.248), transforms.RandomAffine(degrees=(-5, 5), translate=(0, 0.05), scale=(0.9, 1.1)), transforms.RandomResizedCrop((224, 224), scale=(0.35, 1)) ]) val_transforms = transforms.Compose([ transforms.ToTensor(), transforms.Normalize([0.49], [0.248]), ]) class PneumoniaModel(pl.LightningModule): def __init__(self, weight=1): super().__init__() self.model = torchvision.models.resnet18() self.model.conv1 = torch.nn.Conv2d(1, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False) self.model.fc = torch.nn.Linear(in_features=512, out_features=1) self.optimizer = torch.optim.Adam(self.model.parameters(), lr=1e-4) self.loss_fn = torch.nn.BCEWithLogitsLoss(pos_weight=torch.tensor([weight])) self.train_acc = torchmetrics.Accuracy() self.val_acc = torchmetrics.Accuracy() def forward(self, data): return self.model(data) def training_step(self, batch, batch_idx): x_ray, label = batch label = label.float() pred = self(x_ray)[:, 0] loss = self.loss_fn(pred, label) self.log("Train Loss", loss) self.log("Step Train Acc", self.train_acc(torch.sigmoid(pred), label.int())) return loss def training_epoch_end(self, outs): self.log("Train Acc", self.train_acc.compute()) def validation_step(self, batch, batch_idx): x_ray, label = batch label = label.float() pred = self(x_ray)[:, 0] loss = self.loss_fn(pred, label) self.log("Val Loss", loss) self.log("Step Val Acc", self.val_acc(torch.sigmoid(pred), label.int())) return loss def validation_epoch_end(self, outs): self.log("Val Acc", self.val_acc.compute()) def configure_optimizers(self): return [self.optimizer] # 所有运行时逻辑都放到这里 if __name__ == '__main__': # 数据集与加载器初始化 train_dataset = torchvision.datasets.DatasetFolder( "./Processed/train/", loader=load_file, extensions="npy", transform=train_transforms) val_dataset = torchvision.datasets.DatasetFolder( "./Processed/val/", loader=load_file, extensions="npy", transform=val_transforms) # 绘图代码移至主进程执行 fig, axis = plt.subplots(2, 2, figsize=(9, 9)) for i in range(2): for j in range(2): random_index = np.random.randint(0, len(train_dataset)) x_ray, label = train_dataset[random_index] axis[i][j].imshow(x_ray[0], cmap="bone") axis[i][j].set_title(f"Label:{label}, id:{random_index}") plt.show() batch_size = 64 num_workers = 8 train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=batch_size, num_workers=num_workers, shuffle=True, pin_memory=True) val_loader = torch.utils.data.DataLoader(val_dataset, batch_size=batch_size, num_workers=num_workers, shuffle=False, pin_memory=True) print(f"There are {len(train_dataset)} train images and {len(val_dataset)} val images") print("Train class distribution:", np.unique(train_dataset.targets, return_counts=True)) print("Val class distribution:", np.unique(val_dataset.targets, return_counts=True)) # 模型与训练器初始化 model = PneumoniaModel() checkpoint_callback = ModelCheckpoint(monitor='Val Acc', save_top_k=10, mode='max') gpus = 1 trainer = pl.Trainer(gpus=gpus, logger=TensorBoardLogger(save_dir="./logs"), log_every_n_steps=1, callbacks=checkpoint_callback, max_epochs=35) trainer.fit(model, train_loader, val_loader)
额外的PyCharm配置注意事项
如果修改后仍报错,检查PyCharm的运行配置:
- 打开Run/Debug Configurations,找到你的运行脚本
- 取消勾选Emulate terminal in output console选项(该选项会干扰多进程的标准输出重定向)
- 确保使用普通"Run"模式,而非"Run in Python Console"模式
这样调整后,num_workers>0的多进程数据加载就能正常工作,训练速度会显著提升。
内容的提问来源于stack exchange,提问作者Enrique Gil Garcia
相关产品推荐
相关产品推荐

