在训练神经网络时,很多难以排查的问题都来自weight decay的不当使用。一个典型的现象是:模型在训练初期收敛正常,但loss逐渐停滞,某些层的权重范数越来越小,最终几乎归零,表达能力被彻底破坏。这往往是因为对所有参数一刀切地施加了衰减,而其中一些参数——比如归一化层的缩放系数和偏置项——根本不应该参与正则化。本文围绕这个问题展开,讲清楚原理,并给出参数分组与例外设置的完整实践方案。
为什么Decay会把不该归零的参数压向零
权重衰减在每次更新时对参数施加一个比例缩小的操作。以SGD为例,衰减项直接加入梯度:w = w - lr * (grad + wd * w)。可以看出,只要参数值非零,衰减项就持续把它往零的方向拉,与梯度是否为零无关。对于卷积核权重这类高维参数,这种正则化通常是有益的,它能限制模型复杂度、提升泛化能力。
问题出在那些“值本身就承载语义”的参数上。BatchNorm层的weight(缩放系数gamma)和bias(偏移系数beta)通常初始化为1和0,它们的作用是对归一化后的激活做仿射变换。如果对gamma施加衰减,它会被持续拉向零,相当于不断压缩该层特征的幅度,严重时整个分支的信号被衰减殆尽,输出退化为纯归一化结果,模型表达能力大幅受损。
同理,普通的偏置参数也不建议衰减。偏置的作用是调整激活的阈值位置,衰减偏置会强迫模型学出一个接近零的偏置,这与衰减权重所追求的“简化模型”目标并不相关。此外,如果使用了不做归一化的embedding层,其向量也是参数化的语义表示,直接衰减同样会损伤表示质量。
哪些参数应该豁免Decay:分组原则
通用的分组原则可以概括为一句话:凡是维度为一的参数以及归一化层的所有参数,都不参与衰减。具体来说包括:BatchNorm、LayerNorm、GroupNorm、RMSNorm等归一化层的weight和bias;全连接层和卷积层的bias;某些场景下的embedding权重。
为什么用“维度为一”作为判断标准?因为shape为(out_channels,)或(hidden_size,)的一维参数,绝大多数是gamma、beta或偏置,它们控制的是逐通道的平移和缩放,衰减它们没有正则化收益。而二维以上的参数(卷积核、全连接权重矩阵)才是真正需要约束的对象。按维度过滤的好处是不依赖参数命名,即使模型来自第三方库,也能比较稳健地识别。
需要注意一个例外中的例外:如果归一化层的weight被冻结(requires_grad为False),或者你使用了affine=False的归一化层,那么这些层根本没有可学习的仿射参数,过滤逻辑自然不会命中它们,不会产生副作用。但对于embedding,社区存在分歧——做大规模预训练时通常会对embedding做轻度衰减或单独调参,而微调场景往往将其豁免。建议根据任务实验确定,而不是照搬别人的配置。
PyTorch中的参数分组实现
PyTorch的optimizer支持按参数组传入不同的超参数,这是实现例外控制的标准方式。最通用的做法是根据参数维度自动分组,代码如下:
def build_optimizer(model, lr=1e-3, weight_decay=0.01):
decay_params = []
no_decay_params = []
for name, param in model.named_parameters():
if not param.requires_grad:
continue
# 一维参数(bias、归一化层gamma/beta)不衰减
if param.dim() == 1 or name.endswith(".bias"):
no_decay_params.append(param)
else:
decay_params.append(param)
return torch.optim.AdamW([
{"params": decay_params, "weight_decay": weight_decay},
{"params": no_decay_params, "weight_decay": 0.0},
], lr=lr)也可以按参数名称中的关键词过滤,例如匹配bn、norm、bias等字符串。名称匹配的好处是显式、可控,坏处是换模型结构就可能漏配。实践中推荐以维度判断为主、名称匹配为辅,两者结合可以覆盖绝大多数情况。
有一个容易踩的坑:分组时如果传入的是param张量列表,要确保没有参数被同时分到两个组,也没有可训练参数被遗漏。PyTorch对重复参数不会报错,只会导致该参数被更新两次。可以在分组后加一个断言校验总数:
total = len(decay_params) + len(no_decay_params) trainable = sum(1 for p in model.parameters() if p.requires_grad) assert total == trainable, "参数分组存在遗漏或重复"
AdamW与内置衰减的区别及注意事项
传统Adam配合L2正则化的方式是把衰减项加进梯度,这会导致衰减经过自适应学习率的缩放,效果与真正的weight decay并不等价。AdamW将衰减从梯度中解耦出来,直接在参数上执行w = w - lr * wd * w,这才是与原始论文一致的权重衰减实现。因此现代训练基本都使用AdamW,而参数分组在AdamW中同样有效,因为weight_decay是参数组级别的超参数。
另一个经验是衰减系数的取值。微调预训练模型时常用0.01甚至0.0,而大规模预训练可能用0.1。衰减越强,参数被压向零的力道越大,对gamma这类敏感参数的伤害也越快显现。如果你观察到训练中某些层输出幅度持续萎缩,可以打印各参数组的范数来诊断:
for name, param in model.named_parameters():
if param.dim() == 1 and param.requires_grad:
print(name, param.norm().item())最后总结一下排查思路:先确认优化器是否用了AdamW的解耦衰减;再确认归一化层和偏置是否被错误地加入了decay组;训练过程中监控一维参数的范数变化,发现持续下降就说明衰减在侵蚀这些参数。参数分组看似只是几行配置代码,却直接决定了模型能否稳定训练到高精度,值得在项目初期就规范地搭好。