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

修改AdamW将优化器状态移至CPU后训练Loss波动异常排查

优化器状态移存导致Loss波动的问题分析

问题背景

拥有8张80G NVIDIA GPU,训练70B Llama模型时显存不足,无法同时加载模型与优化器状态。因此修改transformers.AdamW的step方法,将优化器状态(exp_avg、exp_avg_sq)平时存储在CPU,仅在参数更新时移至GPU,更新完成后再移回CPU,修改后的代码如下:

@torch.no_grad()
def step(self, closure: Callable = None):
    """
    Performs a single optimization step.

    Arguments:
        closure (`Callable`, *optional*): A closure that reevaluates the model and returns the loss.
    """
    loss = None
    if closure is not None:
        loss = closure()

    for group in self.param_groups:
        for p in group["params"]:
            if p.grad is None:
                continue
            grad = p.grad
            if grad.is_sparse:
                raise RuntimeError("Adam does not support sparse gradients, please consider SparseAdam instead")

            state = self.state[p]

            # State initialization
            if len(state) == 0:
                state["step"] = 0
                # Exponential moving average of gradient values
                state["exp_avg"] = torch.zeros_like(p)
                # Exponential moving average of squared gradient values
                state["exp_avg_sq"] = torch.zeros_like(p)

            exp_avg, exp_avg_sq = state["exp_avg"], state["exp_avg_sq"]
            #modified1: only move the states corresponde to the updating parameters to gpu
            if exp_avg.device != p.device:
                exp_avg = exp_avg.to(p.device)
                exp_avg = exp_avg.to(p.dtype)
            if exp_avg_sq.device != p.device:
                exp_avg_sq = exp_avg_sq.to(p.device)
                exp_avg_sq = exp_avg_sq.to(p.dtype)
            beta1, beta2 = group["betas"]

            state["step"] += 1

            # Decay the first and second moment running average coefficient
            # In-place operations to update the averages at the same time
            exp_avg.mul_(beta1).add_(grad, alpha=(1.0 - beta1))
            exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)
            denom = exp_avg_sq.sqrt().add_(group["eps"])

            step_size = group["lr"]
            if group["correct_bias"]:  # No bias correction for Bert
                bias_correction1 = 1.0 - beta1 ** state["step"]
                bias_correction2 = 1.0 - beta2 ** state["step"]
                step_size = step_size * math.sqrt(bias_correction2) / bias_correction1

            p.addcdiv_(exp_avg, denom, value=-step_size)
            if group["weight_decay"] > 0.0:
                p.add_(p, alpha=(-group["lr"] * group["weight_decay"]))
            #modified2: when updated, the optimizer state will move to cpu to free gpu
            state["exp_avg"] = exp_avg.to('cpu')
            state["exp_avg_sq"] = exp_avg_sq.to('cpu')
    return loss

训练时Loss出现大幅波动,日志如下:

{'loss': 1.4473, 'learning_rate': 1.9993007047883988e-05, 'epoch': 0.05}
{'loss': 1.3078, 'learning_rate': 1.9972037971811802e-05, 'epoch': 0.09}
{'loss': 1.2186, 'learning_rate': 1.9937122098932428e-05, 'epoch': 0.14}
{'loss': 0.9871, 'learning_rate': 1.9888308262251286e-05, 'epoch': 0.19}
{'loss': 0.9528, 'learning_rate': 1.9825664732332886e-05, 'epoch': 0.23}
{'loss': 0.8264, 'learning_rate': 1.9749279121818235e-05, 'epoch': 0.28}
{'loss': 0.8139, 'learning_rate': 1.9659258262890683e-05, 'epoch': 0.33}
{'loss': 0.802, 'learning_rate': 1.955572805786141e-05, 'epoch': 0.38}
{'loss': 0.7674, 'learning_rate': 1.9438833303083677e-05, 'epoch': 0.42}
{'loss': 0.8438, 'learning_rate': 1.9308737486442045e-05, 'epoch': 0.47}
{'loss': 0.8202, 'learning_rate': 1.9165622558699763e-05, 'epoch': 0.52}
{'loss': 0.8353, 'learning_rate': 1.900968867902419e-05, 'epoch': 0.56}
{'loss': 0.8675, 'learning_rate': 1.8841153935046098e-05, 'epoch': 0.61}
{'loss': 0.8664, 'learning_rate': 1.866025403784439e-05, 'epoch': 0.66}
{'loss': 0.9016, 'learning_rate': 1.8467241992282842e-05, 'epoch': 0.7}
{'loss': 0.8926, 'learning_rate': 1.826238774315995e-05, 'epoch': 0.75}
{'loss': 0.7756, 'learning_rate': 1.8045977797666685e-05, 'epoch': 0.8}
{'loss': 0.6959, 'learning_rate': 1.78183148246803e-05, 'epoch': 0.84}
{'loss': 1.2584, 'learning_rate': 1.757971723145453e-05, 'epoch': 0.89}

错误分析与原因

你的修改存在多处关键错误,直接导致Loss波动:

1. Step计数逻辑错误

state["step"] += 1被放在单个参数的循环体内,意味着每处理一个参数,该参数对应的step就会加1。而原始AdamW的设计是:每个优化步骤(一次step调用)中,所有参数的step仅需增加1。

这种错误会引发:

  • 不同参数的step值严重不一致,有梯度的参数被多次累加step,无梯度参数的step始终为0。
  • 基于step的偏差修正(bias_correction1/bias_correction2)计算完全偏离预期,学习率修正值混乱,参数更新步长忽大忽小,直接导致Loss波动。

2. 数据类型精度丢失

将优化器状态移回CPU时,未指定数据类型,默认会转换为CPU支持的默认类型(例如模型用bfloat16训练时,CPU会自动转为float32)。下次更新时又转回模型的dtype,反复类型转换会:

  • 导致exp_avg和exp_avg_sq的精度持续丢失,动量与二阶矩的计算出现偏差。
  • 破坏Adam的自适应学习率机制,梯度平滑效果失效,参数更新出现异常抖动。

3. 分布式训练下的状态一致性问题

8卡训练必然依赖分布式框架(如FSDP/DDP),每个GPU进程独立将优化器状态存到本地CPU,会导致:

  • 不同GPU上对应参数的优化器状态完全独立,无同步机制。
  • 参数更新时,各GPU的动量和二阶矩计算不一致,模型参数在不同卡上出现偏差,反映为Loss大幅波动。

修复建议

  • 修正Step计数:将state["step"] += 1移到参数循环体外,确保每个优化步骤只全局增加一次step;或者保持原始逻辑,但需保证所有参数的step同步更新。
  • 保留数据类型:移回CPU时指定与模型参数一致的dtype,修改为:
    state["exp_avg"] = exp_avg.to('cpu', dtype=p.dtype)
    state["exp_avg_sq"] = exp_avg_sq.to('cpu', dtype=p.dtype)
    
  • 改用成熟方案:放弃手动修改优化器,直接使用Hugging Face accelerate库、PyTorch FSDP或bitsandbytes量化优化器,这些方案已成熟解决大模型训练的显存问题,且无自定义修改的潜在bug。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 23:10:56