學(xué)原理到PyTorch實現(xiàn),OPD在線策略蒸餾實戰(zhàn)指南)
Transformer 相關(guān)的文章里講 KL 散度的不少但絕大多數(shù)都是一句話帶過“損失函數(shù)里加了一個 KL 項”。這對理解模型行為來說遠遠不夠。尤其是在 OPD 這類在線策略蒸餾框架里KL 的方向選擇直接決定了學(xué)生模型是“覆蓋教師分布”還是“聚焦教師分布”訓(xùn)練曲線和生成多樣性都會因此產(chǎn)生明顯差異。這篇文章我打算把反向 KL 這個細節(jié)徹底手撕一遍先對比正向 KL 和反向 KL 的數(shù)學(xué)性質(zhì)再給出 OPD 的框架建模然后用 PyTorch 寫出可運行的反向 KL 損失最后用一維高斯擬合的小實驗觀察 mode-seeking 行為。如果你正在做 Transformer 系列的模型蒸餾、LLM 對齊或生成模型微調(diào)這篇可以直接收藏。需要先說明一點OPD 在技術(shù)語境里最常見的意思是 Online Policy Distillation也就是在線策略蒸餾。它不是一個固定名字的開源倉庫而是一類訓(xùn)練框架。核心設(shè)置是一個凍結(jié)的教師 Transformer一個正在訓(xùn)練的學(xué)生 Transformer每一步用學(xué)生當(dāng)前策略采樣輸出再用教師對數(shù)似然作為訓(xùn)練信號。反向 KL 天然適合這種在線設(shè)置因為它的期望采樣分布就是學(xué)生自己的策略。如果你在其他場景看到的是 Online Preference Distillation 或 Optimal Policy Distillation文中反向 KL 的推導(dǎo)、代碼和調(diào)參思路同樣適用。1. 核心內(nèi)容速覽內(nèi)容維度說明技術(shù)主題Transformer / 生成模型中的反向 KL 散度核心框架OPD在線策略蒸餾Online Policy Distillation前置知識Softmax、交叉熵、自回歸 Transformer 基本結(jié)構(gòu)代碼環(huán)境Python 3.8、PyTorch具體版本以本機環(huán)境為準核心函數(shù)reverse_kl_exact / reverse_kl_gumbel / reverse_kl_reinforce主要應(yīng)用知識蒸餾、RLHF/DPO 風(fēng)格策略約束、生成分布對齊是否涉及 Web API不涉及核心是訓(xùn)練損失函數(shù)是否涉及批量任務(wù)支持訓(xùn)練循環(huán)中按 batch 進行批量蒸餾顯存說明取決于 batch_size、seq_len、vocab_size后文給出估算公式這篇文章不教你怎么部署模型服務(wù)也不涉及推理框架調(diào)用。它解決的是一個更底層的問題當(dāng)你想讓一個學(xué)生 Transformer 去模仿教師 Transformer 的分布時為什么要用“反向 KL”而不是直接用傳統(tǒng)交叉熵或者正向 KL以及如何在 PyTorch 里把它實現(xiàn)成一個穩(wěn)定可用的損失函數(shù)。2. 適用場景與使用邊界2.1 適合誰用這個內(nèi)容適合三類讀者。第一類是做大模型蒸餾的算法工程師。比如有一個 7B 的教師模型想把它壓縮成 2B 或者 1B 的學(xué)生模型OPD 框架下用反向 KL 可以在線獲得教師對“學(xué)生當(dāng)前生成結(jié)果”的反饋比一次性緩存所有教師 logits 的離線蒸餾更貼近學(xué)生當(dāng)前的分布變化。第二類是在做 RLHF、DPO 或偏好對齊的工程師。RLHF 的 KL 懲罰項通常寫作 KL(πθ || πref)這本身就是反向 KL。很多人在看公式時只記住了“加一個 KL 懲罰”但沒有意識到這里的期望采樣自當(dāng)前策略 πθ而不是參考策略 π