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

如何在PyTorch模块内并行运行多个子模块?

解决方案

首先要修正你的初始化代码:不能用普通列表存储子模块,PyTorch无法识别列表中的LocalModel,会导致参数无法被优化器跟踪、无法自动迁移到GPU等问题。必须改用torch.nn.ModuleList,它是PyTorch专门设计的子模块容器,能自动管理其中的所有子模型参数。

修正后的初始化代码

import torch

class GlobalModel(torch.nn.Module):
    def __init__(self, n_local_models):
        super().__init__()
        # 用ModuleList替代普通列表
        self.local_models = torch.nn.ModuleList([LocalModel() for _ in range(n_local_models)])
        
        # 可选:根据LocalModel的输出维度动态初始化线性层(替代占位的100)
        # 假设输入到单个LocalModel的维度为input_dim_per_local,先获取示例输出
        sample_input = torch.randn(1, input_dim_per_local)  # 替换为你实际的单LocalModel输入维度
        local_output_dim = self.local_models[0](sample_input).shape[1]
        self.linear = torch.nn.Linear(n_local_models * local_output_dim, 100)
        
        self.activation = torch.nn.ReLU()

高效并行的forward实现

在GPU环境下,PyTorch会自动异步调度多个子模型的前向计算,用列表推导即可实现并行(无需额外复杂的并行框架)。核心步骤是:拆分输入→并行处理每个子输入→拼接输出→后续层计算。

def forward(self, x):
    # 1. 将输入按特征维度拆分为n_local_models个等长张量
    # dim=1表示按特征维度拆分(假设输入形状为[batch_size, total_feature_dim])
    split_inputs = torch.chunk(x, len(self.local_models), dim=1)
    
    # 2. 并行处理每个子输入:GPU会自动调度多个LocalModel的forward并行执行
    local_outputs = [model(sub_x) for model, sub_x in zip(self.local_models, split_inputs)]
    
    # 3. 拼接所有LocalModel的输出(按特征维度拼接)
    concat_out = torch.cat(local_outputs, dim=1)
    
    # 4. 后续层计算
    out = self.linear(concat_out)
    out = self.activation(out)
    return out

关键说明

  • ModuleList的必要性:它会将所有子模型的参数注册到GlobalModel中,确保优化器能更新这些参数,同时支持model.to(device)一键将所有子模型迁移到GPU/CPU。
  • 并行效率:在GPU上,列表推导中的多个model(sub_x)调用会被PyTorch的异步执行引擎自动并行处理,无需手动启动多线程或使用DataParallel(后者适用于单模型多输入的并行,不适用于多模型)。
  • 动态维度适配:如果需要根据子模型输出动态调整线性层维度,可以通过示例输入获取LocalModel的输出维度,再计算总输入维度(n_local_models * local_output_dim),避免硬编码占位值。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.02 21:30:18