【Bug已解决】LoRA gradients not normalized by input norm → training instability (NaN) 解决方案一、现象长什么样用 LoRA 微调大模型时经常遇到一种诡异的不稳定loss 前几百步正常突然变成nan或者某些层通常是靠后的层、或 embedding 附近的层梯度爆炸而其余层安然无恙。具体表现训练中途loss变nantorch.isfinite(loss)为 Falsemodel.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·xB ∈ ℝ^{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, r4): super().__init__() self.W0 nn.Linear(in_f, out_f, biasFalse) 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, r4) opt torch.optim.Adam(layer.parameters(), lr1e-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], dim0) 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/Infloss{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, r4, eps1e-5): super().__init__() self.W0 nn.Linear(in_f, out_f, biasFalse) 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, keepdimTrue).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_norm1.0) opt.step()即便不改造前向全局梯度裁剪也能拦住单步大梯度避免参数越界成inf。修复 3warmup 适配学习率from torch.optim.lr_scheduler import LinearLR scheduler LinearLR(opt, start_factor0.01, total_iters100) # 前 100 步线性升温让 Adam 的二阶矩先稳定六、解决方案第二层结构性改进改进 1用 LoRA 思想给 A/B 不同学习率LoRA 的核心发现A降维和B升维适合用不同 lrB用更大的 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, eps1e-5): norm x.norm(dim-1, keepdimTrue).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, r4).to(torch.bfloat16) # B 初始化为 0保证训练起点 Δ0不会一开始就引入偏移B0初始化让 LoRA 增量从 0 起步配合输入归一能显著降低早期发散概率。七、解决方案第三层断言 / CI 守护import torch import torch.nn as nn import pytest class LoraLinearNormed(nn.Module): def __init__(self, in_f, out_f, r4, eps1e-5): super().__init__() self.W0 nn.Linear(in_f, out_f, biasFalse) 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, keepdimTrue).clamp_min(self.eps) delta (self.B (self.A (x / norm).T)).T return base delta def _train_step(layer, x, lr1e-2, steps50): opt torch.optim.Adam(layer.parameters(), lrlr) 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, r4) x_small torch.randn(4, 16) * 0.1 x_big torch.randn(4, 16) * 200.0 # 范数爆大 x torch.cat([x_small, x_big], dim0) 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, biasFalse) 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, r4) 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_norm1.0)。warmup 拉满前 100 步线性升温让 Adam 二阶矩稳定。B0 初始化保证 Δ 从 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、B0 初始化结构性改进是用 LoRA 的 A/B 分 lr、把归一做成可插拔包装、加数值健康监测最后用测试守护“归一版抗大范数输入、朴素版会发散、裁剪兜底有效”。把有效步长从输入范数解耦LoRA 训练就能稳定收敛。