已关闭
sync_loss 仅按 cp_world_size 平均,梯度累积场景下日志损失值与实际单步损失不一致 #73
崇理战队创建于  8月13日关闭于  8 天前
崇理战队
8月13日 创建

描述

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 时:

  1. 4 个 micro-batch 的 loss 被追加到 _accumulated_losses["loss"]
  2. sync_loss 对这 4 个损失执行 torch.stack().sum(),得到 4 个 micro-batch 损失的总和
  3. 在 CP 场景下,除以 cp_world_size 得到 CP 平均后的总和
  4. 该值被传递给 _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 倍,影响用户对训练进展的判断。

likedislike
崇崇理战队
8月13日 添加了label:bug
Keilo_W成员
9 天前 评论:

trainer.py:195 每个 micro-batch 的 loss 已在 train_step 中除以 gradient_accumulation_steps , sync_loss 求和后得到的正是 单 micro-batch 平均 loss 。

likedislike
KKeilo_W成员
8 天前 关闭了 issue
KKeilo_W成员
8 天前 issue状态由 TODO 改变为 DONE
ascend-robotascend-robot成员
8 天前 添加了label:resolved