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

如何基于TorchSharp实现数据/模型并行训练?梯度同步方案咨询

手动实现多GPU梯度同步(PyTorch/TorchSharp示例)

一、核心原理

数据并行的核心流程:

  • 多GPU维护完全相同的模型副本:初始时将同一模型复制到所有目标GPU
  • 各GPU独立计算:为每个GPU分配独立的训练batch,分别执行前向传播、损失计算、反向传播(生成当前GPU的参数梯度)
  • 梯度合并与同步:收集所有GPU的梯度,计算平均值/总和后,同步到所有模型的参数梯度上
  • 统一参数更新:用合并后的梯度执行一次优化器更新,再将更新后的参数同步到所有GPU的模型副本

二、PyTorch手动实现步骤与示例

步骤清单

  • 初始化模型并复制到所有目标GPU
  • 为每个GPU分配独立的训练batch
  • 各GPU独立完成前向传播、损失计算、反向传播
  • 收集所有GPU的梯度,计算平均值/总和
  • 将合并后的梯度同步到所有模型副本
  • 执行优化器更新,再将更新后的参数同步到所有GPU

代码示例

import torch
import torch.nn as nn
from torch.utils.data import DataLoader, Dataset

# 1. 定义模拟400K参数规模的模型
class SimpleModel(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc1 = nn.Linear(1024, 512)
        self.fc2 = nn.Linear(512, 256)
        self.fc3 = nn.Linear(256, 10)
        self.fc4 = nn.Linear(256, 128)
        self.fc5 = nn.Linear(128, 64)
        self.fc6 = nn.Linear(64, 10)

    def forward(self, x):
        x = torch.relu(self.fc1(x))
        x = torch.relu(self.fc2(x))
        x = torch.relu(self.fc3(x))
        x = torch.relu(self.fc4(x))
        x = torch.relu(self.fc5(x))
        return self.fc6(x)

# 2. 初始化多GPU模型配置
device_ids = [0, 1]
model = SimpleModel()
# 复制模型到两个GPU
models = [model.to(f'cuda:{id}') for id in device_ids]
# 确保所有模型初始参数完全一致
for m in models[1:]:
    m.load_state_dict(models[0].state_dict())

# 3. 绑定优化器(仅需绑定第一个模型副本)
optimizer = torch.optim.Adam(models[0].parameters(), lr=1e-3)

# 4. 模拟数据集(实际替换为你的1.5e9样本数据集)
class MockDataset(Dataset):
    def __len__(self):
        return 10000
    def __getitem__(self, idx):
        return torch.randn(1024), torch.randint(0,10,(1,))

dataloader = DataLoader(MockDataset(), batch_size=64*len(device_ids), shuffle=True)

# 5. 训练循环
for epoch in range(10):
    for batch_idx, (data, target) in enumerate(dataloader):
        optimizer.zero_grad()
        
        # 将当前batch拆分为对应GPU的子batch
        split_data = torch.chunk(data, len(device_ids))
        split_target = torch.chunk(target, len(device_ids))
        
        losses = []
        # 各GPU独立计算梯度
        for i in range(len(device_ids)):
            current_model = models[i]
            x = split_data[i].to(f'cuda:{device_ids[i]}')
            y = split_target[i].to(f'cuda:{device_ids[i]}').squeeze()
            
            output = current_model(x)
            loss = nn.CrossEntropyLoss()(output, y)
            loss.backward()  # 生成当前GPU的参数梯度
            losses.append(loss.item())
        
        # 梯度同步:收集所有GPU梯度,取平均后同步到所有模型
        with torch.no_grad():
            # 遍历所有参数
            for param_idx, param in enumerate(models[0].parameters()):
                # 收集所有GPU的对应参数梯度
                grads = [param.grad]
                for m in models[1:]:
                    grad = list(m.parameters())[param_idx].grad.to(device_ids[0])
                    grads.append(grad)
                # 计算平均梯度
                avg_grad = torch.mean(torch.stack(grads), dim=0)
                # 将平均梯度同步到每个GPU的模型参数
                for m in models:
                    list(m.parameters())[param_idx].grad = avg_grad.to(m.device)
        
        # 执行参数更新,再同步到所有GPU
        optimizer.step()
        for m in models[1:]:
            m.load_state_dict(models[0].state_dict())
        
        if batch_idx % 10 == 0:
            print(f'Epoch {epoch}, Batch {batch_idx}, Avg Loss: {sum(losses)/len(losses):.4f}')

三、TorchSharp对应实现(C#)

核心逻辑与PyTorch完全一致,仅API语法适配:

代码示例片段

using TorchSharp;
using static TorchSharp.torch;
using static TorchSharp.torch.nn;

// 1. 定义模型
public class SimpleModel : Module<Tensor, Tensor>
{
    private readonly Linear fc1, fc2, fc3, fc4, fc5, fc6;

    public SimpleModel() : base("SimpleModel")
    {
        fc1 = Linear(1024, 512);
        fc2 = Linear(512, 256);
        fc3 = Linear(256, 10);
        fc4 = Linear(256, 128);
        fc5 = Linear(128, 64);
        fc6 = Linear(64, 10);
        
        RegisterComponents();
    }

    public override Tensor forward(Tensor x)
    {
        x = relu(fc1.forward(x));
        x = relu(fc2.forward(x));
        x = relu(fc3.forward(x));
        x = relu(fc4.forward(x));
        x = relu(fc5.forward(x));
        return fc6.forward(x);
    }
}

// 2. 初始化多GPU模型
var deviceIds = new[] { 0, 1 };
var model = new SimpleModel();
var models = deviceIds.Select(id => model.to(Device.CUDA(id))).ToList();
// 同步初始参数
foreach (var m in models.Skip(1))
{
    m.load_state_dict(models[0].state_dict());
}

// 3. 初始化优化器
var optimizer = torch.optim.Adam(models[0].parameters(), lr: 1e-3);

// 4. 训练循环核心逻辑
foreach (var epoch in Enumerable.Range(0, 10))
{
    foreach (var (data, target) in dataloader) // 替换为你的实际DataLoader实现
    {
        optimizer.zero_grad();
        
        // 拆分batch到各GPU
        var splitData = data.chunk(deviceIds.Length);
        var splitTarget = target.chunk(deviceIds.Length);
        
        var losses = new List<double>();
        
        // 各GPU独立计算梯度
        for (int i = 0; i < deviceIds.Length; i++)
        {
            var currentModel = models[i];
            var x = splitData[i].to(Device.CUDA(deviceIds[i]));
            var y = splitTarget[i].to(Device.CUDA(deviceIds[i])).squeeze();
            
            var output = currentModel.forward(x);
            var loss = CrossEntropyLoss().forward(output, y);
            loss.backward();
            losses.Add(loss.item<double>());
        }
        
        // 梯度同步
        using (torch.no_grad())
        {
            var paramList = models[0].parameters().ToList();
            for (int paramIdx = 0; paramIdx < paramList.Count; paramIdx++)
            {
                // 收集所有GPU的对应参数梯度
                var grads = new List<Tensor>();
                grads.Add(paramList[paramIdx].grad);
                foreach (var m in models.Skip(1))
                {
                    var grad = m.parameters().ToList()[paramIdx].grad.to(Device.CUDA(deviceIds[0]));
                    grads.Add(grad);
                }
                // 计算平均梯度
                var avgGrad = torch.stack(grads).mean(0);
                // 同步到所有模型
                foreach (var m in models)
                {
                    m.parameters().ToList()[paramIdx].grad = avgGrad.to(m.Device);
                }
            }
        }
        
        // 更新参数并同步到所有GPU
        optimizer.step();
        foreach (var m in models.Skip(1))
        {
            m.load_state_dict(models[0].state_dict());
        }
    }
}

四、针对你的场景的优化建议

  • 数据加载优化:利用高性能CPU做数据预处理与batch拆分,避免GPU等待数据,最大化GPU利用率
  • 梯度合并策略:根据损失计算逻辑选择梯度平均或求和,若用平均需除以GPU数量
  • 减少同步开销:可采用梯度累积(每N个batch同步一次梯度),但需调整学习率适配
  • 性能验证:训练前测试单GPU与双GPU的batch处理速度,排查数据瓶颈或GPU利用率不足问题

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 21:50:31