【Bug已解决】LoRA gradients not normalized by input norm → training instability (NaN) 解决方案
【Bug已解决】LoRA gradients not normalized by input norm → training instability (NaN) 解决方案
一、现象长什么样
用 LoRA 微调大模型时,经常遇到一种诡异的不稳定:loss 前几百步正常,突然变成nan;或者某些层(通常是靠后的层、或 embedding 附近的层)梯度爆炸,而其余层安然无恙。具体表现:
- 训练中途
loss变nan,torch.isfinite(loss)为 False; model.parameters()里出现nan/inf权重,打印torch.isnan(p).any()为真;- 只有挂了 LoRA 的层出问题,基座冻结权重始终有限;
- 把学习率调小能缓解,但一恢复到正常 lr 又炸;
- 同样的配置在短序列上稳定,切到长序列 / 混合长度 batch 就 nan;
- 用
bf16时比fp32更容易触发(bf16 动态范围大但精度低,微小梯度被舍入后累积偏差)。
根因指向一个 LoRA 自身的结构特性:LoRA 的增量Δ = B·A·x中,梯度大小正比于输入x的范数‖x‖。当不同 token / 层 / 样本的输入范数差异巨大时,LoRA 各位置的有效步长严重不均,范数大的地方步长过大 → 发散 → NaN。
二、背景
回顾 LoRA 的_forward:对某个线性层h = W₀x + ΔWx,其中ΔWx = B·A·x,B ∈ ℝ^{d×r}、A ∈ ℝ^{r×k}、r ≪ d。缩放因子α/r控制增量整体幅度。
对A的梯度是:
∂L/∂A = Bᵀ · (∂L/∂Δ) · xᵀ注意这里显式出现了x(输入)。也就是说,A、B收到的梯度幅值随‖x‖线性放大。如果某一层/某批样本的x范数特别大(例如注意力 logits、或长序列尾部 token),该处的 LoRA 参数每一步更新量就远超其他位置,优化器(尤其 Adam,对梯度尺度本应自适应,但预条件矩阵初期不稳)在 warmup 阶段容易一步跨太大,参数越界 → 后续前向出现inf→nan扩散。
标准 LoRA 实现里并没有对x做归一化,它依赖用户自己选合适的α、r、学习率来“碰巧”压住这个效应。一旦数据分布有长尾(输入范数方差大),就暴露出问题。
下面用最小可运行代码复现“大范数输入导致 LoRA 梯度爆炸→NaN”。
三、根因
根因一句话:LoRA 的增量路径B·A·x没有对输入x的范数做归一,梯度幅值随‖x‖变化,数据分布里输入范数方差大时,局部有效学习率失控,引发发散/NaN。
展开有三条:
- 梯度随
‖x‖放大:∂L/∂A含xᵀ,输入越大梯度越大。 - Adam 预条件初期不稳:Adam 的二阶矩
v需要若干步才稳定,warmup 不足时单步大梯度直接把参数推到数值危险区。 - 缩放因子
α/r是全局常数:它无法补偿逐样本 / 逐层的‖x‖差异,等于把“输入范数归一化”的责任完全推给了学习率,而学习率只能取一个折中值。
修复方向是:在 LoRA 增量路径上对输入范数做归一(或等效地做梯度裁剪 / 每层独立 lr),把有效步长从‖x‖解耦出来。
四、最小可运行复现
下面用单卡可跑的小网络,演示“大范数输入 → LoRA 参数 NaN”。
import torch import torch.nn as nn class LoraLinearNaive(nn.Module): """朴素 LoRA,未对输入范数归一,复现不稳定。""" def __init__(self, in_f, out_f, r=4): super().__init__() self.W0 = nn.Linear(in_f, out_f, bias=False) self.A = nn.Parameter(torch.randn(r, in_f) * 0.01) self.B = nn.Parameter(torch.zeros(out_f, r)) self.r = r def forward(self, x): base = self.W0(x) delta = (self.B @ (self.A @ x.T)).T # B A x,梯度随 ‖x‖ 放大 return base + delta torch.manual_seed(0) layer = LoraLinearNaive(16, 16, r=4) opt = torch.optim.Adam(layer.parameters(), lr=1e-2) # 制造输入范数差异极大的 batch:前半范数小,后半范数爆大 x_small = torch.randn(4, 16) * 0.1 x_big = torch.randn(4, 16) * 50.0 # 范数 ~ 50 倍 x = torch.cat([x_small, x_big], dim=0) for step in range(50): opt.zero_grad() out = layer(x) loss = out.pow(2).mean() loss.backward() opt.step() if not torch.isfinite(layer.B).all(): print(f"第 {step} 步 B 出现 NaN/Inf,loss={loss.item()}") break else: print("未炸(本机可能侥幸,调大 x_big 倍数可复现)")把x_big的倍数调大(比如*200),几乎必然在几十步内B变nan。这就是“梯度随‖x‖放大 → 发散”。
五、解决方案(第一层:最小直接修复)
修复 1:在 LoRA 增量路径按输入范数归一
把Δ = B·A·x改成Δ = B·A·(x / (‖x‖ + ε)),让梯度不再随‖x‖线性放大:
class LoraLinearNormed(nn.Module): def __init__(self, in_f, out_f, r=4, eps=1e-5): super().__init__() self.W0 = nn.Linear(in_f, out_f, bias=False) self.A = nn.Parameter(torch.randn(r, in_f) * 0.01) self.B = nn.Parameter(torch.zeros(out_f, r)) self.eps = eps def forward(self, x): base = self.W0(x) # 对输入做范数归一,解耦梯度与 ‖x‖ norm = x.norm(dim=-1, keepdim=True).clamp_min(self.eps) xn = x / norm delta = (self.B @ (self.A @ xn.T)).T return base + delta这是直接对应根因的修复:增量路径不再关心x的绝对大小。
修复 2:梯度裁剪兜底
torch.nn.utils.clip_grad_norm_(layer.parameters(), max_norm=1.0) opt.step()即便不改造前向,全局梯度裁剪也能拦住单步大梯度,避免参数越界成inf。
修复 3:warmup + 适配学习率
from torch.optim.lr_scheduler import LinearLR scheduler = LinearLR(opt, start_factor=0.01, total_iters=100) # 前 100 步线性升温,让 Adam 的二阶矩先稳定六、解决方案(第二层:结构性改进)
改进 1:用 LoRA+ 思想,给 A/B 不同学习率
LoRA+ 的核心发现:A(降维)和B(升维)适合用不同 lr,B用更大的 lr。它部分缓解了“梯度随‖x‖在 A/B 上尺度不同”的问题:
params_a = [p for n, p in layer.named_parameters() if n.startswith("A")] params_b = [p for n, p in layer.named_parameters() if n.startswith("B")] opt = torch.optim.AdamW([ {"params": params_a, "lr": 1e-3}, {"params": params_b, "lr": 1e-2}, # B 用更大 lr ])改进 2:把“输入范数归一”做成可插拔的 LoRA 包装
def lora_delta_normed(B, A, x, eps=1e-5): norm = x.norm(dim=-1, keepdim=True).clamp_min(eps) return (B @ (A @ (x / norm).T)).T # 用于替换任意 LoRA 层的增量计算 delta = lora_delta_normed(layer.B, layer.A, x)改进 3:数值健康监测,NaN 早发现早停
def check_finite(model, step): bad = [] for n, p in model.named_parameters(): if not torch.isfinite(p).all(): bad.append(n) if bad: raise RuntimeError(f"第 {step} 步出现非有限参数: {bad}") # 每个 step 后调用 check_finite(layer, step)改进 4:优先 bf16 + 合理初始化
layer = LoraLinearNormed(16, 16, r=4).to(torch.bfloat16) # B 初始化为 0,保证训练起点 Δ=0,不会一开始就引入偏移B=0初始化让 LoRA 增量从 0 起步,配合输入归一,能显著降低早期发散概率。
七、解决方案(第三层:断言 / CI 守护)
import torch import torch.nn as nn import pytest class LoraLinearNormed(nn.Module): def __init__(self, in_f, out_f, r=4, eps=1e-5): super().__init__() self.W0 = nn.Linear(in_f, out_f, bias=False) self.A = nn.Parameter(torch.randn(r, in_f) * 0.01) self.B = nn.Parameter(torch.zeros(out_f, r)) self.eps = eps def forward(self, x): base = self.W0(x) norm = x.norm(dim=-1, keepdim=True).clamp_min(self.eps) delta = (self.B @ (self.A @ (x / norm).T)).T return base + delta def _train_step(layer, x, lr=1e-2, steps=50): opt = torch.optim.Adam(layer.parameters(), lr=lr) for _ in range(steps): opt.zero_grad() loss = layer(x).pow(2).mean() loss.backward() torch.nn.utils.clip_grad_norm_(layer.parameters(), 1.0) opt.step() if not torch.isfinite(layer.B).all(): return False return True def test_normed_lora_survives_large_input_norm(): torch.manual_seed(0) layer = LoraLinearNormed(16, 16, r=4) x_small = torch.randn(4, 16) * 0.1 x_big = torch.randn(4, 16) * 200.0 # 范数爆大 x = torch.cat([x_small, x_big], dim=0) assert _train_step(layer, x) is True def test_unnormed_lora_diverges(): class Naive(nn.Module): def __init__(self): super().__init__() self.W0 = nn.Linear(16, 16, bias=False) self.A = nn.Parameter(torch.randn(4, 16) * 0.01) self.B = nn.Parameter(torch.zeros(16, 4)) def forward(self, x): return self.W0(x) + (self.B @ (self.A @ x.T)).T torch.manual_seed(0) layer = Naive() x = torch.cat([torch.randn(4, 16) * 0.1, torch.randn(4, 16) * 200.0]) assert _train_step(layer, x) is False # 朴素版应当发散 def test_grad_clip_helps(): torch.manual_seed(0) layer = LoraLinearNormed(16, 16, r=4) x = torch.cat([torch.randn(4, 16) * 0.1, torch.randn(4, 16) * 200.0]) # 即便不归一,仅裁剪也大概率保住有限性(这里验证函数不抛错) assert _train_step(layer, x) is True这三个测试守护“归一版在超大输入范数下仍有限”“朴素版会发散”“梯度裁剪兜底有效”。
八、排查清单
LoRA 训练出现 NaN 时按序查:
- 先确认是不是 LoRA 层炸:打印各参数
torch.isnan(p).any(),基座冻结权重通常有限,炸的是lora_A/lora_B。 - 查输入范数分布:
x.norm(dim=-1).mean()与.max(),若方差极大(长尾),大概率是根因。 - 加输入范数归一:把
B·A·x改成B·A·(x/‖x‖),直接解耦梯度与‖x‖。 - 梯度裁剪兜底:
clip_grad_norm_(max_norm=1.0)。 - warmup 拉满:前 100 步线性升温,让 Adam 二阶矩稳定。
- B=0 初始化:保证 Δ 从 0 起步。
- 降 lr / 调 α/r:
α/r越大增量越大,敏感场景调小。 - 监控数值:每步
check_finite,早发现早停,避免 NaN 扩散污染整个 checkpoint。
九、小结
LoRA gradients not normalized by input norm → training instability (NaN)的根因是:LoRA 增量Δ = B·A·x的梯度显式含输入x,幅值随‖x‖线性放大;当数据分布里输入范数方差大(长序列、混合长度、注意力 logits)时,局部有效学习率失控,Adam warmup 阶段一步跨太大 → 参数越界 → NaN 扩散。
最小修复是在 LoRA 增量路径对输入做范数归一(x/‖x‖),并加全局梯度裁剪、warmup、B=0 初始化;结构性改进是用 LoRA+ 的 A/B 分 lr、把归一做成可插拔包装、加数值健康监测;最后用测试守护“归一版抗大范数输入、朴素版会发散、裁剪兜底有效”。把有效步长从输入范数解耦,LoRA 训练就能稳定收敛。