PyTorch自定义优化器实战:从零实现ClipAdam
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 性能优化:避免重复
更多推荐


所有评论(0)