戰(zhàn))
強(qiáng)化學(xué)習(xí)這兩年在業(yè)界和學(xué)術(shù)圈的熱度不用多說(shuō)從游戲博弈到機(jī)器人控制再到 ChatGPT 背后的 RLHFReinforcement Learning from Human Feedback人類反饋強(qiáng)化學(xué)習(xí)幾乎處處都能看到它的影子。而不管算法包裝成什么樣只要深入到大模型對(duì)齊那一層就必然會(huì)碰到一個(gè)繞不開的組合Actor-Critic 算法。網(wǎng)上關(guān)于強(qiáng)化學(xué)習(xí)的資料很多但普遍存在兩個(gè)問題要么只講概念不給推導(dǎo)要么直接甩出一堆公式不講直覺。這篇教程會(huì)把演員-評(píng)論家算法作為主線從策略梯度為什么需要評(píng)論家講起先把 TD 誤差的數(shù)學(xué)含義拆開再一步步推導(dǎo) Actor-Critic 的更新公式最后結(jié)合 RLHF 場(chǎng)景說(shuō)明 PPO 在語(yǔ)言模型對(duì)齊中的落地方式并給出一份可以運(yùn)行的 PyTorch 示例代碼。適合剛接觸強(qiáng)化學(xué)習(xí)、希望看懂 RLHF 內(nèi)部原理或者準(zhǔn)備在實(shí)際項(xiàng)目中使用 Actor-Critic 框架的開發(fā)者閱讀。1. 強(qiáng)化學(xué)習(xí)與 RLHF 的關(guān)系1.1 什么是強(qiáng)化學(xué)習(xí)強(qiáng)化學(xué)習(xí)研究的是一個(gè)智能體Agent在與環(huán)境Environment交互過程中如何通過不斷試錯(cuò)來(lái)學(xué)習(xí)最優(yōu)策略Policy的問題。每次交互可以拆解為智能體在當(dāng)前狀態(tài) (s) 下根據(jù)策略 (\pi(a|s)) 選擇一個(gè)動(dòng)作 (a)環(huán)境返回一個(gè)即時(shí)獎(jiǎng)勵(lì) (r)并轉(zhuǎn)移到新狀態(tài) (s)。如此循環(huán)智能體的目標(biāo)不再是最大化某一步的獎(jiǎng)勵(lì)而是最大化長(zhǎng)期累積獎(jiǎng)勵(lì)的期望。這個(gè)過程和人們常說(shuō)的“監(jiān)督學(xué)習(xí)”有本質(zhì)區(qū)別。監(jiān)督學(xué)習(xí)有明確的標(biāo)準(zhǔn)答案模型通過擬合標(biāo)簽來(lái)學(xué)習(xí)輸入輸出映射而強(qiáng)化學(xué)習(xí)沒有標(biāo)準(zhǔn)答案只有獎(jiǎng)勵(lì)信號(hào)。獎(jiǎng)勵(lì)信號(hào)往往稀疏、延遲、存在噪聲這就導(dǎo)致算法必須解決“信用分配”問題到底哪一步動(dòng)作導(dǎo)致了最終的高回報(bào)傳統(tǒng)表格型方法在狀態(tài)空間較小時(shí)可以工作一旦狀態(tài)空間連續(xù)或維度很高就必須借助函數(shù)近似例如神經(jīng)網(wǎng)絡(luò)。深度強(qiáng)化學(xué)習(xí)Deep Reinforcement Learning就是“強(qiáng)化學(xué)習(xí) 深度神經(jīng)網(wǎng)絡(luò)”用深度網(wǎng)絡(luò)來(lái)表示策略、價(jià)值函數(shù)或模型。1.2 從強(qiáng)化學(xué)習(xí)到 RLHFRLHF 嚴(yán)格來(lái)說(shuō)并不是一種全新的強(qiáng)化學(xué)習(xí)算法而是一種“訓(xùn)練范式”。它解決的是“如何讓模型行為符合人類偏好”的問題。以大語(yǔ)言模型為例普通的監(jiān)督微調(diào)只能讓模型學(xué)會(huì)模仿人類給出的答案但無(wú)法真正理解“什么回答更好”。于是研究者設(shè)計(jì)了一條流程先訓(xùn)練一個(gè)獎(jiǎng)勵(lì)模型Reward Model讓它學(xué)會(huì)對(duì)人類偏好打分。再用強(qiáng)化學(xué)習(xí)算法去優(yōu)化語(yǔ)言模型讓模型生成的回答獲得更高的獎(jiǎng)勵(lì)分?jǐn)?shù)。在這個(gè)流程中語(yǔ)言模型本身就是強(qiáng)化學(xué)習(xí)里的 Actor它負(fù)責(zé)生成文本獎(jiǎng)勵(lì)模型輸出的分?jǐn)?shù)就是獎(jiǎng)勵(lì)信號(hào)。為了讓訓(xùn)練穩(wěn)定通常還會(huì)引入一個(gè)評(píng)論家網(wǎng)絡(luò)Critic來(lái)估計(jì)狀態(tài)價(jià)值或動(dòng)作價(jià)值。因此 RLHF 在實(shí)現(xiàn)層面大量使用 Actor-Critic 框架尤其是 PPOProximal Policy Optimization算法。這也是本文把“Actor-Critic”和“RLHF”放在一起講的原因。1.3 容易混淆的概念實(shí)際閱讀資料時(shí)有幾組概念經(jīng)常被混用先在這里做一次區(qū)分概念含義常見誤區(qū)Actor策略網(wǎng)絡(luò)輸入狀態(tài)輸出動(dòng)作分布被誤認(rèn)為只是“生成動(dòng)作的模型”Critic價(jià)值網(wǎng)絡(luò)輸入狀態(tài)輸出價(jià)值估計(jì)被誤認(rèn)為必須和 Actor 共享結(jié)構(gòu)TD 誤差時(shí)序差分誤差用于衡量真實(shí)獎(jiǎng)勵(lì)與估計(jì)之間的偏差被誤認(rèn)為是“優(yōu)勢(shì)函數(shù)”本身優(yōu)勢(shì)函數(shù)某個(gè)動(dòng)作相對(duì)平均水平的優(yōu)勢(shì)A(s,a)Q(s,a)-V(s)被誤認(rèn)為只能用 TD 誤差估計(jì)RLHF用人類反饋訓(xùn)練獎(jiǎng)勵(lì)模型再用強(qiáng)化學(xué)習(xí)優(yōu)化策略被誤認(rèn)為是一種具體算法理解這些概念后再看 Actor-Critic 推導(dǎo)會(huì)順暢很多。2. 策略梯度與 Actor-Critic 的動(dòng)機(jī)2.1 策略梯度為什么直接用獎(jiǎng)勵(lì)會(huì)不穩(wěn)定在策略梯度方法中我們希望最大化期望累積獎(jiǎng)勵(lì)[ J(\theta) \mathbb{E}{\tau \sim \pi\theta} [R(\tau)] ]其中 (\tau (s_0, a_0, r_0, s_1, a_1, r_1, \dots)) 是一條軌跡(R(\tau) \sum_{t0}^{T} \gamma^t r_t)。對(duì)參數(shù) (\theta) 求梯度可以得到著名的策略梯度定理[ \nabla_\theta J(\theta) \mathbb{E}{\tau \sim \pi\theta} \left[ \sum_{t0}^{T} \nabla_\theta \log \pi_\theta(a_t|s_t) \cdot R_t \right] ]這里的 (R_t \sum_{kt}^{T} \gamma^{k-t} r_k) 是軌跡從時(shí)刻 (t) 起的累積折扣回報(bào)。這個(gè)公式看起來(lái)很簡(jiǎn)單但直接使用存在一個(gè)問題方差極大。因?yàn)椴煌壽E之間的回報(bào)差異可能非常大同一個(gè)動(dòng)作在高回報(bào)軌跡和低回報(bào)軌跡中都可能出現(xiàn)算法難以判斷動(dòng)作本身是好是壞。為了降低方差通常會(huì)給獎(jiǎng)勵(lì)減去一個(gè)基線Baseline改成[ \nabla_\theta J(\theta) \mathbb{E} \left[ \sum_t \nabla_\theta \log \pi_\theta(a_t|s_t) \cdot \left( R_t - b(s_t) \right) \right] ]基線 (b(s_t)) 只要不依賴于動(dòng)作 (a_t)就不會(huì)改變梯度的期望但能顯著降低方差。一個(gè)自然的選擇是狀態(tài)價(jià)值函數(shù) (V^\pi(s_t))。于是策略梯度中的“回報(bào)減去基線”就變成了“動(dòng)作價(jià)值減去狀態(tài)價(jià)值”也就是優(yōu)勢(shì)函數(shù) (A^\pi(s_t, a_t) Q^\pi(s_t, a_t) - V^\pi(s_t))。這正是 Actor-Critic 的思想起點(diǎn)。2.2 從 REINFORCE 到 Actor-CriticREINFORCE 是典型的蒙特卡洛策略梯度算法它必須等到一條軌跡結(jié)束后才能計(jì)算累積回報(bào)。這種做法的優(yōu)點(diǎn)是估計(jì)無(wú)偏缺點(diǎn)是方差大、學(xué)習(xí)效率低。如果能在每個(gè)時(shí)間步都獲取一個(gè)“即時(shí)反饋”就可以更快更新。于是評(píng)論家出現(xiàn)了。評(píng)論家網(wǎng)絡(luò)負(fù)責(zé)估計(jì)價(jià)值函數(shù)例如 (V(s)) 或 (Q(s,a))。Actor 網(wǎng)絡(luò)負(fù)責(zé)生成策略。訓(xùn)練時(shí)Actor 根據(jù)評(píng)論家提供的優(yōu)勢(shì)估計(jì)更新策略評(píng)論家則根據(jù)實(shí)際觀察到的獎(jiǎng)勵(lì)和下一狀態(tài)的價(jià)值來(lái)更新自己。兩者交替訓(xùn)練這就是 Actor-Critic 框架。從 REINFORCE 到 Actor-Critic核心變化是把“采樣完整軌跡再計(jì)算回報(bào)”改成“用價(jià)值函數(shù)估計(jì)作為回報(bào)近似”。這種做法引入了偏差因?yàn)閮r(jià)值函數(shù)本身估計(jì)不準(zhǔn)但換來(lái)的是方差大幅下降。在深度強(qiáng)化學(xué)習(xí)實(shí)踐中方差問題往往比偏差問題更致命因此 Actor-Critic 方法成為主流。2.3 為什么需要 CriticCritic 網(wǎng)絡(luò)承擔(dān)了“評(píng)估當(dāng)前策略好壞”的任務(wù)。如果沒有 CriticActor 就只能依賴真實(shí)獎(jiǎng)勵(lì)信號(hào)而真實(shí)獎(jiǎng)勵(lì)通常是稀疏的。例如機(jī)械臂抓取任務(wù)只有抓到物體才給 1 獎(jiǎng)勵(lì)其他時(shí)候都是 0這種稀疏獎(jiǎng)勵(lì)讓策略梯度很難學(xué)習(xí)。Critic 可以通過學(xué)習(xí)狀態(tài)價(jià)值為每一步提供一個(gè)稠密的“預(yù)期收益”信號(hào)即使當(dāng)前沒有真實(shí)獎(jiǎng)勵(lì)也能通過 (r \gamma V(s) - V(s)) 給出一個(gè)更新信號(hào)。這就是時(shí)序差分Temporal DifferenceTD學(xué)習(xí)的價(jià)值。3. TD 誤差概念與公式推導(dǎo)3.1 TD 誤差的定義時(shí)序差分學(xué)習(xí)是強(qiáng)化學(xué)習(xí)中的核心思想之一。它結(jié)合了蒙特卡洛采樣和動(dòng)態(tài)規(guī)劃的思想用一步真實(shí)獎(jiǎng)勵(lì)加上下一狀態(tài)的價(jià)值估計(jì)來(lái)更新當(dāng)前狀態(tài)的價(jià)值估計(jì)。TD 誤差定義為[ \delta_t r_t \gamma V(s_{t1}) - V(s_t) ]其中(r_t) 是時(shí)刻 (t) 執(zhí)行動(dòng)作后獲得的即時(shí)獎(jiǎng)勵(lì)(\gamma) 是折扣因子取值范圍通常為 ([0, 1])(V(s_t)) 是當(dāng)前狀態(tài)價(jià)值的估計(jì)值(V(s_{t1})) 是下一狀態(tài)價(jià)值的估計(jì)值。如果 (V) 是真實(shí)價(jià)值函數(shù)那么根據(jù)貝爾曼方程(r_t \gamma V(s_{t1})) 的期望應(yīng)該等于 (V(s_t))因此 TD 誤差的期望為 0。如果估計(jì)有偏差TD 誤差就反映了這種偏差。在 Actor-Critic 中TD 誤差不僅是評(píng)論家的更新目標(biāo)也可以作為動(dòng)作優(yōu)勢(shì)的近似估計(jì)。3.2 TD 誤差與優(yōu)勢(shì)函數(shù)的關(guān)系優(yōu)勢(shì)函數(shù)定義為[ A^\pi(s_t, a_t) Q^\pi(s_t, a_t) - V^\pi(s_t) ]而動(dòng)作價(jià)值 (Q^\pi(s_t, a_t)) 滿足[ Q^\pi(s_t, a_t) r_t \gamma V^\pi(s_{t1}) ]將上式代入優(yōu)勢(shì)函數(shù)[ A^\pi(s_t, a_t) r_t \gamma V^\pi(s_{t1}) - V^\pi(s_t) \delta_t ]因此 TD 誤差可以被看作優(yōu)勢(shì)函數(shù)的一種單步采樣估計(jì)。當(dāng)然實(shí)際使用中由于價(jià)值函數(shù) (V) 是近似估計(jì)TD 誤差存在偏差但它在每一步都能計(jì)算非常適合在線學(xué)習(xí)。3.3 多步 TD 與 GAE單步 TD 雖然方差低但偏差可能較大尤其是在獎(jiǎng)勵(lì)延遲明顯的任務(wù)中。為了平衡偏差和方差可以使用多步回報(bào)[ A^{(n)}t \sum{k0}^{n-1} \gamma^k r_{tk} \gamma^n V(s_{tn}) - V(s_t) ]更進(jìn)一步GAEGeneralized Advantage Estimation廣義優(yōu)勢(shì)估計(jì)通過對(duì)不同步數(shù)的 TD 誤差做指數(shù)加權(quán)平均得到更靈活的優(yōu)勢(shì)估計(jì)[ A^{GAE}t \sum{l0}^{\infty} (\gamma \lambda)^l \delta_{tl} ]其中 (\lambda \in [0,1]) 控制偏差和方差的權(quán)衡。(\lambda 0) 時(shí)退化為單步 TD(\lambda 1) 時(shí)接近蒙特卡洛估計(jì)。PPO 等現(xiàn)代算法普遍使用 GAE就是因?yàn)樗谙∈瑾?jiǎng)勵(lì)場(chǎng)景下表現(xiàn)更好。關(guān)于 GAE 的推導(dǎo)本質(zhì)上是把多步優(yōu)勢(shì)估計(jì)按照 (\lambda) 做指數(shù)平滑這里不再展開矩陣形式但理解這個(gè)加權(quán)思路對(duì)調(diào)參很有幫助。4. Actor-Critic 算法數(shù)學(xué)推導(dǎo)4.1 目標(biāo)函數(shù)與梯度Actor 的目標(biāo)是最大化期望累積獎(jiǎng)勵(lì)但為了引入基線我們寫成[ \nabla_\theta J(\theta) \mathbb{E}{\pi\theta} \left[ \sum_t \nabla_\theta \log \pi_\theta(a_t|s_t) A(s_t, a_t) \right] ]其中 (A(s_t, a_t)) 代替了原來(lái)的回報(bào) (R_t)。這就是帶基線策略梯度的一般形式。在實(shí)際實(shí)現(xiàn)中期望用采樣近似因此 Actor 的損失函數(shù)一般取負(fù)的代理目標(biāo)[ L_{actor} -\frac{1}{N} \sum_{i1}^{N} \log \pi_\theta(a_i|s_i) \cdot A_i ]如果使用 TD 誤差作為優(yōu)勢(shì)估計(jì)則 (A_i) 用 (\delta_i r_i \gamma V_{\phi}(s_{i1}) - V_{\phi}(s_i)) 代替。但單步 TD 偏差較大所以實(shí)際項(xiàng)目中更常用 GAE 或使用評(píng)論家輸出 (V(s)) 計(jì)算多步優(yōu)勢(shì)。4.2 評(píng)論家的目標(biāo)函數(shù)評(píng)論家的目標(biāo)是讓價(jià)值函數(shù)估計(jì)更準(zhǔn)通常使用均方誤差MSE損失[ L_{critic} \frac{1}{N} \sum_{i1}^{N} \left( V_\phi(s_i) - \hat{R}_i \right)^2 ]其中 (\hat{R}_i) 是目標(biāo)回報(bào)例如[ \hat{R}i r_i \gamma V{\phi_{target}}(s_{i1}) ]需要注意評(píng)論家網(wǎng)絡(luò) (V_\phi) 如果和目標(biāo)回報(bào)使用同一組參數(shù)會(huì)造成自舉偏差的累積。因此很多實(shí)現(xiàn)會(huì)使用一個(gè)目標(biāo)網(wǎng)絡(luò)Target Network或延遲更新機(jī)制。在 PPO 中由于策略更新幅度被限制評(píng)論家網(wǎng)絡(luò)通常直接和 Actor 共享部分底層特征或完全獨(dú)立具體取決于狀態(tài)空間和動(dòng)作空間的復(fù)雜度。4.3 Actor-Critic 更新算法流程一個(gè)標(biāo)準(zhǔn)的 Actor-Critic 循環(huán)可以寫成下面的偽代碼初始化 Actor 網(wǎng)絡(luò) π_θ 和 Critic 網(wǎng)絡(luò) V_φ for 每個(gè)回合: 初始化狀態(tài) s for 每個(gè)時(shí)間步 t: 根據(jù) π_θ(·|s) 采樣動(dòng)作 a 執(zhí)行動(dòng)作 a觀察獎(jiǎng)勵(lì) r 和下一狀態(tài) s 計(jì)算 TD 誤差: δ r γ V_φ(s) - V_φ(s) 存儲(chǔ) (s, a, r, s, δ) 更新 Critic: φ ← φ - α_c * ?_φ (δ)^2 更新 Actor: θ ← θ α_a * ?_θ log π_θ(a|s) * δ s ← s實(shí)際深度強(qiáng)化學(xué)習(xí)中通常不是單步更新而是先收集一批數(shù)據(jù)再用小批量梯度更新類似于經(jīng)驗(yàn)回放。為什么需要批量因?yàn)閱蝹€(gè)樣本的 TD 誤差噪聲太大直接在線更新會(huì)讓網(wǎng)絡(luò)參數(shù)震蕩。4.4 從 Actor-Critic 到 PPOActor-Critic 框架雖然有效但直接使用策略梯度更新 Actor 時(shí)如果學(xué)習(xí)率設(shè)置不當(dāng)一次更新過大策略會(huì)發(fā)生突變導(dǎo)致訓(xùn)練崩潰。PPO 的解決方案是裁剪代理目標(biāo)[ 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] ]其中 (r_t(\theta) \frac{\pi_\theta(a_t|s_t)}{\pi_{\theta_{old}}(a_t|s_t)}) 是新舊策略的比率。通過裁剪PPO 限制每次更新的幅度從而保持訓(xùn)練穩(wěn)定。PPO 仍然使用 Actor-Critic 結(jié)構(gòu)Critic 估計(jì)狀態(tài)價(jià)值A(chǔ)ctor 輸出策略分布。這也是 RLHF 中最常用的強(qiáng)化學(xué)習(xí)算法。5. RLHF 中的 Actor-Critic 應(yīng)用5.1 RLHF 訓(xùn)練流程拆解RLHF 的完整流程可以分成三個(gè)階段監(jiān)督微調(diào)SFT先讓語(yǔ)言模型具備基礎(chǔ)對(duì)話能力。獎(jiǎng)勵(lì)模型訓(xùn)練收集人類對(duì)多個(gè)回答的排序或打分訓(xùn)練一個(gè)獎(jiǎng)勵(lì)模型用來(lái)代替人類的實(shí)時(shí)反饋。強(qiáng)化學(xué)習(xí)優(yōu)化用 PPO 等算法以獎(jiǎng)勵(lì)模型給出的分?jǐn)?shù)作為獎(jiǎng)勵(lì)信號(hào)更新語(yǔ)言模型策略。在第三個(gè)階段Actor 就是語(yǔ)言模型本身它的輸入是提示Prompt輸出是生成的整段文本。Critic 則是一個(gè)價(jià)值網(wǎng)絡(luò)輸入狀態(tài)通常是提示 已生成文本的某種嵌入到當(dāng)前步為止的上下文輸出一個(gè)標(biāo)量?jī)r(jià)值。因?yàn)樯呻A段被建模為馬爾可夫決策過程每一步生成一個(gè) token 相當(dāng)于一個(gè)動(dòng)作。5.2 為什么 RLHF 要用 PPO 而不是直接策略梯度直接使用原始策略梯度在語(yǔ)言模型場(chǎng)景下有幾個(gè)問題語(yǔ)言模型參數(shù)規(guī)模龐大一次采樣一條完整回答成本很高逐 token 的獎(jiǎng)勵(lì)非常稀疏只有整個(gè)回答生成完畢才能得到獎(jiǎng)勵(lì)模型的分?jǐn)?shù)如果單步更新太大模型會(huì)產(chǎn)生語(yǔ)法崩壞、重復(fù)冗余的輸出。PPO 通過重要性采樣和裁剪解決了更新幅度問題同時(shí)利用 Critic 對(duì)每個(gè) token 的狀態(tài)價(jià)值進(jìn)行估計(jì)能夠把“整段獎(jiǎng)勵(lì)”逐 token 分解成優(yōu)勢(shì)信號(hào)從而更新每一步的 token 生成策略。此外PPO 還加入 KL 散度懲罰項(xiàng)防止模型為了獎(jiǎng)勵(lì)模型的高分而偏離原始 SFT 模型太遠(yuǎn)。KL 懲罰的實(shí)際作用可以理解為“不要為了分?jǐn)?shù)犧牲可讀性”。5.3 PPO 在語(yǔ)言模型中的損失函數(shù)完整 PPO 損失通常包含三個(gè)部分Actor 的裁剪策略目標(biāo)Critic 的價(jià)值函數(shù)損失熵正則Entropy Bonus和 KL 懲罰項(xiàng)??梢詫懗蒣 L^{PPO} - L^{CLIP}(\theta) c_1 L^{VF}(\phi) - c_2 \mathcal{H}[\pi_\theta] \beta \cdot \text{KL} ]其中 (L^{VF}) 是評(píng)論家價(jià)值損失(\mathcal{H}) 是策略熵KL 項(xiàng)衡量當(dāng)前策略與參考策略一般是 SFT 模型的分布距離。實(shí)際工程中各項(xiàng)系數(shù) (c_1, c_2, \beta) 都需要針對(duì)具體任務(wù)調(diào)整。一個(gè)常見問題是KL 懲罰加在獎(jiǎng)勵(lì)上還是加在損失函數(shù)上兩種方式都存在。把 KL 懲罰加到實(shí)時(shí)獎(jiǎng)勵(lì)上更直觀PPO 采樣時(shí)每個(gè) token 獎(jiǎng)勵(lì)可以設(shè)置為[ r_t^{total} r_t^{reward_model}(samples) \cdot \mathbb{I}(t T) \beta \cdot \log \frac{\pi_{\theta}(a_t|s_t)}{\pi_{ref}(a_t|s_t)} ]這里只有最后一個(gè) token 收到獎(jiǎng)勵(lì)模型的標(biāo)量獎(jiǎng)勵(lì)其余 token 只收到 KL 懲罰信號(hào)。但實(shí)際實(shí)現(xiàn)中強(qiáng)化學(xué)習(xí)訓(xùn)練時(shí)整個(gè)序列的每個(gè) token 位置都會(huì)計(jì)算 Critic 的價(jià)值因此需要把序列級(jí)別的獎(jiǎng)勵(lì)映射到每一個(gè) token 位置。這也是 RLHF 訓(xùn)練代碼中最容易踩坑的地方。5.4 一個(gè)簡(jiǎn)化的 RLHF 訓(xùn)練偽代碼# 偽代碼僅用于理解流程不直接可運(yùn)行 for prompt in dataset: # 1. Actor 生成回答 response actor.generate(prompt) # 2. 計(jì)算每個(gè) token 的 logprob 和 KL log_probs actor.get_logprobs(prompt, response) ref_log_probs ref_model.get_logprobs(prompt, response) kl log_probs - ref_log_probs # 3. 獎(jiǎng)勵(lì)模型打分整段分?jǐn)?shù) reward_score reward_model(prompt, response) # 4. 將整段獎(jiǎng)勵(lì)分配到每個(gè) token常見做法最后一個(gè) token 得到獎(jiǎng)勵(lì) token_rewards torch.zeros_like(log_probs) token_rewards[-1] reward_score - beta * kl.mean() token_rewards[:-1] - beta * kl[:-1] # 5. 用 GAE 計(jì)算優(yōu)勢(shì) advantages gae(values, token_rewards, masks) # 6. 更新 Actor 和 Critic actor_loss clip_actor_loss(log_probs, old_log_probs, advantages) critic_loss mse_loss(values, returns) loss actor_loss critic_weight * critic_loss - entropy_weight * entropy loss.backward() optimizer.step()注意真實(shí) RLHF 代碼會(huì)比這段復(fù)雜得多包括共享內(nèi)存采樣、微批量更新、動(dòng)態(tài) KL 系數(shù)、長(zhǎng)度歸一化獎(jiǎng)勵(lì)等。理解流程后再去看開源實(shí)現(xiàn)會(huì)輕松很多。6. 最小可運(yùn)行示例PyTorch 實(shí)現(xiàn) Actor-Critic接下來(lái)給出一份可以直接運(yùn)行的 Actor-Critic 示例代碼環(huán)境使用 OpenAI Gym 的 CartPole-v1。這個(gè)任務(wù)相對(duì)簡(jiǎn)單不需要 GPU適合驗(yàn)證算法流程。代碼核心思路是用一個(gè)共享兩層全連接網(wǎng)絡(luò)分別輸出策略分布參數(shù)和價(jià)值標(biāo)量通過 TD 誤差更新。6.1 環(huán)境準(zhǔn)備本示例建議使用以下環(huán)境Python 3.8 及以上PyTorch 1.13 或 2.xgymnasium 0.28 或以上注意新版已從 gym 改名為 gymnasium安裝命令pip install torch gymnasium如果你的環(huán)境已經(jīng)安裝了老版本gym也可以跳過安裝但需要將代碼導(dǎo)入部分改為import gym。版本差異不影響算法邏輯。6.2 完整代碼# 文件路徑actor_critic_cartpole.py import torch import torch.nn as nn import torch.optim as optim import gymnasium as gym import math class ActorCritic(nn.Module): 共享底層特征的 Actor-Critic 網(wǎng)絡(luò)。 Actor 輸出動(dòng)作概率的 logitsCritic 輸出狀態(tài)價(jià)值 V(s)。 def __init__(self, state_dim, action_dim, hidden_dim128): super(ActorCritic, self).__init__() self.fc1 nn.Linear(state_dim, hidden_dim) self.fc2 nn.Linear(hidden_dim, hidden_dim) self.actor_head nn.Linear(hidden_dim, action_dim) self.critic_head nn.Linear(hidden_dim, 1) def forward(self, state): x torch.relu(self.fc1(state)) x torch.relu(self.fc2(x)) policy_logits self.actor_head(x) value self.critic_head(x) return policy_logits, value def get_action(self, state): policy_logits, value self.forward(state) dist torch.distributions.Categorical(logitspolicy_logits) action dist.sample() log_prob dist.log_prob(action) return action.item(), log_prob, value def train(env, agent, optimizer, gamma0.99, max_steps1000): 單回合訓(xùn)練采集軌跡計(jì)算 TD 誤差并更新網(wǎng)絡(luò)。 log_probs [] rewards [] values [] dones [] state, _ env.reset() state torch.tensor(state, dtypetorch.float32) total_reward 0 for step in range(max_steps): action, log_prob, value agent.get_action(state) next_state, reward, terminated, truncated, _ env.step(action) done terminated or truncated log_probs.append(log_prob) values.append(value) rewards.append(reward) dones.append(done) state torch.tensor(next_state, dtypetorch.float32) total_reward reward if done: break # 計(jì)算 TD 誤差和 Actor/Critic 損失 returns [] G 0.0 for i in reversed(range(len(rewards))): G rewards[i] gamma * G * (1 - dones[i]) returns.insert(0, G) returns torch.tensor(returns, dtypetorch.float32) values_tensor torch.cat(values).squeeze() log_probs_tensor torch.stack(log_probs) # Critic 損失價(jià)值估計(jì)與累計(jì)回報(bào)之間的 MSE critic_loss nn.functional.mse_loss(values_tensor, returns) # Actor 損失策略梯度的負(fù)對(duì)數(shù)似然 * 優(yōu)勢(shì)用 TD 誤差近似 advantages returns - values_tensor.detach() actor_loss -(log_probs_tensor * advantages).mean() loss actor_loss critic_loss optimizer.zero_grad() loss.backward() optimizer.step() return total_reward if __name__ __main__: env gym.make(CartPole-v1) state_dim env.observation_space.shape[0] action_dim env.action_space.n agent ActorCritic(state_dim, action_dim) optimizer optim.Adam(agent.parameters(), lr1e-3) episodes 500 for episode in range(episodes): reward train(env, agent, optimizer) if (episode 1) % 20 0: print(fEpisode {episode 1}, total reward: {reward})6.3 代碼說(shuō)明這段代碼雖然簡(jiǎn)單但完整展示了 Actor-Critic 的核心流程。ActorCritic類中actor_head輸出動(dòng)作分布的 logitscritic_head輸出狀態(tài)價(jià)值。get_action通過 Categorical 分布采樣動(dòng)作返回動(dòng)作、動(dòng)作對(duì)數(shù)概率和當(dāng)前狀態(tài)價(jià)值估計(jì)。訓(xùn)練函數(shù)中先與環(huán)境交互收集一整條軌跡然后在軌跡結(jié)束時(shí)計(jì)算累積回報(bào)G并把累積回報(bào)作為 Critic 的回歸目標(biāo)。優(yōu)勢(shì)函數(shù)使用returns - values.detach()這里并沒有直接用 TD 誤差而是用蒙特卡洛回報(bào)近似優(yōu)勢(shì)因?yàn)閱尾?TD 在 CartPole 上也可以運(yùn)行但方差更大。如果你希望觀察 TD 誤差的作用可以把優(yōu)勢(shì)換成rewards gamma * next_value - value注意需要額外計(jì)算下一狀態(tài)價(jià)值。運(yùn)行后如果網(wǎng)絡(luò)正常收斂你會(huì)在 200 到 400 個(gè)回合之間看到回報(bào)穩(wěn)定在 200 左右CartPole 的最大步數(shù)限制通常為 200。如果你的代碼在 100 回合內(nèi)回報(bào)不升反降一般是因?yàn)閷W(xué)習(xí)率過大或優(yōu)勢(shì)計(jì)算錯(cuò)誤。6.4 擴(kuò)展到 PPO 需要考慮什么上面的示例是單步更新的“原始 Actor-Critic”與 PPO 還有差距。要把它改成 PPO需要額外做幾件事保存每步的old_log_prob使用 GAE 計(jì)算優(yōu)勢(shì)對(duì) Actor 損失進(jìn)行裁剪多輪小批量更新而不是每回合只更新一次增加熵正則項(xiàng)鼓勵(lì)探索。如果這篇文章發(fā)布后大家感興趣我可以再單獨(dú)寫一篇“從零實(shí)現(xiàn) PPO”的教程把這三處改造逐一拆開?,F(xiàn)在先把 Actor-Critic 的基礎(chǔ)打好。7. 常見問題與排查思路問題現(xiàn)象常見原因解決思路訓(xùn)練一開始回報(bào)就在下降然后陷入震蕩學(xué)習(xí)率過大Actor 更新步長(zhǎng)過大調(diào)小學(xué)習(xí)率或使用 PPO 裁剪機(jī)制回報(bào)一直很低但 Critic 損失很小Critic 已經(jīng)收斂但 Actor 策略陷入局部最優(yōu)增大熵正則系數(shù)或增加動(dòng)作探索優(yōu)勢(shì)值波動(dòng)非常大梯度爆炸獎(jiǎng)勵(lì)尺度差異過大或 GAE 中的 lambda 設(shè)置不當(dāng)對(duì)獎(jiǎng)勵(lì)做歸一化例如除以獎(jiǎng)勵(lì)均值調(diào)整 gamma 和 lambdaRLHF 訓(xùn)練中模型輸出變得重復(fù)、死板KL 懲罰過小模型過度優(yōu)化獎(jiǎng)勵(lì)模型增大 KL 懲罰系數(shù)或使用動(dòng)態(tài) KL 系數(shù)語(yǔ)言模型生成退化輸出無(wú)意義符號(hào)PPO 更新時(shí)熵正則系數(shù)為 0策略過早確定添加熵正則并監(jiān)控 token 級(jí)別的熵多個(gè)進(jìn)程采樣時(shí)動(dòng)作分布不一致隨機(jī)種子未固定或模型權(quán)重不同步對(duì)齊隨機(jī)種子統(tǒng)一初始化權(quán)重Actor 和 Critic 共享網(wǎng)絡(luò)導(dǎo)致兩者梯度沖突兩個(gè)任務(wù)的梯度方向不一致使用雙頭網(wǎng)絡(luò)但底層分離或使用兩個(gè)獨(dú)立網(wǎng)絡(luò)訓(xùn)練不穩(wěn)定損失出現(xiàn) NaN數(shù)值溢出通常是 log 概率為 0 或梯度爆炸給 log_prob 加 epsilon使用梯度裁剪這些坑在學(xué)術(shù)示例和實(shí)際項(xiàng)目中都很常見。特別是 RLHF 場(chǎng)景問題往往不是算法本身而是獎(jiǎng)勵(lì)設(shè)計(jì)和 KL 約束沒有調(diào)好。建議每訓(xùn)練一段時(shí)間就手動(dòng)查看模型生成文本不要只盯著獎(jiǎng)勵(lì)曲線。8. 最佳實(shí)踐與工程建議8.1 優(yōu)勢(shì)函數(shù)與價(jià)值網(wǎng)絡(luò)的設(shè)計(jì)在 Actor-Critic 框架中Critic 的作用是為 Actor 提供低方差的學(xué)習(xí)信號(hào)。因此價(jià)值網(wǎng)絡(luò)的預(yù)測(cè)準(zhǔn)確性直接影響訓(xùn)練效果。建議在同一個(gè) batch 中Critic 可以多更新幾次而 Actor 每次更新幅度不要太大。在 PPO 中Critic 通常使用與 Actor 相同的 batch 數(shù)據(jù)但損失函數(shù)權(quán)重可以獨(dú)立調(diào)節(jié)。如果 Critic 和 Actor 使用共享底層特征需要小心梯度競(jìng)爭(zhēng)。很多成熟實(shí)現(xiàn)例如 Stable-Baselines3默認(rèn)使用兩個(gè)獨(dú)立的 MLP或共享一個(gè) Encoder但分開 Head。8.2 獎(jiǎng)勵(lì)工程設(shè)計(jì)強(qiáng)化學(xué)習(xí)對(duì)獎(jiǎng)勵(lì)尺度極其敏感。RLHF 中獎(jiǎng)勵(lì)模型輸出的分?jǐn)?shù)范圍可能很大直接在多個(gè)任務(wù)上使用會(huì)不穩(wěn)定。一種常見做法是對(duì)獎(jiǎng)勵(lì)做標(biāo)準(zhǔn)化z-score或者除以獎(jiǎng)勵(lì)標(biāo)準(zhǔn)差。但也不要過度歸一化否則會(huì)損失區(qū)分度。對(duì)于語(yǔ)言模型還有一種有效做法是按回答長(zhǎng)度歸一化獎(jiǎng)勵(lì)避免模型通過生成超長(zhǎng)文本來(lái)刷分。這是工業(yè)界總結(jié)出的寶貴經(jīng)驗(yàn)在學(xué)術(shù)論文中很容易被忽略。8.3 探索與利用的平衡Actor-Critic 的本質(zhì)是“邊探索邊利用”策略熵正則項(xiàng)是控制探索能力的關(guān)鍵。熵太大會(huì)導(dǎo)致策略過于隨機(jī)無(wú)法收斂熵太大會(huì)導(dǎo)致過早收斂到局部最優(yōu)。實(shí)踐中一般將熵系數(shù)設(shè)置為 0.01 到 0.001并在訓(xùn)練后期線性衰減。監(jiān)控每批次策略熵的平均值如果熵突然跌到 0說(shuō)明策略幾乎變成了確定性策略這在許多任務(wù)中是危險(xiǎn)的。8.4 RLHF 工程中的安全與合規(guī)RLHF 的目的是讓模型行為對(duì)齊人類偏好但“人類偏好”本身帶有主觀性必須遵守法律法規(guī)和倫理邊界。在訓(xùn)練獎(jiǎng)勵(lì)模型時(shí)要使用合法合規(guī)的數(shù)據(jù)集避免通過對(duì)抗樣本或惡意提示來(lái)“攻擊”獎(jiǎng)勵(lì)模型。不要試圖讓模型繞過安全限制也不要讓模型在敏感領(lǐng)域產(chǎn)生虛假信息。工程上RLHF 訓(xùn)練需要在受控環(huán)境中進(jìn)行設(shè)置人工審核環(huán)節(jié)對(duì)于人工反饋數(shù)據(jù)要做匿名化處理并確保數(shù)據(jù)采集獲得用戶授權(quán)。8.5 代碼可持續(xù)性強(qiáng)化學(xué)習(xí)實(shí)驗(yàn)代碼比普通訓(xùn)練代碼更容易腐化。建議從一開始就把“環(huán)境配置”“采樣循環(huán)”“模型更新”拆成獨(dú)立模塊。采樣環(huán)境使用向量化環(huán)境例如gymnasium.vector可以提升訓(xùn)練吞吐量。模型定義和損失函數(shù)用獨(dú)立函數(shù)封裝方便切換算法。日志系統(tǒng)記錄每個(gè) step 的獎(jiǎng)勵(lì)、價(jià)值損失、策略熵、KL 散度這能在排錯(cuò)時(shí)省下大量時(shí)間。保存模型時(shí)除了權(quán)重還要保存 optimizer 狀態(tài)、隨機(jī)種子和超參數(shù)否則無(wú)法復(fù)現(xiàn)實(shí)驗(yàn)結(jié)果。8.6 性能優(yōu)化CartPole 這類小環(huán)境用單進(jìn)程訓(xùn)練沒問題但真實(shí) RLHF 任務(wù)通常要并行采樣。語(yǔ)言模型生成文本很慢一般使用 vLLM 等推理框架加速采樣Critic 的價(jià)值網(wǎng)絡(luò)可以用單一的 GPU 作為獨(dú)立服務(wù)。數(shù)據(jù)傳輸建議使用共享內(nèi)存隊(duì)列避免序列化開銷。在更新階段Actor 和 Critic 可以輪流更新也可以同時(shí)更新。PPO 通常采用“先收集大量樣本再多次更新”的方式因此采樣進(jìn)程和訓(xùn)練進(jìn)程分離是必要的。如果你需要在多機(jī)環(huán)境訓(xùn)練還涉及參數(shù)同步和通信開銷建議先單機(jī)多卡跑通再擴(kuò)展。9. 總結(jié)與學(xué)習(xí)建議Actor-Critic 算法的價(jià)值在于它為深度強(qiáng)化學(xué)習(xí)和 RLHF 提供了一套穩(wěn)定的訓(xùn)練框架。理解了 Actor 輸出策略、Critic 輸出價(jià)值、TD 誤差充當(dāng)學(xué)習(xí)信號(hào)這三者之間的關(guān)系就相當(dāng)于打通了強(qiáng)化學(xué)習(xí)的主流脈絡(luò)。在此基礎(chǔ)上繼續(xù)學(xué)習(xí) GAE、PPO、TRPO、SAC 會(huì)更快。RLHF 本質(zhì)上是 Actor-Critic 在語(yǔ)言模型對(duì)齊中的應(yīng)用核心難點(diǎn)并不在策略梯度本身而在于獎(jiǎng)勵(lì)模型的質(zhì)量、KL 約束的調(diào)節(jié)以及大規(guī)模分布式訓(xùn)練工程。如果你想繼續(xù)深入建議按以下路線實(shí)踐先手動(dòng)推導(dǎo)一遍策略梯度定理和 TD 誤差公式運(yùn)行本文的 CartPole 示例修改優(yōu)勢(shì)估計(jì)方式觀察方差變化閱讀 Stable-Baselines3 的 PPO 源碼逐行對(duì)照公式用開源 RLHF 框架如 TRL、DeepSpeed-Chat跑一個(gè)小模型查看訓(xùn)練日志和生成效果。強(qiáng)化學(xué)習(xí)是一門“看公式覺得懂了寫代碼立刻懵”的學(xué)問但多寫幾次代碼多調(diào)幾次參數(shù)后直覺會(huì)逐漸建立起來(lái)。希望這篇教程能成為你理解 Actor-Critic 和 RLHF 的第一塊墊腳石遇到問題也歡迎在評(píng)論區(qū)討論。