跳到主要内容

数据截至 (上游 commit bd2a0fc7c314)

05 · PEFT:LoRA、DoRA 与 QLoRA

这一章讲什么: torchtune 没有依赖 HuggingFace 的 peft 库,torchtune/modules/peft/ 四个文件就是它全部的参数高效微调实现。本章讲透五件事:LoRALinear 这个替换件长什么样、builder 怎么把普通 Linear 换成它、参数冻结靠什么保证没出错、DoRA 比 LoRA 多算哪一步、QLoRA 的 4 bit 基座怎么和适配器拼在一起。

贯穿全章的主线例子: 全章只用一个输入 —— Llama-3.1-8B 的第 0 层注意力里那个 4096×4096 的 q_proj,换成 rank=8alpha=16 的 LoRA 版,配置就取官方的 recipes/configs/llama3_1/8B_lora_single_device.yaml。后面每一节开头的「主线例子 · 第 N 步」都是这同一条线上的一步,每步旁边写出具体的数。


1. 它要解决的小问题

先算一笔账。Llama-3.1-8B 的参数表写死在 torchtune/models/llama3_1/_model_builders.py:122-131:词表 128,256、32 层、隐藏维 4096、前馈中间维 14336。按这张表逐项加起来是 8,030,261,248 个数

全量微调时,这 80.3 亿个数每一个都要配一份梯度(这个数该往哪个方向挪、挪多少),最常用的那个优化器 AdamW 还要再给每一个配两份历史统计量。换句话说,真正吃显存的不是模型本身,是模型的三四倍。 第 4 章讲的 optimizer-in-backward 就是在这一项上抢显存。

参数高效微调(parameter-efficient fine-tuning,业界缩写 PEFT)换了个方向:别动原来那 80.3 亿个数,只训练新加的一小撮。 原权重不训练,就不需要梯度,也不需要优化器状态,那三四倍的账直接归零。

torchtune 把这件事做成了「替换件」:模型结构一个字符不改,只把其中若干个 nn.Linear 换成一个行为兼容、但内部多挂了两片小矩阵的类。换哪些、换成什么、换完谁能训,就是这一章。

2. 思路:在旁边挂一条窄路

2.1 「低秩」到底在省什么

一个 4096×4096 的线性层,是一张 4096 行、4096 列的数表,共 16,777,216 个数。微调想让它变成 W + ΔWLoRA 的赌注是:微调需要的那个 ΔW 其实很「瘦」,不必用一张同样大的数表去表示。

具体做法:把 ΔW 写成两张窄矩阵相乘 —— 一张 A 是 8 行 4096 列,一张 B 是 4096 行 8 列,B @ A 出来正好还是 4096×4096。中间那个 8 就叫秩(rank,代码里的 rank/配置里的 lora_rank),它是这条路最窄处的宽度。

代价与收益一眼可见:

参数量相对 4096×4096
原权重 W16,777,216100%
A(8×4096)32,7680.195%
B(4096×8)32,7680.195%
A + B 合计65,5360.39%,即 1/256

能表示的东西确实变少了 —— 秩为 8 的矩阵撑不出一张满秩 4096×4096 数表能表达的全部变化。LoRA 论文的主张是:微调这件事本来就不需要那么多自由度。torchtune 只是实现者,LoRALinear 的 docstring 直接指向 2021 年那篇原论文(torchtune/modules/peft/lora.py:28)。

2.2 LoRALinear 逐行(主线例子 · 第 1 步)

主线例子 · 第 1 步:q_projnn.Linear(4096, 4096, bias=False) 变成 LoRALinear(in_dim=4096, out_dim=4096, rank=8, alpha=16)。这一层的参数从 16,777,216 个变成 16,777,216 + 65,536 = 16,842,752 个,多了 0.39%。

LoRALinear(torchtune/modules/peft/lora.py:27)的构造函数只做四件事(:82-102):

  1. 先造一个普通 nn.Linear 拿它的权重,再把这份权重重新登记成自己的 weight 参数(:82-94)——所以它的基座权重键名就叫 weight,和被替换掉的那层完全同名,checkpoint 里对得上。
  2. lora_a = nn.Linear(4096, 8, bias=False)lora_b = nn.Linear(8, 4096, bias=False)(:99-100)。这两片是唯一会被训练的东西,业界管它们叫这一层的适配器(adapter)
  3. dropout(训练时随机把一部分输入丢掉,免得模型太依赖某几个通道)是可选的:配置里 lora_dropout 大于 0 才真放一个,否则放 nn.Identity() 占位(:98)。
  4. self.disabled 置 False(:93)——这个开关第 5 节专讲。

前向只有五行(torchtune/modules/peft/lora.py:128-147):

# 示意,非源码(对应 :137-147,略去了量化分支与 bias)
out = F.linear(x, self.weight, self.bias) # 走原来那条路
if self.disabled:
return out # 开关拨到「关」,窄路整条跳过
lora_out = self.lora_a(self.dropout(x)) # x[..., 4096] → [..., 8]
lora_out = (self.alpha / self.rank) * self.lora_b(lora_out) # [..., 8] → [..., 4096]
return out + lora_out

注意 B @ A 从来没被真的乘出来过 —— 那张 4096×4096 的表从头到尾不在显存里(:145 那一行乘的是 lora_b(lora_out),不是两片权重相乘)。输入先被压到 8 维再放回 4096 维,两次矩阵乘的代价是 2 × 4096 × 8,而原来那条路是 4096 × 4096。按 batch 2、序列长 2048 算(合计 4096 个 token;batch 2 取自官方配置的 batch_size: 2,序列长 2048 是为演示设的,那份配置里 max_seq_len 写的是 null):原路 68.7 G 次乘加,窄路 0.27 G 次,多出来的算力开销是 0.39%。省显存没有以「变慢一大截」为代价。

2.3 初始化:B 必须是全零(主线例子 · 第 2 步)

主线例子 · 第 2 步:lora_a 按 Kaiming uniform(一种按这一层输入宽度定随机幅度的常用初始化)随机填,lora_b 全部填 0。此刻 lora_out = (16/8) × B(A(x)) = 0,这一层的输出与替换之前逐位相同。

两个初始化函数就在文件末尾,各三行(torchtune/modules/peft/lora.py:309-320)。B 全零这一条是承重的: 它保证训练的第 0 步模型行为与原模型完全一致,微调是从原模型出发一点点偏离,而不是一上来就被一片随机噪声推歪。

反过来 A 不能也全零 —— 那样 B 的梯度会恒为 0,窄路永远学不动。一头随机、一头置零,是这类结构的标准写法,torchtune 的注释直接给出了它抄的那份参考实现的行号(:111-112)。

2.4 alpha / rank 那个系数

前向里乘的 self.alpha / self.rank(:146)是个纯粹的实用设计:它让你改 rank 时不必重调学习率。 秩越大,B @ A 的输出量级越大;除以 rank 把这份放大抵消掉,alpha 才是你真正想调的那个旋钮。

官方配置里的注释把经验值写在了行尾:lora_rank: 8 # higher increases accuracy and memorylora_alpha: 16 # usually alpha=2*rank(recipes/configs/llama3_1/8B_lora_single_device.yaml)。alpha=16rank=8,系数就是 2。 这个 2 会在第 8 节合并权重时原样再出现一次。

3. LoRA builder:普通 Linear 是怎么被换掉的(主线例子 · 第 3 步)

主线例子 · 第 3 步:配置里那行 _component_: torchtune.models.llama3_1.lora_llama3_1_8b 装配出整个模型 —— 32 层里每层换掉 6 个 Linear,全模型多出 19,660,800 个新参数。

第 2 章讲过 torchtune 的两级 builder:型号 builder 填参数表,组件 builder 拼积木。LoRA 走的是同一套,没有任何「往现成模型里注入适配器」的运行时改写

lora_llama3_1_8b(lora_attn_modules=['q_proj','v_proj','output_proj'],
apply_lora_to_mlp=True, lora_rank=8, lora_alpha=16)
│ _model_builders.py:83 —— 只填 8B 的参数表(:122-131)

lora_llama3_1(...) _component_builders.py:138
│ for _ in range(32):
├──► lora_llama3_attention(...) :276 —— 换 q/k/v/output_proj
├──► lora_llama3_mlp(...) :422 —— 换 w1/w2/w3
└──► 其余积木原样(RMSNorm、RoPE、TransformerSelfAttentionLayer 外壳)

替换的开关就一行:adapter_cls = DoRALinear if use_dora else LoRALinear(:337)。 之后每个投影都是同一个三选一的表达式,下表拿名单内的 q_proj(4096→4096)和名单外的 k_proj(4096→1024)对照(:338-369):

这个投影在 lora_attn_modules 里吗quantize_base造出来的是
在(q_proj)任意adapter_cls(4096, 4096, rank=8, alpha=16, quantize_base=...)
不在(k_proj)Falsenn.Linear(4096, 1024, bias=False)(原样)
不在(k_proj)TrueFrozenNF4Linear(4096, 1024, bias=False)(第 7 节)

配置 ['q_proj','v_proj','output_proj'] + apply_lora_to_mlp=True 落到每一层是 6 个替换件,算出来的新参数是:

被换掉的投影形状新增参数(8 × (in + out))
attn.q_proj4096 → 409665,536
attn.v_proj4096 → 102440,960
attn.output_proj4096 → 409665,536
mlp.w1 / mlp.w34096 → 14336(各一份)147,456 × 2
mlp.w214336 → 4096147,456
每层合计614,400
32 层合计19,660,800

k_proj 不在名单里,保持普通 nn.Linear;apply_lora_to_output: False,所以最后那个 4096→128,256 的输出投影也保持原样(:253-257)。这几个参数量是按源码里那张参数表逐项乘出来的,不是跑起来实测的。

两个值得记住的边角:

  • quantize_base 管不到最终输出投影。 源码里就一行 TODO 挑明了(torchtune/models/llama3_1/_component_builders.py:251),第 7 节算显存时这一项要单独扣出来。
  • QLoRA 不是另一套 builder。 qlora_llama3_1_8b = partial(lora_llama3_1_8b, quantize_base=True)(torchtune/models/llama3_1/_model_builders.py:250)——一行 partial(给一个函数预先钉死某个参数,得到一个新函数)就是全部,没有独立实现。

4. 冻结与三道校验

4.1 谁算适配器参数:模块自己报名(主线例子 · 第 4 步)

主线例子 · 第 4 步:get_adapter_params(model) 扫出 384 个参数名 —— 32 层 × 6 个替换件 × 2 片矩阵。

torchtune 不靠「名字里带 lora 就算」这种字符串规则来定谁是适配器,而是让模块自己声明。AdapterModule(torchtune/modules/peft/_utils.py:20)是个只有一个方法的协议:adapter_params() 返回本模块内适配器参数的名字。LoRALinear 的实现就是返回 ["lora_a.weight", "lora_b.weight"](torchtune/modules/peft/lora.py:125)。

get_adapter_params(_utils.py:39)遍历所有子模块,凡是有这个方法的就照它报的名字去取参数,取到一个就从待办清单里划掉一个,最后用一句断言(条件不成立就当场抛错停下)要求清单必须为空(:62-64):

# 示意,非源码(对应 :57-64)
for n, p in v.named_parameters(recurse=True):
if n in current_adapter_params:
adapter_params.update({f"{k}.{n}": p})
current_adapter_params.remove(n)
assert current_adapter_params == [], f"Adapter params {current_adapter_params} not converted"

这道断言防的是「报了名却不存在」。 比如有人把 lora_a 改名成 lora_A 却忘了同步 adapter_params(),断言当场炸;要是没有它,那片矩阵会被静静地当成基座权重冻掉,训练照跑、loss 照降,只是效果差一截 —— 这种 bug 最难查。源码里 adapter_params() 上方那句 NOTE 说的正是这件事(lora.py:123-124)。

4.2 冻结只有一行(主线例子 · 第 5 步)

主线例子 · 第 5 步:set_trainable_params 把模型里 8,049,922,048 个数中的 8,030,261,248 个标成「不训练」,剩 19,660,800 个可训练 —— 占 0.244%。

# torchtune/modules/peft/_utils.py:82-83
for k, v in model.named_parameters():
v.requires_grad_(k in adapter_params)

关键是它遍历的是全部参数,而不是只把适配器打开。 写成「只 requires_grad_(True) 那 384 个」看着等价,实则漏掉了「其余的必须显式关掉」这一半 —— nn.Parameter 默认就是要训练的。这一行同时做了两件事,所以不存在漏冻的缝。

省下来的是什么,一张表说清(比较对象是同一个模型跑全量微调):

这一项的大小正比于全量微调LoRA缩小
梯度8,030,261,24819,660,800408 倍
优化器状态同上同上408 倍
基座权重(不变)8,030,261,2488,030,261,2481 倍

最后一行就是 QLoRA 存在的理由:LoRA 削掉的是梯度和优化器状态,基座权重那 16.06 GB(bf16,每个数 2 字节)一个字节没少。 那一刀在第 7 节。

多模态模型走的是同一个函数的另一种用法:DeepFusionModel 按「编码器 / 解码器 / 融合层各自训不训」拼出一个参数名集合,末尾调的还是同一个 set_trainable_params(torchtune/modules/model_fusion/_deep_fusion.py:82-95)。传进去的是这个集合而不是适配器字典,那一行照样管用。

4.3 灌权重之后那道 missing / unexpected 校验(主线例子 · 第 5 步续)

主线例子 · 第 5 步续:灌完基座权重,load_state_dict(..., strict=False) 报回 384 个 missing 键;校验器确认这 384 个全是 LoRA 键,放行。

先把这个模型的键数摆出来。每层 9 个基座键(4 个注意力投影 + 3 个前馈投影 + 2 个归一化尺度)、12 个 LoRA 键;加上词嵌入、最终归一化、输出投影三个:

基座键:32 × 9 + 3 = 291 LoRA 键:32 × 12 = 384 合计 675

训练要灌两份权重,而且两份都必须用 strict=False —— 因为任何一份单独看都是不全的:HF 下载来的基座 checkpoint 里没有 LoRA 键,续训用的适配器文件里没有基座键。strict=False 的代价是「少了什么都不吭声」,于是 torchtune 补了一道自己的校验(validate_missing_and_unexpected_for_lora,_utils.py:324):

报回来的东西期望不符合时
base_missing只能是 LoRA 键出现基座键 → RuntimeError: Missing non-LoRA key ...
lora_missing只能是基座键出现 LoRA 键 → RuntimeError: Missing LoRA key ...
base_unexpected / lora_unexpected必须为空非空直接抛(:414-417)

判定靠三个集合求交(:366-378):先按「键名里有没有 loramagnitude」把模型的 675 个键分成 291 + 384 两堆,再看谁落进了不该落的那堆。

主线例子里这一步的两个具体失败长什么样:

  • 基座 checkpoint 少了一个分片,layers.7.mlp.w2.weight 没灌进去 → 它出现在 base_missing 里 → 报 Missing non-LoRA key layers.7.mlp.w2.weight from base model dict不然这一层前馈就是随机初始化的,训练能跑完,模型是废的。
  • 适配器是用 apply_lora_to_mlp=True 训的,续训时配置写成了 False → 模型里根本没有 mlp.w1.lora_a.weight 这类键,它们进 lora_unexpected → 报 Unexpected key loading adapter

4.4 另外两道

recipe 里紧跟着还有两道更粗但更快的:

校验查什么位置
validate_expected_param_dtype那 384 个适配器参数的数值精度是不是配置里写的 bf16torchtune/training/precision.py:153;单机 recipe 调用在 recipes/lora_finetune_single_device.py:470-472
validate_no_params_on_meta_device有没有参数还留在 meta device 上(第 4 章:只有形状、没有存储的占位状态)torchtune/training/_distributed.py:319-333;分布式 recipe 调用在 recipes/lora_finetune_distributed.py:600

第二道在 LoRA 分布式路径上尤其要紧,因为这条路径有个专门的绕行:LoRALinear.to_empty() 被重写成只给 lora_alora_b 分配真存储,基座 weight 故意留在 meta 上(torchtune/modules/peft/lora.py:104-108)。理由是基座权重马上要由 checkpoint 以 assign=True 整个替换掉,提前分配纯属浪费。故意留了一批 meta 参数,就必须有人在最后确认它们都被填上了。

recipe 侧的顺序因此是刻意的(recipes/lora_finetune_distributed.py:566-600):先给适配器分配存储并初始化 → 补建 RoPE 缓存 → 灌基座权重 → DoRA 补算长度 → 三道校验。

5. disabled:同一份权重当两个模型用(主线例子 · 第 6 步)

主线例子 · 第 6 步:进入 with disable_adapter(model),q_proj.disabled 从 False 变 True,前向在第二行就 return out —— 这一层退化成原封不动的 Llama-3.1-8B;退出 with 块变回 False。

这个字段存在的理由,注释写得明明白白(torchtune/modules/peft/lora.py:90-92):在 DPO 里,把带适配器的模型当作正在训练的那个,把关掉适配器的同一份权重当作参考模型。

DPO(直接偏好优化)要同时用到两个模型:一个在训,一个不动,损失比较两者对同一段文本的打分差。朴素做法是把基座模型再加载一份 —— 8B 模型多占 16.06 GB,正好是原模型的一倍。torchtune 的全量 DPO 配方就是这么干的:配置里另有一个 ref_checkpointer,_setup_reference_model 真的再建一个模型(recipes/full_dpo_distributed.py:215:514)。

而 LoRA 训练时基座权重本来就没变过 —— 关掉窄路,剩下的就是参考模型本身,不必再存第二份权重。

disable_adapter(_utils.py:275)是个上下文管理器,遍历所有模块,凡是既有 adapter_params 方法、又有 disabled 字段的就置 True,finally 里再置回 False(:296-312)。两个条件缺一不可,这样普通 nn.Linear 不会被误伤。

DPO recipe 的用法就一行(recipes/lora_dpo_single_device.py:543,分布式版在 recipes/lora_dpo_distributed.py:696):

with torch.no_grad(), disable_adapter(self._model):
reference_chosen_rejected_outputs = self.concatenated_forward(self._model, batch)

DoRALinear 有同名同义的字段和同样的提前返回(torchtune/modules/peft/dora.py:86:170-171),QATLoRALinear 也有(lora.py:261-262),所以这个上下文管理器对三种替换件一视同仁。

判断(我们的,不是源码里的): 用一个布尔字段代替「再开一份参考模型」,省的显存很实在,但它把「参考模型 = 基座权重」焊死在了实现里。想换个模型当参考,LoRA 这条路给不了入口,只能退回全量 DPO 配方那种「另配一个 checkpointer」的写法 —— 而那一份的显存账又回到了两倍。 如果错,会错在: 如果 DPO 的参考模型在实践中几乎总是取基座本身,这就不算约束,只是把常见情形做到了极致。我们没有在 torchtune 之外统计过这个比例。

6. DoRA:多分解出一个「长度」(主线例子 · 第 7 步)

主线例子 · 第 7 步:把配置里的 use_dora 改成 True,q_proj 变成 DoRALinear,比 LoRA 多一个长度为 4096 的向量 magnitude。全模型的可训练参数从 19,660,800 涨到 21,004,288,多 6.8%。

6.1 它和 LoRA 差的那一步

LoRA 直接把 W 改成 W + s·BA(s 就是那个 alpha/rank)。DoRA 的主张是:一个权重矩阵的每一行,可以拆成「朝哪个方向」和「有多长」两部分,这两部分应该分开训。

于是它给每个输出通道配一个可训练的长度标量,4096 个输出通道就是长度 4096 的向量 magnitude(torchtune/modules/peft/dora.py:94);方向那部分仍旧交给 BA 那条窄路。前向的最终形式是(:186-192):

输出 = ( magnitude / ‖W + s·BA‖每行 ) × (W + s·BA) x
└────── 学出来的长度 ──────┘ └── 归一化成纯方向 ──┘

源码把它拆成了两段等价的加法(:188-192),但展开后就是上面这一式。‖·‖每行 那一项被 detach() 掉了(:185)——它只当缩放系数用,不往回传梯度,否则「长度」这件事会被算两遍。

6.2 x_eye:为什么不敢直接写 B @ A

要算 ‖W + s·BA‖,得先有 BA 这张 4096×4096 的表。代码却绕了一大圈(torchtune/modules/peft/dora.py:176-181):

# Can't use raw matmul since FSDP hooks are attached to __call__
# Instead follow the approach in https://github.com/huggingface/peft/pull/1806
x_eye = torch.eye(self.lora_a.weight.shape[1], device=..., dtype=x.dtype) # [4096, 4096] 单位矩阵
lora_weight = self.lora_b(self.lora_a(x_eye)).T # 得到 BA

读法:与其直接乘两片权重,不如喂一个单位矩阵走一遍这两个模块的正常前向 —— 出来的正是 BA

理由就在那行注释里,而且第 4 章讲过它的另一半:FSDP2 把参数切成分片分给各张卡存着,完整参数只在模块的 __call__ 被调用时,由挂在它前面的 hook(到点自动执行的回调)临时聚合出来。直接写 self.lora_b.weight @ self.lora_a.weight 是绕过 __call__ 去摸属性,拿到的是分片、不是完整权重,结果是错的。过一遍单位矩阵,就是为了让那个 hook 有机会开火。

这一手不白拿:

代价具体多少参照物
临时张量单位矩阵 4096×4096 bf16 = 33.5 MB,结果又一份 33.5 MB这一层自己的基座权重也正好 33.5 MB
额外乘加2 × 4096 × 4096 × 8 ≈ 0.27 G 次,与 batch 大小无关batch 2 × 序列 2048 时,基座那条路是 68.7 G 次,占 0.39%

注意最后那句「与 batch 大小无关」: batch 越小,这笔固定开销占比越高 —— 单条序列的推理式前向下,它反而成了主要成本。

6.3 初始化必须晚一步,而且要与 LoRA 起点重合

magnitudetorch.empty(4096) 建出来时是未初始化的(dora.py:94),真正的值由 initialize_dora_magnitude() 补(:117),填的是 ‖W + s·BA‖每行(:137-139)。

因为 B 此刻还是全零,BA 就是零,填进去的其实是 ‖W‖每行 于是第一次前向时 magnitude / ‖W + s·BA‖ = 1,整式退化成 1 × W x —— 和第 2.3 节那个「起点与原模型逐位相同」的性质完全一致,docstring 里那句「its outputs are initially identical to standard LoRA's outputs」说的就是这个(:119-120)。

这带来两个硬约束,都在源码里有对应:

  1. 必须等基座权重灌完才能算。 它依赖 W 的真实数值。recipe 因此把这一步排在 load_from_full_model_state_dict 之后(recipes/lora_finetune_distributed.py:579-587);单机 recipe 更明确,注释直接写「for any adapters that need to be initialized after base weights have been loaded (e.g. DoRA)」(recipes/lora_finetune_single_device.py:444-449)。
  2. 参数还在 meta device 上就得报错。 函数开头三个 is_meta 检查,命中就抛 RuntimeError(dora.py:127-136)——因为 meta 上的张量没有数值,算出来的「每行的长度」全是垃圾数,而这种错不会在训练时暴露。

还有一个只有 DoRA 才有的小麻烦:magnitude 是挂在模块自己身上的参数,不像 lora_a/lora_b 那样躲在子模块里,所以 to_empty() 要额外用 torch.utils.swap_tensors 把它原地换成新设备上的同形张量,并保留 requires_grad(dora.py:103-108)。

7. QLoRA:把冻住的那一大坨压成 4 bit(主线例子 · 第 8 步)

主线例子 · 第 8 步:把配置换成 qlora_llama3_1_8b,q_proj 的基座 weight 变成 4 bit 表示,k_proj(不在适配器名单里)整个换成 FrozenNF4Linear。全模型基座从 16.06 GB 降到约 5.6 GB。

第 4.2 节留下的问题是:LoRA 削掉了梯度和优化器状态,基座权重本身一个字节没少。既然它全程被冻住、永远不更新,那就没有理由用 bf16 存着 —— 用一种更省位数的编码存起来,前向时临时还原成高精度参与计算就行。这就是 QLoRA。

torchtune 用的编码叫 NF4(4 bit NormalFloat,每个数只占 4 个二进制位,是 bf16 的四分之一),实现来自 torchao(PyTorch 官方的低精度与量化库)的 to_nf4,torchtune 侧只负责决定在什么地方调它。

7.1 两种基座层,按「这一层有没有适配器」分

这一层有没有 LoRA基座权重位置
LoRALinear(quantize_base=True)构造时 to_nf4(linear.weight)torchtune/modules/peft/lora.py:83-87;前向走 linear_nf4(:137-140)
没有FrozenNF4Linear构造时先 requires_grad_(False),再 to_nf4 并用 swap_tensors 换掉自己的 weighttorchtune/modules/low_precision/nf4_linear.py:16:46-54:67

FrozenNF4Linear 有一处细节值得记:它先把 requires_grad 关掉再量化(:46-48),而不是反过来 —— 类名里的 Frozen 是构造契约的一部分,不依赖后面 set_trainable_params 那一行补救。它的前向注释也点明了收益的来源:计算在高精度下做,但反向只保存 4 bit 那份用于求梯度,所以不会把省下的显存又吐回去(:57-60)。

按主线例子的配置算这笔账:

哪部分参数量bf16NF4
32 层里的全部注意力与前馈投影6,979,321,85613.96 GB约 3.49 GB
词嵌入 + 输出投影 + 各归一化尺度(不量化)1,050,939,3922.10 GB2.10 GB
合计8,030,261,24816.06 GB约 5.6 GB

这几个 GB 是按位宽乘参数量算出来的,不是实测显存,也没算量化尺度本身占的那点额外空间。第二行不量化的原因见第 3 节那条 TODO。省下的 10.5 GB,比一张 24 GB 显存的消费级显卡的一半还多(显卡容量这个参照物不在源码里,来自通用知识)。

7.2 量化发生在哪一刻:单机和分布式不在同一处

这是 QLoRA 在 torchtune 里最容易看漏的一处分叉。

  • 单机:模型直接在真设备上实例化(recipes/lora_finetune_single_device.py:410-411),构造函数当场量化;随后 load_state_dict 把 bf16 的基座权重拷进这些 4 bit 参数里。
  • 分布式:模型在 meta device 上实例化,构造时无从量化。真正动手的地方在灌权重的函数里 —— 它先检查目标参数的底层张量是不是 NF4,是就用目标参数记着的分块粒度把这份 bf16 权重现场量化,再切成分片装进去(torchtune/training/_distributed.py:399-407)。

同一个文件开头还有一处相关判断:整个模型只要含 NF4 参数,就不走 PyTorch 那套通用的分布式状态字典接口,而是回退到手写路径,注释写明是因为那套接口现在还不支持 NF4(:371-378)。取回权重时同样要手写聚合(_gather_nf4_tensor,:462)。

7.3 保存时的反量化 hook

存 checkpoint 时要写出去的是能被别人加载的普通权重,不是 4 bit 表示。torchtune 的做法是给模型挂一个 state_dict 后置 hook —— state_dict() 是 PyTorch 里「把模型全部权重摊成一个『名字 → 张量』字典」的标准接口,后置 hook 就是在它返回之前插一手:把里面每个 NF4 张量还原成 bf16,并顺手搬到 CPU(torchtune/modules/common_utils.py:24,循环体在 :55-59)。

搬到 CPU 这一下是重点:不搬的话,还原出来的 16.06 GB 会和正在训练的模型挤在同一张卡上,保存瞬间显存翻倍,前面省的全白省。

挂钩子的地方就在 builder 末尾三行 —— quantize_base 为真才挂(torchtune/models/llama3_1/_component_builders.py:268-271),注册函数在 torchtune/modules/common_utils.py:206:235-237换句话说,「怎么存」这件事在拼模型的时候就决定好了,recipe 完全不用知情。

8. 收工:合并回去(主线例子 · 第 9 步)

主线例子 · 第 9 步:训练结束存盘。适配器那 384 个键先被单独摘出来存一份;主权重则执行 W += 2 × B @ A,然后把那 384 个键删掉 —— 键数回到 291,形状与原始 Llama-3.1-8B 一模一样。

两件事在 CheckpointClient 里挨着发生(torchtune/training/checkpointing/_checkpoint_client.py:362-372):

  1. get_adapter_state_dict 按「键名里有 loramagnitude」摘出适配器(torchtune/modules/peft/_utils.py:132),存成单独文件 —— 续训和 PEFT 导出都用它(格式转换见第 3 章)。
  2. get_merged_lora_ckpt(_utils.py:194)就地改写主权重字典。

合并本身是一行 W += (alpha/rank) × B @ A(_utils.py:264-266)——这里那个 2 又出现了,和前向里乘的是同一个数,所以合并后的模型与训练时的模型在数值上等价。合并完删掉 lora_a.weightlora_b.weight(:268-269)。

DoRA 走另一个分支(:244-258),多两步:先算 W + s·BA,再整体乘上 magnitude / ‖W + s·BA‖每行,最后删掉 magnitude这与第 6.1 节前向那一式是同一个公式,只是把逐次前向要算的东西一次性烙进权重里。

合并这个函数自己带一个开关 use_distributed_barriers。打开时,它会在每次矩阵乘之前插一道 dist.barrier()(路障:所有进程都跑到这一行,才一起往下走),把各进程对齐(_utils.py:228-229:247-248)。

但本节这条同步存盘路径压根没打开它 —— :370-372 调用时根本不传这个参数,默认就是 False。原因在几十行之后:非分布式 checkpointer 的情况下,合并只在 0 号进程上跑,而且跑的是已经聚合到这一个进程上的完整权重;别的进程只在函数外面等一道 barrier(:396-401)。函数内部没有第二个进程要对齐,自然用不着。

真正打开这个开关的是另外两条路:异步保存(_checkpoint_client.py:177-182)、从分布式断点续训时的那次合并(:543-547)。两处都写成 use_distributed_barriers=not single_device —— 单机跑仍然传 False。

配置里的 save_adapter_weights_only 决定要不要跳过合并:置 True 就只留那份 384 键的适配器文件,产物从十几 GB 变成几十 MB(recipes/lora_finetune_single_device.py:172:579)。

主线例子全表

发生了什么具体的数 / 状态讲在哪一节
1q_proj 换成 LoRALinear16,777,216 → 16,842,752 个参数(+0.39%)§2.2
2初始化A 随机、B 全零 → 输出与替换前逐位相同§2.3
3builder 拼完整模型32 层 × 6 个替换件,新增 19,660,800§3
4扫出适配器参数384 个键;断言「报了名的都找到了」§4.1
5冻结 + 三道校验可训练 19,660,800 / 8,049,922,048 = 0.244%;base_missing 恰好 384 个 LoRA 键§4.2–4.4
6DPO 关掉适配器disabled = True,前向第二行返回,不必再存第二份权重§5
7换成 DoRA多一个长度 4096 的向量;初始时缩放系数恒为 1§6
8换成 QLoRA基座 16.06 GB → 约 5.6 GB§7
9合并存盘W += 2 × B @ A;键数 675 → 291§8

9. 关键细节 / 坑

  • 换了 LoRA 超参,旧适配器就续不上,而且两种改法报的错不一样。apply_lora_to_mlp 之类的开关,是模型里该有哪些 LoRA 键变了 —— 第 4.3 节那道校验报 Unexpected key loading adapter;只改 rank,键名不变但每片矩阵的形状变了 —— load_state_dict 直接报 size mismatch,连 strict=False 也拦不住。两种都是好事 —— 不报的话你会拿到一个静默错配的模型。
  • get_adapter_state_dict 的过滤规则是字符串匹配(_utils.py:132),所以任何自定义模块都不许在参数名里出现 loramagnitude 字样,否则会被当成适配器摘走。
  • MoE 层的 LoRA 是另一套键名。 MoE(混合专家,一层里并排放好几份前馈网络、每个 token 只走其中几份)的适配器叫 lora_gate_a / lora_down_b 这类名字,单独由 _get_lora_moe_modules 处理(_utils.py:165),合并走另一个分支,而且DoRA 目前不支持 MoE,源码里就一句 TODO(:222)。
  • QAT 与 LoRA 叠加靠事后换件。 QAT(quantization-aware training,训练时就模拟量化后的数值误差,让模型提前适应)在 torchtune 里是这样接上的:swap_lora_linear_with_qat(torchtune/training/quantization.py:205)递归遍历模型,把每个 LoRALinear 就地换成 QATLoRALinear 并搬走权重(:231-238)。换件函数有两条硬拒绝:带 bias 的、以及已经 quantize_base 的,都直接抛错(torchtune/modules/peft/lora.py:280-283)——QAT 和 QLoRA 不能同时用
  • 部分模型族不能导出成 PEFT 格式(即 HuggingFace peft 库那套适配器文件命名,见第 3 章)。Phi3/Phi4、Llama3.2 Vision、Llama4 的适配器只存 torchtune 自己的格式,保存时仅打一条 warning;Llama4 走 LoRA 时更硬:save_adapter_weights_only 不是 True,recipe 在 setup 阶段就抛 ValueError 拒绝开训,连第一步都跑不到(recipes/lora_finetune_distributed.py:306-312)。
  • TrainableParams 这个枚举只服务多模态。 它有 full/lora/frozen 三档(torchtune/modules/peft/lora.py:21-24),供 Llama3.2 Vision 的 builder 按「编码器、解码器、融合层」分别设定;而且混用 full 与 LoRA 的支持被临时移除了,builder 里一句断言拦着(torchtune/models/llama3_2_vision/_model_builders.py:168-172)。

10. 代码地图

主题文件路径符号名
LoRA 替换件本体torchtune/modules/peft/lora.pyLoRALinearLoRALinear.forwardLoRALinear.to_empty_lora_a_init_params_lora_b_init_params
QAT + LoRAtorchtune/modules/peft/lora.pyQATLoRALinearQATLoRALinear.from_lora_linear
QAT 换件入口torchtune/training/quantization.pyswap_lora_linear_with_qat
DoRA 替换件torchtune/modules/peft/dora.pyDoRALinear.forwardinitialize_dora_magnitude_get_weight_normDoRALinear.to_empty
适配器协议与收集torchtune/modules/peft/_utils.pyAdapterModuleget_adapter_paramsget_adapter_state_dict
冻结与校验torchtune/modules/peft/_utils.pyset_trainable_paramsvalidate_missing_and_unexpected_for_loraget_lora_module_names
关闭适配器torchtune/modules/peft/_utils.pydisable_adapter
合并回基座torchtune/modules/peft/_utils.pyget_merged_lora_ckpt_get_lora_modules_get_lora_moe_modules
LoRA builder(型号级)torchtune/models/llama3_1/_model_builders.pylora_llama3_1_8bqlora_llama3_1_8b
LoRA builder(组件级)torchtune/models/llama3_1/_component_builders.pylora_llama3_1lora_llama3_attentionlora_llama3_mlp
量化基座层torchtune/modules/low_precision/nf4_linear.pyFrozenNF4Linear
保存时反量化torchtune/modules/common_utils.pyreparametrize_as_dtype_state_dict_post_hook_register_reparametrize_state_dict_hooks
分布式下的 NF4 灌权重torchtune/training/_distributed.pyload_from_full_model_state_dict_gather_nf4_tensorvalidate_no_params_on_meta_device
LoRA recipe(分布式)recipes/lora_finetune_distributed.pyLoRAFinetuneRecipeDistributed._setup_model
LoRA recipe(单机)recipes/lora_finetune_single_device.pyLoRAFinetuneRecipeSingleDevice._setup_modelsave_checkpoint
DPO 里关适配器recipes/lora_dpo_single_device.py训练循环内的 disable_adapter
存盘时摘适配器 + 合并torchtune/training/checkpointing/_checkpoint_client.pyCheckpointClient._save_checkpoint_sync 内的 adapter 分支
导出 PEFT 键名映射torchtune/models/convert_weights.py_TO_PEFT_KEYStune_to_peft_adapter_weights