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

PyTorch多GPU推理问题:生成图像缺失与GPU负载不均

PyTorch/PyTorch-Lightning多GPU推理图像缺失问题排查

问题描述

我用PyTorch和PyTorch-Lightning做多GPU推理时遇到了问题:推理需要以自回归方式调用两个不同模型,由于自回归步骤计算成本高,我把数据集拆分后分配到多个GPU并行运行来提速。但因为自回归逻辑复杂,没法用标准的PyTorch-Lightning test_step。计算完成后导出生成的图像批次时,发现部分图像缺失。

具体问题:保存的图像里有缺失,对应部分批次的生成结果没被保存。

import os
import hydra
import torch
import torch.distributed as dist
import torch.multiprocessing as mp
from hydra.utils import instantiate
from omegaconf import DictConfig


def ddp_setup(rank: int, world_size: int):
    os.environ["MASTER_ADDR"] = "localhost"
    os.environ["MASTER_PORT"] = "47144"
    dist.init_process_group(backend="nccl", rank=rank, world_size=world_size)
    torch.cuda.set_device(rank)


@hydra.main(version_base=None, config_path="../../configs", config_name="transfert")
def main(cfg: DictConfig) -> None:
    world_size = torch.cuda.device_count()
    mp.spawn(run_da, args=(world_size, cfg), nprocs=world_size)


def run_da(rank: int, world_size: int, cfg: DictConfig) -> None:
    ddp_setup(rank, world_size)

    # define models map location
    map_location = {"cuda:%d" % 0: "cuda:%d" % rank}

    # get eps model of source domain
    source_model = Model.load_from_checkpoint(
        cfg.source_model_path,
        map_location=map_location,
    )
    source_model.eval()
    source_model.to(rank)

    target_model = Model.load_from_checkpoint(
        cfg.target_model_path,
        map_location=map_location,
    )
    target_model.eval()
    target_model.to(rank)

    # load source and target datamodules
    source_datamodule = instantiate(cfg.source)
    source_datamodule.setup()

    test_data_sampler = torch.utils.data.distributed.DistributedSampler(
        source_datamodule.test,
        num_replicas=world_size,
        rank=rank,
    )
    dataloaders = {
        "val": source_datamodule.val_dataloader(sampler=test_data_sampler),
        "test": source_datamodule.test_dataloader(sampler=test_data_sampler),
    }

    for stage, dataloader in dataloaders.items():
        for batch in dataloader:
            x = batch["img"].to(rank)
            cond = batch["cond"].to(rank)

            with torch.no_grad():
                # compute output with several calls with source and target model
                for _ in range(100):
                    x = autoregressive_step(source_model, target_model, x, cond)

            # save output as images
            export_output(x, stage)


if __name__ == "__main__":
    main()

假设

部分GPU接收的批次没被处理/保存,监控显示不同GPU接收的批次数量差异很大。

我参考过多GPU训练教程,但推理场景差别很大(我理解Distributed Data Parallel不适用),想请教代码里有没有明显bug,以及有没有PyTorch/PyTorch-Lightning多GPU推理的优质资源。


代码问题排查

1. DistributedSampler使用错误

  • 你给val_dataloader和test_dataloader共用了同一个针对测试集创建的test_data_sampler,这会导致验证集的数据拆分逻辑完全混乱,不同GPU拿到的验证集批次数量必然异常。正确做法是给验证集单独创建DistributedSampler:
    val_data_sampler = torch.utils.data.distributed.DistributedSampler(
        source_datamodule.val,
        num_replicas=world_size,
        rank=rank,
    )
    dataloaders = {
        "val": source_datamodule.val_dataloader(sampler=val_data_sampler),
        "test": source_datamodule.test_dataloader(sampler=test_data_sampler),
    }
    
  • 遍历dataloader前必须调用sampler.set_epoch(0)(推理时固定epoch为0即可),否则不同GPU可能出现重复采样或漏采数据的情况。

2. 模型加载的潜在问题

  • 用map_location映射设备的方式容易遗漏部分参数,建议改用PyTorch-Lightning支持的device参数直接指定加载到目标GPU:
    source_model = Model.load_from_checkpoint(
        cfg.source_model_path,
        device=f"cuda:{rank}"
    )
    
  • 确认autoregressive_step函数内没有意外将模型切换回训练模式的代码(比如调用了model.train())。

3. 图像导出的冲突问题

  • 如果export_output函数直接向同一目录写入文件,多GPU并行写入会出现文件覆盖或写入失败的情况。必须给每个GPU的输出文件加上rank标识,比如文件名后缀_rank{rank},或者给每个GPU分配单独的子目录。
  • 导出前要先将张量从GPU移到CPU:x = x.cpu(),避免张量在GPU上导致的保存异常。

4. DDP进程清理问题

  • 在run_da函数末尾要调用dist.destroy_process_group(),避免进程残留导致后续运行异常。

多GPU推理资源建议

  • PyTorch官方分布式推理文档:重点关注DistributedDataParallel的推理用法(即使是推理场景,DDP依然适用,只要将模型设为eval()模式并关闭梯度),以及DistributedSampler的规范使用。
  • PyTorch-Lightning多GPU推理指南:可以参考Trainer.predict()方法的自定义逻辑,即使不用test_step,也可以借助Trainer管理分布式进程,减少手动编写DDP的bug。
  • 可以用torch.nn.parallel.DistributedDataParallel包裹模型,推理时关闭梯度,利用DDP的进程管理能力简化分布式逻辑。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 15:02:10