如何基于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
相关产品推荐
相关产品推荐

