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

如何在Huggingface Trainer中配置跨批次MNR损失及多GPU DDP通信

在Transformers Trainer的DDP多GPU环境下实现跨批次MNR损失(单损失计算)

要实现每个GPU独立执行前向传播,收集所有输出后仅计算一次MNR损失的需求,核心是重写Trainer的compute_loss方法,结合PyTorch DDP的通信API完成跨GPU张量收集,最后统一计算损失。以下是具体实现方案:

核心思路

  1. 每个GPU单独执行模型前向,得到当前批次的输出(如query/doc嵌入)
  2. 用DDP的all_gather收集所有GPU的输出及各自的批次大小(处理最后一批次样本数不一致的情况)
  3. 仅在主GPU上计算全局MNR损失,再通过all_reduce同步损失值到所有GPU,保证梯度计算一致性

代码实现

1. 定义MNR损失类

import torch
from torch import nn
from transformers import Trainer, TrainingArguments

class MNRLoss(nn.Module):
    def __init__(self, temperature=0.07):
        super().__init__()
        self.temperature = temperature
        self.cross_entropy = nn.CrossEntropyLoss()

    def forward(self, global_query_embeds, global_doc_embeds):
        # 计算全局query与doc的相似度矩阵
        logits = torch.matmul(global_query_embeds, global_doc_embeds.T) / self.temperature
        # 标签为每个query对应自身的doc(假设query和doc是一一配对的批次)
        labels = torch.arange(logits.size(0), device=logits.device)
        return self.cross_entropy(logits, labels)

2. 重写Trainer类

class CustomMNRTrainer(Trainer):
    def __init__(self, mnr_loss, **kwargs):
        super().__init__(**kwargs)
        self.mnr_loss = mnr_loss

    def compute_loss(self, model, inputs, return_outputs=False):
        # 步骤1:每个GPU独立执行前向传播,获取当前批次的嵌入
        outputs = model(**inputs)
        # 示例:取<CLS>作为query嵌入,第二个token作为doc嵌入(根据你的任务调整)
        query_embeds = outputs.last_hidden_state[:, 0, :]
        doc_embeds = outputs.last_hidden_state[:, 1, :]
        device = query_embeds.device

        # 步骤2:跨GPU收集所有批次的大小和嵌入
        # 先收集每个GPU的批次大小
        local_batch_size = torch.tensor(query_embeds.shape[0], device=device)
        all_batch_sizes = [torch.zeros_like(local_batch_size) for _ in range(self.args.world_size)]
        torch.distributed.all_gather(all_batch_sizes, local_batch_size)
        all_batch_sizes = torch.cat(all_batch_sizes).cpu().tolist()

        # 收集所有GPU的query和doc嵌入
        all_query_embeds = [torch.zeros_like(query_embeds) for _ in range(self.args.world_size)]
        all_doc_embeds = [torch.zeros_like(doc_embeds) for _ in range(self.args.world_size)]
        torch.distributed.all_gather(all_query_embeds, query_embeds)
        torch.distributed.all_gather(all_doc_embeds, doc_embeds)

        # 拼接成全局嵌入(处理最后一批次样本数不一致的情况)
        global_query_embeds = []
        global_doc_embeds = []
        for q_emb, d_emb, bs in zip(all_query_embeds, all_doc_embeds, all_batch_sizes):
            global_query_embeds.append(q_emb[:bs])
            global_doc_embeds.append(d_emb[:bs])
        global_query_embeds = torch.cat(global_query_embeds, dim=0)
        global_doc_embeds = torch.cat(global_doc_embeds, dim=0)

        # 步骤3:仅主GPU计算损失,其他GPU初始化为0
        loss = torch.tensor(0.0, device=device)
        if self.args.local_rank == 0:
            loss = self.mnr_loss(global_query_embeds, global_doc_embeds)

        # 步骤4:同步损失到所有GPU,保证梯度计算一致
        torch.distributed.all_reduce(loss, op=torch.distributed.ReduceOp.SUM)
        loss = loss / self.args.world_size  # 平均损失,避免总损失被放大

        return (loss, outputs) if return_outputs else loss

3. 训练调用示例

# 初始化训练参数(开启DDP)
training_args = TrainingArguments(
    output_dir="./mnr_model",
    per_device_train_batch_size=16,
    num_train_epochs=3,
    logging_dir="./logs",
    logging_steps=10,
    do_train=True,
    # DDP相关设置
    use_cpu=False,
    local_rank=int(os.environ.get("LOCAL_RANK", 0)),
    ddp_find_unused_parameters=False,  # 根据模型结构调整,默认False
)

# 初始化损失和自定义Trainer
mnr_loss = MNRLoss(temperature=0.07)
trainer = CustomMNRTrainer(
    mnr_loss=mnr_loss,
    model=your_model,  # 替换为你的预训练模型
    args=training_args,
    train_dataset=train_dataset,  # 替换为你的训练数据集
)

# 启动训练
trainer.train()

关键注意事项

  • DDP初始化:启动训练时需通过torchrun或python -m torch.distributed.launch开启DDP,确保LOCAL_RANK环境变量正确设置
  • 张量形状一致性:all_gather要求所有GPU传递的张量形状一致,因此代码中通过收集批次大小来截取有效部分,避免最后一批次的形状不匹配问题
  • 损失同步:必须用all_reduce同步损失值,否则非主GPU的梯度会为0,导致训练异常

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.20 22:25:27