前言

在当今的大模型(LLM)时代,强化学习(Reinforcement Learning, RL)已经成为了让 AI 从“只会续写”进化到“听懂指令、符合人类价值观”的关键技术。本文全面总结了强化学习的基础理论推导、RLHF 的三阶段训练流程、PPO 算法详解,以及前沿的 DPO、GRPO 等算法对比,并深入解析了显存计算、LoRA 参数分配、RoPE 位置编码等工程实践细节。

这是一份从理论到实践的完整技术笔记,旨在帮助开发者深入理解大模型背后的核心数学与算法原理。


一、 强化学习基础与公式推导

强化学习的核心是从定义目标建立价值,再到迭代优化的严密逻辑过程。

1.1 马尔可夫决策过程 (MDP) —— 一切推导的起点

强化学习问题被建模为一个 MDP 五元组 ( S , A , P , R , γ ) (S, A, P, R, \gamma) (S,A,P,R,γ)

  • S S S (State): 状态空间(例如棋盘局面)。
  • A A A (Action): 动作空间(例如落子位置)。
  • P P P (Transition Probability): 状态转移概率 P ( s ′ ∣ s , a ) P(s'|s, a) P(ss,a),即在 s s s a a a 后跳到 s ′ s' s 的概率。
  • R R R (Reward): 即时奖励函数 R ( s , a ) R(s, a) R(s,a)
  • γ \gamma γ (Discount Factor): 折扣因子, γ ∈ [ 0 , 1 ] \gamma \in [0, 1] γ[0,1],用于衡量未来的价值折现。

核心目标:最大化累积回报 (Return) G t G_t Gt
G t = R t + 1 + γ R t + 2 + γ 2 R t + 3 + ⋯ = ∑ k = 0 ∞ γ k R t + k + 1 G_t = R_{t+1} + \gamma R_{t+2} + \gamma^2 R_{t+3} + \dots = \sum_{k=0}^{\infty} \gamma^k R_{t+k+1} Gt=Rt+1+γRt+2+γ2Rt+3+=k=0γkRt+k+1
推导技巧:利用递归性质展开:
G t = R t + 1 + γ ( R t + 2 + γ R t + 3 + …   ) = R t + 1 + γ G t + 1 G_t = R_{t+1} + \gamma (R_{t+2} + \gamma R_{t+3} + \dots) = R_{t+1} + \gamma G_{t+1} Gt=Rt+1+γ(Rt+2+γRt+3+)=Rt+1+γGt+1

1.2 贝尔曼方程 (Bellman Equation) —— 核心递归公式

为了评估当前状态“好不好”,引入价值函数。

  • 状态价值函数 V π ( s ) V_\pi(s) Vπ(s):遵循策略 π \pi π 能获得的期望回报。
    V π ( s ) = E π [ G t ∣ S t = s ] V_\pi(s) = \mathbb{E}_\pi [G_t | S_t = s] Vπ(s)=Eπ[GtSt=s]
  • 动作价值函数 Q π ( s , a ) Q_\pi(s, a) Qπ(s,a):在状态 s s s 执行动作 a a a,随后遵循策略 π \pi π 的期望回报。
    Q π ( s , a ) = E π [ G t ∣ S t = s , A t = a ] Q_\pi(s, a) = \mathbb{E}_\pi [G_t | S_t = s, A_t = a] Qπ(s,a)=Eπ[GtSt=s,At=a]

贝尔曼期望方程推导
G t = R t + 1 + γ G t + 1 G_t = R_{t+1} + \gamma G_{t+1} Gt=Rt+1+γGt+1 代入 V π ( s ) V_\pi(s) Vπ(s) 的定义,利用全期望公式展开:
V π ( s ) = E π [ R t + 1 + γ V π ( S t + 1 ) ∣ S t = s ] = ∑ a ∈ A π ( a ∣ s ) ∑ s ′ , r p ( s ′ , r ∣ s , a ) [ r + γ V π ( s ′ ) ] \begin{aligned} V_\pi(s) &= \mathbb{E}_\pi [R_{t+1} + \gamma V_\pi(S_{t+1}) | S_t = s] \\ &= \sum_{a \in A} \pi(a|s) \sum_{s', r} p(s', r | s, a) \left[ r + \gamma V_\pi(s') \right] \end{aligned} Vπ(s)=Eπ[Rt+1+γVπ(St+1)St=s]=aAπ(as)s,rp(s,rs,a)[r+γVπ(s)]

  • 物理意义:当前价值 = (这一步动作的平均即时奖励) + (未来价值的折现)。
  • r r r 的含义:即时奖励(Immediate Reward),指执行动作后环境立即反馈的分数(如赢了+1,输了-1)。

1.3 强化学习的三个阶段 (经典分类)

这是求解 MDP 的三种渐进方法。

阶段 方法 核心特点 是否需要模型 适用场景
1 动态规划 (DP) 上帝视角,解方程 Yes (Model-Based) 已知环境模型,如迷宫地图
2 蒙特卡洛 (MC) 经验主义,玩完再算分 No (Model-Free) 只要能采样即可,如AlphaGo自我对弈
3 时序差分 (TD) 步步为营,边走边算 No (Model-Free) 结合了DP和MC,是Q-Learning核心

二、 大模型 RLHF (三阶段训练流程)

ChatGPT 等大模型的训练通常分为三个标准阶段,被称为“RLHF 之旅”。

2.1 第一阶段:有监督微调 (SFT) —— “学会说话”

  • 目标:让模型从“续写模式”切换到“问答模式”,学会指令遵循。
  • 数据:高质量的 (Prompt, Answer) 对。
  • 损失函数:标准的交叉熵损失 (Cross-Entropy Loss)。
    L S F T ( θ ) = − ∑ t = 1 T log ⁡ P ( x t ∣ x < t ; θ ) L_{SFT}(\theta) = - \sum_{t=1}^{T} \log P(x_t | x_{<t}; \theta) LSFT(θ)=t=1TlogP(xtx<t;θ)
  • 结果:得到 SFT 模型 π S F T \pi^{SFT} πSFT

2.2 第二阶段:奖励模型 (Reward Modeling, RM) —— “学会判卷”

  • 目标:训练一个替身(AI 判卷老师)来模拟人类偏好,解决人工反馈太慢的问题。
  • 数据(Prompt, Answer_Win, Answer_Lose) 三元组。
  • 模型:将 SFT 模型去掉最后的 Linear Head,改为输出一个标量分数。
  • 损失函数:Pairwise Ranking Loss。
    L R M ( ϕ ) = − E ( x , y w , y l ) ∼ D [ log ⁡ σ ( r ϕ ( x , y w ) − r ϕ ( x , y l ) ) ] L_{RM}(\phi) = - \mathbb{E}_{(x, y_w, y_l) \sim D} \left[ \log \sigma \left( r_\phi(x, y_w) - r_\phi(x, y_l) \right) \right] LRM(ϕ)=E(x,yw,yl)D[logσ(rϕ(x,yw)rϕ(x,yl))]
    • 直观理解:最大化好答案得分与坏答案得分的差值。

2.3 第三阶段:强化学习 (PPO) —— “刷分进化”

  • 目标:利用 RM 的打分作为奖励信号,优化生成策略 π θ \pi_\theta πθ,使其生成高分回答。
  • 区别于通用 RL
    1. 约束 (Constraint):必须加 KL 散度惩罚
      R ( x , y ) = r ϕ ( x , y ) − β ⋅ KL ( π θ ( y ∣ x ) ∣ ∣ π S F T ( y ∣ x ) ) R(x, y) = r_\phi(x, y) - \beta \cdot \text{KL}(\pi_\theta(y|x) || \pi^{SFT}(y|x)) R(x,y)=rϕ(x,y)βKL(πθ(yx)∣∣πSFT(yx))
      • 原因:防止模型为了讨好 RM 而输出乱码(Reward Hacking),强迫其保持与 SFT 模型的相似性。
    2. 环境:非随机环境,状态转移是确定的文本拼接。
    3. 动作空间:巨大(词表大小,约 50k-100k),远超一般 RL 任务。

三、 PPO 算法详解

PPO (Proximal Policy Optimization) 是 RLHF 的默认算法,由 OpenAI 提出,核心在于“稳”。

3.1 核心痛点

传统策略梯度(Policy Gradient)步长极其难调:

  • 步长太小 → \to 训练极慢。
  • 步长太大 → \to 策略更新过猛,导致 Policy Collapse(策略崩塌),一旦采出垃圾数据,模型就再也救不回来了。

3.2 核心机制:Clipping (剪切)

PPO 通过引入一个“信任域”(Trust Region),限制每次更新的幅度。

目标函数
L C L I P ( θ ) = E [ min ⁡ ( r t ( θ ) A t , clip ( r t ( θ ) , 1 − ϵ , 1 + ϵ ) A t ) ] L^{CLIP}(\theta) = \mathbb{E} \left[ \min \left( r_t(\theta) A_t, \text{clip}(r_t(\theta), 1-\epsilon, 1+\epsilon) A_t \right) \right] LCLIP(θ)=E[min(rt(θ)At,clip(rt(θ),1ϵ,1+ϵ)At)]

  • r t ( θ ) = π θ ( a t ∣ s t ) π o l d ( a t ∣ s t ) r_t(\theta) = \frac{\pi_\theta(a_t|s_t)}{\pi_{old}(a_t|s_t)} rt(θ)=πold(atst)πθ(atst):概率比率 (Ratio)。衡量新策略比旧策略更倾向于做在这个动作的程度。
  • A t A_t At (Advantage):优势函数。衡量当前动作比平均水平好多少。
  • ϵ \epsilon ϵ:超参数,通常取 0.2。
  • 机制解析
    • 如果动作是好的 ( A t > 0 A_t > 0 At>0),我们要增加概率。但如果 r t > 1.2 r_t > 1.2 rt>1.2(增加了20%以上),则 截断梯度,不再给予奖励。
    • 一句话总结鼓励优化,但严禁贪婪。小步快跑,稳字当头。

四、 进阶算法:DPO 与 GRPO

在 PPO 之后,为了降本增效,出现了两个颠覆性的变体。

4.1 DPO (Direct Preference Optimization) —— “去掉了 Reward Model”

  • 核心思想:利用数学对偶性,证明了奖励函数 r ( x , y ) r(x,y) r(x,y) 可以完全由最优策略 π ∗ \pi^* π 和参考策略 π r e f \pi_{ref} πref 表示。
  • 做法:直接在偏好数据上优化策略,无需训练 RM,也无需 PPO 的采样过程。
  • 损失函数
    L D P O = − log ⁡ σ ( β log ⁡ π θ ( y w ∣ x ) π r e f ( y w ∣ x ) − β log ⁡ π θ ( y l ∣ x ) π r e f ( y l ∣ x ) ) L_{DPO} = - \log \sigma \left( \beta \log \frac{\pi_\theta(y_w|x)}{\pi_{ref}(y_w|x)} - \beta \log \frac{\pi_\theta(y_l|x)}{\pi_{ref}(y_l|x)} \right) LDPO=logσ(βlogπref(ywx)πθ(ywx)βlogπref(ylx)πθ(ylx))
  • 优势:极简、稳定、显存占用小(类似监督学习)。
  • 劣势:上限可能不如 PPO,在复杂逻辑推理任务上表现略逊。

4.2 GRPO (Group Relative Policy Optimization) —— “去掉了 Critic Model”

  • 出处DeepSeek (DeepSeek-Math / DeepSeek-R1)。
  • 背景:PPO 需要一个与 Actor 同等大小的 Critic 模型(价值网络),显存消耗巨大。
  • 核心机制
    1. Group Sampling:对同一个 Prompt 采样一组回答 { y 1 , y 2 , … , y G } \{y_1, y_2, \dots, y_G\} {y1,y2,,yG}
    2. Baseline 计算:利用这组回答的平均分作为基准,代替 Critic 的预测值。
    3. 优势估计
      A i = r i − mean ( { r 1 … r G } ) std ( { r 1 … r G } ) A_i = \frac{r_i - \text{mean}(\{r_1 \dots r_G\})}{\text{std}(\{r_1 \dots r_G\})} Ai=std({r1rG})rimean({r1rG})
  • 优势极度节省显存。不需要 Critic 模型,训练成本减半,使得训练超大参数模型(如 DeepSeek-R1-671B)成为可能。
算法 Critic模型 Reward模型 适用场景
PPO 需要 ✅ 需要 资源充足,追求极致效果
DPO 不需要 不需要 资源受限,快速对齐
GRPO 不需要 ✅ 需要(或规则) 超大模型,推理任务(R1)

五、 工程实践:显存、LoRA 与 RoPE

5.1 显存计算公式 —— “谁吃掉了你的显存?”

训练大模型时,显存杀手往往不是模型参数本身,而是优化器状态

组件 精度 占用公式 7B模型示例 (FP16)
模型权重 FP16 Φ × 2 \Phi \times 2 Φ×2 Bytes 14 GB
梯度 FP16 Φ × 2 \Phi \times 2 Φ×2 Bytes 14 GB
优化器 (AdamW) FP32 (Mixed) Φ × 12 \Phi \times 12 Φ×12 Bytes 84 GB (最大头!)
KV Cache FP16 2 B L H S × 2 2 B L H S \times 2 2BLHS×2 推理时占用,随长度线性增长
  • 优化器状态拆解 (12 Bytes/Param)
    1. Master Weights (4B):FP32 的权重备份。解决 FP16 精度不足导致的下溢出 (Underflow) 问题。
    2. Momentum (4B):动量(一阶矩),模拟物理惯性,加速收敛。
    3. Variance (4B):方差(二阶矩),用于自适应调整学习率(RMSProp原理)。

5.2 LoRA 参数与优化器分配

  • 原理:冻结基座模型权重 W W W,只训练旁路的低秩矩阵 A , B A, B A,B ( W ′ = W + B A W' = W + BA W=W+BA)。
  • 参数量计算
    Params = r × ( d i n + d o u t ) × N l a y e r s \text{Params} = r \times (d_{in} + d_{out}) \times N_{layers} Params=r×(din+dout)×Nlayers
    • 对于 7B 模型 (r=8),LoRA 参数量仅约 4M - 20M
  • 显存优势
    • 全量微调优化器显存:84 GB。
    • LoRA 微调优化器显存:< 0.1 GB
    • 结论:LoRA 让优化器只维护极少量的参数状态,从而在消费级显卡(如 RTX 3090/4090)上实现大模型微调。

5.3 RoPE 位置编码 (Rotary Positional Embedding)

  • 核心思想:通过将词向量在二维复平面上进行旋转来表示位置信息。
  • 数学性质
    Score ( q , k ) = ( R m q ) T ( R n k ) = q T R n − m k \text{Score}(q, k) = (R_m q)^T (R_n k) = q^T R_{n-m} k Score(q,k)=(Rmq)T(Rnk)=qTRnmk
    • 向量点积只与相对距离 ( n − m ) (n-m) (nm) 有关,天然具备相对位置感知能力。
  • 优势
    1. 外推性强:通过线性内插 (Linear Interpolation),可以将位置索引压缩(如把 4096 压到 2048 范围内),让模型处理训练时没见过的超长文本(如 LLaMA 的 Long Context 扩展)。
    2. 实现简单:直接乘旋转矩阵,无需额外的参数学习。

六、 其他关键概念解析

6.1 Function Call vs MCP

  • Function Call (函数调用)
    • 定义:模型的一种能力(类似“手”),使其能输出特定 JSON 指令来调用外部工具。
    • 痛点:每个工具都需要单独写适配代码,维护成本高。
  • MCP (Model Context Protocol)
    • 定义:一种标准协议(类似“USB接口”)。
    • 作用:标准化连接各种数据源(如数据库、Slack、GitHub)。只需开发一次 MCP Server,所有支持 MCP 的模型(如 Claude Desktop, Cursor)都能即插即用。

6.2 RMSNorm (Root Mean Square Normalization)

  • 原理:LayerNorm 的简化版。
  • 区别:去掉了 LayerNorm 中“减均值 (Re-centering)”的操作,只用均方根 (RMS) 进行缩放。
    y = x RMS ( x ) ⋅ γ y = \frac{x}{\text{RMS}(x)} \cdot \gamma y=RMS(x)xγ
  • 优势:计算更快,效果相当,是 LLaMA、Gemma、DeepSeek 等现代大模型的标配。

6.3 RAG Rerank (重排序)

  • Retrieval (粗排)
    • 使用双塔模型 (Bi-Encoder),分别计算 Query 和 Doc 向量。
    • 特点:速度快,但精度低(无法捕捉细粒度交互)。
  • Rerank (精排)
    • 使用交叉编码器 (Cross-Encoder),将 Query 和 Doc 拼接输入 BERT。
    • 特点:速度慢,但精度极高。
    • 作用:解决“关键词匹配但语义不符”的问题(如“不含糖” vs “含糖”),是提升 RAG 系统准确率的关键防线。

七、 高性能计算优化 (Advanced Optimization)

为了支撑大模型(LLM)的超长上下文(Long Context)推理和高效训练,除了算法层面的优化,底层的计算优化至关重要。以下是两项核心技术。

7.1 FlashAttention 加速原理

FlashAttention 是目前大模型训练和推理的标配加速库(如 LLaMA, DeepSeek, Falcon 均默认使用)。

  • 痛点:标准 Attention 的计算复杂度是 O ( N 2 ) O(N^2) O(N2),且显存读写(HBM IO)是主要瓶颈。
  • 核心原理
    1. Tiling (分块计算):将 Q , K , V Q, K, V Q,K,V 矩阵切成小块,放入 GPU 的高速缓存 (SRAM) 中进行计算,减少对显存 (HBM) 的读写次数。
    2. Recomputation (重计算):在前向传播时不存储巨大的 Attention Score 矩阵 ( N × N N \times N N×N),反向传播时根据 Q , K Q, K Q,K 重新计算一遍。虽然多算了一次,但省下的显存读写时间远大于计算时间。
  • 效果
    • 速度:比标准 Attention 快 2-4 倍
    • 显存:显存占用从 O ( N 2 ) O(N^2) O(N2) 降为 线性 O ( N ) O(N) O(N),使得单卡训练 32k/128k 上下文成为可能。

7.2 KV Cache 显存优化

KV Cache 是推理阶段(Inference)显存占用的最大头,甚至可能超过模型权重本身。

  • 痛点:自回归生成(Autoregressive Generation)每生成一个新 token,都要重新计算前面所有 token 的 Key 和 Value,极其浪费计算资源。
  • 原理:将前面计算过的 K , V K, V K,V 矩阵缓存下来,下次只计算新 token 的 Q Q Q 与旧 K , V K, V K,V 的交互。
  • 显存占用公式
    Size K V = 2 × B × L × H × S × Precision \text{Size}_{KV} = 2 \times B \times L \times H \times S \times \text{Precision} SizeKV=2×B×L×H×S×Precision
    • 对于 7B 模型 (FP16),128k 上下文的 KV Cache 可能高达 60GB 以上。
  • 主流优化技术
    1. PagedAttention (vLLM):借鉴操作系统的虚拟内存管理,将 KV Cache 分页存储(非连续显存),彻底解决显存碎片化问题,吞吐量提升 2-4 倍。
    2. MQA / GQA (Grouped Query Attention)
      • MHA (Multi-Head): 每个 Query 头对应独立的 Key/Value 头 (1:1)。
      • MQA (Multi-Query): 所有 Query 头共享同一个 Key/Value 头 (N:1)。
      • GQA (Grouped-Query): 折中方案,每组 Query 共享一个 Key/Value 头 (N:M)。
      • 效果:大幅减少 KV Cache 的大小(LLaMA-2/3, DeepSeek 均采用 GQA)。

总结
本文系统梳理了从 MDP 基础理论DeepSeek-R1 工程实践 的完整知识图谱。理解这些核心算法(PPO/GRPO/DPO)与工程细节(显存/LoRA/RoPE),是掌握大模型微调与对齐技术的必经之路。

Logo

有“AI”的1024 = 2048,欢迎大家加入2048 AI社区

更多推荐