Skip to content

大模型预训练学习笔记

一、Decoder-only 的因果语言模型

1.1 预训练目标

1.1.1 数学形式

给定一个由 Token 组成的本文序列 x=(x1,x2,,xT),其中 xiVV 为词表),Decoder-Only 模型通过链式法则将联合概率分解为条件概率的乘积:

pθ(x1:T)=t=1Tpθ(xtx<t).

其中 θ 为模型全部可学习参数。这是自回归语言模型的核心数学假设:每个位置仅依赖于其左侧的所有历史 token。

训练目标通常是最大化训练语料的似然:

maxθi=1Nt=1Tilogpθ(xt(i)x<t(i))

等价于最小化负对数似然:

L(θ)=i=1Nt=1Tilogpθ(xt(i)x<t(i))

在深度学习实现中,这通常就是 cross entropy loss

推理时,Decoder-only 则把“指令、上下文、回答”全部串成一个 token 序列:

s=[x1,,xm,y1,,yn],pθ(s)=t=1m+npθ(sts<t).

1.1.2 从统计的角度看

对 token 序列 x1:T

pθ(x1:T)=t=1Tpθ(xtx<t).

最大化训练数据的对数似然等价于最小化 next-token 交叉熵:

L(θ)=1Ntoki=1nt=1Tilogpθ(xt(i)x<t(i)).

从统计角度看,这是对真实条件分布 p(xtx<t) 的经验风险最小化。其总体风险满足

Ep[logpθ]=H(p)+DKL(ppθ),

证明:

Ep[logpθ]=i=1nt=1Tip(xt(i)x<t(i))logpθ(xt(i)x<t(i))记为plogpθ=plogpθpp=plogppθplogp

其中:

H(p)=plogp,DKL(ppθ)=plogppθ

所以理想情况下,最小化交叉熵就是缩小模型条件分布与真实条件分布之间的 KL 散度。

1.2 Decoder-only 的整体架构

Decoder-only Transformer 的主体结构并没有脱离原始 Transformer:

EmbeddingDecoder Block1Decoder BlockLLM Head

每个 Decoder Block 通常包含:

Masked Self-Attention+Feed-Forward Network+Residual Connection+Normalization

相对于原始 Transformer Decoder,Decoder-only 的架构:

  • 去掉 Encoder;
  • 去掉 Encoder–Decoder Cross-Attention;
  • 保留带 causal mask 的 Self-Attention;
  • 保留逐位置 FFN、残差连接、归一化;
  • 最后用 LM Head 把隐藏向量映射成词表上的概率分布。

现代模型通常也不再原样使用 2017 Transformer 的配置,而常采用以下组合:

Pre-Norm+RMSNorm+RoPE+SwiGLU+GQA/MQA/MLA.

这些是常见设计,不是 Decoder-only 的数学定义;不同模型会选择不同变体。

1.2.1 一个 Decoder block 的计算

设第 层输入为 H()RT×d。以常见的 Pre-Norm 结构为例:

H~()=H()+Attn(Norm(H())),(Attn是Causal Attn)H(+1)=H~()+FFN(Norm(H~())).

Pre-Norm 中存在接近恒等映射的残差通路,梯度更容易穿过深层网络,通常比 Post-Norm 更适合大规模训练。

1.2.2 输入和输出流程

token idsEmbeddingH(0)L 个 Decoder blocksH(L)NormZWvocablogits.

对位置 t

zt=Wvocabht+b,pθ(xt+1=vxt)=ezt,vu=1Vezt,u.

常见的 weight tying 令输出矩阵与输入词嵌入共享参数:

Wvocab=E.

1.3 关键组件演进

在大模型预训练中,为了提升训练稳定性、收敛速度、数值效率和参数利用率,许多细节组件发生了系统性演进。这些变化看似是“工程细节”,但实际上直接影响大模型能否稳定扩展到几十亿、几百亿甚至更大规模。

1.3.1 归一化位置

1. Post-Norm:原始 Transformer 的结构

Post-Norm 结构,即先经过子层和残差连接,再做 LayerNorm:

xl+1=LN(xl+F(xl))

具体到 Decoder Block 中,可以写成:

x~l=LN(xl+Attn(xl))xl+1=LN(x~l+FFN(x~l))

Post-Norm 的优点是每一层输出都经过归一化,因此表面上看数值范围较稳定。

但问题是:当模型很深时,梯度需要反向穿过多个 LayerNorm 和非线性模块,容易出现训练不稳定:

Lxl=Lxl+1LN(xl+F(xl))xl

由于 LayerNorm 位于残差连接之后,残差路径也要经过 LayerNorm 的变换。这意味着原本应该提供稳定梯度通道的残差连接被 LayerNorm 限制住了。

2. Pre-Norm:现代大模型更常用的结构

Pre-Norm 把归一化放到子层之前:

xl+1=xl+F(LN(xl))

具体到 Decoder Block:

x~l=xl+Attn(LN(xl))xl+1=x~l+FFN(LN(x~l))

这个结构中,残差主路径是:

xlxl+F(LN(xl))

因此至少有一部分梯度可以沿着近似恒等映射传播:

xl+1xlI+other terms

这使得深层网络更容易训练。

Pre-Norm 更稳定,但它也可能带来一个问题:由于每一层都是在残差主干上增加一个较小更新,模型深层部分的表示更新可能偏保守。

所以,从训练角度看,Pre-Norm 的本质作用是:

让残差路径更顺畅,从而改善深层 Transformer 的梯度传播

1.3.2 归一化方法

1. LayerNorm

给定某一层某个 token 的隐藏向量 xRd,LayerNorm 对这个向量的特征维度做归一化:

LN(x)=γxμσ2+ϵ+β

其中:

μ=1di=1dxi,σ2=1di=1d(xiμ)2

(\gamma,\beta\in\mathbb R^d) 是可学习参数。

它的作用是稳定不同层、不同 token 的激活分布。

2. RMSNorm:去掉均值中心化,只保留尺度归一化

RMSNorm 的核心思想是:

不减均值,只根据向量的均方根大小进行缩放。

给定 xRd,定义均方根:

RMS(x)=1di=1dxi2+ϵ

RMSNorm 写作:

RMSNorm(x)=γxRMS(x)

RMSNorm 论文的核心假设是:LayerNorm 中的 re-centering invariance 不是必须的,因此可以只保留 re-scaling invariance,从而减少计算量。


1.3.3 激活函数

1. GELU

GELU,即 Gaussian Error Linear Unit,定义为:

GELU(x)=xΦ(x)

其中 (\Phi(x)) 是标准正态分布的累积分布函数。GELU 最早作为一种高性能神经网络激活函数被提出,后来被 BERT、GPT 等 Transformer 模型广泛使用。

输入 (x) 是否被保留,不是由 (x>0) 这个硬阈值决定,而是由一个连续概率权重决定

2. SiLU

定义为:

SiLU(x)=xσ(x),σ(x)=11+ex

它与 GELU 类似,也是一种“输入乘以门控函数”的形式。

从表达能力角度看:

GELU / SiLU 比 ReLU 更适合大规模 Transformer 的连续表示学习

1.3.4 前馈网络

1. 标准 FFN

Transformer 中的 FFN 是逐 token 作用的多层感知机。

给定隐藏向量 xR1×dmodel,标准 FFN 通常写作:

FFN(x)=ϕ(xW1+b1)W2+b2

其中:W1Rdmodel×dff,W2Rdff×dmodel,通常 dff4dmodel

所以它先升维 dmodeldff,再降维 dffdmodel

直观上,Attention 负责 token 之间的信息交互,而 FFN 负责对每个 token 的表示进行非线性变换

2. GLU:引入门控机制

GLU,即 Gated Linear Unit,基本形式是:

GLU(x)=(xW1)σ(xW2)

其中:

  • (xW_1):内容分支;
  • (\sigma(xW_2)):门控分支;
  • (\odot):逐元素乘法。

从统计建模角度看,可以理解为输入相关的特征选择:

zj=contentj(x)gatej(x)

其中每个通道 (j) 的保留程度由输入 (x) 自适应决定。

3. SwiGLU:现代 LLM 中常见的 FFN 结构

SwiGLU 是 GLU 的一种变体,把 sigmoid gate 换成 SiLU / Swish 型门控:

SwiGLU(x)=(Wax)SiLU(Wbx)

然后再接输出投影:

FFNSwiGLU(x)=W2[(Wax)SiLU(Wbx)]

GLU 变体研究表明,在 Transformer 中使用 GEGLU、SwiGLU 等门控 FFN 变体通常可以改善模型效果;这些变体后来也成为许多大模型的常见组件

4. 标准 FFN 与 SwiGLU 的结构对比

标准 FFN:

xW1xϕ(W1x)W2ϕ(W1x)

SwiGLU FFN:

x{WupxWgatexSiLU(Wgatex)WupxWdown

核心区别是:

SwiGLU 多了一条与输入 x 相关的门控路径
  • 增强通道选择能力
  • 提高非线性表达能力

1.3.5 注意力头

MHA → MQA / GQA
1. MHA

给定输入隐藏状态 XRT×d,对每个注意力头 (h=1,\dots,H),分别计算:

Qh=XWhQ,Kh=XWhK,Vh=XWhV,Qh,Kh,VhRd×dhead

每个头的注意力输出为:

Attnh(Qh,Kh,Vh)=softmax(QhKhdhead+M)VhRT×dhead

最后把所有头拼接:

MHA(X)=Concat(head1,,headH)WORT×d

MHA 的优点是:每个头有独立的 (Q,K,V),可以学习不同的关系模式。

但是在大模型推理阶段,MHA 有一个明显问题:KV Cache 很大

Decoder-only 模型是自回归生成 pθ(xtx<t),生成第 (t) 个 token 时,需要使用之前所有 token 的 Key 和 Value。

为了避免每一步都重新计算历史 token 的 (K,V),推理时会缓存它们:

KV Cache={K1,V1,,Kt1,Vt1}

对于 MHA,每一层、每个注意力头都要保存自己的 (K,V)。

如果模型有:

  • 层数:(L)
  • 注意力头数:(H)
  • 序列长度:(T)
  • 每个头维度:(d_{\text{head}})

那么 KV Cache 的规模大致是:

O(2LHTdhead)=O(2LTd)

当 batch size、上下文长度和层数都较大时,KV Cache 会成为推理显存和带宽瓶颈。


2. MQA

MQA, Multi-Query Attention 的核心思想是:

Query 仍然保留多个头,但所有 Query heads 共享同一组 Key 和 Value。

标准 MHA 是:

Qh=XWhQ,Kh=XWhK,Vh=XWhV

MQA 则在其基础上,

W1K=W2K==WhK,W1V=W2V==WhV

即:

Qh=XWhQ,K=XWK,V=XWV

也就是说:

多个 Query heads+一组共享 KV heads

KV Cache 规模变化:

O(2LHTdhead)O(2LTdhead)

当然,因为所有头共享 KV,模型的表达能力可能下降。


3. GQA

GQA, Grouped-Query Attention 是 MHA 和 MQA 之间的折中方案。

它的思想是:

Query heads 仍然有很多个,但 Key/Value heads 不是 1 个,而是分成若干组,每组 Query heads 共享一组 Key/Value。

例如:

Hq=32,Hkv=8

则每 4 个 Query heads 共享一组 Key/Value。

QA 是当前很多 Decoder-only LLM 的常见选择,因为它在效果和推理效率之间取得平衡。

方法表达能力KV Cache推理效率
MHA
MQA较弱最小
GQA接近 MHA明显降低较快

可以理解为:

GQA 是 MHA 与 MQA 的折中:比 MQA 表达力更强,比 MHA 推理更省

GQA 论文也强调,它可以在接近 MHA 质量的同时达到接近 MQA 的推理速度


1.3.6 参数初始化方式

参数初始化的核心目标是让前向传播和反向传播的尺度保持稳定。

理想情况下,希望:

Var(xl)Var(xl+1)Var(Lxl)Var(Lxl+1)

如果初始化过大:

Fl(xl)

会很大,残差更新过强,训练初期容易不稳定。

如果初始化过小:

Fl(xl)0

模型训练初期接近恒等映射,虽然稳定,但学习可能较慢。

所以初始化本质是在平衡:

稳定性有效更新幅度
1. 标准正态初始化

早期 GPT 类模型常使用类似正态初始化:

WN(0,σ2)
2. Xavier / Glorot 初始化

若输入维度为 (d_{\text{in}}),输出维度为 (d_{\text{out}}),Xavier 初始化目标是。

前向传播方差分析

xRdin,WRdout×din,前向传播为:

y=WxRdout
  • x 的各分量独立,E[x]=0,协方差矩阵 Σx=σx2Idin
  • W 的各元素独立,E[Wij]=0Var(Wij)=σW2Iout
  • Wx 相互独立
Σy=E[yy]=E[WxxW]=EW(Ex(WxxTWT|x))=EW(WΣxWT)=σx2EW(WWT)

计算 E(WWT)

E[(WWT)ij]={dinσW2if i=k0if ikE[WW]=dinσW2Idout

代回得到:

Σy=σx2dinσW2Idout

希望逐层传播后方差不变,即前向传播方差稳定条件 Var(yj)=Var(xi)

dinσW2σx2=σx2σW2=1din

同样的方法,可以得到反向传播方差稳定条件

σW2=1dout

想要同时保持前向和反向方差稳定,Xavier 的做法是取两者的调和平均

1Var(W)=12(11/din+11/dout)=din+dout2Var(W)=2din+dout

常见形式:

WijU(6din+dout,6din+dout)

或:

WijN(0,2din+dout)
3. 残差分支缩放初始化

首先,我们看残差叠加带来的尺度问题。如果每一层都执行:

xl+1=xl+Fl(xl)xL=x0+l=0L1Fl(xl)

如果每个 (F_l(x_l)) 的方差大致相同,假设:

Var(Fl(xl))σ2

如果有 (L) 层,并且各层更新近似独立,则有:

Var(xL)Var(x0)+Lσ2

这说明层数越深,残差累积可能导致激活尺度增大。

为了控制残差分支的输出尺度,可以让某些输出投影权重初始化更小,例如:

WoutN(0,σ22L)

总结

模块早期形式现代常见形式主要目的
归一化位置Post-NormPre-Norm深层训练稳定
归一化方法LayerNormRMSNorm降低计算量,稳定尺度
激活函数ReLUGELU / SiLU更平滑的非线性表达
前馈网络标准 FFNGLU / SwiGLU增强门控与通道选择
注意力头MHAMQA / GQA降低 KV Cache,提高推理效率
残差路径普通残差更重视残差尺度控制防止深层累积失控
参数初始化普通初始化深度相关初始化 / 残差缩放控制前向与反向方差
训练配置Bias + Dropout 常用减少 Bias / 减少 Dropout简化结构,提高大规模训练效率

1.4 MoE稀疏专家模型

MoE,全称 Mixture of Experts,即专家混合模型。在大语言模型中,它通常指:

用多个专家网络共同组成模型,但每个 token 只激活其中少数几个专家

所以 MoE 的核心不是“把模型变小”,而是:

增加总参数量,但保持每次计算只使用一小部分参数

在标准稠密 Transformer 中,每个 token 必须经过全部参数的计算。模型性能随参数量增长而提升(Scaling Laws),但推理 FLOPs、显存、延迟也线性增长。

MoE 的核心目标是解耦两件事:

  • 总参数量(模型容量/知识储备)→ 可以非常大

  • 每个 token 的激活参数量(实际计算量)→ 保持很小

1.4.1 核心架构

在标准 Decoder-only Transformer 中,每一层通常有一个 FFN:

FFN(x)=ϕ(xW1)W2

所有 token 都经过同一个 FFN。

而 MoE 的思想是把一个 FFN 替换成多个专家,即 N 个结构相同的 FFN 子网络,各自拥有独立参数:

E1,E2,,ENEi(x)=ϕ(xWiin)Wiout,i=1,2,,N

然后,对于每个 token 表示 (x),模型通过一个门控网络判断应该交给哪些专家处理:

给定 token 表示 xR1×d,计算所有专家的打分:

s(x)=xWgR1×N,WgRd×N

然后通过一些方法将得分转化为概率权重:

p(x)=gate(s(x))

其中 pi(x) 表示 token x 被分配给专家 i 的权重或概率。

最后计算 MoE FFN:

y=i=1Npi(x)Ei(x)

其中:

  • (E_i(x)):第 (i) 个专家网络;
  • (p_i(x)):第 (i) 个专家的权重;
  • (N):专家总数。

如果所有专家都参与计算,这就是 dense MoE。

但大语言模型中常用的是 sparse MoE,即只激活少数专家,这是由 gate 的方式决定的。

1.4.2 Top-k 稀疏激活

Top-k 稀疏激活:不是所有专家都参与计算,而是只选择得分最高的 (k) 个专家

定义:

Tk(x)=TopK(p(x))

其中:

Tk(x){1,2,,N}

表示 token (x) 选中的 (k) 个专家。

于是 Sparse MoE 的输出为:

MoE(x)=iTk(x)p~i(x)Ei(x)

其中:

p~i(x)=pi(x)jTk(x)pj(x)

表示在选中的 Top-k 专家内部重新归一化后的权重。

所以完整形式是:

MoE(x)=iTopK(softmax(xWg))p~i(x)Ei(x)

总结:

MoE = 大容量参数空间 + 稀疏激活计算

1.4.3 专家容量与 Token 分配

1.4.4 负载均衡与辅助损失

1.4.5 Expert Parallelism 与 All-to-All 通信

1.4.6 MoE 的收益、代价与适用场景

二、Scaling Law

Scaling Law 可以理解为规模定律。在大语言模型预训练中,它关心的是:

模型规模、数据规模、计算量与模型损失之间的经验规律

更具体地说,它试图回答三个问题:

  1. 模型参数 (N) 增大,模型效果会怎么变?
  2. 训练 token 数 (D) 增大,模型效果会怎么变?
  3. 计算预算 (C) 固定时,应该把预算分给更大的模型,还是更多的数据?

通常记:

符号含义
(N)模型参数量,number of parameters
(D)训练 token 数,dataset size
(C)训练计算量,compute,常用 FLOPs 表示
(L)语言模型的预测损失,通常是 cross entropy loss

2.1 核心幂律关系

2.1.1 Scaling Law 的基本形式

Scaling Law 的核心观察是:

语言模型的 loss 会随模型规模、数据规模、计算量呈幂律下降

Kaplan et al. 的 Scaling Laws for Neural Language Models 系统研究了模型大小、数据集大小和训练计算量与语言模型交叉熵损失之间的幂律关系,并指出这些趋势可跨越多个数量级。OpenAI

一个常见的通用形式是:

L(x)=L+Axα

其中:

  • (L(x)):给定规模 (x) 下的 loss;
  • (L_\infty):不可约损失,表示即使规模无限增大也难以消除的部分;
  • (A):常数系数;
  • (\alpha>0):幂律指数;
  • (x):可以是模型参数量 (N)、数据量 (D)、计算量 (C)。

也就是说,loss 的下降不是线性的,而是:

LLxα

2.1.2 三种幂律

通用形式具体表现为三种幂律:

  • 模型规模幂律:当数据充足、训练充分时,增大参数量可降低 loss:
L(N)=L+ANNαN
  • 数据规模幂律:当模型足够大时,增加训练 token 数可降低 loss:
L(D)=L+ADDαD
  • 计算量幂律:训练计算量 C6NDN 为非 embedding 参数量,系数 6 来自前向+反向传播的 FLOPs 估算),loss 随计算预算呈幂律下降:
L(C)=L+ACCαC

所以 Scaling Law 告诉我们:

扩大模型有效,但收益递减

2.1.3 三种瓶颈

实际预训练时,模型效果可能被不同因素限制。

  • 模型瓶颈N 太小,模型容量不足以吸收数据中的信息。此时继续增加数据收益有限,应优先扩大模型。
  • 数据瓶颈D 不足,模型反复看相同数据,容易过拟合。此时继续增大模型不划算,应优先增加高质量 token。
  • 计算瓶颈C 不足,训练步数不够,模型未充分收敛。此时无论模型还是数据多好,都无法达到预期 loss。

2.2 计算最优分配原则

2.2.1 问题:固定计算预算下,模型和数据怎么分配?

Scaling Law 最重要的实际问题是:

给定计算预算 C, 如何选择模型参数量 N 和训练 token 数 D

因为计算量近似满足:

C6ND

所以当 (C) 固定时,预训练规划的核心矛盾是:

是训练更大的模型,还是让较小模型看更多数据?

2.2.2 Kaplan Scaling Law:更重视增大模型

Kaplan et al. 的早期 Scaling Law 研究认为,在计算最优条件下,模型参数量应随计算预算增长得更快,而数据规模增长相对较慢。这个结论影响了早期大模型扩展路径,即更倾向于训练更大的模型。OpenAI

可以粗略理解为:

C↑⇒N 增长较快,D 增长较慢

这条路线对应早期很多超大参数模型,例如 GPT-3 这类模型体现了“扩大参数规模带来 few-shot 能力提升”的思路;OpenAI 对 GPT-3 的介绍也强调,扩大语言模型显著提升了任务无关的 few-shot 性能

2.2.3 Chinchilla Scaling Law:模型和数据应近似等比例增长

后来的 Chinchilla 工作重新研究了计算最优训练问题,结论发生了重要变化:

[ \boxed{ \text{在计算最优条件下,模型规模和训练 token 数应近似等比例增长} } ]

Hoffmann et al. 在 Training Compute-Optimal Large Language Models 中指出,对于 compute-optimal training,模型大小和训练 token 数应当同时扩展;他们训练的 Chinchilla 使用与 Gopher 相同的计算预算,但参数量更小、训练数据约多 4 倍,并取得更好的效果。arXiv

形式上可以写成:

NoptCaDoptCb

Chinchilla 结论中:

ab12

这意味着:

C 增加 k 倍时,N 和 D 都大约增加 k 倍

2.2.4 Chinchilla 的直观含义:很多大模型其实 under-trained

Chinchilla 结论的一个重要启发是:

许多早期大模型参数很多,但训练 token 不够

也就是说,它们不是模型太小,而是数据看得不够多。

如果一个模型参数量非常大,但训练 token 数不足,那么它可能处于 under-trained large model状态。

此时,与其继续增加参数量,不如在相同计算预算下训练一个较小模型,但给它更多 token:

smaller model+more tokens

这可能获得更好的 loss 和下游表现。

2.2.5 计算最优的数学直觉

假设 loss 可以近似分解为:

L(N,D)=L+ANα+BDβ

其中:

  • (\frac{A}{N^\alpha}):模型容量不足带来的误差;

  • (\frac{B}{D^\beta}):数据不足带来的误差。

代入计算约束 C=6ND

L(N,D)=L+ANα+B(6NC)β

所以存在一个优化方式:

Nopt=argminNL(N,D)

这就是计算最优分配的本质。

2.3 对预训练规划的指导

在预训练规划中,最重要的不是单纯追求更大的参数量,而是求解:

(N,D)=argmin6NDCL(N,D)

也就是在固定计算预算下,找到模型规模 (N) 和训练 token 数 (D) 的最优匹配。

2.3.1 从硬件预算估计总计算量

2.3.2 选择模型参数规模和训练 Token 总量

2.3.3 由全局 Batch Size 推算训练步数

2.3.4 规划上下文长度与样本数量

2.3.5 通过小模型试验拟合 Scaling 曲线

在真正训练大模型之前,可以先训练一系列小模型,并使用不同 token 数:

N1<N2<<NmD1<D2<<Dk

得到实验点:

(Ni,Dj,Lij)

然后拟合:

L(N,D)=L+ANα+BDβ

如果拟合稳定,就可以外推到更大规模

2.3.6 为消融实验、评估与故障预留预算

小结:预训练规划的完整推导链条

Scaling Law 对预训练规划的指导可以整理成一条完整链路:

硬件数量、算力、训练时间CtotalCtotalCmainCmain6ND(N,D) 的匹配关系D, BtokensTsteps=DBtokensD, SM=DS(Ni,Dj,Lij)拟合 Scaling Curve(N\*,D\*)

最终,预训练规划不是简单地决定一个模型大小,而是在求解:

(N,D,B,S)=argminL(N,D,B,S)s.t.6NDCmain,memory(N,B,S)GPU memory,time(N,D,B,S)training time.

其中:

  • (N):模型参数量;
  • (D):训练 token 总量;
  • (B):全局 batch token 数;
  • (S):上下文长度;
  • (C_{\text{main}}):主训练计算预算。

三、数据工程

3.1 数据在预训练中的核心地位

3.2 数据来源与类型

3.3 数据规模规划

3.4 数据清洗流水线

3.5 数据配比与混合策略

3.6 Tokenizer

四、完整训练流程

4.1 训练循环基础

可以把一次训练迭代理解为:

读取文本 batch构造输入与标签模型前向传播计算 token-level loss反向传播参数更新

对于 Decoder-only 语言模型,训练目标通常是 next-token prediction

pθ(xt+1x1,x2,,xt)

也就是说,模型看到前面的 token,预测下一个 token。

4.1.1 Batch 读取与样本构造

文本数据首先被 tokenization

对于一段文本,tokenization 就是通过查对应的词表 V(词表是提前构造好的),对每个词找到对应的词表中的位置 xi,从而将整段文本转换成一个整数序列:

xi{0,1,,V1}

Batch 构造

大模型训练时,不一定按照“句子”或“文章”作为基本单位,而是把大量语料拼接成一个很长的 token 序列。

例如我们将一个完整的训练语料拼接成一个长度 N 的 token 序列:

[x1,x2,x3,,xN]

我们通常不会一次性将整个 token 序列全部输入到模型中训练,首先这对显存的要求很高,另外模型的长下文长度不支持。

大多时候,我们是将 token 序列切片为一个个固定长度的片段,记长度为 T,那么就总共有 NT 条长度为 T 的序列。

由于 NT 通常也很多,需要分批次将这些序列输入到模型中,每个批次的大小就是 batch size,记为 B,那么一个 batch 的输入形状就是:

XRB×T

标签右移:构造标签 YRB×T ,通常是 X 向右偏移一位(即预测下一个 token)

4.1.2 前向传播与 Logits

将 token ID embedding:

XRB×TEmbedding()H(0)RB×T×d

经过模型主体:

H(0)L 个 Decoder blocksH(L)NormZWvocablogitsRB×T×V.

这里不展开 Decoder-only 内部结构,比如 self-attention、causal mask、MLP、残差连接等,因为这些属于架构部分。

4.1.3 Token-level Loss

单个 token 的交叉熵损失

如果第 (b) 条样本第 (t) 个位置的真实下一个 token 是 (y_{t}^{b}),那么该位置的 loss 为:

tb=logpt,ytbb

其中:

ptb=[pt,1b,pt,2b,,pt,Vb]

表示模型给第 (b) 条样本第 (t) 个位置预测的概率分布,pt,ytbb 就表示表示模型给真实 token 分配的概率。

单个 token 的交叉熵损失

一个 batch 中有 (B\times T) 个预测位置,因此总 loss 通常取平均:

L(θ)=1BTb=1Bt=1Tb,t

代入交叉熵形式:

L(θ)=1BTb=1Bt=1Tlogpθ(yb,txb,t)

Padding token 的处理

有时候,batch 中不同文本长度不同,可能需要 padding,而 padding 位置不应该参与 loss

因此通常会引入 mask:

mb,t={1,该位置是真实 token0,该位置是 padding

loss 改为:

L(θ)=b=1Bt=1Tmb,tlogpθ(yb,txb,t)b=1Bt=1Tmb,t

也就是说,只对真实 token 求平均,不让 padding 影响训练。

在很多大规模预训练场景中,数据会被拼接和截断成固定长度,所以 padding 可能较少;但在微调、对话数据训练或变长样本训练中,padding mask 很常见。

4.1.4 反向传播与梯度

最基础的梯度下降更新为:

θk+1=θkηθL(θk)

其中:

  • (k):第 (k) 次参数更新;
  • (\eta):学习率;
  • (\nabla_\theta \mathcal{L}(\theta_k)):当前 batch 上计算得到的梯度。

梯度指向 loss 上升最快的方向,所以要沿着负梯度方向更新。

一次前向传播可以抽象为:

XHOL

其中:

  • (X):输入 token;
  • (H):hidden states;
  • (O):logits;
  • (\mathcal{L}):loss。

反向传播使用链式法则:

Lθ=LOOHHθ

现代深度学习框架会自动构建计算图,框架会自动沿着计算图从 loss 反向传播到每一层参数。

4.1.5 梯度累积与梯度裁剪

1. 梯度累积

为什么需要梯度累积?

大模型训练中,为了训练稳定,我们往往需要较大的 Global Batch Size(如 256、512 甚至更大)。但受限于 GPU 显存,单卡一次能塞进去的 Micro Batch Size 可能只有 4 或 8。

梯度累积的核心思想是:用多次小 batch 的前向/反向传播,模拟一次大 batch 的效果

原理

假设 Global Batch Size = 32,单卡最多放 Micro Batch Size = 8,那么累积步数 k=32/8=4

具体流程:

  • 取第 1 个 micro batch,前向传播,计算 loss,反向传播得到梯度 1
  • 取第 2 个 micro batch,前向传播,计算 loss,反向传播得到梯度 2
  • 取第 3 个 micro batch,前向传播,计算 loss,反向传播得到梯度 3
  • 取第 4 个 micro batch,前向传播,计算 loss,反向传播得到梯度 4
  • 将 4 次梯度求平均
¯=14(1+2+3+4)
  • ¯ 更新参数
  1. PyTorch 中 loss.backward() 默认将梯度累加param.grad 上而非覆盖。
  2. 标准训练中,我们在每次 backward() 前调用 optimizer.zero_grad() 清零梯度,因此每次 param.grad 只保留当前 batch 的梯度。
  3. 而梯度累积正是利用这一累加机制,故意在 k*k* 次 backward() 之间不调用 zero_grad(),使多次梯度在 param.grad 上叠加,最后统一更新参数。
2. 梯度裁剪

大模型训练中可能出现梯度过大,导致参数更新不稳定。

常见做法是梯度裁剪,有两种主流方法:

方法一:按范数裁剪(Gradient Clipping by Norm)

计算所有参数的梯度范数( g=θL),如果超过阈值 C,就按比例缩小:

ifg2>C:ggCg2

PyTorch 中常见写法:

python
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)

方法二:按值裁剪(Gradient Clipping by Value)

直接将每个梯度元素截断到 [τ,τ] 范围内:

giclip(gi,τ,τ)=max(τ,min(τ,gi))

直觉理解:对梯度的每个分量独立地"削峰",可能会改变梯度方向。

python
torch.nn.utils.clip_grad_value_(model.parameters(), clip_value=1.0)

主流大模型(GPT、LLaMA 等)的预训练中,几乎都使用按范数裁剪,典型配置:

  • 阈值 τ=1.0
  • 在每个 micro batch 的 backward() 之后、optimizer.step() 之前执行
  • 配合梯度累积时,通常在最后一次 backward 之后统一裁剪

在 PyTorch 风格的训练循环中,一步训练通常是:

optimizer.zero_grad()

logits = model(input_ids)

loss = cross_entropy(
    logits.view(-1, vocab_size),
    labels.view(-1)
)

loss.backward()

torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)

optimizer.step()

4.2 优化器

4.2.1 SGD

最基础的优化器是随机梯度下降,简称 SGD。

假设第 (t) 步参数为 (\theta_t),当前 batch 上的梯度为:

gt=θL(θt)

SGD 的更新规则是:

θt+1=θtηgt

为什么叫随机梯度下降?

理论上的总体目标函数是:

L(θ)=E(X,Y)D[(fθ(X),Y)]

也就是对整个数据分布求期望。

但实际训练时,我们不会每一步都用完整数据集计算梯度,因为其计算量很大,解决方法是从数据集中采样一个 mini-batch:

Bt={(Xi,Yi)}i=1BD

然后用 mini-batch loss 近似总体 loss:

L^t(θ)=1Bi=1B(fθ(Xi),Yi)

对应的梯度:

gt=θL^t(θt),E(gt)=θLt(θt)

是总体梯度无偏估计。"无偏"意味着平均来看方向是对的,但每一次具体采样都会有随机波动——这就是"随机"二字的来源

SGD 的问题

SGD 简单,但在大模型训练中通常不够稳定。

主要问题有三个:

问题含义
梯度噪声大每个 batch 只是总体梯度的近似
不同参数尺度不同有些参数梯度大,有些参数梯度小
收敛路径震荡在高曲率方向容易来回摆动

所以现代大模型训练通常不直接使用普通 SGD,而是使用带动量或自适应学习率的优化器。

4.2.2 Momentum

Momentum 的思想是:

不要只看当前 batch 的梯度,还要参考过去若干步的平均方向。

定义速度变量 (v_t):

vt=βvt1+(1β)gt

其中:

  • (g_t):当前梯度;
  • (v_t):梯度的指数滑动平均;
  • (\beta\in[0,1)):动量系数,常见取值如 (0.9)。

参数更新为:

θt+1=θtηvt

可以把 Momentum 理解为:

当前更新方向=当前梯度+历史梯度惯性

vt 展开后为:

vt=(1β)gt+(1β)βgt1+(1β)β2gt2+

也就是历史梯度的加权平均,越近的梯度权重越大。

SGD 的问题Momentum 的改善
batch 梯度噪声大历史平均可以降低噪声
路径震荡高频震荡方向会被抵消
有效方向前进慢一致方向会被加速

4.2.3 Adam

Adam(Adaptive Moment Estimation) 可以看作结合了两类信息:

  1. 一阶矩估计:梯度的滑动平均,类似 Momentum;
  2. 二阶矩估计:梯度平方的滑动平均,用来调整每个参数的步长。

设当前梯度为:

gt=θL(θt)

Adam 维护一阶矩(梯度方向的平均)

mt=β1mt1+(1β1)gt

其中:

  • (m_t):梯度的一阶指数滑动平均;
  • (\beta_1):一阶矩衰减系数,常见取值 (0.9)。

Adam 还维护二阶矩:梯度平方的平均

vt=β2vt1+(1β2)gt2

其中:

  • (v_t):梯度平方的指数滑动平均;
  • (\beta_2):二阶矩衰减系数,常见取值 (0.999);
  • gt2 是逐元素平方

如果某个参数长期梯度很大,那么对应的 (v_{t,i}) 会变大。

Adam 更新时会用 (\sqrt{v_{t,i}}) 去缩放梯度:

mt,ivt,i+ϵ

误差修正

由于 Adam 初始化时,(m_0=0, v_0=0),因此在训练初期,(m_t) 和 (v_t) 会被系统性低估:

mt=β1mt1+(1β1)gt=i=1t1β1ti(1β1)gi+(1β1)gt=i=1tβ1ti(1β1)gi

假设各步梯度 gi 独立同分布,且 E[gi]=E[g] (即真实梯度):

E[mt]=E[g](1β1)i=1tβ1ti=E[g](i=1tβ1tii=1tβ1t+1i)=E[g](i=0tβ1tii=0t1β1ti)=(1β1t)E[g]

同理:

E[vt]=(1β2t)E[g2]

为了修正这种偏差,Adam 使用:

m^t=mt1β1t,v^t=vt1β2t

这其实是一个无偏化操作:

E[mt^]=E[g],E[vt^]=E[g2]

参数更新为:

θt+1=θtηm^tv^t+ϵ

Adam 的一轮完整更新可以写成

gt=θL(θt)mt=β1mt1+(1β1)gtvt=β2vt1+(1β2)gt2m^t=mt1β1tv^t=vt1β2tθt+1=θtηm^tv^t+ϵ

Adam 的直观解释:

更新方向=梯度的一阶滑动平均梯度尺度的滑动估计方向由 mt 决定,尺度由 vt 调整

4.2.4 AdamW

现代大语言模型训练中,更常用的是 AdamW,而不是原始 Adam。

AdamW 的关键改动是:

把 weight decay 从梯度更新中解耦出来。

权重衰减

权重衰减的目标是限制参数规模,防止参数无限变大。

常见形式是在目标函数中加入 (L_2) 正则项:

Lreg(θ)=L(θ)+λ2θ22

对它求梯度:

θLreg(θ)=θL(θ)+λθgtreg=gt+λθt

在 SGD 中,加入 (L_2) 正则后:

θt+1=θtη(gt+λθt)=(1ηλ)θtηgt

可以看到,参数被乘上了一个衰减因子1ηλ

但是在 Adam 中,更新是:

θt+1=θtηm^tv^t+ϵ

如果把 (\lambda\theta_t) 加进梯度里,由于 (m_t) 和 (v_t) 是依赖梯度的,它也会进入 (m_t) 和 (v_t),然后被自适应缩放。这会导致 λθt 不再是一个单纯的参数衰减项,而是被 Adam 的二阶矩调整过。

AdamW 的更新方式

AdamW 将梯度更新和权重衰减分开。

先做 Adam 更新:

θtmt,vt

同时单独加入权重衰减:

θt+1=θtηm^tv^t+ϵηλθt

其中:

  • (\lambda):weight decay 系数;
  • ((1-\eta\lambda)\theta_t):直接对参数做衰减;
  • Adam 部分只处理数据 loss 的梯度。
Adam + L2AdamW(解耦权重衰减)
正则项是否进入 $ m_t, v_t $✅ 是❌ 否
衰减是否被自适应学习率缩放✅ 是(不均匀)❌ 否(等比例)
大参数的衰减效果偏弱正常
小参数的衰减效果偏强正常

在大语言模型预训练中,常见优化器配置格式大致如下:

optimizer = torch.optim.AdamW(
    model.parameters(),
    lr=learning_rate,
    betas=(0.9, 0.95),
    eps=1e-8,
    weight_decay=0.1
)

4.2.5 不参与权重衰减的参数

在 AdamW 中,权重衰减表示在每次参数更新时,AdamW 会额外把参数往 0 的方向拉一小步:

θt(1ηtλ)θt

这对很多权重矩阵是有益的,但并不是所有参数都应该参与 weight decay。

weight decay 的本质是限制参数规模:

θ2

它适合用于普通权重矩阵,例如:

WQ, WK, WV, WO, WMLP

这些参数通常具有较高维度,并且控制模型的主要表示能力。适当约束它们的范数,可以起到正则化作用。

但是某些参数的功能不是“学习复杂映射”,而是控制尺度、平移或特殊结构。对这些参数做 weight decay,可能会破坏它们的作用。

通常不参与 weight decay 的参数:

参数类型是否 weight decay原因
bias 参数偏置项只负责平移,不适合用范数约束
LayerNorm / RMSNorm 权重归一化层参数控制尺度,衰减会干扰归一化效果
BatchNorm 参数类似归一化尺度参数,通常不衰减
embedding 特殊参数视情况有些实现会衰减,有些不会
普通 Linear 权重矩阵主要学习映射关系,适合正则化

在大语言模型中,最常见的规则是:

decay: Linear weightsno decay: bias + norm weights

4.2.6 优化器状态的显存开销

在大模型训练中,显存不只被模型参数占用。

训练时显存主要来自:

参数+梯度+优化器状态+激活值+临时计算缓存

其中,AdamW 的优化器状态非常占显存。

现代大模型训练通常使用 FP16 或 BF16 做前向和反向。FP16/BF16 每个数占 2 bytes,但为了训练稳定,优化器状态通常仍然用 FP32 保存。

一种常见估算是:

对象精度显存
模型参数BF16 / FP16(2N) bytes
梯度BF16 / FP16(2N) bytes
FP32 master weightsFP32(4N) bytes
Adam 一阶矩 (m)FP32(4N) bytes
Adam 二阶矩 (v)FP32(4N) bytes

合计 16N bytes。这说明在不包括激活值和临时计算缓存的情况下,优化器状态(含 master weights)占总量的 75%(12N16N

为降低 AdamW 的优化器状态显存,常见方法包括:

方法核心思想
ZeRO Stage 1切分 optimizer states
ZeRO Stage 2进一步切分 gradients
ZeRO Stage 3切分 parameters、gradients、optimizer states
optimizer offload把优化器状态放到 CPU 或 NVMe
8-bit optimizer用低精度保存优化器状态
Adafactor减少二阶矩存储

具体方法后面再研究。

4.3 学习率调度

在训练语言模型时,学习率 learning rate, LR 通常不是固定不变的,而是会随着训练步数动态调整。

学习率调度,就是设计一个函数:

ηt=f(t)

使得学习率随着训练过程合理变化。

  • 训练初期:学习率太大容易不稳定

  • 训练后期:学习率太大不利于收敛


4.3.1 Learning Rate Warmup

1. Warmup 是什么?

Learning Rate Warmup 指的是:在训练初期,不直接使用较大的学习率,而是让学习率从一个很小的值逐步增加到目标最大值。

设 warmup 步数为 (T_{\text{warmup}}),最大学习率为 (\eta_{\max}),则线性 warmup 可以写成:

ηt=ηmaxtTwarmup,0tTwarmup
2. 为什么需要 Warmup?

训练刚开始时,模型参数通常是随机初始化的。此时:

θ0Random Initialization

模型输出非常不稳定,loss 也通常较大。对应的梯度可能存在较强波动:

θL(θ0)

如果一开始就使用较大的学习率,参数更新可能过大,从而引起训练不稳定。Warmup 的作用是让模型在训练初期先进行较温和的更新。


4.3.2 Constant、Linear 与 Cosine Decay

Warmup 结束后,学习率通常进入 decay 阶段。常见策略包括:

  1. Constant Learning Rate
  2. Linear Decay
  3. Cosine Decay
1. Constant Learning Rate

Constant 学习率指训练过程中学习率保持不变:

ηt=η

如果带 warmup,则通常是:

ηt={ηmaxtTwarmup,tTwarmupηmax,t>Twarmup
2. Linear Decay

Linear Decay 指学习率在 warmup 后从最大值线性下降到最小值。

设总训练步数为 (T),warmup 步数为 (T_{\text{warmup}}),最小学习率为 (\eta_{\min}),则整体形式为:

ηt={ηmaxtTwarmup,0tTwarmupηmax(ηmaxηmin)tTwarmupTTwarmup,Twarmup<tT

先大步搜索,再小步收敛

3. Cosine Decay

Cosine Decay 指 warmup 后学习率按照余弦曲线下降:

ηt=ηmax12(ηmaxηmin)[1+cos(πtTwarmupTTwarmup)]

整体形式为:

ηt={ηmaxtTwarmup,0tTwarmupηt=ηmax12(ηmaxηmin)[1+cos(πtTwarmupTTwarmup)],Twarmup<tT

笔记插图

4. Constant、Linear、Cosine 对比
策略公式特点优点缺点适用场景
Constant(\eta_t=\eta)简单,调试方便后期可能震荡小实验、调试
Linear Decay线性下降稳定、简单后期可能过早变小BERT 类训练、常规微调
Cosine Decay余弦下降平滑、常用、后期稳定多一个调度超参数大模型预训练、视觉模型训练

代码理解:

python
from transformers import get_scheduler

scheduler = get_scheduler(
    optimizer = optimizer,
    name = 'cosine',
    num_warmup_steps = 1000,
    num_training_steps = 10000,
    lr_min_ratio = 1e-2
)

for step, batch in enumerate(dataloader):
    loss = model(batch).loss
    loss.backward()
    optimizer.step()
    scheduler.step()
    optimizer.zero_grad()

4.3.3 Peak Learning Rate 与 Minimum Learning Rate

Peak LR:最大学习率 ηmax

Minimum LR:最小学习率 ηmin

Peak LR 与 Batch Size 的关系

一般来说,batch size 越大,可以使用的 peak LR 往往越大。

因为大 batch 下梯度估计更稳定,当 batch size (B) 增大时,梯度方差减小:

gt=1Bi=1Bθi(θt)Var(gt)=1B2Var(i=1Bθi(θt))=1BVar(θi(θt))

因此可以承受更大的学习率。但这不是无限成立。过大的 batch 和学习率仍可能损害泛化或导致训练不稳定。

Peak LR 与 Minimum LR 的整体关系

在大模型训练中,minimum LR 常常设置为 peak LR 的一个比例,例如:

ηmin=0.1ηmax

也有些训练直接让学习率衰减到 0。


4.3.4 大模型实践中的经验设置

预训练:linear warmup + cosine decay(0ηmaxηmin

  • 初期稳定,中期强学习,后期收敛

微调:linear warmup + linear decay,或 constant with warmup

  • 步数少,peak LR 远小于预训练(ηmaxftηmaxpt),防止破坏已有表示

4.4 Batch

4.4.1 Micro Batch、Global Batch 与梯度累积

在大模型训练中,Batch Size 不是一个单一概念。实际训练时至少要区分三个量:

Bmicro,Bglobal,A

其中:

名称含义
Micro Batch Size单张 GPU 一次前向/反向真正处理的样本数
Gradient Accumulation Steps梯度累积步数
Global Batch Size一次参数更新实际等价使用的总样本数
1. Micro Batch Size

Micro Batch Size 是指每张 GPU 在一次 forward/backward 中处理的样本数。

假设一张 GPU 一次只能放下:

Bmicro=4

那么每次前向传播时,这张 GPU 只处理 4 条样本。

2. Gradient Accumulation

显存有限时,不能直接把很大的 Batch 放进 GPU,于是可以用 梯度累积 模拟大 Batch。

普通训练是:

θt+1=θtηθLB(θt)

其中一次 forward/backward 后立刻更新参数。

梯度累积的做法是:将一个大 Batch 分小批次,连续做多次 forward/backward,但暂时不更新参数,而是把梯度加起来。

假设累积 (A) 次,每个 micro batch 的 loss 为:

L(a)(θ)g(a)=θL(a)(θ),a=1,,A

累积后的梯度为:

g=1Aa=1Ag(a)

然后再进行参数更新:

θt+1=θtηg

梯度累积的伪代码可以理解为:

python
optimizer.zero_grad()

for step, batch in enumerate(dataloader):
    loss = model(batch)
    loss = loss / accumulation_steps
    loss.backward()

    if (step + 1) % accumulation_steps == 0:
        optimizer.step()
        scheduler.step()
        optimizer.zero_grad()
3. Global Batch Size

Global Batch Size 是指一次 optimizer step 实际使用的总 batch size。

Bglobal=Bmicro×Ngpu×A
  • GPU 数量为 (N_{\text{gpu}});
  • 每张 GPU 的 micro batch size 为 (B_{\text{micro}});
  • 梯度累积步数为 (A)。

4.4.2 Batch Size 与梯度噪声

假设 N 是全量训练数据总量,在大模型训练中这个数据量一般很大,实际训练不会每一步都用全部数据,而是采样一个 mini-batch(其实就是上面介绍的 Global Batch Size:B=Bmicro×Ngpu×A):

B={i1,i2,,iB}

于是 mini-batch 梯度为:

gB(θ)=1BiBi(θ)

它是完整梯度的随机无偏估计 E[gB(θ)]=L(θ)

如果假设单样本梯度的方差为:

Var(i(θ))=σ2

那么 mini-batch 平均梯度的方差近似为:

Var(gB)=σ2B

即,Batch Size 越大,梯度噪声越小。大 batch 看似梯度更平稳,但有一个隐含代价:

steps per epoch=NB

同样遍历一遍数据,参数更新次数变少了。相当于模型"看"了同样多的数据,但"学"的次数变少了。

  • 小 Batch 靠噪声探索,但梯度不稳定且并行效率低;

  • 大 Batch 靠稳定加速,但参数更新次数少可能导致收敛变慢,因此常配合更大的学习率。


4.4.3 Critical Batch Size

当 batch size 小于某个范围时,增大 batch size 可以明显提升吞吐和稳定性。

但超过某个临界值后,继续增大 batch size 的收益下降:

B>Bcritical

此时:

  • 梯度噪声已经很小,现有的梯度已经在很大程度上代表真实的优化方向,如果再增加每个 Batch 的样本,只是重复已有的信息;
  • 相应的需要调大学习率,过大的 LR scaling 可能导致训练不稳定;
  • 现在的大模型预训练通常是大规模的分布式训练,增大 Batch 会提高显存的需求,分布式训练的通信开销随规模非线性增长,代价很大。

4.5 训练的监控与 Checkpoint 管理

大模型预训练通常持续数天甚至数周,训练成本很高,因此不能只关注最终 loss。完整训练系统还需要解决两个问题:

  1. 如何判断训练是否正常
  2. 如何在训练中断后准确恢复训练状态

4.5.1 监控指标

一个典型训练监控面板可能包含:

指标主要作用
Train Loss判断模型是否持续学习
Validation Loss判断泛化性能
Learning Rate检查学习率调度
Gradient Norm检测梯度爆炸/异常
Tokens/s衡量训练吞吐
Step Time检测训练速度异常
GPU Utilization判断计算资源利用率
GPU Memory监测 OOM 风险
NaN/Inf检测数值稳定性问题

现代大模型训练非常关注 tokens/sec

Throughput=Bglobal×LTstep
  • (B_{\mathrm{global}}):global batch size;

  • (L):sequence length;

  • (T_{\mathrm{step}}):一个 optimizer step 的耗时

如果 loss 正常,但 throughput 突然下降,则可能不是模型问题,而是:

  • DataLoader 速度下降;
  • 网络通信瓶颈;
  • GPU 等待数据;
  • Checkpoint 保存阻塞;
  • 集群节点异常;
  • 数据 Shuffle 或 I/O 出现瓶颈。

4.5.2 Checkpoint 管理

若在模型训练的过程中,我们通过监控指标发现有训练不稳定的情况,如果问题严重,我们就需要暂停训练,排查问题并及时修改。完成之后我们要如何恢复训练呢?这时就需要 Checkpoint。

Checkpoint 可以理解为:

训练过程在某一个时间点的完整快照。

最简单的 checkpoint 只有模型参数 θt,但这通常不足以真正恢复训练。

一个完整 checkpoint 通常至少包含:

Ct={θt,Ot,St,t,Rt,Mt}.
  • (\theta_t):模型参数;
  • (O_t):optimizer state;
  • (S_t):scheduler state;
  • (t):当前 global step;
  • (R_t):随机数状态;
  • (M_t):其他训练元数据。
Model+Optimizer+Scheduler+Step+RNG+Data State

五、分布式训练与工程优化

5.1 分布式训练的核心动因

5.2 基础并行策略

5.3 混合精度训练

5.4 显存优化

5.5 主流训练框架

5.6 大规模训练工程实践

六、评估(验证与反馈闭环)

6.1预训练阶段在线评估

6.2 基准测试

6.3 生成质量评估

6.4 评估的反馈作用