尧图网站建设 尧图网络
  • 首页
  • 关于我们
  • 服务项目
  • 案例展示
  • 建站流程
  • 资讯中心
  • 联系我们
首页/资讯中心/详情

【Bug已解决】LoRA gradients not normalized by input norm → training instability (NaN) 解决方案

【Bug已解决】LoRA gradients not normalized by input norm → training instability (NaN) 解决方案
📅 发布时间:2026/7/25 4:15:07

【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。

展开有三条:

  1. 梯度随‖x‖放大:∂L/∂A含xᵀ,输入越大梯度越大。
  2. Adam 预条件初期不稳:Adam 的二阶矩v需要若干步才稳定,warmup 不足时单步大梯度直接把参数推到数值危险区。
  3. 缩放因子α/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 时按序查:

  1. 先确认是不是 LoRA 层炸:打印各参数torch.isnan(p).any(),基座冻结权重通常有限,炸的是lora_A/lora_B。
  2. 查输入范数分布:x.norm(dim=-1).mean()与.max(),若方差极大(长尾),大概率是根因。
  3. 加输入范数归一:把B·A·x改成B·A·(x/‖x‖),直接解耦梯度与‖x‖。
  4. 梯度裁剪兜底:clip_grad_norm_(max_norm=1.0)。
  5. warmup 拉满:前 100 步线性升温,让 Adam 二阶矩稳定。
  6. B=0 初始化:保证 Δ 从 0 起步。
  7. 降 lr / 调 α/r:α/r越大增量越大,敏感场景调小。
  8. 监控数值:每步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 训练就能稳定收敛。

相关新闻

  • Unity Asset Bundle分析器:透视资源包,优化游戏性能与包体
  • 计算机毕业设计之运动场馆预约系统设计与实现
  • ComRAG框架:工业级问答系统的动态检索增强生成技术

最新新闻

  • 智能写作工具链:学术专著效率提升实战指南
  • 基于ResNet50的考研资料图片分类实践与优化
  • 美的风尊三代Pro空调选购指南:能效、智能与舒适体验全解析
  • C++ AI模型部署性能调优:内存、SIMD与多线程实战技巧
  • 深入解析Go语言cgo:连接Go与C/C++的桥梁机制与实践指南
  • PSO优化CNN-LSTM混合模型在时间序列预测中的应用

日新闻

  • 从国家条件到买方清单,深入理解 ABAP CDS 单值过滤器派生
  • 2026 年当下,齐齐哈尔专业的不锈钢闸门批发厂家哪个好,揭秘!这个工业“铁门”如何实现成本翻倍的效率提升? - 行业甄选官
  • 2026阳极氧化加工厂推荐:从设备规模看硬质氧化技术的成熟应用推荐百正机械 - 栗子测评

周新闻

  • SaaS软件行业GEO实践:AI搜索时代的品牌可见性与获客新路径
  • 什么是PCTFE?医药高端包装的“防潮王牌“材料
  • 【JVM调优实战】16-可视化利器-JConsole-VisualVM-JMC

月新闻

  • 2026年6月公司网站搭建最新热门渠道测评:四大低成本/零代码平台对比+避坑
  • 【Linux】Linux arm 编译QT程序,出现expected “}“报错
  • 【MATLAB例程】四基站二维AOA定位与距离辅助增强对比仿真。基于角度观测和测距修正的固定目标平面定位精度分析

关于尧图

  • 公司简介
  • 团队介绍
  • 企业文化
  • 荣誉资质

服务项目

  • 定制开发
  • 电商建站
  • UI 设计
  • 运维服务

快速链接

  • 案例展示
  • 建站流程
  • 常见问题
  • 资讯中心

联系方式

  • 📍北京市朝阳区互联网产业园 A 座 10 层
  • 📞400-888-8888
  • ✉️contact@rkmt.cn
  • 🕐周一至周日 9:00-21:00

© 2024 北京尧图网络科技有限公司 版权所有 | 京 ICP 备 XXXXXXXX 号