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

基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 15:54:54