当前位置: 首页 > news >正文

【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 附近的层)梯度爆炸,而其余层安然无恙。具体表现:

  • 训练中途lossnantorch.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·xB ∈ ℝ^{d×r}A ∈ ℝ^{r×k}r ≪ d。缩放因子α/r控制增量整体幅度。

A的梯度是:

∂L/∂A = Bᵀ · (∂L/∂Δ) · xᵀ

注意这里显式出现了x(输入)。也就是说,AB收到的梯度幅值随‖x‖线性放大。如果某一层/某批样本的x范数特别大(例如注意力 logits、或长序列尾部 token),该处的 LoRA 参数每一步更新量就远超其他位置,优化器(尤其 Adam,对梯度尺度本应自适应,但预条件矩阵初期不稳)在 warmup 阶段容易一步跨太大,参数越界 → 后续前向出现infnan扩散。

标准 LoRA 实现里并没有对x做归一化,它依赖用户自己选合适的αr、学习率来“碰巧”压住这个效应。一旦数据分布有长尾(输入范数方差大),就暴露出问题。

下面用最小可运行代码复现“大范数输入导致 LoRA 梯度爆炸→NaN”。

三、根因

根因一句话:LoRA 的增量路径B·A·x没有对输入x的范数做归一,梯度幅值随‖x‖变化,数据分布里输入范数方差大时,局部有效学习率失控,引发发散/NaN。

展开有三条:

  1. 梯度随‖x‖放大∂L/∂Axᵀ,输入越大梯度越大。
  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),几乎必然在几十步内Bnan。这就是“梯度随‖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 训练就能稳定收敛。

http://www.jsqmd.com/news/1258590/

相关文章:

  • CNN-LSTM-Attention混合模型在时序数据分类中的应用
  • 毛绒挂件手工精致款推荐哪些品牌?2026年测评 - 科技焦点
  • 2026 年现阶段,宜春正规的K9 级球墨铸铁管生产商深度剖析,颠覆想象:这管管铸铁如何让水管寿命翻倍? - 领域鉴赏官
  • RAG架构下小模型性能优化实战指南
  • VMware虚拟机多系统安装:Windows XP多实例部署与优化指南
  • 四步本地部署Dify:开源AI应用平台,实现私有化与深度定制
  • Codex实战指南:用AI生成自动化脚本,提升开发效率
  • 2026成都美术集训画室精选测评,艺考家庭择校避坑攻略! - 资讯报道
  • 毛绒挂件水果造型款怎么挑?2026年品牌推荐 - 科技焦点
  • 计算机毕业设计之运动场馆预约系统设计与实现
  • Stats.js前端性能监控实战指南
  • ComRAG框架:工业级问答系统的动态检索增强生成技术
  • 2026年电动扫地车厂家大盘点,看看谁更胜一筹 - 品牌排行榜
  • 2026 年当下,苍溪口碑好的陶粒订制厂家格局重塑与选型新思路,别再花冤枉钱!这小东西如何颠覆你的装修成本?-沈氏陶粒 - 企业官方推荐【认证】
  • 3步完成macOS软件管理革命:Applite让Homebrew变得如此简单
  • DAF-YOLO算法在工地安全监控中的创新应用
  • 视觉大语言模型技术演进与跨模态应用实践
  • 告别繁琐点击:京东自动化脚本终极指南,轻松管理日常任务
  • 婴儿玩具布艺类怎么挑?2026年品牌推荐 - 科技焦点
  • 国内版NotebookLM实测:音视频转图文笔记的AI工具有哪些选择
  • 计算机毕业设计之在线点餐系统的设计与实现
  • 基于Dify与MCP协议构建岗位专属AI副驾:从工作流到Claude/Cursor集成
  • AI驱动全域智能营销的技术架构与实战
  • 天心大师:关注AI时代的互联网心理困境,百问百答
  • AI规模化困境与Anthropic Skills模块化解决方案
  • 2026青岛漏水检测维修本地口碑榜TOP5权威推荐-专业仪器精准测漏-正规防水补漏公司推荐:卫生间/厨房/屋顶/阳台/外墙渗漏水检测师傅上门 - 安佳防水
  • 毛绒挂件小众设计款值得买吗?2026年品牌测评 - 科技焦点
  • AI提示词优化指南:提升大模型交互效率300%
  • Kubernetes集群部署预检实战指南
  • 北京华恒智信破解发电企业人员冗余、忙闲不均管理痛点