基于PyTorch的多GPU多时序数据生成器并行训练方法问询
多GPU多生成器并行训练实现方案
需求背景
- 拥有16个时间相关数据生成器,每个生成器固定batch size为1,仅支持逐个生成数据
- 配备8块GPU,需按“每2个生成器对应1块GPU”的方式分配,以多进程方式并行训练
- 核心逻辑:单GPU内合并对应两个生成器的batch数据,计算损失;汇总所有GPU损失后执行反向传播
实现代码
import torch import torch.multiprocessing as mp def train_on_gpu(gpu_id, dataloaders, model, criterion, losses_queue): device = torch.device(f'cuda:{gpu_id}') model = model.to(device) # 配对遍历两个生成器的数据 for (input1, target1), (input2, target2) in zip(*dataloaders): # 合并两个batch并转移到对应GPU x = torch.cat([input1, input2], dim=0).to(device) target = torch.cat([target1, target2], dim=0).to(device) # 前向传播计算损失 out = model(x) loss = criterion(out, target) # 将损失存入进程安全队列 losses_queue.put((gpu_id, loss.item())) def main(): # 假设已初始化16个数据生成器:dataloader1 ~ dataloader16 dataloaders = [dataloader1, dataloader2, dataloader3, dataloader4, dataloader5, dataloader6, dataloader7, dataloader8, dataloader9, dataloader10, dataloader11, dataloader12, dataloader13, dataloader14, dataloader15, dataloader16] # 按GPU分配生成器:每2个生成器对应1块GPU dataloader_pairs = {i: (dataloaders[2*i], dataloaders[2*i+1]) for i in range(8)} # 初始化模型与损失函数 model = YourModel() # 替换为你的模型类 criterion = YourLoss() # 替换为你的损失函数 # 进程安全队列,用于收集各GPU的损失值 losses_queue = mp.Queue() processes = [] # 启动每个GPU对应的训练进程 for gpu_id in range(8): proc = mp.Process( target=train_on_gpu, args=(gpu_id, dataloader_pairs[gpu_id], model, criterion, losses_queue) ) proc.start() processes.append(proc) # 等待所有进程完成当前轮次训练 for proc in processes: proc.join() # 汇总所有GPU的损失 losses = {gpu: 0.0 for gpu in range(8)} while not losses_queue.empty(): gpu_id, loss_val = losses_queue.get() losses[gpu_id] += loss_val # 计算总损失并反向传播 total_loss = torch.tensor(sum(losses.values()) / 8, requires_grad=True) total_loss.backward() # 此处可添加优化器更新、参数同步等后续步骤 # optimizer.step() if __name__ == '__main__': mp.set_start_method('spawn') # PyTorch多进程推荐使用spawn启动方式 main()
关键注意事项
- 生成器进程兼容性:若生成器涉及文件读取或外部资源,需确保资源在子进程中可正常访问,必要时可在子进程内重新初始化生成器
- 模型参数同步:若需多轮训练,需在每轮训练后同步各进程的模型参数,可结合
torch.nn.DistributedDataParallel简化同步逻辑 - 显存控制:每块GPU处理合并后batch size为2的数据,需根据模型体量调整,避免显存溢出
- 队列安全:使用
multiprocessing.Queue确保进程间损失数据的安全传递,避免竞态问题
内容的提问来源于stack exchange,提问作者Tommy Yu
相关产品推荐
相关产品推荐

