数据截至 (上游 commit 1fe27b1b53f3)
01 · 补丁机制:怎么接管 transformers
这章回答:Unsloth 没有 fork transformers,它是怎么让你的 HF 模型「不知不觉」跑上手工 kernel 的?答案是一套分三个时机点火的 monkey-patch 体系。
1. 小问题
要换掉别人库里「类的行为」,常规做法是继承 + 让用户改用你的子类。但 Unsloth 的目标恰好是
用户代码一行不改——你还是 AutoModelForCausalLM、还是 SFTTrainer。那就只剩一条路:
运行时改别人的类(monkey-patch)。难的是三件事:
- 什么时候改(太早太晚都出错);
- 改哪些点(改少了快不起来,改多了必崩);
- 上游升级了怎么办(补丁点漂了)。
2. 思路 / 直觉
Unsloth 把补丁按「点火时机」分成三层:
| 时机 | 补什么 | 为什么在这个时机 |
|---|---|---|
import unsloth 时 | 环境变量、设备探测、日志/后端开关 | 必须在 transformers 被 import 之前落定 |
from_pretrained() 时 | 模型类的 forward 整组替换(pre_patch)、Trainer 训练循环重写 | 此刻才知道要补哪个模型族 |
get_peft_model() 后 | 逐层装 LoRA 快路径(patch_peft_model) | 只有 LoRA 权重存在后才知道能不能走快路径 |
一句话:越具体的补丁,越晚点火。
3. import 时机:必须先导入 unsloth
import unsloth 做的第一件事不是定义函数,而是抢在 transformers 前面布置环境:
unsloth/__init__.py:17设os.environ["UNSLOTH_IS_PRESENT"] = "1",标记自己在场;unsloth/__init__.py:20-37若 transformers 尚未被 import,主动把USE_TF/USE_FLAX设为"0"——因为 transformers 4.x 只要检测到装了 TensorFlow/Flax 就会去 import 它们, 而 Unsloth 用不上,导入即炸。注释原话:“It reads these variables once at its own import, so this has to land first.”
这就是「import unsloth 必须写在 import transformers 之前」的全部原因:它的补丁有一
半是环境变量,事后补没用。
4. pre_patch():把模型类的 forward 整组换掉
加载模型时,from_pretrained 先调 model_patcher.pre_patch()(unsloth/models/llama.py:2458)。
以 Llama 为例,FastLlamaModel.pre_patch(llama.py:2314-2347)做的事就是一连串类属性赋值:
# 示意,非源码 —— pre_patch 的核心动作
LlamaAttention.forward = LlamaAttention_fast_forward # 注意力前向
LlamaDecoderLayer.forward = LlamaDecoderLayer_fast_forward # 整层前向
LlamaModel.forward = LlamaModel_fast_forward # 模型主干
LlamaForCausalLM.forward = CausalLM_fast_forward(...) # 顶层(含融合 CE)
PeftModelForCausalLM.forward = PeftModel_fast_forward # PEFT 包装器也换
对应真实代码在 unsloth/models/llama.py:2326-2332。改的是类,不是实例:此后创建的
所有 Llama 模型——包括你自己 AutoModelForCausalLM.from_pretrained 出来的——都自动是快
版本。这解释了为什么补丁是「全局」的,也是边界(同进程里想跑原生 HF 就得 unpatch)。
5. 最难缠的补丁:exec 生成 RoPE 缩放的 __init__
RoPE 缩放(yarn/llama3/longrope 等)需要改 LlamaAttention.__init__,但 HF 的 __init__
源码随版本漂移,手写一个固定的替代品必然过时。Unsloth 的解法在
patch_llama_rope_scaling(unsloth/models/_utils.py:3084-3120):
inspect.getsource(attention_module.__init__)把上游当前的__init__源码抽成字符串;- 字符串级改写:换函数名、在末尾插入「按 config 选 rotary 模块」的逻辑;
exec(function, globals())编译后赋 回LlamaAttention.__init__(llama.py:2324-2325)。
直觉:不跟上游的代码,跟上游的「当前源码」——版本漂移被运行时读源码吸收了。失败模式
也诚实:抽不到源码(说明已被补过)就直接返回 (None, None) 跳过(_utils.py:3110-3113)。
6. 最粗暴的补丁:重写 HF Trainer 的训练循环
from_pretrained 里还有一段更大的 exec:Unsloth 要改 Trainer._inner_training_loop
(加自己的日志横幅、禁掉 TPU 分支、接入编译缓存重置),做法是(llama.py:2957-3044):
# 示意,非源码 —— 流程骨架
src = inspect.getsource(Trainer._inner_training_loop) # 抽源码
src = src.replace(original_debug, debug_info) # 正则/字符串级改写
src = src.replace("is_torch_tpu_available()", "False") # 关掉 TPU 分支
exec(src, globals()) # 重新编译
Trainer._inner_training_loop = _fast_inner_training_loop
真实代码里它还会先把 transformers.trainer 模块里所有在源码字符串中「出现过的名字」收集
起来 exec("from transformers.trainer import (...)") 造好命名空间(llama.py:2967-2974)——
这样重写后的函数体能找到原来的依赖。完成后再有自我校验:patch_peft_model 会检查
Trainer._inner_training_loop.__name__ == "_fast_inner_training_loop",不满足直接抛错
(llama.py:3815-3817)。
关键细节: 它保留 Trainer._original_training_loop(llama.py:2960),二次进入时用
原件再改——补丁可重复点火,不会把改过的版本当原料继续改。
7. 模块级替换与反悔通道
除了类方法,还有「模块里的类」级替换。例如 RMSNorm:
patch_rms_layernorm(unsloth/kernels/rms_layernorm.py:277-286):transformers.models.llama.modeling_llama.LlamaRMSNorm = Unsloth_LlamaRMSNorm;- 对应的
unpatch_rms_layernorm(rms_layernorm.py:289-298)把原类赋回去。
子类继承原类(Unsloth_LlamaRMSNorm(LlamaRMSNorm),rms_layernorm.py:261),只覆写
forward 一行去调 Triton kernel——构造签名、权重属性完全不动,上游代码无感知。
8. 坑与失败模式
- 上游签名漂移 = 补丁失效。 整套机制赌的是「被替换方法的调用契约不变」;transformers
升级改了
forward签名,快版本就要跟着改。源码里大量 issue 链接(如llama.py:2335-2339引 transformers#27931)就是这种耦合的考古层。 - 执行顺序敏感。
exec类补丁抽不到源码时静默跳过,可能造成「模型加载了但没加速」; Unsloth 用运行时校验(如第 6 节的__name__检查)兜底。 - 全局污染。 补丁影响整个 Python 进程;想在同一进程对比「原生 vs 加速」,只能靠 unpatch 函数逐个反悔。
- 更深的补丁在
unsloth_zoo。 编译器级重写(unsloth_compile_transformers)、offloaded 梯度检查点等不在这个仓库,只在_utils.py:155-172处 import——本仓库是调度层。