深色模式
继续预训练 Continued Pretraining
摘要:当基座模型「不认识领域词汇/事实」时,SFT 力不从心,应先做 Continued Pretraining(CPT,又称 mid-training / domain-adaptive pretraining)。本文讲清 CPT 与 SFT 的边界、数据混合与回放防遗忘、学习率与调度、分词器扩展,以及评估与回滚。适用于医疗/法律/代码/金融等垂直领域适配。版本:
transformers4.40+、分布式见infra.md([版本相关])。
适用版本与前提
- 框架:
transformers4.40+、分布式训练用 DeepSpeed/FSDP(见infra.md) - 数据:无标注领域语料(清洗、去重后的文本),量级通常数十亿 token
- 模型示例:Llama-3.1-8B 作为起始 checkpoint
核心概念:CPT 教的是「分布与知识」
三阶段谱系的定位
- 预训练(PT):从随机初始化在万亿 token 上学的通用底座。
- CPT(继续预训练 / mid-training):从已训练底座在领域/新鲜语料上继续下一词预测,教领域词汇、事实分布、推理模式。OLMo 2 等已把它作为 SFT 前的独立阶段。
- SFT:在 CPT 产出的「领域底座」上教指令遵循与行为。
- 经验法则:不认识词 → CPT;认识但不会用 → SFT;会用但品味不对 → 偏好对齐(DPO/RLHF)。
CPT 与基础预训练的差异:
| 维度 | 基础预训练 | CPT |
|---|---|---|
| 起点 | 随机初始化 | 已训练底座 |
| 数据 | 广而杂 | 目标领域为主 + 回放 |
| Token 预算 | 万亿级 | 十亿~数百亿 |
| 学习率 | 峰值 ~1e-4~1e-3 | 更低,约基座峰值的 10%–30% |
| 风险 | 欠拟合 | 灾难性遗忘 |
架构与原理
灾难性遗忘是头号风险
纯领域数据继续训练会侵蚀通用能力(如 MMLU 下滑)。缓解手段:混入 5%–30% 原始预训练分布作为回放(replay)、使用更低学习率、限制训练 token 数,并每几百步跑通用评测(MMLU/HellaSwag 等)监控回归。研究(Ibrahim et al., 2024 等)表明即便是 5% 的通用回放也能显著减少遗忘。
生产实践
数据配比与调度
- 典型混合:领域 60% / 通用回放 25% / 高质量 curated 10% / 指令相邻 5%。纯领域 = 高领域质量但低通用留存;60/40 通常是多数场景最佳平衡。
- 退火(annealing):训练末期(最后 5%–10% token)把 LR 激进降到接近 0,并切到纯高质量数据,可带来 ~1%–3% 基准提升(LLaMA-3 在最后 40M token 用 curated 数据退火)。
- 越短越好:在领域指标达标且不造成大通用回退的前提下,尽量短训;过度训练会过拟合风格伪影。
- 分词器覆盖:领域专有词(如医学术语、化学式)若被切成很多子词,考虑扩展词表并初始化新 embedding(平均子词嵌入或随机),但会增加复杂度。
操作步骤:CPT 训练脚本
python
# cpt_train.py —— 领域继续预训练(因果 LM 续训)
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, TrainingArguments, Trainer
from datasets import load_dataset
model_id = "meta-llama/Llama-3.1-8B"
model = AutoModelForCausalLM.from_pretrained(
model_id, torch_dtype=torch.bfloat16, attn_implementation="flash_attention_2",
)
tokenizer = AutoTokenizer.from_pretrained(model_id)
# 领域语料(示例:代码/医学/法律),需先清洗去重
dataset = load_dataset("json", data_files="domain_clean.jsonl", split="train")
args = TrainingArguments(
output_dir="./cpt-out",
per_device_train_batch_size=2,
gradient_accumulation_steps=16,
learning_rate=2e-5, # 远低于预训练峰值(~3e-4),约 1/10~1/15
lr_scheduler_type="cosine",
warmup_steps=500,
max_steps=50000, # 数十亿 token 量级,非百万级
bf16=True,
gradient_checkpointing=True, # 省显存,约 -30% 速度
logging_steps=10,
save_steps=1000,
save_total_limit=3,
)
trainer = Trainer(model=model, args=args, train_dataset=dataset)
trainer.train()1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
学习率是第一旋钮
CPT 最常见的错误是用 SFT/预训练同量级学习率继续训,直接导致灾难性遗忘。CPT 学习率应比基座峰值低 10–100 倍(常用 1e-5 ~ 5e-5)。DAPT 实践建议 2e-5,显著低于 SFT 的 2e-4。[具体取值依模型/数据,需实测]
验证
bash
# 1) 训练/验证 perplexity(领域语料上应下降,通用语料上不应大幅上升)
# 2) 周期性跑通用基准,监控回归
# lm_eval --model hf --model_args pretrained=./cpt-out,dtype=bfloat16 \
# --tasks mmlu,hellaswag,arc_challenge --num_fewshot 5
# [未实测具体分数;对比 CPT 前底座,MMLU 跌幅建议 < 3–5%]
# 3) 领域基准(如 PubMedQA / 法律 QA)应上升1
2
3
4
5
6
7
2
3
4
5
6
7
回滚与清理
CPT 回滚
- CPT 产出完整底座 checkpoint,体积大;保留 CPT 前快照,若回归严重可回退并重调混合比。
- 删除旧 checkpoint 前确认下游 SFT/对齐任务不依赖;大目录跨 NFS 删除注意 IO 风暴。
- 数据快照与混合比必须版本化——改了语料 = 换了实验,不可复现。
故障排查
- 通用能力崩塌:提高回放比例(到 25%–30%)、降学习率、缩短训练;检查是否纯领域数据。
- 领域指标不涨:数据质量/去重不足、领域覆盖不够;先做小代理模型验证混合比。
- OOM:开梯度检查点、降 micro batch、上 ZeRO-2/3(见
infra.md)。 - 基准污染:领域语料可能含评测题,CPT 前需重新去污染(即使原始语料已去过重)。
安全与合规
领域数据的特殊风险
- 安全回归:医学/安全/法律等领域数据可能含不安全操作指引,CPT 后必须重跑安全评测再宣布完成。
- PII / 机密:领域语料(病历、合同、代码仓)常含敏感信息,需脱敏与访问控制,等同 SFT 数据要求(见
sft-data.md)。 - 版权/许可:爬虫语料与受版权文本用于训练可能触发合规问题;确认数据来源授权。
- 越权:CPT 集群多租户需隔离存储与网络,避免跨团队读取领域语料与产出底座。
成本与性能(估算,[未实测])
| 规模 | 配置 | 示例时长 | 单价假设 | 估算 |
|---|---|---|---|---|
| CPT 7B 10B token | 8× A100 80GB | 数十小时 | ~$16/h | $数百–千 |
| CPT 70B 数十B token | 多节点 | 数百小时 | ~$数十/h | $数千–万 |
成本备注
CPT 比 SFT 贵(token 量级大、常全量权重更新),但远小于从零预训练。Diminishing returns 明显:10–100B 领域 token 后边际收益快速下降,质量与多样性比纯粹堆量更重要。利用率监控同 infra.md。[时长/单价为估算,非实测报价]
参考资料
- Open LLM Training Wiki — Continued pretraining / mid-training
- Continual Pre-training of Language Models(Ke et al., 2022, arXiv:2302.03241)
- Simple and Scalable Strategies to Continually Pre-train Large Language Models(Ibrahim et al., 2024, arXiv:2403.08763)
- OLMo 2(形式化 mid-training 阶段, arXiv:2501.00656)
- Code Llama(Llama 2 上继续预训练代码, arXiv:2308.12950)
- Domain-Adaptive Pretraining(Gururangan et al., 2020, DAPT)