数据截至 (上游 commit 1fe27b1b53f3)
04 · 无降级 LoRA 快路径
这章回答:Unsloth 宣称的「更快且零精度损失」在数学上怎么成立。答案是:不近似任何运算, 只是把 LoRA 的前向反向用闭合公式手工重排,并顺手把 4bit 反量化塞进矩阵乘里。
1. 小问题
PEFT 的 LoRA 层是「两个模块叠着跑」:冻结的 base Linear(4bit 量化权重)+ 可训练的
lora_A/lora_B。默认走 PyTorch autograd,三个浪费:
- 中间激活全要存:autograd 为反向保留每个算子的输入,MLP 这种三层链存一堆张量;
- 4bit 权重每层前向反量化一次,反向再来一遍,临时张量反复分配;
- kernel launch 碎:X@A、@B、缩放、相加,全是小 launch。
2. 思路 / 直觉
LoRA 的前向是线性的:out = X@(W + s·A@B)ᵀ。它的反向有闭合形式——比如 MLP 里
dC/dA = X.T @ (…) @ B.T、dC/dB = A.T @ X.T @ (…)。既然如此:
- 不建 autograd 图,写一个
autograd.Function,反向里用addmm_(融合乘加,原地累 加)按公式直接算出dA, dB, dX; - 反向需要的激活只有
X, e, g(两个投影输出),其余中间量重算或根本不需要; - 4bit 权重在反向需要
W时临时反量化、用完即删(fast_lora.py:193-202)。
「无降级」的底气就在这:每一步都是初等代数恒等变形,没有任何数值近似;kernel 头部的注释
把完整推导都写出来了(unsloth/kernels/fast_lora.py:29-64)。
3. 零件一:matmul_lora —— 一次调用完成「反量化 + XW + XAB」
matmul_lora(unsloth/kernels/utils.py:1128-1168)是所有 LoRA 快路径的公共矩阵乘:
# 示意,非源码 —— matmul_lora 在做什么
W = fast_dequantize(W, W_quant) # 4bit → bf16,临时张量
out = X @ W.t() # 冻结主干
if A is not None: # LoRA 增量
out.addmm_(X @ A.t(), B.t(), alpha = s) # out += s·(XA)B
del W # 量化态用完即释放
要点:
W_quant是 bitsandbytes 的quant_state,由get_lora_parameters从 PEFT 层里一并取 出(utils.py:335-397,含 FP8 scale 的兼容分支);addmm_把「乘 + 加 + 缩放」融成一次调用,不产生新张量;- FP8 权重走另一条
fp8_linear路径(utils.py:1153-1154)。
4. 零件二:LoRA_MLP —— 一整个 MLP 的手工 autograd
LoRA_MLP(unsloth/kernels/fast_lora.py:28-229)把 SwiGLU MLP(gate/up/down 三个 LoRA
投影 + 激活)包成一个 Function:
前向(:91-112):三个 matmul_lora 分别算 e = X@G'、g = X@U'、i = h@W',激活
函数用 02 章 的 swiglu_fg_kernel。存给反向的只有 A/B 矩阵和
X, e, g 三个张量(:110)。
反向(:116-229),按注释里的推导执行:
# 示意,非源码 —— LoRA_MLP.backward 骨架
DW = matmul_lora(dY, downW.t(), ...) # dY 穿过 down 投影(含 LoRA)
h, df, de = swiglu_DWf_DW_dfg_kernel(DW, e, g) # 一个 kernel 吐三个梯度
d_downA.addmm_(h.t(), dY @ downB.t(), alpha = downS, beta = 0) # dC/dA
d_downB.addmm_(downA.t() @ h.t(), dY, alpha = downS, beta = 0) # dC/dB
upW = fast_dequantize(upW.t(), upW_quant) # 反量化,算 dX 用
dX = torch.matmul(df, upW.t(), out = X) # inplace:直接写回 X 的内存
...
真实代码见 fast_lora.py:156-204。注意两处「抠门」:out = X 把 dX 写回输入缓冲区
(:194,ctx.inplace 控制);fast_dequantize 出的临时权重 del 释放(:195)。
返回值顺序与 forward 的参数一一对应,不需要梯度的位置全是 None(:209-229)。
QKV 侧同理:LoRA_QKV(fast_lora.py:335-541)一次前向出 Q、K、V,反向按
dC/dAq = X.T @ D(Wq) @ B.T 这组对称公式算六个 LoRA 梯度(推导注释 :348-364);
O 投影是单权重版 LoRA_W(:574)。
5. 装配:patch_peft_model 什么时候才换上快路径
快路径不是无条件替换。FastLlamaModel.patch_peft_model(unsloth/models/llama.py:3752-3975)
的流程:
-
按
model_type选 MLP 变体(SwiGLU / GeGLU 近似,llama.py:3770-3790,未实现的模型族 直接NotImplementedError); -
门槛检查:只有
lora_dropout == 0 and bias == "none"(llama.py:3848)且各投影lora_A存在、无 bias、无 DoRA magnitude 向量(:3861-3870)才替换; -
替换动作本身是实例级赋值:
mlp_module.forward = types.MethodType(_apply_lora_mlp, mlp_module)(llama.py:3878);layer.self_attn.apply_qkv = apply_lora_qkv(llama.py:3901);layer.self_attn.apply_o = apply_lora_o(llama.py:3919)。
-
不满足条件的层打印 warning 走 PEFT 默认路径,并汇总「patch 了 N 层中 n_qkv/n_o/n_mlp」 (
llama.py:3924-3929)。
注意 apply_qkv/apply_o 这两个钩子:快版注意力前向(01 章 的
LlamaAttention_fast_forward)调的是 self.apply_qkv(self, hidden_states)
(llama.py:719),默认指向 original_apply_qkv(llama.py:156-160),有 LoRA 时被换成
apply_lora_qkv——注意力代码不区分有没有 LoRA,差的那部分全在这两个函数指针里。
6. 与 PEFT 的关系
Unsloth 不重新实现 LoRA 的「定义」(配置、权重初始化、adapter 管理全用 PEFT 的,
get_peft_model 内部照常调 PEFT),它替换的是 LoRA 的「执行」。所以:
- 你的
peft_config、adapter 存取、多 adapter 语义原样保留; - DPO 这类「禁用 adapter 跑参考模型」的场景,
get_lora_parameters在disable_adapters/merged时返回A=B=None,matmul_lora自动退化为纯X@W(utils.py:371-373)——快路径与「关掉 LoRA」同一份代码。
7. 坑
- 带 bias 的模型(如 Qwen)走不了 QKV 快路径,只有 warning 提示;这是「无降级」的另一
面:不愿意为带 bias 的变体再推一套公式,宁可退回慢路径也不近似(
llama.py:3903-3912)。 - 反向里的 dtype 提升:A/B 在反向被
.to(dtype)+ 转置(fast_lora.py:138-154), fp16 老卡(T4)会被抬成 fp32——省显存与兼容性的权衡写死在代码里。 - 手工 autograd 怕上游改结构:MLP/QKV 的数学绑定死 SwiGLU/GQA 结构,新架构(新激活、
MLA 等)都要重新推导——模型族列表(
llama.py:3770-3790)就是已推导清单。