已关闭
sync_loss 仅按 cp_world_size 平均,梯度累积场景下日志损失值与实际单步损失不一致 #73
崇理战队创建于 8月13日关闭于 8 天前
Keilo_W
9 天前 评论:
9 天前 评论:
trainer.py:195 每个 micro-batch 的 loss 已在 train_step 中除以 gradient_accumulation_steps , sync_loss 求和后得到的正是 单 micro-batch 平均 loss 。


8 天前 添加了label:resolved
描述
BaseTrainer.sync_loss方法在 Context Parallel 启用时,将累积的损失按cp_world_size平均后返回用于日志记录。但在梯度累积(gradient accumulation)场景下,返回的损失值是多个 micro-batch 的累积和除以cp_world_size,而非用户通常期望的「单 micro-batch 平均损失」,导致日志中的损失值与实际单步损失存在gradient_accumulation_steps倍数的差异。问题代码
# 文件: fsdp_turbo/training/trainer.py (第 197-241 行) @classmethod def sync_loss(cls, accumulated_losses: dict) -> dict: from fsdp_turbo.distributed.parallel_state import get_parallel_state parallel_state = get_parallel_state() cp_enabled = parallel_state.is_cp_enable() if cp_enabled: cp_group = parallel_state.get_cp_group() cp_world_size = parallel_state.get_cp_group_size() result = {} for name, losses in accumulated_losses.items(): if not losses: result[name] = None continue summed = torch.stack(losses).sum() # ← 累积所有 micro-batch 的损失 if cp_enabled: synced = summed.clone() torch.distributed.all_reduce(synced, op=torch.distributed.ReduceOp.SUM, group=cp_group) synced.div_(cp_world_size) result[name] = synced else: result[name] = summed return result分析
在训练循环中,
sync_loss的调用时机如下:# trainer.py 训练循环 (第 102-129 行) for step, batch in enumerate(self.dataloader): loss, aux_loss = self.train_step(batch) self._accumulated_losses["loss"].append(loss.detach()) # ... if (step + 1) % self.config.run.gradient_accumulation_steps == 0: synced = self.sync_loss(self._accumulated_losses) loss = synced["loss"] # ... self._on_log_step(loss, grad_norm, aux_loss)当
gradient_accumulation_steps = 4时:loss被追加到_accumulated_losses["loss"]sync_loss对这 4 个损失执行torch.stack().sum(),得到 4 个 micro-batch 损失的总和cp_world_size得到 CP 平均后的总和_on_log_step用于日志输出用户通常期望日志中的
loss表示「单个 micro-batch 的平均损失」,但实际得到的是「gradient_accumulation_steps个 micro-batch 的累积损失除以cp_world_size」。这导致:
gradient_accumulation_steps = 1时,日志值正确gradient_accumulation_steps > 1时,日志值是实际单步损失的gradient_accumulation_steps倍gradient_accumulation_steps时会发现日志损失值不成比例地变化,产生困惑建议修复
在
sync_loss中除以gradient_accumulation_steps,使返回值为单 micro-batch 的平均损失:@classmethod def sync_loss(cls, accumulated_losses: dict, gradient_accumulation_steps: int = 1) -> dict: # ... for name, losses in accumulated_losses.items(): if not losses: result[name] = None continue summed = torch.stack(losses).sum() if cp_enabled: synced = summed.clone() torch.distributed.all_reduce(synced, op=torch.distributed.ReduceOp.SUM, group=cp_group) synced.div_(cp_world_size) result[name] = synced / gradient_accumulation_steps # ← 新增 else: result[name] = summed / gradient_accumulation_steps # ← 新增 return result同时在调用处传入
gradient_accumulation_steps:synced = self.sync_loss(self._accumulated_losses, self.config.run.gradient_accumulation_steps)影响范围
影响所有使用梯度累积(
gradient_accumulation_steps > 1)且启用 Context Parallel 的训练任务。日志中的损失值将是实际值的gradient_accumulation_steps倍,影响用户对训练进展的判断。