Axolotl继续预训练实战:流式加载喂饱TB级领域语料
【免费下载链接】axolotlGo ahead and axolotl questions项目地址: https://gitcode.com/GitHub_Trending/ax/axolotl
Axolotl 是 Hugging Face 生态里的 LLM 训练框架,除了 LoRA 微调,它最被低估的能力是继续预训练:用流式数据集加载边读边训,不落地 token 化缓存,普通 GPU 服务器就能让通用模型吃下领域语料。这篇把参数一个个摊开讲清楚。
先想清楚:你的语料装不进内存吗?
选错路径后面全白干,所以第一件事是判断语料规模。Axolotl 给了两条路:
| 对比项 | 非流式(type: completion) | 流式(pretraining_dataset) |
|---|---|---|
| 语料规模 | 能装进内存 | TB 级,装不下 |
| token 化时机 | 训练前一次性完成 | 训练中按需计算 |
| 长文本处理 | 超出sequence_len就截断 | 多段拼接成定长序列 |
| 适用场景 | 小语料、反复迭代同一份数据 | 大语料、开箱就训 |
数据格式两边通用,JSONL 每行一个text字段即可,领域术语和原始句式尽量保留,别过度清洗。
用 pretrain.yaml 搭一条能跑的流式预训练管线
官方示例 examples/streaming/pretrain.yaml 是最小的可用起点,关键片段:
pretraining_dataset: - path: HuggingFaceFW/fineweb-edu # 换成你的领域语料,本地路径也行 type: pretrain text_column: text split: train streaming_multipack_buffer_size: 10000 # 打包缓冲:越大越省算力,越吃内存 shuffle_merged_datasets: true # 用缓冲窗口混洗,避免语料顺序泄露 sequence_len: 1024 sample_packing: true pretrain_multipack_attn: true # 切断打包样本间的交叉注意力 attn_implementation: flash_attention_2 # 样本打包依赖它 micro_batch_size: 1 gradient_accumulation_steps: 8 learning_rate: 5e-4 max_steps: 1000 # 流式模式必填:框架算不出数据集总长 output_dir: ./outputs/my-domain-pretrain两个容易踩的点:
max_steps在流式模式是必填的,因为框架没法推断流式数据集的总长度;- 语料偏小、或要反复实验,直接用非流式路径更省事,流式文档里也建议小数据集走
axolotl preprocess离线 token 化。
sequence_len没有万能值,按语料定:
| 语料类型 | 建议序列长度 |
|---|---|
| 网页、通用短文本 | 1024~2048 |
| 领域文档、长报告 | 4096 起步 |
样本打包会互相串味,靠 pretrain_multipack_attn 隔离
sample_packing把多条短文本拼进一条定长序列,padding 少、GPU 利用率高。代价是多段文本会互相"看见",继续预训练里这属于注意力泄漏——一段法律条文不该影响另一段医学文本的生成概率。
pretrain_multipack_attn: true就是干这个的:配 Flash Attention 时用 cu_seqlens 告诉内核每段的边界,根本不需要显式构造 4D mask。
打包细节可以看 多打包(样本打包)文档。
每步吃多少token,先算再定训练量
流式模式下别凭感觉填步数。一步消耗:
tokens_per_step = sequence_len × micro_batch_size × gradient_accumulation_steps × GPU数按上面示例:1024 × 1 × 8 × 1 = 8192 tokens/步,跑 1000 步约 820 万 tokens。继续预训练想看到领域效果,语料量至少要在千万 tokens 以上,据此反推max_steps。学习率上,全新预训练和继续预训练差一个量级,官方示例用 5e-4 是配 135M 小模型的;在 7B~8B 级别的基座上继续预训练,参考 Llama-3 完整微调示例,2e-5 这类低学习率更稳。
显存不够时按这个优先级压
| 手段 | 配置 | 代价 |
|---|---|---|
| 梯度检查点 | gradient_checkpointing: true | 约 20% 算力换大笔激活显存 |
| 混合精度 | bf16: auto | 几乎无感 |
| 页式优化器 | optimizer: paged_adamw_8bit | 优化器状态显存减半 |
| 4bit 量化 | load_in_4bit: true等 | 继续预训练一般不建议,会损失基座能力 |
原则:继续预训练默认动全参数,别上 LoRA;量化是最后的退路。
中断了怎么续:auto-resume-from-checkpoints
流式训练一跑就是小时级,断电、OOM 都不是意外。示例里配了save_steps: 250、save_total_limit: 3控制检查点,续训用 CLI 参数(注意不是网上流传的--auto-resume):
axolotl train config.yaml --auto-resume-from-checkpoints它会自动找到output_dir下最新的检查点接着跑,不用手动填路径。
最后几个收在口袋里的提醒:
- 流式目前只支持单个数据集,多源混合得先离线合并;
- 验证集不会被流式化,始终全量加载,别指望它省显存;
- 启动后先看 loss 是否平稳下降,再检查数据管道有没有成为瓶颈(GPU 利用率长期偏低通常就是流式吞吐不够);
- 第一次跑可以临时把步数调小,验证检查点能存能续,再放满
max_steps。
【免费下载链接】axolotlGo ahead and axolotl questions项目地址: https://gitcode.com/GitHub_Trending/ax/axolotl
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考