如何在CPU上并行运行多个PyTorch推理任务?
关于PyTorch CPU多进程推理的问题解答
1. torch.multiprocessing.Pool是不是最优方案?
是的,这是当前场景下的推荐方案。Python的GIL会限制多线程的CPU并行效率,而多进程能完全绕过GIL,非常适合CPU密集型的推理任务。torch.multiprocessing在原生Python多进程基础上做了适配,支持张量的共享内存传递,避免不必要的数据拷贝,完美匹配你这种任务完全独立的场景。
2. 如何部署N个CPU任务并收集结果?
核心思路是将数据集拆分为N份(对应N个CPU核心),每个进程处理自己的子集,最后汇总所有进程的结果。以下是两种可行实现方式:
方式一:使用torch.multiprocessing.Pool(简洁高效)
import torch import torch.multiprocessing as mp from torch.utils.data import Subset, DataLoader from torchvision.datasets import CIFAR10 from torchvision.transforms import ToTensor # 替换成你自己的模型定义 class YourModel(torch.nn.Module): def __init__(self): super().__init__() self.conv_layers = torch.nn.Sequential( torch.nn.Conv2d(3, 16, kernel_size=3), torch.nn.ReLU(), torch.nn.MaxPool2d(2) ) self.fc_layers = torch.nn.Sequential( torch.nn.Linear(16*15*15, 128), torch.nn.ReLU(), torch.nn.Linear(128, 10) ) def forward(self, x): x = self.conv_layers(x) x = x.flatten(1) x = self.fc_layers(x) return x def worker_func(subset_indices, dataset, model_path): # 每个进程独立加载模型(内存充足时无需共享) model = YourModel() model.load_state_dict(torch.load(model_path)) model.eval() # 创建当前进程的DataLoader subset = Subset(dataset, subset_indices) loader = DataLoader(subset, batch_size=64, num_workers=2) # 推理并计算当前子集的准确率 correct = 0 total = 0 with torch.no_grad(): for imgs, labels in loader: outputs = model(imgs) _, preds = torch.max(outputs, 1) total += labels.size(0) correct += (preds == labels).sum().item() return correct, total if __name__ == "__main__": # 初始化CIFAR10测试集 dataset = CIFAR10(root="./data", train=False, transform=ToTensor(), download=True) # 设置进程数(建议等于CPU核心数) num_processes = mp.cpu_count() model_path = "./trained_model.pth" # 替换成你的模型路径 # 拆分数据集索引 total_samples = len(dataset) indices_per_process = [list(range(i, total_samples, num_processes)) for i in range(num_processes)] # 启动进程池(PyTorch推荐用spawn避免fork的潜在问题) mp.set_start_method('spawn') with mp.Pool(num_processes) as pool: # 向每个进程传递任务参数 results = pool.starmap(worker_func, [(idxs, dataset, model_path) for idxs in indices_per_process]) # 汇总所有进程的结果 total_correct = sum(c for c, t in results) total_samples = sum(t for c, t in results) overall_accuracy = total_correct / total_samples print(f"整体准确率: {overall_accuracy:.4f}")
方式二:手动创建mp.Process(更灵活)
如果需要更精细的进程控制,可以手动创建进程并通过队列收集结果:
import torch import torch.multiprocessing as mp from torch.utils.data import Subset, DataLoader from torchvision.datasets import CIFAR10 from torchvision.transforms import ToTensor class YourModel(torch.nn.Module): def __init__(self): super().__init__() self.conv_layers = torch.nn.Sequential( torch.nn.Conv2d(3, 16, kernel_size=3), torch.nn.ReLU(), torch.nn.MaxPool2d(2) ) self.fc_layers = torch.nn.Sequential( torch.nn.Linear(16*15*15, 128), torch.nn.ReLU(), torch.nn.Linear(128, 10) ) def forward(self, x): x = self.conv_layers(x) x = x.flatten(1) x = self.fc_layers(x) return x def worker_func(subset_indices, dataset, model_path, result_queue): model = YourModel() model.load_state_dict(torch.load(model_path)) model.eval() subset = Subset(dataset, subset_indices) loader = DataLoader(subset, batch_size=64) correct = 0 total = 0 with torch.no_grad(): for imgs, labels in loader: outputs = model(imgs) _, preds = torch.max(outputs, 1) total += labels.size(0) correct += (preds == labels).sum().item() result_queue.put((correct, total)) if __name__ == "__main__": dataset = CIFAR10(root="./data", train=False, transform=ToTensor(), download=True) num_processes = mp.cpu_count() model_path = "./trained_model.pth" result_queue = mp.Queue() total_samples = len(dataset) indices_per_process = [list(range(i, total_samples, num_processes)) for i in range(num_processes)] mp.set_start_method('spawn') processes = [] for idxs in indices_per_process: p = mp.Process(target=worker_func, args=(idxs, dataset, model_path, result_queue)) processes.append(p) p.start() # 等待所有进程结束 for p in processes: p.join() # 收集结果 total_correct = 0 total_samples = 0 while not result_queue.empty(): c, t = result_queue.get() total_correct += c total_samples += t overall_accuracy = total_correct / total_samples print(f"整体准确率: {overall_accuracy:.4f}")
3. 是否需要手动处理torch.device?
不需要额外手动处理。纯CPU推理场景下,PyTorch默认会将模型和张量放在CPU设备上,每个进程的运行环境是独立的,无需显式指定torch.device('cpu')(当然显式指定也没问题,不影响运行)。如果你的模型之前在GPU上训练,加载时只要不调用.cuda(),模型会自动切换到CPU运行。
关键注意事项
- 必须在
if __name__ == "__main__":代码块下启动多进程,避免Windows系统下的进程启动异常。 - 推荐使用
spawn作为进程启动方式,PyTorch对其兼容性更好,能避免fork带来的模型状态紊乱问题。 - 若内存资源紧张,可使用
model.share_memory()让所有进程共享模型权重;但你内存充足的情况下,每个进程独立加载模型更简单,无需额外操作。
内容的提问来源于stack exchange,提问作者edmz
相关产品推荐
相关产品推荐

