Skip to content

Latest commit

 

History

12 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

lora-demo

一次从零到有、每一行都讲得清的 LoRA 微调实验。

  • 模型:Qwen3-4B(默认 Qwen/Qwen3-4B-Instruct-2507)
  • 方法:裸 HuggingFace peft + 手写 training loop,不用 Trainer,不用 SFTTrainer,不用 LLaMA-Factory / ms-swift
  • 任务:中文口语 → 固定 schema 的紧凑 JSON(日程指令解析)
  • 硬件:单卡 bf16,实测显存峰值 15.45 GB → 24GB 卡(4090 / A10 / 3090)足够 (主结果是在 H20 95GiB 上跑的,但那张卡的余量完全没用上)
  • 耗时:训练 1.3 分钟,全套流程含下权重约半小时

这不是一个"跑通就行"的模板。仓库里的每个非显然决定都在注释里写了为什么, 每条技术断言都是我在 transformers 5.15.1 / peft 0.20.0 上实际跑出来的, 包括几条推翻了我自己最初直觉的(见 docs/pitfalls.md)。

一句话结论

同一个任务、同一个 base 模型,三个条件(详见 §3):

prompt 权重 逐字段全对 平均 prompt token
A 极简(17 字) base 0.0% 37.7
B 完整 spec + 3-shot(950 字) base 68.0% 564.7
C 极简(和 A 逐字符相同) base + LoRA 99.7% 37.7

A → C 的 prompt 逐字符相同,所以那 99.7 个点只能来自权重。 C vs B 说明:同样的规则,写进权重比写进 prompt 又准又省 —— 省下的 527 token/次不是一次性的,是每一次推理都在省。

这就是 LoRA 的收益形状:把 prompt 里每次都一样的那部分,从 context 搬进权重。 不是"让模型变强"。


0. 为什么要这么做,而不是找个教程照抄

网上 LoRA 教程绝大多数是这个形态:

trainer = SFTTrainer(model=model, train_dataset=ds, peft_config=LoraConfig(r=16))
trainer.train()

三行跑通,然后你什么都没学到。更糟的是,一旦效果不好,你完全不知道该查哪儿: 是 rank 太小?数据太少?lr 不对?label masking 错了?chat template 不匹配? 这五个可能性里,只有前两个是网上教程会讨论的,而实际最常出错的是后三个。

所以这个仓库反过来做:把 Trainer 藏起来的东西全摊开, 每一步都配一个只做检查、不做训练的脚本,让你在训练之前就确认自己没错。

还有一个现实原因:transformers 已经是 5.x,有破坏性变更。 2024/2025 年写的 LoRA 教程在今天大面积跑不通(torch_dtype 改名、 apply_chat_template 返回类型变了……)。这仓库的版本是写死并实测过的。


1. 学习路径

第 0 级(可选但强烈推荐,笔记本 CPU 就能跑,约 30 分钟)

先手写一遍 LoRA。推荐 Sebastian Raschka《Build a Large Language Model (From Scratch)》 的附录 E:

它用纯 PyTorch 手写 LoRA:一个 LoRALayer 类,一个递归函数把模型里所有 nn.Linear 换成 Linear + LoRA 旁路。跑完你会亲眼看到 ΔW = BA 是怎么变成 forward 里两行代码的。之后你用 peft 时,get_peft_model() 就不再是黑盒。

跳过也行 —— 本仓库的 scripts/02_inspect_lora.py 会把 peft 内部结构 拆开打给你看,能补上大部分认知。

第 1~7 级(GPU 机器,本仓库)

步骤 脚本 干什么 训练?
0 00_check_env.py 环境自检:CUDA / 显存 / 版本 / 权重从哪下 否
1 01_make_dataset.py 合成数据集,检查泄漏与分布 否
1b 01b_download_model.py 单独把权重下好,和训练解耦 否
2 02_inspect_lora.py 解剖 LoRA,验证 4 条理论事实 否
3 03_inspect_batch.py 摊开一个 batch,看 label masking 边界 否
4 04_train_lora.py 手写 training loop 训练 是
5 05_eval.py 三条件对照评测 否
6 06_merge_and_infer.py 合并 adapter,验证旁挂==合并 否
7 run_ablations.sh + 07_collect_results.py 消融实验与汇总 是

第 2、3 步不要跳。 它们不训练任何东西,几十秒跑完,但能在你烧掉一小时 GPU 之前把 90% 的错误挡住。


2. 完整命令序列

git clone https://github.com/structDream/lora-demo && cd lora-demo

# --- 环境 ---
# GPU 机器上通常已有能用的 torch,不要重装它(会被换成 CPU 版)
python scripts/00_check_env.py          # 先看这个的结论再决定装什么

# 隔离 venv:如果这台机器上跑着 sglang / vllm,这一步不是"最佳实践"而是必需 ——
# 它们硬钉 transformers 的具体版本,直接装会把推理服务弄坏(第 14 条)。
# --system-site-packages 是为了复用系统那份已编译好 CUDA 的 torch。
python -m venv .venv --system-site-packages
source .venv/bin/activate

pip install -r requirements.txt         # 注意 requirements 里故意没有 torch

# 若 00 报了 torchao 版本过低(推理机上很常见,sglang 钉 0.9.0),补这一句。
# 它没有运行时依赖,不会动你的 torch,且只装进 venv(第 17 条)
pip install "torchao==0.16.0"

# 装完再跑一次 00,确认隔离成立:
#   venv 内  transformers 5.15.1 且路径在 .venv 下、torch 路径**不在** .venv 下
#   deactivate 后全局仍是原版本、推理框架仍能 import
python scripts/00_check_env.py

# --- 数据(秒级) ---
python scripts/01_make_dataset.py

# --- 权重(7.5GB,单独下,失败了重跑就行) ---
# 直连 huggingface.co 就不用管下面这行。
# 网络受限(如中国大陆、公司内网)时走镜像;Xet 会由脚本自动关掉(第 15 条)
export HF_ENDPOINT=https://hf-mirror.com
python scripts/01b_download_model.py --model Qwen/Qwen3-4B-Instruct-2507

# 下完之后**不需要**再设 HF_HUB_OFFLINE —— 脚本检测到权重已在缓存就自动切
# offline 并打印一行提示。(不这么做的后果:from_pretrained 即使缓存齐全也要
# 联网校验一次;如果你的网络是"丢包"而不是"拒连",它会静默卡到 socket 超时,
# 实测 5 分钟不输出任何东西。第 18 条)
#
# 想更保险可以自己 export,效果一样,脚本会尊重你已设的值:
#   export HF_HUB_OFFLINE=1

# --- 训练前检查(不训练,一两分钟) ---
python scripts/02_inspect_lora.py  --model Qwen/Qwen3-4B-Instruct-2507
python scripts/03_inspect_batch.py --model Qwen/Qwen3-4B-Instruct-2507

# --- 先跑一个冒烟测试,确认端到端不炸(约 0.2 分钟) ---
python scripts/04_train_lora.py --model Qwen/Qwen3-4B-Instruct-2507 \
    --limit 64 --epochs 1 --batch-size 8 --out out/smoke

# --- 正式训练(实测 1.3 分钟,显存峰值 15.45 GB) ---
python scripts/04_train_lora.py \
    --model Qwen/Qwen3-4B-Instruct-2507 \
    --r 16 --targets all --lr 2e-4 --epochs 2 \
    --batch-size 8 --grad-accum 1 \
    --out out/r16

# --- 三条件评测(全量 300 条,约 5 分钟) ---
python scripts/05_eval.py --model Qwen/Qwen3-4B-Instruct-2507 --adapter out/r16/best

# --- 合并 + 泛化检查 + 延迟基准 ---
python scripts/06_merge_and_infer.py --model Qwen/Qwen3-4B-Instruct-2507 --adapter out/r16/best

# --- 评测(三条件对照,这是收官) ---
python scripts/05_eval.py \
    --model Qwen/Qwen3-4B-Instruct-2507 \
    --adapter out/r16-all/best \
    --data data/test.jsonl

# --- 合并与部署 ---
python scripts/06_merge_and_infer.py \
    --model Qwen/Qwen3-4B-Instruct-2507 \
    --adapter out/r16-all/best \
    --merge-to merged/lora-demo-4b

# --- 消融(2~3 小时,想清楚再跑) ---
bash scripts/run_ablations.sh Qwen/Qwen3-4B-Instruct-2507
python scripts/07_collect_results.py

显存不够就调 --batch-size / --grad-accum / 加 --grad-checkpointing, 00_check_env.py 会按你的卡直接给出建议值。


3. 这个实验到底想让你看见什么

任务是"中文口语 → 固定 schema JSON"。schema 是故意设计得难猜的: 字段名是 act/subj/who/d/t0/len/pri 这种任意缩写,act 的取值是 EVT_NEW 这种自造枚举,还有"不同 act 要把哪些字段清零"这种纯约定规则。

然后设三个条件对照:

prompt 权重
A 极简(一句话,17 字) base
B 完整 spec + 3-shot(950 字) base
C 极简(和 A 逐字符相同) base + LoRA
  • A → C 的差 = LoRA 真正学到的东西。 prompt 完全一样,所以差值只能来自权重。
  • C vs B = LoRA 和"把规则写进 prompt"谁强。

PROMPT_MINIMAL 是一个共享常量,A 和 C 从同一处 import; render_prompt() 是唯一的 chat template 入口,训练和评测都走它。 这两条纪律从代码层面消灭了"prompt 不一致"这个隐形变量。

参考结果

主结果:Qwen3-4B-Instruct-2507 + 1200 条训练 + 全量 300 条测试, 单卡 bf16(我用的是 H20 95GiB,但峰值只吃 15.45 GB),r=16、lr 2e-4、2 epoch、batch 8。 训练 1.3 分钟,dev loss 5.1441 → 0.0003,显存峰值 15.45 GB。

指标 A base+极简 B base+spec C LoRA+极简
裸 JSON 可解析率 92.7% 100.0% 100.0%
键集合正确率 0.0% 100.0% 100.0%
act 分类正确率 0.0% 83.7% 100.0%
逐字段全对(EXACT) 0.0% 68.0% 99.7%
平均 prompt token 37.7 564.7 37.7

各字段单独正确率(C 只错在 1 条的 d 上):

act subj who d t0 len pri
B 83.7% 93.7% 94.0% 96.7% 99.3% 99.7% 93.3%
C 100% 100% 100% 99.7% 100% 100% 100%

四个读法:

  1. A 的 0% 不是模型笨,是它没有任何理由猜中你的 schema。 它输出的是完全合理的 JSON,只是键名/枚举/格式全不是你要的: {"action":"delete","event":"产品需求评审","date":"明天","notes":"已取消…"}。 注意它的可解析率有 92.7% —— "会写 JSON"和"会写你的 JSON"是两件事, 只看解析率会得出完全错误的结论。

  2. B 是这张表里信息量最大的一列:键集合 100%,EXACT 只有 68%。 spec 全写进 prompt 了,键、格式、枚举它都照做了, 但 act 只有 83.7%、pri 只有 93.3% —— 那些"哪种 act 要清零哪些字段"、 "'有空再说'该给 pri=0 还是 1"的细则,它读到了但没稳定遵守。 典型错例:pri 期望 0 却给了 1;"记得…叫我"该判 RMD_NEW 却判成 EVT_NEW。 这就是 prompt 的天花板:context 里的规则是"建议",权重里的规则才是"本能"。

  3. C 用 37.7 个 prompt token 拿到 99.7%,B 用 564.7 个才 68%。 省下的 527 token/次不是一次性的,是每一次推理都在省。 按每天 10 万次调用算,一天省 5270 万 prefill token。

  4. 规模会改变结论的强度,但不改变结论。 我先在 Mac 上用 Qwen3-0.6B 跑通同一套流程(400 条训练 / 40 条测试),B 的 EXACT 只有 5%; 换到 4B 就涨到 68% —— 模型越大越会读 spec。 但 A→C 的差值(0% → 99.7%)在两个规模上都成立, 而这个差值才是本实验要证的东西(A 和 C 的 prompt 逐字符相同, 所以差值只能来自权重)。

这就是 LoRA 的收益形状:把 prompt 里每次都一样的那部分,从 context 搬进权重。 不是"让模型变强"。


4. 训练之前必须确认的 4 条理论事实

scripts/02_inspect_lora.py 会逐条验证并打 PASS/FAIL。这些不是装饰, 是"我真的懂了"的可执行检查:

验证 内容 为什么重要
1 B 初始化全 0,A 随机非零 两个都设 0 则梯度恒为 0,永远学不动
2 挂上 LoRA 那一刻,输出与 base 逐位相同 ΔW = B@A = 0。LoRA 是严格增量,不会破坏预训练能力
3 第一步 ∂L/∂A 精确为 0,只有 B 在动 因为 ∂L/∂A = Bᵀ·(∂L/∂ΔW),B=0 则整项为 0。B 离开 0 后 A 才活过来
4 旁挂 vs merge_and_unload() argmax 逐位一致 这是"训练时旁挂保精度,部署时合并零开销"的依据

验证 3 是我在写这仓库时实测出来的,很少有教程提。它让"A 随机 / B 全零"这个 不对称初始化的全部后果变得可见。

验证 4 有个坑值得单独说:它的判据不能是绝对阈值。 我第一版写 diff < 5e-2,在 Mac 的 fp32 上实测 7.3e-05 很宽裕, 换到 bf16 上就变成 3.1e-01 直接 FAIL —— 量级差 5300 倍。 原因是合并要把 fp32 的 ΔW 存回 bf16 的 W,而 bf16 在 |W|~0.02 处的量化间隔 约 6e-5、ΔW 元素才 9e-4,这一"存"就吃掉了 ΔW 的约 2.7%。 所以判据改成「相对误差 + argmax 逐位一致」,后者在 bf16 下依然 14/14。 详见 docs/pitfalls.md 第 16 条。

顺带一个实测结论:合并带来的加速不是一个固定百分比,它取决于你有多 compute-bound。 同一张卡、batch=8、纯前向、warmup 后取 p50:

序列长度 旁挂 合并 合并省下
44 token 59.26 ms 36.23 ms 38.9%
571 token 448.27 ms 343.20 ms 23.4%

短序列下前向被 kernel 启动开销主导(36 层 × 7 个 target module = 252 次额外的 小 kernel 启动),所以省得多;序列一长,真正的矩阵乘开始主导,那 252 次固定 开销被摊薄,比例就掉下来。别拿短序列的数字做容量规划。


5. 踩坑清单

完整 19 条含实测证据在 docs/pitfalls.md, 其中第 7、16 条是推翻我自己判断的(留着比正确结论更有用)。 这里是最要命的九条:

  1. label masking 必须手写,官方 API 对 Qwen3 是废的。 apply_chat_template(..., return_assistant_tokens_mask=True) 在 Qwen3 上 返回全 0 mask(模板里没有 {% generation %} 块)。信了它又没看 warning, labels 会全部被屏蔽,训练等于没训。

  2. thinking mode 不关掉,所有指标归零。 Qwen3-4B / 0.6B 是 hybrid 模型,默认走思维链,会把生成预算全烧在 <think> 里,一个 JSON 字符都吐不出来。本仓库默认 enable_thinking=False (对非 thinking 模型是安全的空操作),且训练/评测一致性会被校验。

  3. lr 要 1e-4 量级,不是 2e-5。 拿全量微调的 lr 训 LoRA 是最常见的失败方式: loss 缓慢下降"看起来在学",评测分数几乎不动。消融里有这一组。

  4. 生成时必须左 padding。 右 padding 会让 PAD 挡在 prompt 和生成位置之间, batch 越大结果越差。训练右 padding 没问题,两者不一样。

  5. transformers 5.x 改了 API。 torch_dtype → dtype; apply_chat_template(tokenize=True) 现在返回 BatchEncoding 而非 list[int] —— 旧教程里的 len(prompt_ids) 在 5.x 上会静默算成 2。

  6. 合成数据必须跨集合去重。 "不同 random seed 当然不重复"是错的: 本仓库第一版就在 dev 上撞了 34 条 train 见过的句子。撞上的部分测的是 背诵而不是泛化。01_make_dataset.py 会显式报告。

  7. 别往推理机的全局环境 pip install。 sglang / vllm 硬钉 transformers 的具体版本,直接装会装成功但把推理服务弄坏,pip 只在末尾打一行容易 漏掉的 ERROR。而且那行 ERROR 报不全:实测它只提了 transformers, 没提被一起顶掉的 huggingface_hub(0.36 → 1.28,大版本破坏)。 用 python -m venv .venv --system-site-packages。

  8. 旧版 torchao 会让 LoRA 一个模块都挂不上。 peft 的 is_torchao_available() 对 <0.16.0 是 raise 而非 return False, 而 dispatcher 对每个 target module 都会调它。报错里一个 LoRA 字样都没有。

  9. 低精度环境里,写死的数值容差不可移植。 我的验证 4 阈值在 fp32 上余量 很大,换到 bf16 上直接 FAIL —— 代码没错,判据错了。判据要贴住你真正在乎 的东西(这里是 argmax 是否改变),而不是贴住某台机器上的数值。


6. 消融实验

run_ablations.sh 里 13 组,每组只改一个变量。详细读法见 docs/ablations.md。最值得亲手量的三组:

  • rank 扫描 r ∈ {4,8,16,32,64} —— 看从哪儿开始饱和。 本任务是固定格式映射,"内在秩"应该很低,预期 r=8 就够。 这是 LoRA 核心假设(任务需要的更新是低秩的)的直接检验。
  • attn vs mlp vs all —— 原论文只挂 attention,QLoRA 说要挂全部线性层, 而 2024/2025 两项独立实验说收益主要来自 MLP、attention 加上去几乎没有额外贡献 (参数量对齐后 attention-only 仍明显更差,见 docs/external-evidence.md 第 2 节)。 所以这里有三档:如果 mlp 和 all 打平,可训练参数能砍掉约一半。 脚本里带了参数量对齐的公平对照(attn-r45 vs all-r16,都是 33M), 不对齐的话你测出的差异有一部分只是"参数多了"。 对齐的 r 要按 Σ(in+out) 算:16 × 57344/20480 = 44.8。 我第一版按模块个数算成 7/4 → r=28,结果只有 62.5% —— FFN 的矩阵宽得多,按个数数会严重低估。教训见 docs/ablations.md。
  • lr 2e-5 vs 1e-4 vs 3e-4 —— 亲眼看到第 3 条坑。

7. 目录结构

loralab/                  可复用模块(都有大段 why 注释)
  schema.py               任务 schema 定义 + 判分(为什么故意设计得难猜)
  synth.py                数据合成(为什么不用 GPT 造数据)
  prompts.py              三种 prompt 条件 + 唯一的 chat template 入口
  data.py                 ★ label masking / collate / 手算 loss —— 最该读的文件
  modeling.py             加载 base + 挂 LoRA(target_modules 的取舍)
  jsonl.py                读写
scripts/                  按 0~7 编号,逐步执行
docs/
  pitfalls.md             踩坑清单 19 条(每条都附实测证据,含 2 条推翻我自己的)
  external-evidence.md    外部文献 + 六大框架默认值对照(每条带 URL)
  ablations.md            消融怎么设计、怎么读
  next-steps.md           从这个玩具任务到真实项目

复现说明

数据集不入库(.gitignore 挡掉了 data/),因为它是合成的: 01_make_dataset.py 里 seed 写死,跑一次就能得到和上表完全相同的 1200/300/300 条。 out/、merged/ 同理不入库。所以 clone 下来第一步就是跑 01, 而不是去找数据文件。


8. 下一步

这个任务是刻意选的"最容易看出效果"的类型。真实项目会难得多, 从这里到那里缺什么,写在 docs/next-steps.md。


参考

License

MIT,见 LICENSE。

About

No description, website, or topics provided.

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages