数据截至 (上游 commit 5779b17b9a67)
03 · LoRA 层实现与权重合并
这一章讲什么: 注入完成后,每个目标层的真实长相——
LoraLayer的数据结构、forward 的四个分支、merge/unmerge/merge_and_unload三件套怎么把 ΔW 并回基座,以及量化基座(QLoRA 场景)是怎么被 dispatcher 链接住的。
1. 它要解决的小问题
一个 LoRA 层不只是「W 旁边加两个小矩阵」那么简单,它同时要支撑四件事:
- 多适配器:同一层挂多个 LoRA,运行时切换谁生效。
- 可逆:能合并进基座(推理提速),也能拆出来(继续训练或换适配器)。
- 量化基座:基座可能是 bnb 4bit 的,不能当普通
nn.Linear处理。 - 变体:DoRA、VeLoRA 等 LoRA 衍生方法要复用同一层骨架。
PEFT 的答案是把这层设计成一个带字典的包装器。
2. 数据结构:一层 = 原层 + 按名索引的参数字典
LoraLayer.__init__(src/peft/tuners/lora/layer.py:128)建立的全部状态:
self.base_layer = base_layer # 原层,被包住
self.r = {}
self.lora_alpha = {}
self.scaling = {}
self.lora_dropout = nn.ModuleDict({})
self.lora_A = nn.ModuleDict({}) # adapter 名 → A 矩阵
self.lora_B = nn.ModuleDict({}) # adapter 名 → B 矩阵
self.merged_adapters = [] # 已合并的 adapter 名(栈)
self.lora_magnitude_vector = torch.nn.ModuleDict() # DoRA 用
(src/peft/tuners/lora/layer.py:129-145,有删节)两个设计决策值得记住:
- 一切按 adapter 名索引。 加第二个适配器不用新建层,往字典里再插一套即可——这就是第 2 章「已包过的层只调
update_layer不再重包」能成立的原因。 - 类属性
adapter_layer_names/other_param_names声明哪些名字属于适配器(src/peft/tuners/lora/layer.py:110-126)。冻结、存档、状态检查都靠这份声明做结构化判定,而不是猜名字。
所有适配层(不止 LoRA)的公共协议在 BaseTunerLayer(src/peft/tuners/tuners_utils.py:1808):get_base_layer、set_adapter、enable_adapters、merge、unmerge——方法想接入 PEFT 生态,实现这套接口即可。
3. forward:四个分支
Linear.forward(src/peft/tuners/lora/layer.py:1035)开头先做两件事:校验参数、从 kwargs 里弹出 PEFT 私货(adapter_names 等,防止透传给基座层报错)。然后按状态分四路:
| 分支 | 条件 | 行为 |
|---|---|---|
| 禁用 | disable_adapters | 若已合并先 unmerge() 还原,然后纯基座前向 |
| 混合 batch | 传了 adapter_names | 走 _mixed_batch_forward,见 §4 |
| 已合并 | merged | 纯基座前向——合并后零开销就体现在这 |
| 正常 | 默认 | 基座前向 + 逐个叠加活跃适配器的支路 |
正常分支的核心循环(src/peft/tuners/lora/layer.py:1052-1063):
for active_adapter in self.active_adapters:
if active_adapter not in lora_A_keys:
continue
...
result = result + lora_B(lora_A(dropout(x))) * scaling
多个活跃适配器是直接相加的——多 LoRA 同时生效的语义就是支路求和。dtype 也有讲究:输入先被 cast 到 A 矩 阵的 dtype,最后结果再 cast 回基座输出的 dtype(src/peft/tuners/lora/layer.py:1061、:1074),这让 fp32 适配器可以骑在 4bit 基座上。
4. 一个 batch 里混用不同适配器
推理服务场景:同一批请求里,有的要用「客服 LoRA」、有的要用「代码 LoRA」、有的要裸基座。_mixed_batch_forward(src/peft/tuners/lora/layer.py:816)的做法:
按 adapter_names 把 batch 行分组
├── "__base__" 组:跳过,只保留基座输出
└── 每组:取子 batch 过该适配器支路,按行号 += 回 result
即每个适配器只算自己那几行,而不是全 batch 逐适配器串行。特殊名字 "__base__" 表示该行不要任何适配器。
5. 合并三件套:merge / unmerge / merge_and_unload
5.1 层级的 merge 与 unmerge
Linear.merge(src/peft/tuners/lora/layer.py:915)把 get_delta_weight 算出的 (α/r)·B@A 加进基座权重,并把 adapter 名压入 merged_adapters 栈(src/peft/tuners/lora/layer.py:980)。unmerge(src/peft/tuners/lora/layer.py:982)反向:弹栈、把同样的 delta 减掉——加减对称,所以可反复合并/拆分。
两条安全设计:
safe_merge=True先克隆再检查。 在副本上加 delta、torch.isfinite全量检查通过才落笔;发现 NaN 直接报错「adapter seems to be broken」,基座权重保持原样(src/peft/tuners/lora/layer.py:936-950)。- CPU 上半精度合并先升 fp32。 一些 CPU 的 fp16/bf16 matmul 很慢,所以 CPU 设备上算 delta 时先
.float(),算完再转回去(src/peft/tuners/lora/layer.py:1016-1031)。
5.2 模型级的 merge_and_unload 与 unload
层级 merge 只改权重,层还在(包装还在)。要彻底还原成普通模型 ,用 BaseTuner.merge_and_unload(src/peft/tuners/tuners_utils.py:738)或 unload(src/peft/tuners/tuners_utils.py:776)。两者都走 _unload_and_optionally_merge(src/peft/tuners/tuners_utils.py:675):
遍历模块树(跳过 tuner 层内部)
每个 BaseTunerLayer:
merge=True → 先 target.merge() 再替换回 base_layer
merge=False → 直接替换回 base_layer(丢弃适配器)
最后 del model.peft_config # 防止再次 get_peft_model 时误报
一个处理得很细的边角:合并后如果输入/输出 embedding 已经分叉(比如训练时 resize 过词表),会把 config 里的 tie_word_embeddings 改成 False 并发警告(src/peft/tuners/tuners_utils.py:714-732)——不然后续保存/加载会按「权重共享」的假设处理两份其实已经不同的权重。
6. 量化基座:dispatcher 链
QLoRA 场景里基座层不是 nn.Linear,而是 bnb 的 Linear4bit。创建适配层时 _create_new_module(src/peft/tuners/lora/model.py:394)按顺序问一串 dispatcher「这个 target 你接不接」,第一个返回非 None 的赢(src/peft/tuners/lora/model.py:445-449):
| 顺序 | dispatcher | 接住什么 |
|---|---|---|
| 1 | 用户自定义 _custom_modules | 实验性自定义层 |
| 2 | dispatch_bnb_8bit / dispatch_bnb_4bit | bitsandbytes 量化层(src/peft/tuners/lora/bnb.py:288、:571) |
| 3 | eetq / aqlm / awq / gptq / hqq / inc / torchao / megatron / transformer_engine | 各自生态的量化层 |
| 4 | dispatch_default(src/peft/tuners/lora/layer.py:2656) | 普通 nn.Linear/Embedding/ConvNd/MultiheadAttention |
全都不接就抛错并列出支持的层类型(src/peft/tuners/lora/model.py:451-457)。顺序就是优先级:默认实现必须排最后,注释里明说了这一点(src/peft/tuners/lora/model.py:395-396)。
7. 变体系统:DoRA 们怎么复用同一层
LoRA 有一堆衍生方法(DoRA、Arrow、VeLoRA……),它们改的是 forward/merge 的细节,不改层的骨架。PEFT 用 LoraVariant 协议(src/peft/tuners/lora/layer.py:54)把它们插件化:
- 层类用
lora_variants属性声明自己支持哪些变体组合,key 是配置字段名的排序元组。Linear支持 8 种:("use_dora",)→DoraLinearVariant等(src/peft/tuners/lora/layer.py:900-913)。 resolve_lora_variant(src/peft/tuners/lora/layer.py:178)读 LoraConfig 的字段元数据,算出当前激活的变体组合,查表实例化;空元组()就是 vanilla LoRA。- forward/merge 里凡是看到
active_adapter in self.lora_variant就委托给变体对象,否则走 vanilla 路径(如前向src/peft/tuners/lora/layer.py:1062-1072)。
这套机制让「加一个 LoRA 变体」变成「写一个 Variant 类 + 在字典里登记」,不用动层代码。
8. 关键细节与坑
merge_and_unload不是原地操作,返回值才是新模型——文档字符串里用感叹号强调了这一点(src/peft/tuners/tuners_utils.py:747)。lora_bias=True+ 基座无 bias = 无法合并。 merge 时会抛RuntimeError(src/peft/tuners/lora/layer.py:954-958);建层时其实已有PeftWarning预警(src/peft/tuners/lora/layer.py:242-247)。- 换适配器权重不重建设备/结构:
hotswap_adapter。 热替换同名适配器的权重(如 A/B 测试新 checkpoint),目前只支持 LoRA(src/peft/utils/hotswap.py:613)。 - adapter 默认会升到 fp32。
_cast_adapter_dtype把 fp16/bf16 的适配器权重转成 fp32 求训练稳定(src/peft/tuners/tuners_utils.py:624-634);想要「纯」半精度训练需显式关掉autocast_adapter_dtype。 - 合并后想再训练要先
unmerge;forward里merged分支直接跳过支路计算,不会帮你自动拆(src/peft/tuners/lora/layer.py:1046-1047)。 - 多适配器线性组合的另一种形态:
add_weighted_adapter(src/peft/tuners/lora/model.py:681)把多个已训好的适配器按权重融合成一个新适配器(SVD/线性拼接等方式),和第 3 节「多活跃适配器求和」是不同的东西——前者产出新权重,后者是运行时叠加。
9. 代码地图(本章涉及)
| 主题 | 文件路径 | 符号名 |
|---|---|---|
| 层数据结构与建层 | src/peft/tuners/lora/layer.py | LoraLayer.__init__、LoraLayer.update_layer |
| 前向四分枝 | 同上 | Linear.forward、Linear._mixed_batch_forward |
| 合并/拆分 | 同上 | Linear.merge、Linear.unmerge、Linear.get_delta_weight |
| 模型级卸载 | src/peft/tuners/tuners_utils.py | BaseTuner.merge_and_unload、BaseTuner._unload_and_optionally_merge |
| 量化 基座分发 | src/peft/tuners/lora/model.py、src/peft/tuners/lora/bnb.py | LoraModel._create_new_module、dispatch_bnb_4bit |
| 变体协议 | src/peft/tuners/lora/layer.py、src/peft/tuners/lora/variants.py | LoraVariant、LoraLayer.resolve_lora_variant |
| 热插拔 | src/peft/utils/hotswap.py | hotswap_adapter |