You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.01 21:47:07