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

PyTorch中如何让nn.ModuleList内的独立nn.Module实例真正并行运行

解决方案

一、向量化批量计算(GPU/CPU都高效,优先推荐)

循环处理之所以慢,是因为没有利用PyTorch底层的并行计算能力。针对多个独立模型处理对应输入切片的场景,向量化改造是最优选择,能一次性完成所有计算。

方法1:用torch.vmap自动并行(PyTorch 1.12+)

torch.vmap是专门用来批量执行独立操作的工具,无需手动合并参数,直接让多个模型的计算并行化:

import torch
import torch.nn as nn

class FullyConnectedNetwork(nn.Module):
    def __init__(self):
        super(FullyConnectedNetwork, self).__init__()
        self.fc1 = nn.Linear(20, 10)
        self.fc2 = nn.Linear(10, 1)
    
    def forward(self, x):
        x = self.fc1(x)
        x = self.fc2(x)
        return x

class ParallelFCN(nn.Module):
    def __init__(self, n):
        super(ParallelFCN, self).__init__()
        self.models = nn.ModuleList([FullyConnectedNetwork() for _ in range(n)])
        # 把所有模型的参数整理成批量格式
        self.fc1_w = torch.stack([m.fc1.weight for m in self.models])  # shape: (n, 10, 20)
        self.fc1_b = torch.stack([m.fc1.bias for m in self.models])    # shape: (n, 10)
        self.fc2_w = torch.stack([m.fc2.weight for m in self.models])  # shape: (n, 1, 10)
        self.fc2_b = torch.stack([m.fc2.bias for m in self.models])    # shape: (n, 1)

    def forward(self, x):
        # 拆分输入为n个20维切片,堆叠成批量格式
        x_slices = x.chunk(len(self.models), dim=1)
        x_batch = torch.stack(x_slices)  # shape: (n, batch_size, 20)
        
        # 定义单个模型的计算逻辑
        def single_model_calc(x_slice, fc1_w, fc1_b, fc2_w, fc2_b):
            x = nn.functional.linear(x_slice, fc1_w, fc1_b)
            x = nn.functional.linear(x, fc2_w, fc2_b)
            return x
        
        # vmap自动在n的维度上并行计算
        outputs = torch.vmap(single_model_calc)(x_batch, self.fc1_w, self.fc1_b, self.fc2_w, self.fc2_b)
        # 调整形状后返回
        return outputs.permute(1, 0, 2).squeeze(-1)

# 示例
n = 400
model = ParallelFCN(n)
x = torch.randn(32, 20*n)  # batch_size=32,总输入维度20*400
output = model(x)
print(output.shape)  # 输出 (32, 400)

方法2:手动合并线性层(兼容低版本PyTorch)

把所有独立模型的线性层参数合并成大的权重矩阵,通过两次矩阵运算完成所有计算,兼容性更强:

import torch
import torch.nn as nn

class ParallelFCN(nn.Module):
    def __init__(self, n):
        super(ParallelFCN, self).__init__()
        self.n = n
        # 先创建临时模型,提取参数
        temp_models = [FullyConnectedNetwork() for _ in range(n)]
        # 合并fc1的权重和偏置
        self.fc1_weight = nn.Parameter(torch.cat([m.fc1.weight for m in temp_models], dim=0))  # (10n, 20)
        self.fc1_bias = nn.Parameter(torch.cat([m.fc1.bias for m in temp_models], dim=0))      # (10n,)
        # 合并fc2的权重和偏置
        self.fc2_weight = nn.Parameter(torch.cat([m.fc2.weight for m in temp_models], dim=0))  # (n, 10)
        self.fc2_bias = nn.Parameter(torch.cat([m.fc2.bias for m in temp_models], dim=0))      # (n,)

    def forward(self, x):
        bs = x.size(0)
        # 把输入拆成(batch_size, n, 20)的格式
        x_splits = x.view(bs, self.n, 20)
        # 批量计算fc1输出:(batch_size, n, 10)
        fc1_out = torch.einsum('bni,ni->bnj', x_splits, self.fc1_weight.view(self.n, 10, 20)) + self.fc1_bias.view(self.n, 10)
        # 批量计算fc2输出:(batch_size, n)
        fc2_out = torch.einsum('bnj,nj->bn', fc1_out, self.fc2_weight.view(self.n, 10)) + self.fc2_bias
        return fc2_out

class FullyConnectedNetwork(nn.Module):
    def __init__(self):
        super(FullyConnectedNetwork, self).__init__()
        self.fc1 = nn.Linear(20, 10)
        self.fc2 = nn.Linear(10, 1)
    
    def forward(self, x):
        x = self.fc1(x)
        x = self.fc2(x)
        return x

二、多进程处理(仅适合CPU计算场景)

如果模型在CPU上运行,且向量化优化后速度仍不够,可以用多进程并行处理。GPU场景不推荐,会导致显存占用飙升,反而降低效率:

import torch
import torch.nn as nn
from torch.multiprocessing import Pool, set_start_method

class FullyConnectedNetwork(nn.Module):
    def __init__(self):
        super(FullyConnectedNetwork, self).__init__()
        self.fc1 = nn.Linear(20, 10)
        self.fc2 = nn.Linear(10, 1)
    
    def forward(self, x):
        x = self.fc1(x)
        x = self.fc2(x)
        return x

# 单个模型的计算函数,供多进程调用
def process_single_model(args):
    model, x_slice = args
    return model(x_slice)

class ParallelFCN(nn.Module):
    def __init__(self, n):
        super(ParallelFCN, self).__init__()
        self.models = nn.ModuleList([FullyConnectedNetwork() for _ in range(n)])
        # 进程数建议不超过CPU核心数
        self.pool = Pool(processes=min(n, 8))

    def forward(self, x):
        x_slices = x.chunk(len(self.models), dim=1)
        args_list = [(self.models[i], x_slices[i]) for i in range(len(self.models))]
        # 多进程并行计算
        outputs = self.pool.map(process_single_model, args_list)
        return torch.cat(outputs, dim=1)

if __name__ == '__main__':
    # Windows系统需要设置启动方式
    try:
        set_start_method('spawn')
    except RuntimeError:
        pass
    n = 400
    model = ParallelFCN(n)
    x = torch.randn(32, 20*n)
    output = model(x)
    print(output.shape)

注意:多进程会复制多份模型参数,内存开销大,仅适合模型小、CPU资源充足的场景。

三、关键注意事项

  • 优先用向量化方案:GPU的核心优势就是并行计算,向量化能最大化利用算力,比多线程/多进程效率高得多。
  • 多进程仅适用于CPU:GPU上用多进程会增加显存占用和调度成本,得不偿失。
  • 如果用多GPU,可以考虑将模型分散到不同GPU处理不同切片,但实现复杂度高,单GPU的向量化方案足以应对n=400的场景。

内容的提问来源于stack exchange,提问作者Peyman

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 08:48:13