[LLM 1/10] 继续预训练:把新知识真正灌进泰语 LLM
一个泰语能力还不错的大语言模型,往往对你所在组织的专业知识"一无所知"—— 不懂泰国的政府法规,不懂你这个行业的术语,更没见过公司内部文档。 这篇文章要讲的是最直接的解法:继续预训练(Continue Pretraining,CPT), 从公式一直讲到能在免费 Colab 上大约 15 分钟真正跑完的代码。
Open in Colab01_continue_pretraining.ipynb
1. 问题(Problem statement)
设想你拿 Qwen3-0.6B 去问一句:"按照泰国总理府的规定,什么情况下可以采用特定方式采购?" 模型会非常自信地回答你,而且答错——因为它压根没见过足够多的泰国政府公文。
很多人试图用几千条问答对做 fine-tuning 来补救,然后发现根本没用。 原因在于:SFT 教的是"回答的形式",不是"知识"本身。 如果知识从来就不在模型权重里, 教它用正确的语气去回答,只不过是让它在胡说八道的时候更加理直气壮而已。
新知识进入模型有三条路,选错路正是大多数 LLM 项目失败的根源:
| 方法 | 适合的场景 | 推理时的开销 |
|---|---|---|
| RAG | 变动频繁、需要标注来源的知识 | 每次都要检索 + prompt 变长 |
| 继续预训练 | 大量且相对稳定的专业领域知识 | 无(知识已在权重里) |
| SFT | 格式、语气、答案结构 | 无 |
这篇文章走的是第二条路。
2. 我们要做什么(Solution)
我们会拿一个 base 模型(尚未经过 instruction tuning), 用与预训练阶段完全相同的 objective——预测下一个词——在我们关心领域的泰语原始文本上继续训练。 没有标签,没有问答对,只有纯文本。
但这篇文章的核心不是"训完变强了",而是你为此付出的代价:
CPT 是用牺牲通用能力来换取领域上的精准。 它是一笔交易,不是免费的午餐,而这笔交易的"汇率"由一个叫 replay ratio 的数字单独控制。
模型忘掉原本会做的事情,这个现象叫作灾难性遗忘(catastrophic forgetting)。 我们不会空口谈它,而是要把它量化成数字,再去找一个你能接受的平衡点。
3. 公式(Equation)
3.1 CPT 的 objective
- = 第 个位置的 token
- = 它前面的全部 token
- = 模型预测出的概率
这个式子和预训练时一模一样,唯一变的是数据。 这正是 CPT 不需要标签的原因——文本自己就是答案。
3.2 Perplexity:我们的度量单位
翻译成人话就是:"平均而言,模型正在多少个选项之间犹豫。" PPL = 20 表示大约在 20 个词之间摇摆,PPL = 5 则笃定得多。困惑度越低越好。
3.3 本章最重要的公式——Replay Mixing
就是每个 batch 中领域数据所占的比例。
- → 纯领域数据 → 领域能力涨得最快,同时也忘得最快
- → 一半一半 → 涨得慢一些,但遗忘少得多
不要把这个值写死,去扫一遍它的取值,然后挑一个你能接受的点。
4. 把公式画出来(Visualize)
困惑度到底在告诉我们什么
Figure 1.1PPL 是 loss 的指数——同样降低 5 个单位,含义会因起点不同而天差地别
右边这张图是大家最容易忽略的地方:如果有人说"困惑度降了 5 个点",却不告诉你起点是多少, 这句话基本没有信息量——因为 80 → 75 只是改善了 6%,而 10 → 5 是改善了 50%。
replay ratio 所控制的那笔交易
Figure 1.2replay mixing 公式所刻画出的权衡形状(示意机制的插图,并非实测结果——真实测量见第 8 节)
5. 准备环境(Environment)
打开 Colab,选择 Runtime → Change runtime type → T4 GPU(免费额度就够用)。
Colab 的 T4 是 Turing 架构(SM 7.5),它不支持 bfloat16,也不支持 FlashAttention-2。
但 Qwen3-0.6B 的 config.json 里写的是 torch_dtype: bfloat16。
所以如果你照着网上大多数教程写 torch_dtype="auto",代码要么直接崩,要么慢得离谱。
在这个系列里,我们每次都会把它写死:
torch_dtype=torch.float16 # 不是 bfloat16
attn_implementation="sdpa" # 不是 flash_attention_2
fp16=True # 在 TrainingArguments 里(不是 bf16=True)
本系列每个 notebook 的第一个 cell 都会把这行打印出来,让你亲眼看到:
cap = torch.cuda.get_device_capability(0)
print("compute capability:", cap) # T4 = (7, 5)
print("native bf16:", cap[0] >= 8) # T4 -> False
print("torch says :", torch.cuda.is_bf16_supported()) # T4 -> True(把 emulation 也算上了!)
is_bf16_supported() 在 T4 上会骗你较新的 torch 在 T4 上返回 True,因为它把**模拟(emulation)**也算作支持——而模拟比 fp16 慢得多。
请改为判断 compute capability ≥ 8.0(Ampere 及以上)。这是真正在 Colab 上跑才发现的 bug。
显存预算有多少,又都花到哪去了
Figure 1.3全参数训练时显存的构成,按 Qwen3-0.6B config 中的真实数值计算(596M 参数)
注意 optimizer state 占的显存比模型本身还多——Adam 要为每个参数各存一份 和 。
换成 adamw_bnb_8bit 能省下 3.3 GB,这意味着你可以把 batch size 或序列长度再往上加不少。
自己动手玩一下显存预算——调调参数,看看什么时候会 OOM:
- Weights1.13 GiB
- Gradients19.25 MiB
- Optimizer state115.50 MiB
- Activations170.00 MiB
- KV cache—
It fits.This run needs 1.43 GiB and leaves 14.57 GiB of headroom on a free Colab T4.
6. 准备数据(Data)
我们使用 pythainlp/thaigov-v2-corpus-22032023——这是泰国政府的新闻与公文语料库(公有领域),
它扮演的角色是"模型见得远远不够的专业领域知识"。这类政府公文用词高度固定、术语密集,
和任何一个组织的内部知识库在性质上是一样的,你完全可以把它换成自己的语料。
另外再用一份通用泰语文本作为 replay data,用来抑制遗忘。
from datasets import load_dataset
domain = load_dataset("pythainlp/thaigov-v2-corpus-22032023", split="train")
domain = domain.shuffle(seed=42).select(range(8000))
Packing:别让 padding 吃掉你的预算
如果把每篇文档都 pad 到同样长度,你会有海量算力浪费在 <pad> 上。
正确的做法是把所有文档首尾相接,再切成等长的 512 token 块。
def pack(examples, block_size=512):
ids = []
for text in examples["context"]:
ids.extend(tokenizer(text + tokenizer.eos_token).input_ids)
n = (len(ids) // block_size) * block_size
return {"input_ids": [ids[i:i+block_size] for i in range(0, n, block_size)]}
大多数模型的 tokenizer 主要是在英语数据上训练出来的。 泰语文本因此会被切得更碎——同一句话消耗的 token 可能是英语的 2–3 倍。 结果就是 API 费用更高、context 更快被填满、训练也更慢。notebook 里会把这个数字实测给你看。
中文读者对这件事应该并不陌生:中文同样是被切得更碎的一方,一个汉字常常就要占掉一到两个 token, 所以下面所有关于"token 效率"的讨论,对中文语料同样成立。
7. 核心代码(Main code)
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, TrainingArguments, Trainer
model = AutoModelForCausalLM.from_pretrained(
"Qwen/Qwen3-0.6B-Base", # 要 base,不要 instruct —— CPT 必须从 base 开始
torch_dtype=torch.float16, # T4 没有 bf16
attn_implementation="sdpa", # T4 没有 FlashAttention-2
).cuda()
model = model.float() # 训练前必须转成 fp32 —— 见下方提示框
args = TrainingArguments(
output_dir="cpt-out",
per_device_train_batch_size=2,
gradient_accumulation_steps=8, # effective batch = 16
num_train_epochs=1,
learning_rate=2e-5, # 比 SFT 低 10 倍 —— 见下方警告
lr_scheduler_type="cosine",
warmup_steps=50,
optim="adamw_bnb_8bit", # 省下 3.3 GB
gradient_checkpointing=True,
max_grad_norm=1.0, # 防止 fp16 数值爆炸
fp16=True, # 不是 bf16
logging_steps=10,
)
fp16=True 并不意味着权重是 fp16,它指的是混合精度:矩阵乘法在 fp16 中进行,
但主权重(master weights)必须保持 fp32,因为优化器要加上非常小的量(lr = 2e-5),
而 fp16 的精度根本表示不了这么小的数。
如果你把模型以 fp16 加载,然后直接用 fp16=True 训练全部参数,就会看到:
ValueError: Attempting to unscale FP16 gradients.
因为 max_grad_norm=1.0 要求在裁剪梯度前先反缩放(unscale),而这些梯度是 fp16 的。
正确做法是训练前调用 model.float(),评估时再转回 fp16。
(第 2 篇和第 4 篇用 LoRA 时,只需要转换 adapter 的参数。)
如果你拿 learning_rate=2e-4(大家做 LoRA 时常用的值)来做全参数 CPT,
几百个 step 之内你就会把模型的能力抹掉。
CPT 需要的学习率大约比 SFT 低 10–50 倍,因为我们动的是每一个权重。
8. 结果(Results)
notebook 会在训练前后各测三项指标,并写入 results.json:
- 领域 held-out PPL —— 应当明显下降(这是我们花钱买到的东西)
- 通用 held-out PPL —— 应当有所上升(这是我们付出的代价)
- 来自 KobEval-TH 评测集的 TH-KNOW accuracy,附带 Wilson 95% CI
如果测试集只有 100 道题,95% 置信区间的宽度大约是 ±10 个点。 也就是说,"78% 对比 74%"通常和随机波动区分不开。 没有 CI 的 accuracy 数字不是实验结果,只是传闻——第 9 章我们会深入讲这件事。
Promptอธิบายว่าทำไมท้องฟ้าถึงเป็นสีฟ้า แบบสั้น ๆbase
sft
Showing the built-in sample.
9. 对比(Comparison)
notebook 在同一份数据上训练了三种配置,好让这笔交易变成看得见的数字:
| 模型 | 领域 PPL ↓ | 通用 PPL ↓ | TH-DOMAIN | 训练时间 |
|---|---|---|---|---|
| Base(未训练) | 4.83 | 5.88 | 27.3% | — |
| CPT,λ = 1.0(纯领域) | 4.07(−0.76) | 6.72(+0.85) | — | 8.0 分钟 |
| CPT,λ = 0.5(含 replay) | 4.29(−0.53) | 4.98(−0.90) | 36.4% | 8.0 分钟 |
results.json。
读懂这张表就是本章的核心:
- λ = 1.0 在领域 PPL 上最好(4.07),但通用 PPL 变差,从 5.88 升到 6.72 —— 这就是被量化出来的灾难性遗忘,而不是空口断言。
- λ = 0.5 在领域上让出一点(4.29),通用 PPL 反而变好 到 4.98。 replay 在这里不只是防止遗忘,还让模型整体的泰语建模能力更强了。
27.3% → 36.4% 看着令人振奋,但两者的 Wilson 95% 置信区间是 13.2–48.2 与 19.7–57.0 —— 几乎完全重叠。在 n=22 的规模下,这只是迹象,不是结论。
真正扎实的证据是 PPL,因为它建立在数万个 token 之上,而不是 22 道题。 要让 TH-DOMAIN 具备结论性,题目数量需要提升到数百 —— 这正是第 9 篇要讲的内容。
10. 小结(Summary)
- CPT 用与预训练相同的 objective 把知识写进权重,不需要任何标签
- 它永远是一笔交易:领域上的精准,是拿丢掉的通用能力换来的
- replay ratio 就是调节汇率的旋钮——去扫,别猜
- 足够低的学习率是"做 CPT"和"毁掉模型"之间的分界线
- 每一个数字都必须带上置信区间
我们只用了大约 8,000 篇文档,而 OpenThaiGPT 那种量级的真实 CPT 用的是百亿 token 级别的数据, 两者相差约 6 个数量级(order of magnitude)。
这个实验确实能证明**"机制"和"权衡关系"的存在, 但它并不会产出一个可用于生产的更好的模型**。请不要拿这个结果去宣称你做出了更强的泰语模型。 你真正得到的是"每个旋钮各自在做什么"的理解,而这份理解是可以迁移到真实规模的工作上的。
下一章: SFT 与 LoRA——当模型已经有了知识,我们该怎么教它回答, 以及为什么只训练 1.7% 的参数,效果就能逼近全参数训练。
参考文献(References)
- Gururangan et al. (2020). Don't Stop Pretraining: Adapt Language Models to Domains and Tasks — 本章所遵循的领域自适应预训练方法的源头
- Ibrahim et al. (2024). Simple and Scalable Strategies to Continually Pre-train Large Language Models — 让 CPT 不至于毁掉模型的 replay 与学习率策略
- Gupta et al. (2023). Continual Pre-Training of Large Language Models: How to (re)warm your model? — 继续预训练时学习率 warmup 为何如此关键
- Luo et al. (2023). An Empirical Study of Catastrophic Forgetting in Large Language Models During Continual Fine-tuning — 对灾难性遗忘的系统性量化
- Kaplan et al. (2020). Scaling Laws for Neural Language Models — scaling laws——"8,000 篇文档远远不够"的依据
- Hoffmann et al. (2022). Training Compute-Optimal Large Language Models — Chinchilla:算力最优的数据与参数配比
- Yuenyong et al. (2025). OpenThaiGPT 1.6 and R1: Thai-Centric Open Source and Reasoning Large Language Models — 真实规模的泰语 CPT,可与本章的小实验对照
- Lowphansirikul et al. (2021). WangchanBERTa: Pretraining transformer-based Thai Language Models — 泰语预训练模型的先行者及其语料处理
本系列的文章、代码与 notebook 均以 CC BY-NC-SA 4.0 授权 —— 可自由使用与改编,须署名、限非商业用途,并以相同方式共享。文中引用的第三方模型与数据集仍适用各自的许可证。
