LLaMA-Factory实战:训练提速117%省显存50%
【免费下载链接】LlamaFactoryUnified Efficient Fine-Tuning of 100+ LLMs & VLMs (ACL 2024)项目地址: https://gitcode.com/GitHub_Trending/ll/LlamaFactory
凌晨两点,你把 7B 全量微调任务丢上去,40 分钟后屏幕上弹出一行 CUDA out of memory。这种等待和爆显存,在 LLaMA-Factory 里可以用一套组合拳解决:Liger Kernel 融合算子、选择性梯度检查点、DeepSpeed ZeRO-3。读完本文,你能在三步之内拿到一份跑得动的提速配置,并判断它在你自己的场景里值多少。
代价锚点:不优化会怎样
先看一个真实量级的场景:Llama-2-7B 跑 56k 超长序列微调,FlashAttention-2 路线在 24GB 卡上根本放不下。官方记录里,启用优化后的方案把同样的任务塞进了 24GB 显存,这就是"优化前跑不完、优化后跑得动"的差距。
短序列场景差距更直观一点:同样一张卡,不开检查点时 4096 tokens 全量微调就 OOM,开了之后同配置直接跑通,只是单步耗时略涨。代价是时间,收益是任务能完成。
方案全景:三层各管一件事
整个加速链路是流水线式的,三层各切一块问题:
- 前向层:Liger Kernel 把 RoPE、RMSNorm、SwiGLU、交叉熵等默认实现换成融合 CUDA kernel,减少 kernel 启动和全局显存往返;
- 反向层:选择性梯度检查点,只对真正要算梯度的层做"前向不存激活、反向重算",冻结层不浪费时间;
- 跨卡层:ZeRO-3 把参数和 optimizer 状态切分到多卡,单卡只持有一份。
输入是普通训练配置,输出是同一个模型,区别是每小时能多跑多少 step、以及多大的 batch 塞得进显存。
数据先行:速度对比与显存对比
在解释原理前,先把有出处的数字摆出来:
| 方案 | 训练速度 | 显存占用 | 7B 跑 56k 长序列 |
|---|---|---|---|
| FlashAttention-2(基线) | 100% | 100% | 24GB 卡放不下 |
| LLaMA-Factory 优化方案 | 快 117% | 省 50% | Llama-2-7B-56k 在 24GB 内跑通 |
数据来源:README 更新日志(24/04/16),unsloth 长序列训练方案,Llama-2-7B 56k 序列场景,详见仓库 wiki 的性能对比页
两点说明,避免误读:
- 这组 117%/50% 是长序列场景的官方口径,短序列 LoRA 拿不到同样的倍数,主要收益来自交叉熵和检查点部分;
- Liger Kernel 单独的官方对比数字仓库没有公布,它的价值体现在机制层,下面拆开讲。
算子融合:一次 kernel 启动省一次显存往返
Liger Kernel 说白了就是"换一套更快的算子实现":把几个相邻的小操作合并成一个 CUDA kernel,省掉中间的读写显存。
最容易量化的是交叉熵。默认实现会先把完整的 logits 张量(batch × 序列长 × 词表大小)在显存里摆出来再算损失,而融合交叉熵是分块现算现丢,这个大张量从头到尾不落显存。Qwen 7B 词表约 15 万,一条 2048 tokens 的样本这个张量就有几百 MB 量级,batch 一放大就是 GB 级——这就是最后一层最大的显存吃客。
仓库按模型类型分发对应实现,支持 Llama、Qwen2/2.5-VL/3、Gemma、GLM4、Mistral 等家族:
# src/llamafactory/model/model_utils/liger_kernel.py model_type = getattr(config, "model_type", None) if model_type == "qwen2_vl": from liger_kernel.transformers import apply_liger_kernel_to_qwen2_vl as apply_liger_kernel elif model_type == "qwen3_moe": from liger_kernel.transformers import apply_liger_kernel_to_qwen3_moe as apply_liger_kernel ... apply_liger_kernel(**kwargs)完整分发逻辑见 liger_kernel.py。注意代码里有一个自动降级:训练阶段需要保留完整 logits 时,融合交叉熵会被关掉(fused_linear_cross_entropy: False),换成普通交叉熵——好处缩水,但正确性优先。
选择性梯度检查点:冻结层不浪费重算
梯度检查点是一笔交易:前向不存中间激活,反向时重算,用算力换显存。50% 的显存节约主要就来自这条链路。
LLaMA-Factory 做了两处增强。第一处,只对含可训练参数的层启用检查点,冻结层正常前向,不做无用功:
# src/llamafactory/model/model_utils/checkpointing.py has_grad = any(param.requires_grad for param in module.parameters()) if has_grad: return gradient_checkpointing_func(func, *args, **kwargs) else: return func(*args, **kwargs)第二处,unsloth 风格的检查点会把层输入异步搬到 CPU 内存再算,反向时搬回来:显存直接变成"GPU 显存 + 内存"两级,长序列场景靠它才能塞进 24GB。完整实现见 checkpointing.py。
ZeRO-3 分片:把 optimizer 状态摊到 N 张卡
单卡显存放不下的另一大块是 optimizer 状态(Adam 下约等于参数量的 2 倍)。ZeRO-3 把参数、梯度、状态都切成 1/N 摊到每张卡,用的时候再临时聚合:
// examples/deepspeed/ds_z3_config.json(节选) "zero_optimization": { "stage": 3, "contiguous_gradients": true, "stage3_gather_16bit_weights_on_model_save": true }bf16 也默认自动启用("bf16": {"enabled": "auto"}),半精度训练同时省一半激活显存。完整配置在 ds_z3_config.json。
三步启用 Liger Kernel
第一步:装依赖
pip install -r requirements/liger-kernel.txt版本要求liger-kernel>=0.6.3,文件在 requirements/liger-kernel.txt。验证:
python -c "import liger_kernel; print(liger_kernel.__version__)"能打印出 0.6.3 以上版本号即成功。
第二步:改配置
在训练 yaml 里加两行关键参数(以 qwen3_lora_sft.yaml 为底改):
model_name_or_path: Qwen/Qwen3-7B-Instruct stage: sft finetuning_type: full enable_liger_kernel: true # 关键:开启融合算子 bf16: true deepspeed: examples/deepspeed/ds_z3_config.json注意:梯度检查点默认就是开的(disable_gradient_checkpointing默认 False),不用额外配置。
第三步:启动并确认
git clone https://link.gitcode.com/i/22cff475469f8adcacf2a58b951984c0 cd LlamaFactory && llamafactory-cli train 你的配置.yaml启动日志里看到这两行就说明全部生效:
Liger kernel has been applied to the model.Gradient checkpointing enabled.
然后对比同数据同硬件的 tokens/sec,即可确认提速幅度。
边界与坑:什么时候别开
- 模型不在支持列表:只会打印
Current model does not support liger kernel.警告并回退普通训练,不会崩,但你也拿不到收益。开跑前对照 liger_kernel.py 里的分发列表确认。 - 需要 logits 的阶段:如奖励建模等需要完整输出的 stage,融合交叉熵自动关闭,显存收益明显缩水,这是代码里的有意降级。
- NPU 用户:非 Ascend 910 的 NPU 会自动关掉 swiglu 和融合交叉熵两个开关,收益比 CUDA 小不少。
- 高频报错 1:
ImportError且模型是 gpt_oss —— 装的 liger-kernel 版本太老,升到新版即可。 - 高频报错 2:开了 ZeRO-3 但保存的权重分片异常 —— 确认配置里保留了
stage3_gather_16bit_weights_on_model_save: true,否则存档不完整。 - 预期管理:别拿 56k 长序列的 117% 去要求 2048 短序列 LoRA,短序列瓶颈在数据加载和通信,收益主要来自交叉熵那部分。
路线图:ROCm 与 NPU
项目正在把优化面铺到更多硬件:docker/docker-rocm/下已有 AMD GPU 的现成 Dockerfile 和 compose 配置,docker/docker-npu/覆盖昇腾场景,Megatron 并行训练则有 examples/megatron/ 的完整示例。如果你在做模型适配,算子相关代码集中在 src/llamafactory/model/model_utils/,性能问题直接提 issue 反馈。
显存和速度从来不是二选一,核心动作只有一个:别让显存替你保管用不到的东西。觉得有用的话,先收藏这篇,下一篇拆解多模态场景下图像 token 爆炸的显存控制思路。
【免费下载链接】LlamaFactoryUnified Efficient Fine-Tuning of 100+ LLMs & VLMs (ACL 2024)项目地址: https://gitcode.com/GitHub_Trending/ll/LlamaFactory
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考