1. 项目概述:为什么你需要亲手写一个优化器,而不是直接调用 torch.optim

在 PyTorch 项目里敲下 optimizer = torch.optim.Adam(model.parameters(), lr=1e-3) 的那一刻,你其实已经把模型收敛的“方向盘”交给了别人——不是交给人,而是交给了一个封装得严丝合缝、文档写得滴水不漏、但内部逻辑对你完全黑盒的预设实现。我带过二十多个工业级训练项目,从轻量级时序预测到百层视觉Transformer微调,几乎每次遇到收敛震荡、梯度爆炸、学习率衰减失灵、或者需要在参数空间做非标准约束(比如让某几组权重始终满足和为1)时,第一反应不再是调参,而是打开 torch/optim/ 目录下的源码文件夹,点开 adam.py ,然后默默新建一个 my_adam.py 。这不是炫技,是刚需。

“Custom Optimizer”这个词听起来像学术论文里的玩具实验,但真实场景中它解决的是三类硬问题: 第一,算法层面的定制化 ——比如你在复现一篇顶会论文,它提出了一种带二阶动量修正的 Adam 变体,PyTorch 官方还没收录; 第二,工程层面的轻量化 ——你部署在边缘设备上,不能容忍 torch.optim.lr_scheduler 那套冗余的状态管理; 第三,调试与可解释性需求 ——你想在每步更新后打印某个参数组的梯度范数变化率,或者记录历史动量向量的 L2 衰减曲线,官方优化器根本不提供 hook 接口。这五个步骤不是教你怎么“造轮子”,而是教你如何把轮子拆开、看清轴承怎么咬合、再按自己车轴尺寸重装回去。核心关键词就三个: PyTorch 自定义优化器、 torch.optim.Optimizer 基类、 step() 方法重载 。如果你正在跑一个关键训练任务,卡在 loss 不降、grad nan、或者想加个简单约束却找不到入口,这篇就是为你写的——它不讲理论推导,只讲你打开 IDE 后该敲哪几行、为什么这么敲、以及敲错哪一行会导致整个训练崩掉。

2. 核心设计思路与基类原理:为什么必须继承 Optimizer ,而不是写个普通函数

2.1 优化器的本质不是“计算公式”,而是“状态容器 + 更新引擎”

很多人初学时有个误解:优化器 = 更新公式。比如 Adam 就是那几个带 beta 参数的指数滑动平均公式。于是他们写个函数:

def simple_adam_step(params, grads, m, v, lr=1e-3, beta1=0.9, beta2=0.999):
    m = beta1 * m + (1 - beta1) * grads
    v = beta2 * v + (1 - beta2) * grads ** 2
    return params - lr * m / (torch.sqrt(v) + 1e-8)

这函数能跑通,但立刻会撞上三堵墙:
第一堵墙:参数与状态分离失控 。PyTorch 模型有成百上千个参数张量( weight , bias , embedding.weight …),每个都需要独立维护 m v 。你不可能手动为每个参数创建两个同形张量并传进函数——这违背了 PyTorch 的自动参数管理哲学。
第二堵墙: zero_grad() 失效 。PyTorch 的 model.zero_grad() 是靠遍历 model.parameters() 并清空 .grad 属性实现的。你的函数如果绕过 .grad 直接操作梯度值,下次 backward() 就会累加而非覆盖,loss 瞬间爆炸。
第三堵墙: state_dict() 无法保存/加载 。训练中断要 resume?你得把所有 m , v , step_count 全部手动打包进字典,还要确保加载时形状对齐——而 PyTorch 的 checkpoint 机制默认只认 optimizer.state_dict() 返回的标准结构。

所以, torch.optim.Optimizer 基类存在的根本意义,是强制你把“状态”和“参数”绑定在一起,并提供统一的生命周期管理接口 。它内部维护一个 self.state 字典,键是参数张量对象(注意:是对象引用,不是名字!),值是该参数专属的状态字典(如 {'step': 0, 'exp_avg': tensor(...), 'exp_avg_sq': tensor(...)} )。当你调用 optimizer.step() ,它自动遍历 self.param_groups 中所有参数组,对每个参数 p ,取出 self.state[p] ,执行你的更新逻辑,再把新参数值写回 p.data 。这个设计看似繁琐,实则是稳定性的基石。

2.2 为什么 param_groups 是不可绕过的抽象层?

看一眼 torch.optim.Adam 的初始化签名:

Adam(params, lr=0.001, betas=(0.9, 0.999), eps=1e-08, weight_decay=0, amsgrad=False)

你以为 params 就是 model.parameters() ?错了。 params 实际接受三种形式:

  • 单个参数迭代器(最常见)
  • 参数列表( [p1, p2, p3]
  • 参数组列表( [{'params': [p1, p2], 'lr': 1e-4}, {'params': [p3], 'lr': 1e-3}]

这个设计解决了真实项目中的刚性需求: 不同参数需要不同学习率、不同 weight decay、甚至不同更新规则 。比如 BERT 微调时,通常对 classifier 层用 2e-5,对 encoder 层用 5e-5;又比如在对比学习中, projection head 的 weight decay 设为 0,而 backbone 设为 1e-4。如果你写个裸函数,就得手动拆分参数、分别调用、再合并结果——而 Optimizer 基类通过 param_groups 把这个过程标准化了: self.param_groups 是一个列表,每个元素是一个字典,包含 'params' (参数列表)、 'lr' 'weight_decay' 等键。你在 step() 里只需两层循环:

for group in self.param_groups:          # 遍历参数组
    for p in group['params']:            # 遍历组内参数
        if p.grad is None: continue      # 跳过无梯度参数(如冻结层)
        grad = p.grad.data               # 获取梯度张量
        # 在这里写你的更新逻辑
        p.data.add_(update_tensor)       # 原地更新参数

提示: p.data 是参数张量的数值部分, p 本身是 Parameter 对象(继承自 Tensor ),包含 grad requires_grad 等属性。直接操作 p.data 是安全的,因为 backward() 只影响 p.grad ,不影响 p.data 的值。

2.3 state 字典的键必须是参数对象,不是字符串名

这是新手踩坑最多的地方。有人试图这样写:

# ❌ 错误示范:用参数名当键
self.state['encoder.weight'] = {'step': 0, 'exp_avg': ...}
# 后续取值时:
state = self.state[p.name]  # p.name 不存在!Parameter 没有 name 属性

PyTorch 的 state 字典键必须是参数张量对象本身(即 p 这个变量),因为:

  • 参数对象在内存中是唯一标识,即使两个张量数值相同,对象引用也不同;
  • model.named_parameters() 返回的名字只是调试用的标签,训练过程中参数对象不会变,但名字可能因模型重构而失效;
  • Optimizer 内部通过 id(p) 或弱引用管理状态,确保 p 被 GC 时状态也能清理。

所以正确写法永远是:

# ✅ 正确:用参数对象 p 作为键
if p not in self.state:
    self.state[p] = {}
    self.state[p]['step'] = 0
    self.state[p]['exp_avg'] = torch.zeros_like(p.data, memory_format=torch.preserve_format)
# 后续使用:
state = self.state[p]
state['step'] += 1

torch.preserve_format 这个参数很关键——它让新张量继承原参数的内存布局(如 channel-last 格式),避免因格式不一致导致 CUDA kernel 报错。我在一个 ResNet50+channel-last 的训练中,漏掉这个参数, exp_avg 初始化为默认的 channel-first,结果 add_() 操作直接报 RuntimeError: expected input and weight to have same layout ,debug 了三小时才发现是这里。

3. 五步实操详解:从零构建一个带梯度裁剪与自适应学习率的 Adam 变体

我们来构建一个真实项目中高频使用的优化器: ClipAdam ——它在标准 Adam 基础上增加两项能力:

  • 每步自动对梯度做全局裁剪( torch.nn.utils.clip_grad_norm_ 的等效逻辑);
  • 学习率随全局 step 动态缩放(类似 warmup + decay,但由优化器自身管理,不依赖外部 scheduler)。

这个优化器在训练不稳定模型(如 GAN、RL 策略网络)时非常实用,避免手动在 train_step() 里插 clip_grad_norm_ scheduler.step() ,减少出错概率。

3.1 第一步:继承 Optimizer 并初始化 param_groups defaults

import torch
from torch.optim import Optimizer

class ClipAdam(Optimizer):
    def __init__(self, params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8,
                 weight_decay=0, amsgrad=False, max_grad_norm=1.0,
                 warmup_steps=1000, total_steps=100000):
        # 1. 构建 defaults 字典:存储所有超参数的默认值
        # 这些值会被复制到每个 param_group 的字典中,供 step() 读取
        defaults = dict(
            lr=lr, betas=betas, eps=eps,
            weight_decay=weight_decay, amsgrad=amsgrad,
            max_grad_norm=max_grad_norm,
            warmup_steps=warmup_steps,
            total_steps=total_steps
        )
        # 2. 调用父类初始化:解析 params,生成 param_groups,注入 defaults
        super(ClipAdam, self).__init__(params, defaults)
        
        # 3. 初始化全局计数器(不属于任何参数组,所以放 self 上)
        self._step = 0

这里的关键细节:

  • defaults 字典必须包含所有 step() 中要用到的超参数,否则 group['lr'] 会报 KeyError
  • super().__init__() 会自动处理 params 的三种输入格式,并为每个参数组设置 group['lr'] 等字段;
  • self._step 是全局 step 计数器,不是 per-parameter 的,所以不能放 self.state[p] 里,必须挂 self 上。

注意: warmup_steps total_steps 是调度参数,不是优化算法参数,但放在这里能让优化器完全自治。很多团队把 scheduler 和 optimizer 分开管理,结果 resume 时 scheduler.last_epoch optimizer.state['step'] 不同步,loss 曲线突变。把它们收归优化器, state_dict() 里自然就包含 '_step' ,resume 时一并加载,彻底解决同步问题。

3.2 第二步:重载 state_dict() load_state_dict() ,支持断点续训

PyTorch 默认的 state_dict() 只序列化 self.state self.param_groups ,但我们的 self._step 是全局计数器,必须显式加入:

def state_dict(self):
    # 1. 调用父类方法获取基础 state_dict
    state_dict = super(ClipAdam, self).state_dict()
    # 2. 手动添加自定义字段
    state_dict['_step'] = self._step
    return state_dict

def load_state_dict(self, state_dict):
    # 1. 提取自定义字段
    _step = state_dict.pop('_step', 0)
    # 2. 调用父类方法加载基础状态
    super(ClipAdam, self).load_state_dict(state_dict)
    # 3. 恢复自定义字段
    self._step = _step

这个操作看似简单,但有两个致命陷阱:

  • 陷阱一: state_dict.pop() 必须带默认值 。如果 checkpoint 是旧版本(没存 _step ), pop('_step') 会抛 KeyError ,训练直接中断。所以必须写 state_dict.pop('_step', 0)
  • 陷阱二: super().load_state_dict() 必须在 pop 之后调用 。因为父类方法会校验 state_dict 是否有多余键,如果你先调用它, '_step' 会被视为非法键而报错。

我在一个医疗影像分割项目中,因为没加默认值,线上服务自动 reload checkpoint 时崩溃,P0 故障。后来加了日志发现是这里,从此所有自定义字段都强制加默认值。

3.3 第三步:实现 step() 主逻辑——状态初始化、梯度裁剪、自适应 LR 计算

这是核心中的核心。我们分四小步写:

3.3.1 遍历参数组,初始化每个参数的状态
def step(self, closure=None):
    loss = None
    if closure is not None:
        loss = closure()

    # 1. 全局 step +1
    self._step += 1

    # 2. 遍历每个参数组
    for group in self.param_groups:
        # 解包超参数(从 group 字典中取,不是 defaults)
        lr = group['lr']
        beta1, beta2 = group['betas']
        eps = group['eps']
        weight_decay = group['weight_decay']
        amsgrad = group['amsgrad']
        max_grad_norm = group['max_grad_norm']
        warmup_steps = group['warmup_steps']
        total_steps = group['total_steps']

        # 3. 遍历组内每个参数
        for p in group['params']:
            if p.grad is None:
                continue
            grad = p.grad.data

            # 4. 初始化该参数的状态(如果第一次见)
            if p not in self.state:
                self.state[p] = {}
                state = self.state[p]
                state['step'] = 0
                state['exp_avg'] = torch.zeros_like(p.data, memory_format=torch.preserve_format)
                state['exp_avg_sq'] = torch.zeros_like(p.data, memory_format=torch.preserve_format)
                if amsgrad:
                    state['max_exp_avg_sq'] = torch.zeros_like(p.data, memory_format=torch.preserve_format)
            else:
                state = self.state[p]

            # 5. 状态 step +1(per-parameter)
            state['step'] += 1

注意 state['step'] self._step 的区别:前者是每个参数自己的更新次数(用于 bias correction),后者是全局训练步数(用于 LR 调度)。Adam 的 bias correction 公式 m_hat = m / (1 - beta1^step) 必须用 state['step'] ,否则所有参数共享一个 step,小 learning rate 的参数会严重欠校正。

3.3.2 梯度裁剪:在更新前对 grad 做 in-place 归一化
            # 6. 梯度裁剪:计算全局梯度 norm(L2)
            # 注意:这里要对整个 group 的 grad 做 global norm,不是单个 p
            # 但我们是在 per-parameter 循环里,所以先收集所有 grad 到 list
            # ✅ 正确做法:在 group 循环外预计算 global norm
            # ❌ 错误做法:在 p 循环里单独算 norm —— 会得到错误的 global norm

等等,这里发现一个经典误区!上面代码注释里指出的问题正是多数人栽跟头的地方。 梯度裁剪必须基于整个 param_group 的梯度拼接计算 global norm,而不是单个参数的 norm 。所以我们要重构逻辑:把裁剪提到 group 循环内、p 循环外:

        # 在 group 循环内,p 循环前,收集本组所有梯度
        grads = []
        for p in group['params']:
            if p.grad is not None:
                grads.append(p.grad.data.view(-1))  # 展平
        if len(grads) == 0:
            continue
        # 拼接所有梯度为一个大向量
        grad_vec = torch.cat(grads)
        # 计算 global L2 norm
        total_norm = torch.norm(grad_vec, p=2)
        # 计算裁剪系数
        clip_coef = max_grad_norm / (total_norm + 1e-6)
        if clip_coef < 1:
            # 对组内每个参数的梯度做 in-place 缩放
            for p in group['params']:
                if p.grad is not None:
                    p.grad.data.mul_(clip_coef)

这个 mul_() 是 in-place 操作,直接修改 p.grad.data ,后续 step() 里的 grad = p.grad.data 就拿到裁剪后的值。 1e-6 是防除零,比 eps=1e-8 稍大,避免浮点精度问题导致 clip_coef 溢出。

3.3.3 计算自适应学习率:warmup + cosine decay
            # 7. 计算当前 step 对应的自适应 lr
            # 公式:lr * (min(step, warmup_steps) / warmup_steps) * 
            #       (0.5 * (1 + cos(pi * (step - warmup_steps) / (total_steps - warmup_steps))))
            # 但 step 可能 > total_steps,需 clamp
            if self._step <= warmup_steps:
                # warmup 阶段:线性增长
                adaptive_lr = lr * float(self._step) / float(max(1, warmup_steps))
            else:
                # decay 阶段:cosine decay
                progress = float(self._step - warmup_steps) / float(max(1, total_steps - warmup_steps))
                progress = min(progress, 1.0)  # clamp to [0,1]
                adaptive_lr = lr * 0.5 * (1.0 + math.cos(math.pi * progress))
            # 注意:adaptive_lr 是标量,后续乘 grad 时会自动 broadcast

这里用 math.cos 而不是 torch.cos ,因为 self._step 是 Python int, math 函数更快且无 device 问题。 max(1, ...) 是防 warmup_steps=0 total_steps=warmup_steps 导致除零。

3.3.4 执行 Adam 更新:bias correction 与参数更新
            # 8. 执行 Adam 更新(标准逻辑,加 weight_decay)
            state = self.state[p]
            exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

            if weight_decay != 0:
                grad = grad.add(p.data, alpha=weight_decay)

            # 更新一阶矩估计
            exp_avg.mul_(beta1).add_(grad, alpha=1 - beta1)
            # 更新二阶矩估计
            exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1 - beta2)

            if amsgrad:
                max_exp_avg_sq = state['max_exp_avg_sq']
                torch.max(max_exp_avg_sq, exp_avg_sq, out=max_exp_avg_sq)
                denom = max_exp_avg_sq.sqrt().add_(eps)
            else:
                denom = exp_avg_sq.sqrt().add_(eps)

            # bias correction
            bias_correction1 = 1 - beta1 ** state['step']
            bias_correction2 = 1 - beta2 ** state['step']
            step_size = adaptive_lr / bias_correction1
            denom = denom / math.sqrt(bias_correction2)

            # 参数更新:p.data = p.data - step_size * exp_avg / denom
            p.data.addcdiv_(exp_avg, denom, value=-step_size)

addcdiv_() 是 PyTorch 原生的 p.data -= step_size * exp_avg / denom 的高效实现,比 p.data = p.data - ... 少一次内存分配。 value=-step_size 是标量乘数, addcdiv_ 会自动 broadcast。

3.4 第四步:添加调试钩子——在 step() 结束时打印关键指标

为了验证优化器是否按预期工作,我们在 step() 最后加一段调试输出(生产环境注释掉即可):

        # 9. 调试:打印本 group 的梯度 norm 和 lr(仅第一个 group)
        if group == self.param_groups[0]:
            grads = [p.grad.data.view(-1) for p in group['params'] if p.grad is not None]
            if grads:
                grad_vec = torch.cat(grads)
                grad_norm = torch.norm(grad_vec, p=2).item()
                print(f"[ClipAdam] Step {self._step:6d} | "
                      f"Grad Norm: {grad_norm:.4f} | "
                      f"LR: {adaptive_lr:.6f} | "
                      f"Param Group Size: {len(group['params'])}")

这段代码放在 for group in self.param_groups: 循环的末尾,只对第一个参数组打印,避免日志刷屏。 grad_norm 应该在 warmup 阶段缓慢上升,decay 阶段逐渐下降; LR 应该从 0 线性升到 lr ,再平滑降到接近 0。我在调试一个语音合成模型时,发现 grad_norm 在 step=500 时突然跳到 100+,立刻定位到是 weight_decay 符号写反了( add 写成 sub ),这种实时反馈比看 loss 曲线快十倍。

3.5 第五步:完整类定义与使用示例

把以上所有片段组合起来,就是完整的 ClipAdam 类:

import math
import torch
from torch.optim import Optimizer

class ClipAdam(Optimizer):
    def __init__(self, params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8,
                 weight_decay=0, amsgrad=False, max_grad_norm=1.0,
                 warmup_steps=1000, total_steps=100000):
        if not 0.0 <= lr:
            raise ValueError("Invalid learning rate: {}".format(lr))
        if not 0.0 <= eps:
            raise ValueError("Invalid epsilon value: {}".format(eps))
        if not 0.0 <= betas[0] < 1.0:
            raise ValueError("Invalid beta1 parameter: {}".format(betas[0]))
        if not 0.0 <= betas[1] < 1.0:
            raise ValueError("Invalid beta2 parameter: {}".format(betas[1]))
        if not 0.0 <= weight_decay:
            raise ValueError("Invalid weight_decay value: {}".format(weight_decay))
        if not 0.0 <= max_grad_norm:
            raise ValueError("Invalid max_grad_norm: {}".format(max_grad_norm))

        defaults = dict(
            lr=lr, betas=betas, eps=eps,
            weight_decay=weight_decay, amsgrad=amsgrad,
            max_grad_norm=max_grad_norm,
            warmup_steps=warmup_steps,
            total_steps=total_steps
        )
        super(ClipAdam, self).__init__(params, defaults)
        self._step = 0

    def __setstate__(self, state):
        super(ClipAdam, self).__setstate__(state)
        for group in self.param_groups:
            group.setdefault('amsgrad', False)

    def state_dict(self):
        state_dict = super(ClipAdam, self).state_dict()
        state_dict['_step'] = self._step
        return state_dict

    def load_state_dict(self, state_dict):
        _step = state_dict.pop('_step', 0)
        super(ClipAdam, self).load_state_dict(state_dict)
        self._step = _step

    def step(self, closure=None):
        loss = None
        if closure is not None:
            loss = closure()

        self._step += 1

        for group in self.param_groups:
            lr = group['lr']
            beta1, beta2 = group['betas']
            eps = group['eps']
            weight_decay = group['weight_decay']
            amsgrad = group['amsgrad']
            max_grad_norm = group['max_grad_norm']
            warmup_steps = group['warmup_steps']
            total_steps = group['total_steps']

            # 收集本组所有梯度用于 global norm
            grads = []
            for p in group['params']:
                if p.grad is not None:
                    grads.append(p.grad.data.view(-1))
            if len(grads) == 0:
                continue
            grad_vec = torch.cat(grads)
            total_norm = torch.norm(grad_vec, p=2)
            clip_coef = max_grad_norm / (total_norm + 1e-6)
            if clip_coef < 1:
                for p in group['params']:
                    if p.grad is not None:
                        p.grad.data.mul_(clip_coef)

            # 计算自适应 lr
            if self._step <= warmup_steps:
                adaptive_lr = lr * float(self._step) / float(max(1, warmup_steps))
            else:
                progress = float(self._step - warmup_steps) / float(max(1, total_steps - warmup_steps))
                progress = min(progress, 1.0)
                adaptive_lr = lr * 0.5 * (1.0 + math.cos(math.pi * progress))

            # 更新每个参数
            for p in group['params']:
                if p.grad is None:
                    continue
                grad = p.grad.data

                if p not in self.state:
                    self.state[p] = {}
                    state = self.state[p]
                    state['step'] = 0
                    state['exp_avg'] = torch.zeros_like(p.data, memory_format=torch.preserve_format)
                    state['exp_avg_sq'] = torch.zeros_like(p.data, memory_format=torch.preserve_format)
                    if amsgrad:
                        state['max_exp_avg_sq'] = torch.zeros_like(p.data, memory_format=torch.preserve_format)
                else:
                    state = self.state[p]

                state['step'] += 1

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                if weight_decay != 0:
                    grad = grad.add(p.data, alpha=weight_decay)

                exp_avg.mul_(beta1).add_(grad, alpha=1 - beta1)
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1 - beta2)

                if amsgrad:
                    max_exp_avg_sq = state['max_exp_avg_sq']
                    torch.max(max_exp_avg_sq, exp_avg_sq, out=max_exp_avg_sq)
                    denom = max_exp_avg_sq.sqrt().add_(eps)
                else:
                    denom = exp_avg_sq.sqrt().add_(eps)

                bias_correction1 = 1 - beta1 ** state['step']
                bias_correction2 = 1 - beta2 ** state['step']
                step_size = adaptive_lr / bias_correction1
                denom = denom / math.sqrt(bias_correction2)

                p.data.addcdiv_(exp_avg, denom, value=-step_size)

            # 调试打印(生产环境请注释)
            if group == self.param_groups[0]:
                grads = [p.grad.data.view(-1) for p in group['params'] if p.grad is not None]
                if grads:
                    grad_vec = torch.cat(grads)
                    grad_norm = torch.norm(grad_vec, p=2).item()
                    print(f"[ClipAdam] Step {self._step:6d} | "
                          f"Grad Norm: {grad_norm:.4f} | "
                          f"LR: {adaptive_lr:.6f} | "
                          f"Param Group Size: {len(group['params'])}")

        return loss

使用方式和原生优化器完全一致

import torch
import torch.nn as nn

model = nn.Sequential(nn.Linear(10, 5), nn.ReLU(), nn.Linear(5, 1))
criterion = nn.MSELoss()
# ✅ 直接替换,无需改其他代码
optimizer = ClipAdam(
    model.parameters(),
    lr=1e-3,
    warmup_steps=500,
    total_steps=10000,
    max_grad_norm=0.5
)

for epoch in range(10):
    for x, y in dataloader:
        optimizer.zero_grad()
        y_pred = model(x)
        loss = criterion(y_pred, y)
        loss.backward()
        optimizer.step()  # 自动完成裁剪、LR调度、更新

4. 常见问题排查与实战避坑指南:那些文档里不会写的细节

4.1 问题速查表:训练异常时的定位路径

现象 最可能原因 快速验证方法 修复方案
loss 不下降, grad_norm 持续为 0 p.grad None ,参数未参与计算图 print([p.grad is None for p in model.parameters()]) 检查 forward() 是否用了 with torch.no_grad(): ,或 loss 是否未对 model 输出求导
loss 爆炸( inf / nan ), grad_norm 极大 梯度裁剪未生效,或 max_grad_norm 设太小导致裁剪系数 >1e6 step() 中打印 clip_coef max_grad_norm 1.0 改为 5.0 ,观察 clip_coef 是否 < 1
state_dict() 加载后 loss 曲线突变 self._step 未正确 load_state_dict print(optimizer._step) 加载前后 确保 load_state_dict() pop('_step') 带默认值,且在 super() 调用之后赋值
GPU 显存暴涨,OOM exp_avg / exp_avg_sq 未用 memory_format=torch.preserve_format print(p.data.stride(), state['exp_avg'].stride()) 在初始化状态时显式传入 memory_format 参数
多卡 DDP 训练报错 Expected all tensors to be on the same device state 中的张量未随 p.data 自动 move 到 GPU print(p.data.device, state['exp_avg'].device) if p not in self.state: 分支中,用 torch.zeros_like(p.data, ...) ,它会自动继承 p.data 的 device

这张表来自我处理过的 37 个线上故障案例。最常被忽略的是最后一项: torch.zeros_like(p.data) 会自动匹配 p.data device dtype ,但如果你写 torch.zeros(p.data.shape, dtype=p.data.dtype).to(p.data.device) ,在 DDP 下可能因 to() 触发隐式 copy 导致 device 不一致。

4.2 关于 torch.no_grad() 的深度陷阱

很多教程说“在 step() 里操作 p.data 是安全的,因为不涉及 autograd”。这句话 只对一半 。看这个例子:

# ❌ 危险操作:在 no_grad 块里更新 p.data,但 p.grad 仍被 backward 修改
with torch.no_grad():
    p.data = p.data - lr * p.grad  # 这里没问题
# 但下一轮 backward() 仍会往 p.grad 里累加!
# 如果你忘了 zero_grad(),p.grad 就是上一轮的残留 + 新梯度 → 爆炸

正确姿势永远是:

# ✅ 安全操作:zero_grad() 必须在 backward() 之前,且与 optimizer 绑定
optimizer.zero_grad()  # 这行会清空所有 p.grad
loss.backward()        # 这行会重新填充 p.grad
optimizer.step()       # 这行用 p.grad 更新 p.data

optimizer.zero_grad() 的源码就是遍历 self.param_groups ,对每个 p 执行 p.grad = None p.grad.zero_() 。所以它和你的自定义优化器是强耦合的——你不能用裸 model.zero_grad() ,必须用 optimizer.zero_grad() ,否则 self.state[p] 里的 exp_avg 还在,但 p.grad 已清空,下一步 step() 就会 grad is None 跳过,状态停滞。

4.3 weight_decay 的两种实现方式:L2 正则 vs Decoupled Decay

PyTorch 官方 Adam weight_decay L2 正则 :它把 weight_decay * p.data 加到梯度上(见 step() grad = grad.add(p.data, alpha=weight_decay) )。但论文《Decoupled Weight Decay Regularization》指出,AdamW 才是正确的解耦方式: weight_decay 应直接作用于参数更新,而不是污染梯度。如果你要实现 AdamW,只需把 weight_decay 移到更新后:

# ✅ AdamW 风格:decoupled weight decay
p.data.mul_(1 - lr * weight_decay)  # 先对参数做 decay
p.data.addcdiv_(exp_avg, denom, value=-step_size)  # 再做 Adam 更新

这个细节影响极大。我在一个蛋白质结构预测项目中,用 L2 版本 weight_decay=0.01 ,val loss 一直卡在 1.2;换成 AdamW 风格, weight_decay=0.01 ,val loss 降到 0.85。原因是 L2 正则在 Adam 的自适应学习率下会失效——大 learning rate 的参数被过度惩罚,小 learning rate 的参数惩罚不足。

4.4 性能优化:避免重复

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐