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

PyTorch Lightning双GPU DDP模式下WandbLogger仅记录2张预测图的问题

解决DDP双GPU下Wandb仅记录部分预测图像的问题

问题根源

在DDP分布式模式中,每个GPU对应独立的运行进程。PyTorch Lightning的WandbLogger默认仅让**主进程(rank=0)**的日志上传至Wandb,非主进程的log_image调用不会实际生效。你的双GPU各自处理2个样本,非主进程的2组预测结果对应的日志完全没被上传,因此最终只看到主进程的部分结果。

解决方案

推荐两种可靠的处理方式,优先选第二种,逻辑更清晰:

方案1:在Batch级别收集所有进程结果后由主进程记录

修改on_predict_batch_end方法,先通过分布式通信收集所有GPU的当前Batch结果,再仅让主进程执行日志记录:

def on_predict_batch_end(self, outputs, batch, batch_idx):
    # 收集所有GPU进程的当前Batch输出
    gathered_outputs = self.all_gather(outputs)
    
    # 仅主进程执行日志上传
    if self.trainer.is_global_zero:
        # 遍历所有进程的结果
        for gt, pred, axis, idx_slice in gathered_outputs:
            gt = gt.squeeze().cpu().numpy()
            pred = pred.squeeze().cpu().numpy()
            if axis != 0:
                gt = np.rot90(gt, 2)
                gt = np.flip(gt, 1)
                pred = np.rot90(pred, 2)
                pred = np.flip(pred, 1)
            self.logger.log_image(
                key="eval_output",
                images=[gt, pred],
                caption=[
                    f"gt_data_{axis}_{idx_slice}",
                    f"output_model_{axis}_{idx_slice}",
                ],
            )

方案2:在预测Epoch结束时统一收集所有结果并记录

这种方式避免在Batch级别频繁做分布式通信,代码结构更清晰:

  1. 确保predict_step返回所需的预测结果:
def predict_step(self, batch, batch_idx):
    # 此处编写你的预测逻辑,生成gt, pred, axis, idx_slice
    # 若不需要梯度可添加 .detach()
    return gt, pred, axis, idx_slice
  1. 实现predict_epoch_end方法,统一收集所有进程的结果并由主进程记录:
def predict_epoch_end(self, predict_outputs):
    # 收集所有GPU进程的所有Batch结果
    all_results = self.all_gather(predict_outputs)
    
    if self.trainer.is_global_zero:
        # 展开嵌套的结果结构(适配多进程+多Batch的输出)
        for batch_results in all_results:
            for result in batch_results:
                gt, pred, axis, idx_slice = result
                gt = gt.squeeze().cpu().numpy()
                pred = pred.squeeze().cpu().numpy()
                if axis != 0:
                    gt = np.rot90(gt, 2)
                    gt = np.flip(gt, 1)
                    pred = np.rot90(pred, 2)
                    pred = np.flip(pred, 1)
                self.logger.log_image(
                    key="eval_output",
                    images=[gt, pred],
                    caption=[
                        f"gt_data_{axis}_{idx_slice}",
                        f"output_model_{axis}_{idx_slice}",
                    ],
                )

关键细节

  • self.trainer.is_global_zero:PyTorch Lightning提供的便捷属性,用于判断当前进程是否为主进程,等价于torch.distributed.get_rank() == 0。
  • self.all_gather():LightningModule内置的分布式通信方法,能将所有进程的张量结果同步到主进程,确保主进程拿到所有GPU的预测数据。
  • 禁止非主进程调用Wandb日志:非主进程的WandbLogger默认处于禁用状态,调用日志方法不会上传任何内容,反而会浪费资源。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 18:57:45