跳到主要内容

[LLM 1/10] 继续预训练:把新知识真正灌进泰语 LLM

· 阅读需 14 分钟
Kobkrit Viriyayudhakorn
CEO, iApp Technology

一个泰语能力还不错的大语言模型,往往对你所在组织的专业知识"一无所知"—— 不懂泰国的政府法规,不懂你这个行业的术语,更没见过公司内部文档。 这篇文章要讲的是最直接的解法:继续预训练(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

LCPT(θ)=ExDdomain[t=1xlogpθ(xtx<t)]\mathcal{L}_{\text{CPT}}(\theta) = -\mathbb{E}_{x\sim\mathcal{D}_{\text{domain}}}\left[\sum_{t=1}^{|x|}\log p_\theta(x_t \mid x_{<t})\right]
  • xtx_t = 第 tt 个位置的 token
  • x<tx_{<t} = 它前面的全部 token
  • pθp_\theta = 模型预测出的概率

这个式子和预训练时一模一样,唯一变的是数据。 这正是 CPT 不需要标签的原因——文本自己就是答案。

3.2 Perplexity:我们的度量单位

PPL(D)=exp ⁣(1NiLCPT(x(i)))\text{PPL}(\mathcal{D}) = \exp\!\left(\frac{1}{N}\sum_{i}\mathcal{L}_{\text{CPT}}(x^{(i)})\right)

翻译成人话就是:"平均而言,模型正在多少个选项之间犹豫。" PPL = 20 表示大约在 20 个词之间摇摆,PPL = 5 则笃定得多。困惑度越低越好。

3.3 本章最重要的公式——Replay Mixing

Dmix=λDdomain+(1λ)Dgeneral\mathcal{D}_{\text{mix}} = \lambda\,\mathcal{D}_{\text{domain}} + (1-\lambda)\,\mathcal{D}_{\text{general}}

λ\lambda 就是每个 batch 中领域数据所占的比例。

  • λ=1.0\lambda = 1.0 → 纯领域数据 → 领域能力涨得最快,同时也忘得最快
  • λ=0.5\lambda = 0.5 → 一半一半 → 涨得慢一些,但遗忘少得多

不要把这个值写死,去扫一遍它的取值,然后挑一个你能接受的点。

4. 把公式画出来(Visualize)

困惑度到底在告诉我们什么

cross-entropy loss 与 perplexity 之间的关系,以及从不同起点降低 5 点 PPL 各自意味着什么cross-entropy loss 与 perplexity 之间的关系,以及从不同起点降低 5 点 PPL 各自意味着什么

Figure 1.1PPL 是 loss 的指数——同样降低 5 个单位,含义会因起点不同而天差地别

右边这张图是大家最容易忽略的地方:如果有人说"困惑度降了 5 个点",却不告诉你起点是多少, 这句话基本没有信息量——因为 80 → 75 只是改善了 6%,而 10 → 5 是改善了 50%。

replay ratio 所控制的那笔交易

随着 lambda 增大,领域困惑度下降而通用困惑度上升的曲线,并标出 Pareto frontier随着 lambda 增大,领域困惑度下降而通用困惑度上升的曲线,并标出 Pareto frontier

Figure 1.2replay mixing 公式所刻画出的权衡形状(示意机制的插图,并非实测结果——真实测量见第 8 节)

5. 准备环境(Environment)

打开 Colab,选择 Runtime → Change runtime type → T4 GPU(免费额度就够用)。

大多数 LLM notebook 在免费版 Colab 上翻车的死穴

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。

显存预算有多少,又都花到哪去了

按 weights、gradients、fp32 master、Adam states 和 activations 拆分显存占用的柱状图,并对比 adamw fp32 与 8-bit按 weights、gradients、fp32 master、Adam states 和 activations 拆分显存占用的柱状图,并对比 adamw fp32 与 8-bit

Figure 1.3全参数训练时显存的构成,按 Qwen3-0.6B config 中的真实数值计算(596M 参数)

注意 optimizer state 占的显存比模型本身还多——Adam 要为每个参数各存一份 mmvv。 换成 adamw_bnb_8bit 能省下 3.3 GB,这意味着你可以把 batch size 或序列长度再往上加不少。

自己动手玩一下显存预算——调调参数,看看什么时候会 OOM:

596.0M parameters, derived from config.json
Weight dtype
Run mode
Trades about 30% more compute for a large drop in activation memory.
weights: 1.13 GiBgradients: 19.25 MiBoptimizer: 115.50 MiBactivations: 170.00 MiB16 GB — Colab T40481216GiB
  • Weights1.13 GiB
  • Gradients19.25 MiB
  • Optimizer state115.50 MiB
  • Activations170.00 MiB
  • KV cache
Total VRAM1.43 GiB14.57 GiB to spare
Trainable params10.1M1.69%
KV cache per token112 KiB2 x 28 x 8 x 128
Full context KV4.38 GiB41.0K tok

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 陷阱

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

  1. 领域 held-out PPL —— 应当明显下降(这是我们花钱买到的东西)
  2. 通用 held-out PPL —— 应当有所上升(这是我们付出的代价)
  3. 来自 KobEval-TH 评测集的 TH-KNOW accuracy,附带 Wilson 95% CI
为什么任何时候都要给置信区间

如果测试集只有 100 道题,95% 置信区间的宽度大约是 ±10 个点。 也就是说,"78% 对比 74%"通常和随机波动区分不开。 没有 CI 的 accuracy 数字不是实验结果,只是传闻——第 9 章我们会深入讲这件事。

Prompt
Promptอธิบายว่าทำไมท้องฟ้าถึงเป็นสีฟ้า แบบสั้น ๆ

base

Thai 18%41 tokens
The sky appears blue because of Rayleigh scattering. ท้องฟ้า is blue เพราะ light scatter ครับ. Shorter wavelengths scatter more than longer ones.

sft

Thai 99%78 tokens
ท้องฟ้าเป็นสีฟ้าเพราะแสงอาทิตย์กระทบกับโมเลกุลของอากาศแล้วเกิดการกระเจิงแบบเรย์ลี ซึ่งแสงสีน้ำเงินที่มีความยาวคลื่นสั้นกว่าจะกระเจิงได้มากกว่าแสงสีแดง เราจึงมองเห็นท้องฟ้าเป็นสีฟ้าครับ

Showing the built-in sample.

9. 对比(Comparison)

notebook 在同一份数据上训练了三种配置,好让这笔交易变成看得见的数字:

模型领域 PPL ↓通用 PPL ↓TH-DOMAIN训练时间
Base(未训练)4.835.8827.3%
CPT,λ = 1.0(纯领域)4.07(−0.76)6.72(+0.858.0 分钟
CPT,λ = 0.5(含 replay)4.29(−0.53)4.98(−0.90)36.4%8.0 分钟
在 Colab T4(sm_75,14.56 GB)上实测 —— 显存峰值 10.50 GB,Qwen3-0.6B-Base, 每轮 100 个优化步。所有数字均来自 notebook 自动写出的 results.json

读懂这张表就是本章的核心:

  • λ = 1.0 在领域 PPL 上最好(4.07),但通用 PPL 变差,从 5.88 升到 6.72 —— 这就是被量化出来的灾难性遗忘,而不是空口断言。
  • λ = 0.5 在领域上让出一点(4.29),通用 PPL 反而变好 到 4.98。 replay 在这里不只是防止遗忘,还让模型整体的泰语建模能力更强了。
TH-DOMAIN 上升了,但还不能下结论

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 λ\lambda 就是调节汇率的旋钮——去扫,别猜
  • 足够低的学习率是"做 CPT"和"毁掉模型"之间的分界线
  • 每一个数字都必须带上置信区间
这个实验的局限

我们只用了大约 8,000 篇文档,而 OpenThaiGPT 那种量级的真实 CPT 用的是百亿 token 级别的数据, 两者相差约 6 个数量级(order of magnitude)。

这个实验确实能证明**"机制""权衡关系"的存在, 但它并不会产出一个可用于生产的更好的模型**。请不要拿这个结果去宣称你做出了更强的泰语模型。 你真正得到的是"每个旋钮各自在做什么"的理解,而这份理解是可以迁移到真实规模的工作上的。

下一章: SFT 与 LoRA——当模型已经有了知识,我们该怎么教它回答, 以及为什么只训练 1.7% 的参数,效果就能逼近全参数训练。

参考文献(References)

  1. Gururangan et al. (2020). Don't Stop Pretraining: Adapt Language Models to Domains and Tasks — 本章所遵循的领域自适应预训练方法的源头
  2. Ibrahim et al. (2024). Simple and Scalable Strategies to Continually Pre-train Large Language Models — 让 CPT 不至于毁掉模型的 replay 与学习率策略
  3. Gupta et al. (2023). Continual Pre-Training of Large Language Models: How to (re)warm your model? — 继续预训练时学习率 warmup 为何如此关键
  4. Luo et al. (2023). An Empirical Study of Catastrophic Forgetting in Large Language Models During Continual Fine-tuning — 对灾难性遗忘的系统性量化
  5. Kaplan et al. (2020). Scaling Laws for Neural Language Models — scaling laws——"8,000 篇文档远远不够"的依据
  6. Hoffmann et al. (2022). Training Compute-Optimal Large Language Models — Chinchilla:算力最优的数据与参数配比
  7. Yuenyong et al. (2025). OpenThaiGPT 1.6 and R1: Thai-Centric Open Source and Reasoning Large Language Models — 真实规模的泰语 CPT,可与本章的小实验对照
  8. Lowphansirikul et al. (2021). WangchanBERTa: Pretraining transformer-based Thai Language Models — 泰语预训练模型的先行者及其语料处理

本系列的文章、代码与 notebook 均以 CC BY-NC-SA 4.0 授权 —— 可自由使用与改编,须署名、限非商业用途,并以相同方式共享。文中引用的第三方模型与数据集仍适用各自的许可证。