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

Torch:单GPU并行运行多同架构异权重模型的实现方案问询

PyTorch下多同架构异权重小模型并行训练方案

针对你提出的两个需求,以下是基于PyTorch的具体实现方案:

一、多个模型在同一minibatch数据上并行运行

这种场景下,我们可以利用PyTorch的torch.vmap(向量化映射)工具,或者手动合并模型参数,实现批量前向/反向传播,避免循环调用单个模型带来的开销。

方案1:使用torch.vmap(推荐,简洁高效)

vmap可以自动将单模型的运算映射到批量维度,完美适配多同架构模型的并行计算场景。

示例代码:

import torch
import torch.nn as nn
from torch.func import vmap, grad

# 定义基础模型架构
class SmallModel(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc1 = nn.Linear(10, 20)
        self.fc2 = nn.Linear(20, 2)
    
    def forward(self, x):
        x = torch.relu(self.fc1(x))
        return self.fc2(x)

# 创建N个不同权重的模型实例
num_models = 8
models = [SmallModel().to('cuda') for _ in range(num_models)]

# 提取所有模型的参数,整理成带批量维度的张量(num_models, ...)
params = [torch.stack([getattr(m, name) for m in models]) for name, _ in models[0].named_parameters()]

# 定义单模型的前向计算逻辑
def model_forward(params, x):
    temp_model = SmallModel()
    for param, p in zip(temp_model.parameters(), params):
        param.data = p
    return temp_model(x)

# 用vmap包装实现多模型并行计算
batch_x = torch.randn(32, 10).to('cuda')  # 共享的minibatch数据
parallel_forward = vmap(model_forward, in_dims=(0, None))  # 参数按第0维并行,数据共享
outputs = parallel_forward(params, batch_x)
# outputs形状:(num_models, 32, 2),对应每个模型在该batch上的输出

# 批量反向传播示例
labels = torch.randint(0, 2, (32,)).to('cuda')
def loss_fn(params, x, y):
    pred = model_forward(params, x)
    return nn.CrossEntropyLoss()(pred, y)

# 计算所有模型的损失和梯度
parallel_loss = vmap(loss_fn, in_dims=(0, None, None))(params, batch_x, labels)
parallel_grad = vmap(grad(loss_fn), in_dims=(0, None, None))(params, batch_x, labels)

# 更新每个模型的参数
for i, model in enumerate(models):
    for param, g in zip(model.parameters(), parallel_grad):
        param.data -= 1e-3 * g[i]

方案2:手动合并参数与数据堆叠

如果不想依赖torch.func,可以手动合并模型参数为带批量维度的张量,同时复制数据到对应维度,通过矩阵运算实现并行:

# 延续上述SmallModel定义
num_models = 8
models = [SmallModel().to('cuda') for _ in range(num_models)]

# 合并参数:例如fc1.weight形状变为(num_models, 20, 10)
fc1_weights = torch.stack([m.fc1.weight for m in models])
fc1_biases = torch.stack([m.fc1.bias for m in models])
fc2_weights = torch.stack([m.fc2.weight for m in models])
fc2_biases = torch.stack([m.fc2.bias for m in models])

batch_x = torch.randn(32, 10).to('cuda')
# 复制数据到模型维度:(num_models, 32, 10)
x_repeated = batch_x.unsqueeze(0).repeat(num_models, 1, 1)

# 并行前向计算
h = torch.relu(torch.bmm(x_repeated, fc1_weights.transpose(1,2)) + fc1_biases.unsqueeze(1))
outputs = torch.bmm(h, fc2_weights.transpose(1,2)) + fc2_biases.unsqueeze(1)
# outputs形状:(num_models, 32, 2)

# 反向传播后拆分参数更新即可

二、多个模型在不同minibatch数据上并行运行

这种场景下,我们可以将数据拆分为对应模型的小batch,结合vmap实现并行计算,充分利用GPU显存和算力。

方案:使用torch.vmap处理多数据batch

import torch
import torch.nn as nn
from torch.func import vmap, grad

class SmallModel(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc1 = nn.Linear(10, 20)
        self.fc2 = nn.Linear(20, 2)
    
    def forward(self, x):
        x = torch.relu(self.fc1(x))
        return self.fc2(x)

num_models = 8
models = [SmallModel().to('cuda') for _ in range(num_models)]
params = [torch.stack([getattr(m, name) for m in models]) for name, _ in models[0].named_parameters()]

# 准备多个不同的minibatch:形状(num_models, batch_size_per_model, 10)
batch_size_per_model = 4
multi_batch_x = torch.randn(num_models, batch_size_per_model, 10).to('cuda')
multi_batch_labels = torch.randint(0, 2, (num_models, batch_size_per_model)).to('cuda')

# 定义单模型单数据batch的损失计算逻辑
def model_loss(params, x, y):
    temp_model = SmallModel()
    for param, p in zip(temp_model.parameters(), params):
        param.data = p
    pred = temp_model(x)
    return nn.CrossEntropyLoss()(pred, y)

# 并行计算所有模型的损失和梯度
parallel_loss = vmap(model_loss)(params, multi_batch_x, multi_batch_labels)
parallel_grad = vmap(grad(model_loss))(params, multi_batch_x, multi_batch_labels)

# 更新每个模型参数
for i, model in enumerate(models):
    for param, g in zip(model.parameters(), parallel_grad):
        param.data -= 1e-3 * g[i]

替代方案:DataLoader拆分+批量运算

如果数据从DataLoader加载,可以一次性取出num_models * batch_size_per_model条数据,拆分为num_models个小batch后,按上述方式并行计算,避免循环加载数据的开销。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.24 18:37:15