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

如何在多GPU环境下基于PyTorch Lightning计算检索类指标?

多GPU环境下PyTorch Lightning计算多验证集检索类指标的解决方案

针对多GPU(2台及以上)、多验证数据集场景,计算torchmetrics.RetrievalNormalizedDCG这类依赖全局组数据的检索指标,核心是跨设备聚合所有批次的预测、标签和组数据后,再统一计算指标,以下是可落地的实现方案:

实现步骤与代码示例

1. 初始化数据存储结构

在LightningModule中,用字典存储不同验证数据集的中间数据,每个数据集对应三个列表(预测值、标签、组ID):

import torch
import pytorch_lightning as pl
from torchmetrics.retrieval import RetrievalNormalizedDCG

class RetrievalModel(pl.LightningModule):
    def __init__(self):
        super().__init__()
        # 初始化检索指标(无需在step中更新,最后统一计算)
        self.ndcg = RetrievalNormalizedDCG(k=10)
        # 存储多验证集的中间数据:key为数据集索引/名称,value为(preds, targets, groups)列表
        self.val_data = {}

    def setup(self, stage=None):
        if stage == "validate":
            # 根据验证数据集数量初始化存储结构
            num_val_datasets = len(self.trainer.datamodule.val_dataloaders())
            for idx in range(num_val_datasets):
                # 可替换为数据集名称,比如self.trainer.datamodule.val_datasets[idx].name
                self.val_data[idx] = {"preds": [], "targets": [], "groups": []}

2. 验证步骤收集批次数据

在validation_step中,将当前批次的预测、标签、组数据移至CPU并添加到对应数据集的存储中,通过dataloader_idx参数区分不同验证集:

def validation_step(self, batch, batch_idx, dataloader_idx=0):
        preds, targets, groups = batch  # 假设你的batch返回这三个字段
        # 将张量移至CPU,避免GPU内存占用
        self.val_data[dataloader_idx]["preds"].append(preds.detach().cpu())
        self.val_data[dataloader_idx]["targets"].append(targets.detach().cpu())
        self.val_data[dataloader_idx]["groups"].append(groups.detach().cpu())

3. 验证epoch结束时跨设备聚合并计算指标

重写validation_epoch_end,使用PyTorch分布式API聚合所有设备的数据,仅在主设备计算并记录指标:

def validation_epoch_end(self, outputs):
        # 遍历每个验证数据集
        for dataloader_idx, data in self.val_data.items():
            # 拼接当前设备的所有批次数据
            preds = torch.cat(data["preds"], dim=0)
            targets = torch.cat(data["targets"], dim=0)
            groups = torch.cat(data["groups"], dim=0)

            # 跨设备聚合数据:all_gather收集所有设备的张量
            if self.trainer.world_size > 1:
                preds = torch.cat(torch.distributed.all_gather(preds), dim=0)
                targets = torch.cat(torch.distributed.all_gather(targets), dim=0)
                groups = torch.cat(torch.distributed.all_gather(groups), dim=0)

            # 仅在主设备计算指标(避免多设备重复计算)
            if self.trainer.is_global_zero:
                # 计算RetrievalNDCG,必须传入groups参数指定检索组
                ndcg_score = self.ndcg(preds, targets, groups=groups)
                # 记录指标,区分不同验证集
                self.log(f"val/{dataloader_idx}/ndcg", ndcg_score, sync_dist=False)
                print(f"Validation Dataset {dataloader_idx} NDCG@10: {ndcg_score.item()}")

        # 重置存储结构,避免下一个epoch数据污染
        self._reset_val_data()

    def _reset_val_data(self):
        for key in self.val_data.keys():
            self.val_data[key]["preds"] = []
            self.val_data[key]["targets"] = []
            self.val_data[key]["groups"] = []

关键细节说明

  • 跨设备聚合:使用torch.distributed.all_gather确保所有设备的批次数据被收集到主设备,拼接后得到全局的检索组数据,解决组跨批次、跨设备拆分的问题。
  • 多验证集区分:通过dataloader_idx参数识别当前处理的验证数据集,将数据存储到对应字典条目,实现多验证集的独立计算。
  • 内存优化:将批次数据移至CPU存储,避免GPU内存随着验证批次累积而溢出;epoch结束后重置存储结构,避免内存泄漏。
  • 指标计算注意点:RetrievalNormalizedDCG等检索类指标必须传入groups参数,否则无法正确划分检索任务组,导致指标计算错误。

内容的提问来源于stack exchange,提问作者Ford O.

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.19 16:53:27