如何在Huggingface Trainer中配置跨批次MNR损失及多GPU DDP通信
在Transformers Trainer的DDP多GPU环境下实现跨批次MNR损失(单损失计算)
要实现每个GPU独立执行前向传播,收集所有输出后仅计算一次MNR损失的需求,核心是重写Trainer的compute_loss方法,结合PyTorch DDP的通信API完成跨GPU张量收集,最后统一计算损失。以下是具体实现方案:
核心思路
- 每个GPU单独执行模型前向,得到当前批次的输出(如query/doc嵌入)
- 用DDP的
all_gather收集所有GPU的输出及各自的批次大小(处理最后一批次样本数不一致的情况) - 仅在主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
相关产品推荐
相关产品推荐

