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

PyTorch Lightning多GPU训练时,如何在CPU上正确计算验证指标?

PyTorch Lightning多GPU训练:验证步骤CPU计算指标的正确姿势

问题场景

在PyTorch Lightning多GPU训练时,验证步骤中仅单块GPU执行验证操作,无法在CPU上正确计算全局指标,推测和cpu()的使用方式有关。原代码片段如下:

def validation_step(self, batch, batch_idx):
    # batch size = 1
    assert batch["image"].shape[0] == 1

    labels, trimaps, dataset_name, image_name = batch["alpha"], batch["trimap"], batch["dataset_name"], batch["image_name"]

    output, guidance_map = self.shared_step(batch, batch_idx)
    
    label = labels.squeeze().cpu().numpy()*255.0
    trimap = trimaps.squeeze().cpu().numpy()*128
    pred = output.squeeze().cpu().numpy()*255.0
    
    metrics_unknown, metrics_all = self.compute_four_metrics(
        pred, label, trimap)

问题根源

不是cpu()函数本身的问题,而是多GPU(如DDP)模式下,每个GPU进程会独立执行validation_step,原代码中每个进程单独转numpy计算指标,但没有做进程间的数据同步或结果聚合,最终只有单个进程的结果被保留,导致看起来只有一块GPU在工作。

解决方案

核心思路

  1. 优先用PyTorch张量完成指标计算(无需转CPU),再通过Lightning的self.log()自动同步结果
  2. 若必须转CPU计算,需先收集所有GPU进程的张量数据,再在主进程统一处理,避免重复计算

修改后的代码示例

def validation_step(self, batch, batch_idx):
    assert batch["image"].shape[0] == 1

    labels, trimaps = batch["alpha"], batch["trimap"]
    output, guidance_map = self.shared_step(batch, batch_idx)
    
    # 先保持张量状态,避免过早转CPU
    label = labels.squeeze() * 255.0
    trimap = trimaps.squeeze() * 128
    pred = output.squeeze() * 255.0
    
    # 收集所有GPU进程的结果,确保拿到全局验证数据
    pred_all = self.all_gather(pred)
    label_all = self.all_gather(label)
    trimap_all = self.all_gather(trimap)
    
    # 仅在主进程执行CPU转换和指标计算,避免重复工作
    if self.trainer.is_global_zero:
        pred_np = pred_all.cpu().numpy()
        label_np = label_all.cpu().numpy()
        trimap_np = trimap_all.cpu().numpy()
        metrics_unknown, metrics_all = self.compute_four_metrics(pred_np, label_np, trimap_np)
        
        # 同步并记录指标,sync_dist=True确保多GPU结果聚合
        self.log("val/metrics_unknown", metrics_unknown, sync_dist=True)
        self.log("val/metrics_all", metrics_all, sync_dist=True)
    else:
        # 非主进程需占位log,避免Lightning抛出进程同步异常
        self.log("val/metrics_unknown", torch.tensor(0.0, device=pred.device), sync_dist=True)
        self.log("val/metrics_all", torch.tensor(0.0, device=pred.device), sync_dist=True)

关键注意点

  • self.all_gather():用于收集所有GPU进程的张量,将分散在不同GPU的数据聚合为一个全局张量
  • self.trainer.is_global_zero:判断当前进程是否为主进程,仅在主进程处理CPU转换和指标计算,减少冗余操作
  • self.log()的sync_dist=True:自动完成多GPU指标的同步与聚合,确保最终日志显示全局计算结果
  • 如果你的compute_four_metrics支持直接处理PyTorch张量,完全可以跳过转CPU步骤,直接在GPU上计算后调用self.log,效率更高

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 19:48:39