Skip to content
alarik.me
Go back

OPD 深度解析(二):SFT 与策略梯度推导

本文转载自:微信公众号「菜菜Vibe Coding」

本文改编自「菜菜Vibe Coding」公众号系列文章,原文发布于 2026-06-11。这是系列第二篇,讲述 SFT、强化学习策略梯度的基础知识和公式推导,为后续的知识蒸馏打下基础。

本文分为 3 个部分:

  1. 数学基础、信息论基础和 KL 散度系列第一篇
  2. 大模型微调(SFT)、强化学习策略梯度
  3. OPD 的介绍和遇到的问题

Table of contents

Open Table of contents

第四章:语言模型与自回归分解

这一章建立整个推导的数学语言——符号和分解方式。理解自回归分解才能理解后面为什么 sequence-level log-ratio 可以写成 token-level log-ratio 之和,为什么 softmax 归一化是 policy-gradient 推导的基石。

4.1 符号定义

符号含义为什么需要它
xxprompt(输入提示)模型的输入,所有生成都以此为起点
y=(y1,,yT)y = (y_1, \dots, y_T)回答序列模型的输出,是一串 token
ct=(x,y<t)c_t = (x, y_{<{t}})tt 步的 prefix模型在第 tt 步看到的所有已知信息——prompt + 已生成的 t1t-1 个 token
πθ\pi_\theta学生模型(参数为 θ\theta我们要训练的模型,参数 θ\theta 是被优化的对象
qqπT\pi_T教师模型不被训练,提供目标信号
V\mathcal{V}词表(如 50000 个词)所有可能 token 的集合,softmax 在词表上做

4.2 自回归分解——为什么要这样分解?

问题:模型生成整句话的概率 πθ(yx)\pi_\theta(y|x) 是什么?

为什么不能直接算? 一个序列 y=(y1,y2,,yT)y = (y_1, y_2, \dots, y_T) 的概率不是一次性的——模型是”逐步选择”下一个 token 的。每一步的选择依赖于之前已经生成的所有 token。

分解过程

模型生成第 1 个 token:

P(y1x)=πθ(y1c1),c1=xP(y_1|x) = \pi_\theta(y_1|c_1), \quad c_1 = x

模型生成第 2 个 token(需要看第 1 个):

P(y2x,y1)=πθ(y2c2),c2=(x,y1)P(y_2|x, y_1) = \pi_\theta(y_2|c_2), \quad c_2 = (x, y_1)

模型生成第 tt 个 token(需要看前 t1t-1 个):

P(ytx,y<t)=πθ(ytct),ct=(x,y<t)P(y_t|x, y_{<{t}}) = \pi_\theta(y_t|c_t), \quad c_t = (x, y_{<{t}})

链式法则(概率论基本规则):联合概率 = 条件概率之积:

P(y1,y2,,yT)=P(y1)P(y2y1)P(yTy<T)P(y_1, y_2, \dots, y_T) = P(y_1) \cdot P(y_2|y_1) \cdot \dots \cdot P(y_T|y_{<{T}})

应用到语言模型:

πθ(yx)=t=1Tπθ(ytct)\pi_\theta(y|x) = \prod_{t=1}^{T} \pi_\theta(y_t|c_t)

为什么要取 log? 乘法在数学上不方便:

取 log 后乘法变加法,解决这些问题:

logπθ(yx)=t=1Tlogπθ(ytct)\log \pi_\theta(y|x) = \sum_{t=1}^{T} \log \pi_\theta(y_t|c_t)
这个分解是后面所有推导的基础
  • Sequence-level log-ratio logπθ(y)q(y)\log\frac{\pi_\theta(y)}{q(y)} 可以分解成 token-level log-ratio 之和 trt\sum_t r_t
  • Sequence-level score gradient logπθ(y)\nabla\log\pi_\theta(y) 可以分解成 token-level score 之和 tgt\sum_t g_t
  • SFT loss 是对每个 token 的 logprob 求和
  • 如果没有这个分解,我们无法把高方差的 sequence-level 目标降到 token-level

Teacher 同理

logq(yx)=t=1Tlogq(ytct)\log q(y|x) = \sum_{t=1}^{T} \log q(y_t|c_t)

两者的差——log-ratio

logπθ(yx)q(yx)=t=1T[logπθ(ytct)logq(ytct)]=t=1Trt\log\frac{\pi_\theta(y|x)}{q(y|x)} = \sum_{t=1}^{T} \left[\log\pi_\theta(y_t|c_t) - \log q(y_t|c_t)\right] = \sum_{t=1}^{T} r_t

这就是为什么 log-ratio 可以按 token 分解——自回归分解让我们能把整条序列的 “student vs teacher” 差异,拆成每个位置上的差异之和。

4.3 Softmax 保证归一——为什么这是整个推导的基石?

每个时间步,模型对词表中每个 token 输出一个实数分数 zvz_v(称为 logit),然后通过 softmax 转成概率:

πθ(vct)=ezv(ct)uVezu(ct)\pi_\theta(v|c_t) = \frac{e^{z_v(c_t)}}{\sum_{u \in \mathcal{V}} e^{z_u(c_t)}}

为什么需要 softmax? 模型的输出 zvz_v 是任意实数(可正可负,没有上界),不能直接当概率用。Softmax 的功能:

证明归一性

vVπθ(vct)=vVezvuezu=vezvuezu=1\sum_{v \in \mathcal{V}} \pi_\theta(v|c_t) = \sum_{v \in \mathcal{V}} \frac{e^{z_v}}{\sum_{u} e^{z_u}} = \frac{\sum_{v} e^{z_v}}{\sum_{u} e^{z_u}} = 1

分母和分子是同一个东西(对所有词表元素求和),所以比值 = 1。

从单步归一到联合分布归一

每个时间步的概率都归一(和为 1),所以联合概率也归一:

yπθ(yx)=y1y2yTtπθ(ytct)=1\sum_y \pi_\theta(y|x) = \sum_{y_1} \sum_{y_2} \dots \sum_{y_T} \prod_t \pi_\theta(y_t|c_t) = 1
为什么归一性是 policy-gradient 推导的基石?

在第 10 章的推导中,关键步骤是:

Eyπθ[logπθ(y)]=yπθ(y)logπθ(y)=(yπθ(y))=1=0\mathbb{E}_{y \sim \pi_\theta}[\nabla\log\pi_\theta(y)] = \sum_y \pi_\theta(y) \nabla\log\pi_\theta(y) = \nabla\left(\sum_y \pi_\theta(y)\right) = \nabla 1 = 0

如果 yπθ(y)1\sum_y \pi_\theta(y) \neq 1,这一步就不成立,整个 policy-gradient 推导就崩了。

Softmax 是归一性的保障——它让每步概率和为 1,进而让联合概率和为 1,进而让梯度推导的第二项消零成立。

PyTorch 中 softmax 的实现

import torch
import torch.nn.functional as F
# 假设模型输出的logits,shape [B, L, V]
logits = torch.randn(2, 10, 50000) # batch=2, seq_len=10, vocab=50000
# softmax转概率
probs = F.softmax(logits, dim=-1) # shape [2, 10, 50000]
# 验证归一性
print(probs.sum(dim=-1)) # 每个位置的概率和,应该全是1.0
# log_softmax(更稳定,避免数值问题)
log_probs = F.log_softmax(logits, dim=-1)
# 验证: exp(log_probs) ≈ probs
print(torch.exp(log_probs) - probs) # 应该接近0

第五章:SFT(监督微调)

SFT 是大模型训练的起点。理解 SFT 的数学本质(它到底在优化什么),才能理解后面 KD、OPD 为什么需要改进、改进了什么。

5.1 什么是 SFT?为什么要做 SFT?

大模型训练的三个阶段

为什么需要 SFT? 预训练后的模型能”说话”,但不会”对话”——它可能续写一段新闻,但不会回答你的问题。SFT 教它:

类比:预训练像学会了中文语法,SFT 像学会了”怎么写商务邮件”——语法对了,但格式和内容风格需要专门训练。

5.2 SFT Loss 的完整推导——从目标到公式

训练数据(x,ydata)(x, y^{data}),其中 xx 是 prompt,ydata=(y1data,,yTdata)y^{data} = (y_1^{data}, \dots, y_T^{data}) 是人类写好的标准回答。

目标:让模型尽可能生成和标准答案一样的回答。

第一步:写出目标

模型对标准答案的生成概率(用自回归分解):

πθ(ydatax)=t=1Tπθ(ytdatactdata)\pi_\theta(y^{data}|x) = \prod_{t=1}^{T} \pi_\theta(y_t^{data}|c_t^{data})

我们希望这个概率尽可能大

为什么最大化概率? 如果模型认为标准答案的概率很高,说明模型”认同”这个答案——它倾向于生成同样的回答。这正是我们想要的。

第二步:取 log,乘法变加法

logπθ(ydatax)=t=1Tlogπθ(ytdatactdata)\log \pi_\theta(y^{data}|x) = \sum_{t=1}^{T} \log \pi_\theta(y_t^{data}|c_t^{data})

取 log 的原因(第四章已讲):乘法数值不稳定、求导复杂,加法更方便。

第三步:从最大化变成最小化

机器学习习惯写最小化问题(因为优化器是做梯度下降的)。最大化 logπ\log \pi 等价于最小化 logπ-\log \pi

LSFT=t=1Tlogπθ(ytdatactdata)\mathcal{L}_{\mathrm{SFT}} = -\sum_{t=1}^{T} \log \pi_\theta(y_t^{data}|c_t^{data})

这就是 SFT Loss(也叫 NLL Loss,Negative Log-Likelihood)。

第四步:逐项理解

每个时间步 tt 的 loss 是:

LSFT,t=logπθ(ytdatactdata)\mathcal{L}_{\mathrm{SFT},t} = -\log \pi_\theta(y_t^{data}|c_t^{data})

直觉:就像考试——每道题都有标准答案,你答对了(给了高概率)就不扣分,答错了(给了低概率)就扣分。

5.3 SFT 与信息论的联系——SFT 到底在优化什么?

关键问题:SFT loss 看起来只是”让概率变大”,但它的数学本质是什么?

推导:每个时间步,标准答案的真实分布是 one-hot——只有 ytdatay_t^{data} 这个 token 概率为 1,其他所有 token 概率为 0:

Pdata(vct)={1if v=ytdata0otherwiseP_{\mathrm{data}}(v|c_t) = \begin{cases} 1 & \text{if } v = y_t^{data} \\ 0 & \text{otherwise} \end{cases}

SFT loss 就是这个 one-hot 分布和 student 分布的交叉熵:

LSFT,t=H(Pdata,PS)=vPdata(v)logPS(v)=logPS(ytdatact)\mathcal{L}_{\mathrm{SFT},t} = H(P_{\mathrm{data}}, P_S) = -\sum_v P_{\mathrm{data}}(v) \log P_S(v) = -\log P_S(y_t^{data}|c_t)

为什么交叉熵只剩一项? 因为 PdataP_{\mathrm{data}} 是 one-hot——只有 v=ytdatav = y_t^{data}Pdata(v)=1P_{\mathrm{data}}(v) = 1,其他都是 0。0 乘任何数都是 0,所以求和只剩一项。

再用第二章的公式 H(P,Q)=H(P)+DKL(PQ)H(P,Q) = H(P) + D_{\mathrm{KL}}(P \| Q)

LSFT,t=H(Pdata)+DKL(PdataPS)\mathcal{L}_{\mathrm{SFT},t} = H(P_{\mathrm{data}}) + D_{\mathrm{KL}}(P_{\mathrm{data}} \| P_S)

One-hot 分布的熵 H(Pdata)=1log10log0=0H(P_{\mathrm{data}}) = -1 \cdot \log 1 - 0 \cdot \log 0 = 0(完全确定,没有不确定性):

LSFT,t=0+DKL(PdataPS)=DKL(PdataPS)\mathcal{L}_{\mathrm{SFT},t} = 0 + D_{\mathrm{KL}}(P_{\mathrm{data}} \| P_S) = D_{\mathrm{KL}}(P_{\mathrm{data}} \| P_S)

结论

SFT=在数据 prefix 上,最小化 student 与 one-hot 分布的 forward KL\text{SFT} = \text{在数据 prefix 上,最小化 student 与 one-hot 分布的 forward KL}
这个结论的意义
  • SFT 的目标是让 student 完全模仿标准答案(one-hot = 100% 确定只有一个正确答案)
  • KL 方向是 forward(DKL(PdataPS)D_{\mathrm{KL}}(P_{\mathrm{data}} \| P_S)),因为标准答案说”只有这个 token 是对的”,student 不能给其他 token 太高概率
  • 这是一种硬监督——每个位置只有一个正确答案,没有”软信号”

5.4 SFT 的根本局限——Exposure Bias

问题:SFT 在数据的前缀上训练:

ct=(x,y<tdata)c_t = (x, y_{<{t}}^{data})

但推理时,模型在自己生成的前缀上运行:

ct=(x,y<tstudent)c_t = (x, y_{<{t}}^{student})

一旦 student 在某个位置生成了和标准答案不同的 token,后续的前缀就进入了训练时从未见过的区域。模型在这个陌生区域没有受过训练,容易继续犯错——错误连锁放大。

这就是 Exposure Bias——训练看到的”世界”和推理看到的”世界”不一致。

SFT 像只在阳光明媚的公路上练车。考试时遇到雨天、夜路,你就慌了——因为你从未在这些条件下练过。

5.5 SFT 的 PyTorch 实现

import torch
import torch.nn.functional as F
def sft_loss(student_logits, labels, ignore_index=-100):
"""计算SFT loss(交叉熵/NLL)
Args:
student_logits: student模型输出的logits, shape [B, L, V]
labels: 标准答案的token id, shape [B, L]
ignore_index: 忽略的label值(如padding token)
Returns:
平均loss,scalar
"""
# shift: logits的第t步对应labels的第t+1步(因为logits预测下一个token)
shift_logits = student_logits[:, :-1, :].contiguous() # [B, L-1, V]
shift_labels = labels[:, 1:].contiguous() # [B, L-1]
# 交叉熵 = NLL = -log P_S(y_t^data | c_t^data)
loss = F.cross_entropy(
shift_logits.view(-1, shift_logits.size(-1)), # [B*(L-1), V]
shift_labels.view(-1), # [B*(L-1)]
ignore_index=ignore_index
)
return loss
# 使用示例
logits = torch.randn(4, 128, 50000) # batch=4, seq_len=128, vocab=50000
labels = torch.randint(0, 50000, (4, 128)) # 标准答案token ids
loss = sft_loss(logits, labels)
print(f"SFT loss: {loss.item():.4f}")

第六章:强化学习基础——策略梯度

本章是整个 OPD 推导的”引擎”——理解策略梯度,才能理解第 10 章 OPD 的梯度推导为什么是那个形式、为什么第二项能消零、为什么可以加 baseline。

6.1 什么是 RL?为什么要引入 RL?

SFT 的局限:SFT 只能告诉模型”标准答案是什么”,但不能告诉模型”你自己的回答好不好”。

类比

为什么需要 RL? 因为很多场景没有标准答案:

SFT 只能学一种固定答案,RL 让模型自己探索并发现更好的策略。

核心要素

6.2 策略梯度推导——为什么要这样推导?

目标:最大化期望奖励

J(θ)=Eyπθ[R(y)]=yπθ(y)R(y)J(\theta) = \mathbb{E}_{y \sim \pi_\theta}[R(y)] = \sum_y \pi_\theta(y) \cdot R(y)

问题πθ(y)\pi_\theta(y) 是一个概率值(经过 softmax),我们不能直接对它求导——因为概率值和参数 θ\theta 的关系是通过 softmax 的复杂映射,直接求导极其困难。

策略梯度的核心思路:不直接对 πθ(y)\pi_\theta(y) 求导,而是通过 log-derivative trickπθ(y)\nabla\pi_\theta(y) 转化成更容易处理的形式。

推导过程

第一步:写出目标函数

J(θ)=yπθ(y)R(y)J(\theta) = \sum_y \pi_\theta(y) \cdot R(y)

第二步:对 θ\theta 求导

J=yπθ(y)R(y)\nabla J = \sum_y \nabla\pi_\theta(y) \cdot R(y)

为什么只对 πθ\pi_\theta 求导? 因为 R(y)R(y) 是 reward,它不依赖 θ\theta(reward 是外部给的,不是模型参数决定的)。所以 R(y)R(y)θ\theta 的导数为 0。

第三步:用 log-derivative trick

πθ(y)=πθ(y)logπθ(y)\nabla\pi_\theta(y) = \pi_\theta(y) \cdot \nabla\log\pi_\theta(y)

为什么这个 trick 成立? 因为 πθ(y)=elogπθ(y)\pi_\theta(y) = e^{\log\pi_\theta(y)},对 ef(θ)e^{f(\theta)} 求导 = effe^{f} \cdot \nabla f,所以 πθ(y)=πθ(y)logπθ(y)\nabla\pi_\theta(y) = \pi_\theta(y) \cdot \nabla\log\pi_\theta(y)

为什么需要这个 trick? 直接对 softmax 输出的概率 πθ(y)=ezθ(y)uezθ(u)\pi_\theta(y) = \frac{e^{z_\theta(y)}}{\sum_u e^{z_\theta(u)}} 求导需要用商的求导法则,还要处理分母(所有词表的指数之和),极其复杂。Log-derivative trick 避开了这个困难——它把 πθ\nabla\pi_\theta 转化成 πθlogπθ\pi_\theta \cdot \nabla\log\pi_\theta,后者只需要对 log 概率求导,简单得多。

第四步:代入

J=yπθ(y)logπθ(y)R(y)=Eyπθ[R(y)logπθ(y)]\nabla J = \sum_y \pi_\theta(y) \cdot \nabla\log\pi_\theta(y) \cdot R(y) = \mathbb{E}_{y \sim \pi_\theta}[R(y) \cdot \nabla\log\pi_\theta(y)]

结果的含义

J=Eyπθ[R(y)logπθ(y)]\nabla J = \mathbb{E}_{y \sim \pi_\theta}[R(y) \cdot \nabla\log\pi_\theta(y)]
组成部分含义
logπθ(y)\nabla\log\pi_\theta(y)Score function:参数怎么调才能让模型更多生成这条轨迹
R(y)R(y)Reward:这条轨迹好不好
乘积好轨迹 → 提高概率;坏轨迹 → 降低概率

就像生活中的学习:做了一件事,得了高分(reward 高),以后多做(提高概率);得了低分,以后少做(降低概率)。

为什么这是采样估计? 因为 Eyπθ\mathbb{E}_{y \sim \pi_\theta} 的意思是”从模型分布中采样 yy“,我们让模型生成几条轨迹,算每条的 R(y)logπθ(y)R(y) \cdot \nabla\log\pi_\theta(y),取平均。不需要遍历所有可能的 yy

6.3 加 baseline——为什么可以减?减了有什么好处?

原始策略梯度方差很高。可以减去一个 baseline bb

J=Eyπθ[(R(y)b)logπθ(y)]\nabla J = \mathbb{E}_{y \sim \pi_\theta}[(R(y) - b) \cdot \nabla\log\pi_\theta(y)]

为什么可以减 baseline? 逐步证明:

Eyπθ[blogπθ(y)]\mathbb{E}_{y \sim \pi_\theta}[b \cdot \nabla\log\pi_\theta(y)]

bb 提出来(常数):

=bEyπθ[logπθ(y)]= b \cdot \mathbb{E}_{y \sim \pi_\theta}[\nabla\log\pi_\theta(y)]

展开期望:

=byπθ(y)logπθ(y)= b \cdot \sum_y \pi_\theta(y) \nabla\log\pi_\theta(y)

用 log-derivative trick 反向写 πθ(y)logπθ(y)=πθ(y)\pi_\theta(y) \nabla\log\pi_\theta(y) = \nabla\pi_\theta(y)

=byπθ(y)= b \cdot \sum_y \nabla\pi_\theta(y)

求和与梯度交换:

=b(yπθ(y))= b \cdot \nabla\left(\sum_y \pi_\theta(y)\right)

归一性 yπθ(y)=1\sum_y \pi_\theta(y) = 1

=b1=b0=0= b \cdot \nabla 1 = b \cdot 0 = 0
又用到了归一性!

yπθ(y)=1\sum_y \pi_\theta(y) = 11=0\nabla 1 = 0。这和第 10 章 OPD 推导中”第二项消零”用的是同一个性质。

减 baseline 有什么好处?

常见 baseline 选择:reward 的均值 b=E[R]b = \mathbb{E}[R],或用另一个模型(critic/reference model)的 logprob 作为 baseline。

PyTorch 实现

import torch
def policy_gradient_loss(logprobs, rewards):
# 策略梯度: -reward * logprob
# 最大化期望奖励 = 最小化 -reward * logprob
loss = -rewards * logprobs
return loss
def policy_gradient_with_baseline(logprobs, rewards, baseline):
# advantage = reward - baseline
advantages = rewards - baseline
loss = -advantages * logprobs
return loss
# 示例: PPO中的baseline就是reference model的logprob
# advantage = reward + ref_logprob - logprob
# 这和OPD中的 sampled-token advantage: A_t = logq - logpi 是同一个思路

6.4 策略梯度与 OPD 的联系

策略梯度是 OPD 推导的”模板”

在第 10 章,OPD 的目标函数 J(θ)=Eyπθ[logπθ(y)logq(y)]J(\theta) = \mathbb{E}_{y \sim \pi_\theta}[\log\pi_\theta(y) - \log q(y)] 中:

  • “reward” R(y)R(y) 被替换成了 log-ratio logπθ(y)logq(y)\log\pi_\theta(y) - \log q(y)
  • 推导过程完全一样:乘积求导 → log-derivative trick → 归一性消零
  • 最终形式:J=E[(logπθlogq)logπθ]\nabla J = \mathbb{E}[(\log\pi_\theta - \log q) \cdot \nabla\log\pi_\theta]

唯一的区别:reward 从外部给的 R(y)R(y) 变成了内部计算的 log-ratio。这意味着 OPD 不需要外部 reward 函数,teacher 的 logprob 就是 reward。


Share this post:

Previous Post
OPD 深度解析(一):概率、熵与 KL 散度
Next Post
数学公式与代码高亮测试