强化学习笔记
强化学习
前言
在当今的大模型(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(s′∣s,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π[Gt∣St=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π[Gt∣St=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]=a∈A∑π(a∣s)s′,r∑p(s′,r∣s,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=1∑TlogP(xt∣x<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:
- 约束 (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(πθ(y∣x)∣∣πSFT(y∣x))- 原因:防止模型为了讨好 RM 而输出乱码(Reward Hacking),强迫其保持与 SFT 模型的相似性。
- 环境:非随机环境,状态转移是确定的文本拼接。
- 动作空间:巨大(词表大小,约 50k-100k),远超一般 RL 任务。
- 约束 (Constraint):必须加 KL 散度惩罚。
三、 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(at∣st)πθ(at∣st):概率比率 (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(yw∣x)πθ(yw∣x)−βlogπref(yl∣x)πθ(yl∣x)) - 优势:极简、稳定、显存占用小(类似监督学习)。
- 劣势:上限可能不如 PPO,在复杂逻辑推理任务上表现略逊。
4.2 GRPO (Group Relative Policy Optimization) —— “去掉了 Critic Model”
- 出处:DeepSeek (DeepSeek-Math / DeepSeek-R1)。
- 背景:PPO 需要一个与 Actor 同等大小的 Critic 模型(价值网络),显存消耗巨大。
- 核心机制:
- Group Sampling:对同一个 Prompt 采样一组回答 { y 1 , y 2 , … , y G } \{y_1, y_2, \dots, y_G\} {y1,y2,…,yG}。
- Baseline 计算:利用这组回答的平均分作为基准,代替 Critic 的预测值。
- 优势估计:
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({r1…rG})ri−mean({r1…rG})
- 优势:极度节省显存。不需要 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):
- Master Weights (4B):FP32 的权重备份。解决 FP16 精度不足导致的下溢出 (Underflow) 问题。
- Momentum (4B):动量(一阶矩),模拟物理惯性,加速收敛。
- 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)=qTRn−mk- 向量点积只与相对距离 ( n − m ) (n-m) (n−m) 有关,天然具备相对位置感知能力。
- 优势:
- 外推性强:通过线性内插 (Linear Interpolation),可以将位置索引压缩(如把 4096 压到 2048 范围内),让模型处理训练时没见过的超长文本(如 LLaMA 的 Long Context 扩展)。
- 实现简单:直接乘旋转矩阵,无需额外的参数学习。
六、 其他关键概念解析
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)是主要瓶颈。
- 核心原理:
- Tiling (分块计算):将 Q , K , V Q, K, V Q,K,V 矩阵切成小块,放入 GPU 的高速缓存 (SRAM) 中进行计算,减少对显存 (HBM) 的读写次数。
- 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 以上。
- 主流优化技术:
- PagedAttention (vLLM):借鉴操作系统的虚拟内存管理,将 KV Cache 分页存储(非连续显存),彻底解决显存碎片化问题,吞吐量提升 2-4 倍。
- 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),是掌握大模型微调与对齐技术的必经之路。
更多推荐

所有评论(0)