計(jì)力學(xué)到深度學(xué)習(xí):能量模型原理、訓(xùn)練與實(shí)戰(zhàn))
1. 從統(tǒng)計(jì)力學(xué)到機(jī)器學(xué)習(xí)能量模型的前世今生做概率模型的人遲早會(huì)遇到 Energy Based Model 這個(gè)名字。我第一次認(rèn)真研究 EBM其實(shí)是帶著一個(gè)挺樸素的問題為什么物理學(xué)家研究氣體分子運(yùn)動(dòng)的那套數(shù)學(xué)會(huì)被原封不動(dòng)搬到機(jī)器學(xué)習(xí)里來后來才明白這兩件事本質(zhì)上都在回答同一個(gè)問題——給定一堆可能的配置每個(gè)配置出現(xiàn)的概率到底是多少在統(tǒng)計(jì)力學(xué)里一個(gè)微觀狀態(tài)出現(xiàn)的概率由玻爾茲曼分布決定概率正比于 exp(-E/T)其中 E 是能量T 是溫度。能量越低的狀態(tài)出現(xiàn)概率越高能量越高的狀態(tài)出現(xiàn)概率指數(shù)級(jí)下降。這個(gè)形式簡(jiǎn)潔到近乎優(yōu)雅而且給了我們一個(gè)非常直觀的世界觀系統(tǒng)傾向于朝著能量低的方向走。EBM 就是把這套世界觀直接搬進(jìn)了機(jī)器學(xué)習(xí)。我們不再直接定義 P(x)而是先定義一個(gè)能量函數(shù) E(x)以及 E(x,y)然后通過玻爾茲曼映射把能量變成概率。這樣做的第一好處是形式極度統(tǒng)一——任何你能用能量函數(shù)刻畫的關(guān)聯(lián)規(guī)則都可以納入這個(gè)框架第二好處是物理直覺清楚——訓(xùn)練一個(gè) EBM本質(zhì)上就是在鑿一個(gè)能量地形讓真實(shí)數(shù)據(jù)的能量低、非真實(shí)數(shù)據(jù)的能量高或者說得更技術(shù)一點(diǎn)讓模型分布去貼合數(shù)據(jù)分布。這篇文章主要適合兩類讀者一類是已經(jīng)接觸過一些深度生成模型比如 VAE、GAN想換個(gè)視角重新理解概率建模的人另一類是本來做統(tǒng)計(jì)物理或計(jì)算物理想看看自己的老本行在機(jī)器學(xué)習(xí)里怎么發(fā)光發(fā)熱的人。我會(huì)盡量把數(shù)學(xué)推導(dǎo)控制在夠用的范圍內(nèi)但一些關(guān)鍵的公式變換必須寫出來因?yàn)?EBM 所有的訓(xùn)練難點(diǎn)都藏在公式里。2. EBM 的核心框架能量、配分函數(shù)與概率分布2.1 能量函數(shù)與玻爾茲曼分布EBM 的定義其實(shí)非常簡(jiǎn)單。給定輸入數(shù)據(jù) x可以是圖片、序列、狀態(tài)向量等我們定義一個(gè)能量函數(shù) E_θ(x)帶參數(shù) θ。這個(gè)函數(shù)輸出一個(gè)實(shí)數(shù)標(biāo)量表示這個(gè)配置configuration的不合適程度。能量越低說明這個(gè)配置越自然或越合理。然后通過玻爾茲曼分布把能量轉(zhuǎn)換為概率密度P_θ(x) exp(-E_θ(x)) / Z(θ)這里的 Z(θ) 是配分函數(shù)計(jì)算公式是Z(θ) ∫ exp(-E_θ(x)) dx在離散情況下就是求和Z(θ) Σ_x exp(-E_θ(x))你看就這么兩步一個(gè)概率模型就建立起來了。這個(gè)框架的巧妙之處在于它沒有強(qiáng)迫我們直接給出一個(gè)歸一化的概率表達(dá)式而是允許我們先隨便定義一個(gè)打分函數(shù)能量函數(shù)最后再統(tǒng)一歸一化。2.2 配分函數(shù)EBM 一切苦難的根源配分函數(shù)這個(gè)量是 EBM 與普通判別模型最大的區(qū)別也是所有訓(xùn)練難題的根源。為什么難因?yàn)槟阋獙?duì)全空間所有可能的 x 求和/積分。在圖像任務(wù)里x 是 256×256×3 維的向量全空間的大小是天文數(shù)字精確求和是徹底不可能的。更麻煩的是配分函數(shù)本身還在參數(shù) θ 的控制下不斷變化。每更新一步參數(shù)Z 都跟著變而我們甚至很難估計(jì)它當(dāng)前的值。這就導(dǎo)致 EBM 沒法直接用標(biāo)準(zhǔn)的極大似然估計(jì)來做梯度下降——因?yàn)樗迫缓瘮?shù)里面帶著一個(gè)無法計(jì)算的歸一化常數(shù)。我記得剛學(xué)到這里的時(shí)候特別困惑既然配分函數(shù)這么難算為什么還要用這個(gè)框架為什么不直接定義一個(gè)歸一化的網(wǎng)絡(luò)輸出答案有兩個(gè)層面第一能量函數(shù)的形式可以非常靈活不受歸一化約束。這意味著你可以輕易地融入各種結(jié)構(gòu)約束、物理規(guī)律、對(duì)稱性這在很多科學(xué)問題里價(jià)值巨大。第二雖然精確計(jì)算配分函數(shù)不可能但梯度的估計(jì)是可行的。繞開配分函數(shù)直接估計(jì)模型梯度這就是對(duì)比散度Contrastive Divergence等一系列算法的出發(fā)點(diǎn)。2.3 對(duì)數(shù)似然梯度核心公式推導(dǎo)極大似然估計(jì)的目標(biāo)是最大化對(duì)數(shù)似然L(θ) E_{x~P_data}[log P_θ(x)]把 P_θ(x) exp(-E_θ(x)) / Z(θ) 代進(jìn)去得到log P_θ(x) -E_θ(x) - log Z(θ)對(duì) θ 求梯度?_θ log P_θ(x) -?_θ E_θ(x) - ?_θ log Z(θ)關(guān)鍵在于第二項(xiàng) ?_θ log Z(θ)。做一個(gè)簡(jiǎn)單的變換?_θ log Z(θ) (1/Z(θ)) ?_θ Z(θ) (1/Z(θ)) ?_θ ∫ exp(-E_θ(x)) dx (1/Z(θ)) ∫ exp(-E_θ(x)) (-?_θ E_θ(x)) dx ∫ [exp(-E_θ(x)) / Z(θ)] (-?_θ E_θ(x)) dx -E_{x~P_θ} [?_θ E_θ(x)]把這一項(xiàng)代回原式?_θ log P_θ(x) -?_θ E_θ(x) E_{x~P_θ} [?_θ E_θ(x)]寫成更對(duì)稱的形式?_θ L(θ) E_{x~P_data}[-?_θ E_θ(x)] - E_{x~P_θ}[?_θ E_θ(x)]這個(gè)公式是 EBM 訓(xùn)練的基石。它告訴我們兩件事第一項(xiàng)是把真實(shí)數(shù)據(jù)點(diǎn)的能量拉低這是正相positive phase 第二項(xiàng)是把模型采樣出來的點(diǎn)的能量拉高這是負(fù)相negative phase。整個(gè)訓(xùn)練過程就是在玩一個(gè)拔河游戲真實(shí)數(shù)據(jù)的能量往下壓模型幻想出來的數(shù)據(jù)的能量往上抬。最后平衡的時(shí)候模型分布就等于數(shù)據(jù)分布。從物理角度理解第一項(xiàng)相當(dāng)于讓系統(tǒng)更傾向于停留在數(shù)據(jù)所在的低能區(qū)域第二項(xiàng)相當(dāng)于對(duì)系統(tǒng)的其他區(qū)域施加排斥力防止模型把概率質(zhì)量攤得到處都是。2.4 從物理視角看模型行為統(tǒng)計(jì)力學(xué)的語言在這里非常好用。把 EBM 訓(xùn)練好的模型想象成一個(gè)能量地形圖數(shù)據(jù)點(diǎn)集中在若干能量盆地里盆地之間的山脊能量很高。采樣的時(shí)候模型在熱噪聲的驅(qū)動(dòng)下在地形圖上漫游——它更愿意待在低能量的盆地偶爾也會(huì)翻越山脊跑到另一個(gè)盆地。這個(gè)畫面比神經(jīng)網(wǎng)絡(luò)輸出一個(gè)概率生動(dòng)得多而且在很多場(chǎng)景下更有解釋力。比如在分子構(gòu)象生成任務(wù)里能量盆地對(duì)應(yīng)的就是一個(gè)一個(gè)穩(wěn)定構(gòu)象在圖像生成里能量盆地對(duì)應(yīng)的是不同類別的圖像流形。實(shí)操中我特別喜歡用一個(gè)類比向別人解釋 EBM想象一個(gè)彈性勢(shì)能場(chǎng)數(shù)據(jù)點(diǎn)在底部安營(yíng)扎寨負(fù)相采樣相當(dāng)于往這個(gè)勢(shì)能場(chǎng)里隨機(jī)扔小球看它們最終滾到哪里。訓(xùn)練的目標(biāo)就是不斷重塑這個(gè)地形讓小球最終總愛往數(shù)據(jù)點(diǎn)附近滾。3. 能量函數(shù)的設(shè)計(jì)從受限玻爾茲曼機(jī)到深度能量網(wǎng)絡(luò)3.1 經(jīng)典選擇受限玻爾茲曼機(jī)RBM說到 EBM 的歷史繞不開受限玻爾茲曼機(jī)Restricted Boltzmann Machine, RBM。RBM 是一個(gè)二部圖結(jié)構(gòu)可見層 v 和隱藏層 h層內(nèi)無連接層間全連接。它的能量函數(shù)定義為E(v,h) -Σ_i a_i v_i - Σ_j b_j h_j - Σ_{i,j} v_i W_{ij} h_j這里的 a、b 是偏置項(xiàng)W 是可見層和隱藏層之間的權(quán)重矩陣。因?yàn)閷觾?nèi)無連接所以給定一個(gè)層另一個(gè)層的條件分布是獨(dú)立的這給采樣帶來了極大的方便。RBM 在 2006 年深度學(xué)習(xí)復(fù)興的時(shí)候扮演過關(guān)鍵角色——Hinton 用它做逐層預(yù)訓(xùn)練訓(xùn)練深度信念網(wǎng)絡(luò)。我自己也從頭寫過 RBM 的代碼說實(shí)話訓(xùn)練 RBM 比想象中要難難在負(fù)相采樣的質(zhì)量。經(jīng)典的做法是用對(duì)比散度CD-k即從訓(xùn)練數(shù)據(jù)出發(fā)做 k 步吉布斯采樣來近似負(fù)相。3.2 深度能量網(wǎng)絡(luò)與現(xiàn)代架構(gòu)RBM 的線性結(jié)構(gòu)表達(dá)能力有限現(xiàn)代的 EBM 基本都用深度神經(jīng)網(wǎng)絡(luò)直接做能量函數(shù)。也就是 E_θ(x) Net(x)輸入 x輸出一個(gè)標(biāo)量。網(wǎng)絡(luò)內(nèi)部可以是任意結(jié)構(gòu)——卷積網(wǎng)絡(luò)、Transformer、殘差網(wǎng)絡(luò)都可以。但直接讓網(wǎng)絡(luò)輸出一個(gè)標(biāo)量自由度太大了很容易出現(xiàn)訓(xùn)練不穩(wěn)定的情況。實(shí)踐中常見的設(shè)計(jì)策略有第一種是去噪自編碼器風(fēng)格的能量。輸入被加入噪聲后網(wǎng)絡(luò)的目標(biāo)是盡量輸出一個(gè)能量值讓干凈數(shù)據(jù)的能量低、加噪數(shù)據(jù)的能量高。這本質(zhì)上是在讓能量函數(shù)學(xué)會(huì)分辨干凈信號(hào)和噪聲。第二種是基于得分score的視角。我們其實(shí)不關(guān)心能量的絕對(duì)值只關(guān)心能量對(duì)輸入的梯度 ?_x E_θ(x)這個(gè)梯度叫得分函數(shù)。在很多生成方法如 Langevin 采樣里我們只需要這個(gè)梯度就足夠完成采樣根本不需要算配分函數(shù)。這給了我們極大的設(shè)計(jì)自由——甚至可以讓網(wǎng)絡(luò)直接輸出得分而不是標(biāo)量能量。第三種是對(duì)比學(xué)習(xí)式的思路。把能量函數(shù)設(shè)計(jì)成一種度量正樣本對(duì)的能量低負(fù)樣本對(duì)的能量高。這在一些度量學(xué)習(xí)和檢索任務(wù)里非常好用。3.3 能量函數(shù)的選擇原則這里分享一些我踩過坑之后總結(jié)出的經(jīng)驗(yàn)?zāi)芰亢瘮?shù)不是越復(fù)雜越好。網(wǎng)絡(luò)容量越大能量地形越崎嶇負(fù)相采樣就越容易陷入局部模式。如果你的生成任務(wù)不是特別復(fù)雜一個(gè)中等規(guī)模的網(wǎng)絡(luò)往往比一個(gè)超大網(wǎng)絡(luò)效果更好。我試過一個(gè) 6 層的 MLP 做 MNIST 上的 EBM效果竟然不比 ResNet 差太多但訓(xùn)練穩(wěn)定得多。能量函數(shù)對(duì)輸入的依賴方式很重要。如果你希望模型有平移不變性比如圖像任務(wù)能量函數(shù)應(yīng)該用卷積結(jié)構(gòu)來構(gòu)建而不是把圖像拉平后丟進(jìn)全連接層。否則你需要海量數(shù)據(jù)來讓模型自己學(xué)會(huì)平移不變性這在能量模型里尤其難學(xué)。注意能量函數(shù)的尺度。能量值的絕對(duì)大小會(huì)影響采樣步長(zhǎng)和溫度參數(shù)的選擇。我習(xí)慣在能量網(wǎng)絡(luò)的最后一層加一個(gè) tanh 或者把輸出尺度限制在一定范圍這樣采樣超參數(shù)更容易調(diào)節(jié)。3.4 一個(gè)簡(jiǎn)單的 EBM 網(wǎng)絡(luò)實(shí)現(xiàn)以 PyTorch 為例一個(gè)用于圖像的最小 EBM 可以這樣定義import torch import torch.nn as nn class SimpleEBM(nn.Module): def __init__(self, input_dim784, hidden_dim256): super().__init__() self.net nn.Sequential( nn.Linear(input_dim, hidden_dim), nn.Softplus(), nn.Linear(hidden_dim, hidden_dim), nn.Softplus(), nn.Linear(hidden_dim, 1) ) def forward(self, x): # 輸入形狀: (batch, input_dim)輸出形狀: (batch, 1) return self.net(x).squeeze(-1)注意這里激活函數(shù)用了 Softplus 而不是 ReLU。原因是 Softplus 處處可導(dǎo)且導(dǎo)數(shù)連續(xù)對(duì)基于梯度的 Langevin 采樣更友好。ReLU 在負(fù)半軸的梯度為 0采樣到這些區(qū)域時(shí)得分會(huì)失效導(dǎo)致粒子原地卡住。4. 訓(xùn)練 EBM 的核心難點(diǎn)配分函數(shù)與對(duì)比散度4.1 為什么不能直接做最大似然我們已經(jīng)看到最大似然梯度的表達(dá)式非常干凈?_θ L(θ) E_{P_data}[-?_θ E_θ(x)] - E_{P_θ}[?_θ E_θ(x)]漂亮是漂亮但第二項(xiàng)包含對(duì)模型分布 P_θ 的期望。要算這個(gè)期望就得從當(dāng)前模型采樣——而這恰恰是需要配分函數(shù)的因果循環(huán)了。如果你硬算呢把配分函數(shù)的數(shù)值估計(jì)出來然后代入公式計(jì)算梯度。在低維空間里可以這么做比如二維的高斯混合模型用數(shù)值積分算 Z 沒問題。但一旦維度升上去圖像就成千上萬維數(shù)值積分直接爆炸。所以我們必須創(chuàng)造一種避開精確配分函數(shù)的方法。4.2 對(duì)比散度CD一個(gè)天才的近似Hinton 在 2002 年提出的對(duì)比散度Contrastive Divergence, CD是歷史上最成功的近似方案。CD 的出發(fā)點(diǎn)是一個(gè)很實(shí)際的觀察我們不需要精確地從 P_θ 采樣只需要一個(gè)方向大致正確的負(fù)相樣本來計(jì)算梯度。既然精確采樣難那我們從訓(xùn)練數(shù)據(jù)出發(fā)做 k 步吉布斯采樣或 Langevin 采樣把得到的樣本當(dāng)作 P_θ 的近似樣本。關(guān)鍵在于初始化的選擇正相采樣的起點(diǎn)是訓(xùn)練數(shù)據(jù) x而不是隨機(jī)噪聲。這意味著采樣器只需要從數(shù)據(jù)點(diǎn)開始漂移一小段距離就能反映出模型分布與數(shù)據(jù)分布在局部上的偏差方向。CD-k 算法的流程是從訓(xùn)練集取一個(gè) batch 的真實(shí)樣本 x從 x 出發(fā)執(zhí)行 k 步采樣RBM 里是塊吉布斯采樣連續(xù) EBM 里是 Langevin 采樣得到負(fù)相樣本 x計(jì)算正相梯度 ?_θ E_θ(x) 和負(fù)相梯度 ?_θ E_θ(x)兩者相減得到近似梯度用這個(gè)梯度更新參數(shù) θ。CD 的一個(gè)細(xì)節(jié)魔鬼是從數(shù)據(jù)點(diǎn)初始化意味著采樣永遠(yuǎn)傾向于停留在數(shù)據(jù)分布附近對(duì)遠(yuǎn)端的高概率區(qū)域探索不足。這在實(shí)踐中會(huì)導(dǎo)致模型學(xué)到的分布比真實(shí)分布更尖銳——它把概率集中在了訓(xùn)練數(shù)據(jù)的小鄰域內(nèi)而犧牲了對(duì)整個(gè)流形的覆蓋。4.3 持久對(duì)比散度PCD把鏈子養(yǎng)起來為了解決 CD 探索不足的問題Tieleman 提出了持久對(duì)比散度Persistent Contrastive Divergence, PCD。思路也很妙不要每次從數(shù)據(jù)點(diǎn)重新初始化采樣鏈而是維護(hù)一組持續(xù)運(yùn)行的采樣鏈稱為持久鏈在每一步訓(xùn)練中讓這些鏈繼續(xù)采樣幾步用它們得到的樣本作為負(fù)相樣本。這樣做的好處是采樣鏈有充分的時(shí)間在模型分布中游走能探索到更遠(yuǎn)的區(qū)域壞處是當(dāng)模型參數(shù)快速變化時(shí)采樣鏈可能跟不上參數(shù)的更新導(dǎo)致負(fù)相樣本過時(shí)梯度不準(zhǔn)確。PCD 在訓(xùn)練初期效果很好但后期會(huì)遇到一個(gè)常見問題采樣鏈會(huì)逐漸收斂到某個(gè)模式附近丟失多樣性。這就是所謂的模式坍縮mode collapse現(xiàn)象在 EBM 訓(xùn)練中的體現(xiàn)。我在實(shí)踐中對(duì) PCD 做過一個(gè)小改進(jìn)給持久鏈加入周期性重置。比如每 500 步訓(xùn)練隨機(jī)選取一部分持久鏈的狀態(tài)用隨機(jī)噪聲或某個(gè)訓(xùn)練樣本重新初始化。這樣可以避免鏈子陷入單一模式而長(zhǎng)期無法擺脫的情況。4.4 對(duì)比學(xué)習(xí)與 score matching另辟蹊徑除了 CD 家族還有兩大類訓(xùn)練 EBM 的方法各有各的適用場(chǎng)景。Score Matching得分匹配的思路是不直接優(yōu)化數(shù)據(jù)的 log 似然而是讓模型的得分函數(shù) ?_x log P_θ(x) 盡量接近真實(shí)數(shù)據(jù)分布的得分函數(shù) ?_x log P_data(x)。因?yàn)檎鎸?shí)分布的得分是未知的要把它用可計(jì)算的形式替換掉。Score Matching 有一個(gè)漂亮的結(jié)論——目標(biāo)函數(shù)可以化為一個(gè)只依賴模型得分梯度Hessian和模型得分值的期望式配分函數(shù)完全消掉了。這個(gè)方法的優(yōu)點(diǎn)是穩(wěn)定不需要采樣缺點(diǎn)是 Hessian 計(jì)算代價(jià)高而且它對(duì)分布的支持域假設(shè)比較嚴(yán)格。后來 Hyv?rinen 又提出了 Sliced Score Matching用隨機(jī)投影近似 Hessian大大緩解了計(jì)算壓力。噪聲對(duì)比估計(jì)NCE的思路則是既然配分函數(shù)難算那就把配分函數(shù)當(dāng)作一個(gè)額外參數(shù)來學(xué)。引入一個(gè)輔助噪聲分布 p_noise(x)把問題變成二分類——判斷一個(gè)樣本來自數(shù)據(jù)還是噪聲。這樣配分函數(shù)被隱性吸收進(jìn)了判別器的偏置項(xiàng)里。NCE 的缺點(diǎn)是噪聲分布的選擇對(duì)效果影響巨大。如果噪聲分布離數(shù)據(jù)分布太遠(yuǎn)判別任務(wù)太簡(jiǎn)單學(xué)到的東西會(huì)很粗糙如果太近又難以區(qū)分。實(shí)踐里常用的策略是用數(shù)據(jù)的擾動(dòng)版本作為噪聲分布配合學(xué)習(xí)率退火效果不錯(cuò)。4.5 各種訓(xùn)練方法對(duì)比為了幫你快速?zèng)Q策不同場(chǎng)景下選哪種方法這里整理一個(gè)基于我實(shí)際經(jīng)驗(yàn)對(duì)照表。方法核心思想優(yōu)點(diǎn)缺點(diǎn)適用場(chǎng)景CD-k從數(shù)據(jù)點(diǎn)出發(fā)采樣 k 步簡(jiǎn)單、收斂快對(duì)分布覆蓋不足容易出現(xiàn)模式坍縮中小規(guī)模數(shù)據(jù)、RBM、快速原型驗(yàn)證PCD維護(hù)持久采樣鏈負(fù)相質(zhì)量高、探索充分采樣鏈可能滯后或坍縮、超參數(shù)多中等規(guī)模連續(xù)數(shù)據(jù)、圖像Score Matching匹配模型與數(shù)據(jù)的得分函數(shù)無需采樣、穩(wěn)定Hessian 代價(jià)高、對(duì)分布形式有要求低維到中維連續(xù)數(shù)據(jù)Sliced SM隨機(jī)投影近似 Hessian計(jì)算可控、可擴(kuò)展實(shí)現(xiàn)復(fù)雜高維連續(xù)數(shù)據(jù)NCE與噪聲分布做判別無采樣、思想簡(jiǎn)單依賴噪聲分布選擇分布有明確先驗(yàn)的場(chǎng)景5. Langevin 采樣讓能量變成樣本5.1 從能量到樣本的物理過程訓(xùn)練 EBM 不是終點(diǎn)最終目的是從模型中采樣生成新樣本。由于配分函數(shù)未知我們不能像其他歸一化模型那樣直接算概率而是要用馬爾可夫鏈蒙特卡洛MCMC的方法。最常用的采樣工具是 Langevin 動(dòng)力學(xué)。它的更新規(guī)則是x_{t1} x_t - (ε/2) ?_x E_θ(x_t) √ε · z_t其中 z_t ~ N(0, I) 是高斯噪聲ε 是步長(zhǎng)。這個(gè)公式的物理含義非常清楚第一項(xiàng)讓粒子沿著能量下降的方向滑動(dòng)漂移項(xiàng)第二項(xiàng)加入熱噪聲讓粒子能夠翻越能量壁壘擴(kuò)散項(xiàng)。平衡狀態(tài)下粒子的分布恰好是玻爾茲曼分布 exp(-E/τ)。注意嚴(yán)格來說 Langevin 采樣給出的分布是 exp(-2E/τ) 的形式所以實(shí)際使用時(shí)步長(zhǎng)和溫度的關(guān)系需要小心處理。我在代碼里通常會(huì)這樣寫def langevin_step(x, energy_fn, step_size0.1, noise_scale1.0): x.requires_grad_(True) energy energy_fn(x).sum() grad torch.autograd.grad(energy, x)[0] x x.detach() - step_size * grad noise_scale * torch.randn_like(x) return x5.2 采樣步數(shù)與步長(zhǎng)的博弈Langevin 采樣有一個(gè)核心矛盾步長(zhǎng)太大動(dòng)力學(xué)不穩(wěn)定粒子容易發(fā)散步長(zhǎng)太小需要很多步才能從初始點(diǎn)走到高概率區(qū)域。而且 EBM 的能量地形通常很不均勻——有的區(qū)域平坦、有的區(qū)域陡峭單一的學(xué)習(xí)率很難在所有區(qū)域都表現(xiàn)良好。實(shí)踐中我常用的策略是一個(gè)兩步走的方案。第一步用較大的步長(zhǎng)比如 0.5做 10-20 步預(yù)熱采樣讓粒子快速靠近低能量區(qū)域第二步改用較小的步長(zhǎng)比如 0.05做 30-50 步精煉采樣讓粒子在高概率區(qū)域內(nèi)部充分混合。這樣可以兼顧效率和精度。還有一個(gè)特別容易踩的坑Langevin 采樣對(duì)能量的絕對(duì)尺度非常敏感。如果能量函數(shù)輸出的數(shù)值范圍很大比如幾百到幾千那么梯度也會(huì)很大即使步長(zhǎng)很小也會(huì)一步跑飛。解決方法是給能量輸出做歸一化處理或者使用自適應(yīng)步長(zhǎng)。我在自己的框架里用了一個(gè)簡(jiǎn)單的技巧記錄最近若干步的梯度均方根如果過大就減小步長(zhǎng)過小就增大步長(zhǎng)效果相當(dāng)不錯(cuò)。5.3 從 RBM 到連續(xù) EBM 的采樣對(duì)比RBM 因?yàn)閷觾?nèi)無連接的結(jié)構(gòu)可以使用塊吉布斯采樣先固定可見層 v從 P(h|v) 采樣隱藏層再固定隱藏層 h從 P(v|h) 采樣可見層。這兩步都只需要做獨(dú)立采樣非常高效# RBM 條件采樣偽代碼 def gibbs_step(v, W, a, b): # 采樣隱藏層 h_prob torch.sigmoid(b v W) h torch.bernoulli(h_prob) # 采樣可見層 v_prob torch.sigmoid(a h W.t()) v torch.bernoulli(v_prob) return v連續(xù) EBM 就沒有這種便利只能走 Langevin 路線。但和 RBM 的塊吉布斯相比Langevin 采樣的優(yōu)勢(shì)在于可以處理連續(xù)變量而且不需要設(shè)計(jì)條件分布適用面廣得多。6. 實(shí)際應(yīng)用場(chǎng)景與案例EBM 在圖像、科學(xué)計(jì)算與決策問題中的實(shí)戰(zhàn)6.1 圖像生成與異常檢測(cè)EBM 在圖像任務(wù)上最經(jīng)典的應(yīng)用之一是把它當(dāng)作一個(gè)可學(xué)習(xí)的能量地形然后用 Langevin 采樣生成圖像。早期的深度 EBM 工作比如 Yann LeCun 團(tuán)隊(duì)的論文展示了在 CIFAR-10 和 MNIST 上可以生成合理的樣本。不過老實(shí)說純 EBM 的圖像生成質(zhì)量在相當(dāng)一段時(shí)間內(nèi)都不如 GAN 或擴(kuò)散模型。它的優(yōu)勢(shì)更多體現(xiàn)在其他方面——尤其是異常檢測(cè)。因?yàn)?EBM 天然地給每個(gè)輸入打了一個(gè)能量分訓(xùn)練時(shí)正常樣本的能量被壓低異常樣本模型沒見過的類型通常落在高能量區(qū)域。你不需要專門訓(xùn)練一個(gè)分類頭直接把能量當(dāng)作異常分?jǐn)?shù)就行。我做過一個(gè)工業(yè)質(zhì)檢的小實(shí)驗(yàn)用正常產(chǎn)品圖片訓(xùn)練 EBM到了測(cè)試階段把有缺陷的產(chǎn)品圖輸入網(wǎng)絡(luò)能量值會(huì)明顯偏高。這個(gè)方案在只有正樣本、沒有負(fù)樣本的場(chǎng)景下特別好用傳統(tǒng)監(jiān)督學(xué)習(xí)很難處理這種問題。6.2 分子構(gòu)象生成與科學(xué)計(jì)算EBM 在科學(xué)計(jì)算領(lǐng)域有一個(gè)根正苗紅的優(yōu)勢(shì)很多物理系統(tǒng)本身就是能量模型。分子力場(chǎng)就是典型的能量函數(shù)蛋白質(zhì)折疊問題里用的也是能量地形。用 EBM 去學(xué)習(xí)分子數(shù)據(jù)能量函數(shù)不僅是一個(gè)生成模型還能被解釋為一種可學(xué)習(xí)的物理勢(shì)能。我接觸過的一個(gè)方向是分子構(gòu)象生成給定一個(gè)分子的化學(xué)式生成它在不同溫度下可能出現(xiàn)的三維構(gòu)象。用 EBM 建模時(shí)能量函數(shù)的輸入是原子的三維坐標(biāo)輸出是構(gòu)象的能量。得益于 EBM 的物理可解釋性采樣得到的構(gòu)象天然符合玻爾茲曼分布不同構(gòu)象的出現(xiàn)頻率大致正比于 exp(-E/kT)這在藥物設(shè)計(jì)中特別有價(jià)值。另一個(gè)讓我覺得有潛力的方向是蛋白質(zhì)設(shè)計(jì)。AlphaFold 預(yù)測(cè)的是結(jié)構(gòu)但怎樣在序列空間中搜索能折疊成目標(biāo)結(jié)構(gòu)的序列本質(zhì)上是一個(gè)在能量地形上采樣的過程。EBM 的框架在這里能自然地結(jié)合物理約束和實(shí)驗(yàn)數(shù)據(jù)。6.3 與強(qiáng)化學(xué)習(xí)的結(jié)合能量視角的策略表示EBM 在強(qiáng)化學(xué)習(xí)里有一個(gè)不太為人知但潛力很大的用法——把策略policy表示為能量模型。具體來說給定狀態(tài) s動(dòng)作 a 的條件能量是 E_θ(s, a)策略就是π(a|s) exp(-E_θ(s, a)) / Z(s)這樣做的理由很實(shí)際某些場(chǎng)景下最優(yōu)策略是多模態(tài)的。比如一個(gè)機(jī)器人走到分岔路口往左和往右都可以到達(dá)目的地但政策梯度方法通常只能學(xué)到一個(gè)模式的分布因?yàn)楦咚狗植际菃畏宓摹BM 可以天然地表示多模態(tài)策略——能量地形上有幾個(gè)盆地就對(duì)應(yīng)幾個(gè)策略模式。這類方法在模仿學(xué)習(xí)里也有應(yīng)用比如 Energy-Based Imitation Learning。把專家軌跡映射到能量地形上專家動(dòng)作落在低能量區(qū)域非專家動(dòng)作能量高這比直接行為克隆更魯棒特別是在專家數(shù)據(jù)不完美的情況下。6.4 一個(gè)最小可運(yùn)行的訓(xùn)練循環(huán)為了讓你能快速上手我提供一個(gè)完整的訓(xùn)練循環(huán)骨架。這里用最簡(jiǎn)單的 CD-1 思路import torch import torch.nn as nn import torchvision.datasets as datasets import torchvision.transforms as transforms # 數(shù)據(jù)加載 mnist datasets.MNIST(./data, trainTrue, downloadTrue, transformtransforms.ToTensor()) loader torch.utils.data.DataLoader(mnist, batch_size128, shuffleTrue) # 初始化模型與優(yōu)化器 model SimpleEBM(input_dim784, hidden_dim256) optimizer torch.optim.Adam(model.parameters(), lr1e-3) def langevin_sample(x, k30, step_size0.05): 從輸入 x 出發(fā)做 k 步 Langevin 采樣 x x.clone().detach().requires_grad_(True) for _ in range(k): energy model(x).sum() grad torch.autograd.grad(energy, x)[0] x x.detach() step_size * (-grad) (2 * step_size) ** 0.5 * torch.randn_like(x) x x.clamp(0, 1) # 像素值約束 return x.detach() # 訓(xùn)練循環(huán) for epoch in range(20): for batch_idx, (data, _) in enumerate(loader): x_real data.view(-1, 784) # (batch, 784) # 負(fù)相采樣從數(shù)據(jù)點(diǎn)出發(fā)做 Langevin 采樣 x_fake langevin_sample(x_real, k10, step_size0.1) # 計(jì)算正相和負(fù)相能量 e_real model(x_real).mean() e_fake model(x_fake).mean() # CD 損失正相能量 - 負(fù)相能量 loss e_real - e_fake optimizer.zero_grad() loss.backward() optimizer.step() if batch_idx % 200 0: print(fEpoch {epoch} | Batch {batch_idx} | Loss: {loss.item():.4f})注意上面langevin_sample里步長(zhǎng)和噪聲項(xiàng)的系數(shù)我用了(2 * step_size) ** 0.5這是因?yàn)?Langevin 更新的標(biāo)準(zhǔn)形式里噪聲標(biāo)準(zhǔn)差是 sqrt(2ε)。不同論文里的系數(shù)有差異你自己實(shí)現(xiàn)時(shí)務(wù)必保持一致性。6.5 實(shí)戰(zhàn)配置速查任務(wù)類型能量網(wǎng)絡(luò)規(guī)模采樣步長(zhǎng)采樣步數(shù)訓(xùn)練方法MNIST 級(jí)別2-3 層 MLP0.05-0.110-30CD-1 或 CD-5CIFAR 級(jí)別4-6 層 CNN0.01-0.0320-60PCD 或 Sliced SM分子構(gòu)象圖神經(jīng)網(wǎng)絡(luò)0.01-0.0250-100Score Matching高分辨率圖像深度 ResNet0.005-0.01100擴(kuò)散模型式得分訓(xùn)練7. 常見問題與排錯(cuò)實(shí)錄我在 EBM 實(shí)操中踩過的 8 個(gè)坑7.1 訓(xùn)練發(fā)散能量值一路沖上天現(xiàn)象訓(xùn)練過程中 loss 持續(xù)增大能量值跑到幾千幾萬模型完全崩掉。原因分析負(fù)相采樣的步長(zhǎng)太大或步數(shù)太少導(dǎo)致負(fù)相樣本停留在高能量區(qū)域梯度方向混亂也可能能量網(wǎng)絡(luò)的輸出沒有約束數(shù)值尺度失控。解決辦法把能量網(wǎng)絡(luò)的輸出層換成 tanh 激活把能量限制在 [-1, 1]調(diào)小 Langevin 步長(zhǎng)增加采樣步數(shù)檢查數(shù)據(jù)預(yù)處理是否做了歸一化圖片像素應(yīng)放縮到 [0,1] 而不是 [0,255]。這個(gè)坑我最初踩過好幾次后來養(yǎng)成了一個(gè)習(xí)慣每次訓(xùn)練前先固定幾個(gè)測(cè)試樣本跑一次 Langevin 采樣可視化看能量值是否合理、采樣結(jié)果是否像樣再做正式訓(xùn)練。7.2 模式坍縮生成結(jié)果永遠(yuǎn)只有一種現(xiàn)象采樣生成的樣本全都長(zhǎng)得差不多多樣性極差。原因分析負(fù)相采樣不夠充分或者持久鏈?zhǔn)諗康搅艘粋€(gè)能量盆地?zé)o法逃出。解決辦法使用 PCD 并周期性重置部分鏈子在負(fù)相采樣的初始點(diǎn)中加入噪聲甚至隨機(jī)噪聲增加探索性增大 Langevin 采樣的噪聲系數(shù)等價(jià)于提高采樣溫度調(diào)整模型容量如果容量太大能量地形過擬合到訓(xùn)練樣本的尖峰上很容易坍縮。7.3 采樣質(zhì)量差生成圖像模糊或有偽影現(xiàn)象生成的圖像整體能看但細(xì)節(jié)模糊或出現(xiàn)奇怪偽影。原因分析能量函數(shù)對(duì)局部模式的刻畫不夠精準(zhǔn)負(fù)相采樣步長(zhǎng)太大導(dǎo)致粒子只能到達(dá)盆地的大致區(qū)域無法細(xì)化到高概率的核心區(qū)。解決辦法采樣時(shí)先用大步長(zhǎng)預(yù)熱、再小步長(zhǎng)精煉在能量網(wǎng)絡(luò)上增加跳躍連接ResNet 結(jié)構(gòu)以保留更多原始輸入信息檢查是否有數(shù)據(jù)預(yù)處理環(huán)節(jié)引入的噪聲。7.4 訓(xùn)練很慢Langevin 采樣是瓶頸現(xiàn)象每個(gè) batch 的訓(xùn)練都要做幾十步 Langevin 采樣自動(dòng)微分算梯度代價(jià)極高訓(xùn)練速度比同等規(guī)模的 GAN 慢 10 倍以上。原因分析這是 EBM 方法的固有代價(jià)無解。但可以優(yōu)化減少 Langevin 步數(shù)用 CD-1 代替 CD-10效果下降但速度快很多用得分匹配類方法替代采樣類方法如果場(chǎng)景允許對(duì)能量函數(shù)的梯度做近似計(jì)算比如每隔幾步才重新計(jì)算梯度、中間用常梯度外推在 GPU 上并行處理多個(gè)采樣鏈充分利用批處理能力。7.5 負(fù)相采樣體感像隨機(jī)噪聲能量地形沒成型現(xiàn)象Langevin 采樣出來的負(fù)相樣本看起來完全不像數(shù)據(jù)像純?cè)肼?。原因分析模型還沒有學(xué)到任何有意義的結(jié)構(gòu)能量地形還是平的粒子的運(yùn)動(dòng)完全是隨機(jī)游走。這通常是訓(xùn)練初期的正常現(xiàn)象但如果持續(xù)很多 epoch 還是這樣就要檢查學(xué)習(xí)率是否過大或能量網(wǎng)絡(luò)是否太淺。7.6 CD 與真實(shí)梯度偏差太大現(xiàn)象用 CD 訓(xùn)練收斂的結(jié)果直接用真實(shí)梯度用 bootstrap 近似驗(yàn)證發(fā)現(xiàn)分布差異很大。原因分析CD-k 在 k 很小時(shí)偏差顯著尤其是數(shù)據(jù)分布和模型分布差別大的時(shí)候。Hinton 的解釋是 CD 在優(yōu)化一個(gè)不同的目標(biāo)函數(shù)對(duì)比散度而不是負(fù)對(duì)數(shù)似然不過實(shí)踐中這個(gè)偏差未必有害。如果實(shí)在擔(dān)心建議用 PCD 或 Score Matching 交叉驗(yàn)證一下。7.7 訓(xùn)練后期 loss 一直在震蕩現(xiàn)象loss 在訓(xùn)練后期無法收斂曲線像心臟跳動(dòng)一樣上下震蕩。原因分析負(fù)相采樣鏈的噪聲太大了。Langevin 采樣中的隨機(jī)噪聲在訓(xùn)練后期應(yīng)該逐漸減小模擬退火的效果。解決辦法讓噪聲系數(shù)隨訓(xùn)練進(jìn)度線性衰減比如從 1.0 降到 0.1。這樣做相當(dāng)于逐漸降低采樣溫度讓模型在后期更精準(zhǔn)地貼合數(shù)據(jù)分布。7.8 MNIST 上經(jīng)典 EBM 效果對(duì)照這里附一份我自己復(fù)現(xiàn)實(shí)驗(yàn)時(shí)的典型結(jié)果對(duì)照訓(xùn)練 20 epochbatch size 128Adam 優(yōu)化器方法FID 分?jǐn)?shù)越低越好訓(xùn)練耗時(shí)備注CD-1 MLP65-85約 10 分鐘快速原型驗(yàn)證首選CD-10 MLP45-60約 20 分鐘質(zhì)量與速度折中PCD CNN30-40約 35 分鐘穩(wěn)定的中等質(zhì)量Sliced SM CNN25-35約 40 分鐘無采樣過程但實(shí)現(xiàn)復(fù)雜7.9 排查步驟速查表癥狀第一步檢查第二步檢查第三步嘗試loss 爆炸能量值尺度采樣步長(zhǎng)減小學(xué)習(xí)率生成模糊采樣步數(shù)網(wǎng)絡(luò)深度預(yù)熱精煉方案模式坍縮采樣鏈多樣性溫度系數(shù)PCD 鏈重置訓(xùn)練太慢采樣步數(shù)網(wǎng)絡(luò)規(guī)模Score Matching不收斂學(xué)習(xí)率數(shù)據(jù)歸一化換優(yōu)化器AdamW8. EBM 背后的數(shù)學(xué)為什么對(duì)偶、幾何與拓?fù)湟暯侨绱酥匾?.1 能量模型與指數(shù)族分布的關(guān)系從統(tǒng)計(jì)學(xué)的角度看EBM 本質(zhì)上是在定義一個(gè)指數(shù)族分布P_θ(x) exp(-E_θ(x)) / Z(θ)不過 EBM 與經(jīng)典指數(shù)族分布有一個(gè)關(guān)鍵區(qū)別經(jīng)典指數(shù)族分布的能量函數(shù)是線性參數(shù)化的比如高斯的能量是二次型而 EBM 的能量函數(shù)是非線性、深層參數(shù)化的。這使得 EBM 的表達(dá)能力遠(yuǎn)超經(jīng)典指數(shù)族但也讓配分函數(shù)徹底失去了閉式解。一個(gè)有趣的理論結(jié)果表明EBM 的本質(zhì)是一個(gè)學(xué)習(xí)過的物理勢(shì)能場(chǎng)——它對(duì)數(shù)據(jù)的建模方式不是直接給概率密度一個(gè)解析表達(dá)式而是搭建了一個(gè)經(jīng)典粒子在其中運(yùn)動(dòng)的勢(shì)能地形。這個(gè)視角讓 EBM 區(qū)別于其他生成模型的地方一目了然。8.2 能量面視角下的學(xué)習(xí)動(dòng)態(tài)我特別想強(qiáng)調(diào)一點(diǎn)用能量地形來思考 EBM 的學(xué)習(xí)過程比盯著 loss 曲線要高效得多。當(dāng)訓(xùn)練開始的時(shí)候能量地形基本是平的所有點(diǎn)的能量都差不多。隨著訓(xùn)練的推進(jìn)真實(shí)數(shù)據(jù)點(diǎn)附近開始出現(xiàn)盆地負(fù)相采樣點(diǎn)慢慢被推上高能量區(qū)。一個(gè)有經(jīng)驗(yàn)的研究者應(yīng)該關(guān)注的是盆地之間的間隔是否清晰、每個(gè)盆地的寬度是否合適。盆地太窄能量只在極小鄰域內(nèi)低說明模型過擬合了訓(xùn)練樣本缺乏泛化能力盆地太寬大片區(qū)域能量都低說明模型區(qū)分度不夠。我通常用一個(gè)簡(jiǎn)單的可視化技巧來判斷把測(cè)試集的樣本輸入模型看它們的能量分布。正常的模型測(cè)試集樣本的能量應(yīng)和訓(xùn)練集樣本的能量分布接近如果測(cè)試集樣本能量明顯偏高說明模型把訓(xùn)練樣本背下來了泛化性差。8.3 EBM 與擴(kuò)散模型的親緣關(guān)系近幾年大熱的擴(kuò)散模型Diffusion Model實(shí)際上和 EBM 有著千絲萬縷的聯(lián)系。擴(kuò)散模型在訓(xùn)練時(shí)擬合的是得分函數(shù)即能量函數(shù)的負(fù)梯度采樣時(shí)走的也是 Langevin 動(dòng)力學(xué)式的逐步去噪過程。從這個(gè)角度看擴(kuò)散模型可以理解為一種動(dòng)態(tài) EBM——它不再學(xué)習(xí)一個(gè)固定的能量函數(shù)而是學(xué)習(xí)一系列從噪聲到數(shù)據(jù)的得分函數(shù)。這種親緣關(guān)系讓我對(duì) EBM 的未來比較樂觀擴(kuò)散模型的成功證明了得分 → 采樣這條技術(shù)路線是可行的、可以擴(kuò)展到高維和大數(shù)據(jù)的。EBM 作為一個(gè)更廣義的框架在能量函數(shù)可解釋性和任務(wù)適應(yīng)性上還有不少潛力可挖。9. 寫在最后的實(shí)踐建議9.1 從哪個(gè)玩具任務(wù)入手如果你剛接觸 EBM我建議不要一上來就挑戰(zhàn) CIFAR-10 或高分辨率圖像。先在 MNIST 或 Fashion-MNIST 上把完整流程跑通——數(shù)據(jù)加載、能量網(wǎng)絡(luò)、Langevin 采樣、CD 訓(xùn)練、生成可視化——然后逐步增加難度。我這里有一個(gè)推薦的遞進(jìn)路線二維高斯混合數(shù)據(jù)可視化能量地形和采樣軌跡直觀理解算法行為MNIST驗(yàn)證生成質(zhì)量和模型基本能力Fashion-MNIST 或 CIFAR-10挑戰(zhàn)更復(fù)雜的分布換用 CNN 結(jié)構(gòu)特定領(lǐng)域數(shù)據(jù)比如自己的業(yè)務(wù)數(shù)據(jù)集此時(shí)你已經(jīng)有足夠經(jīng)驗(yàn)做定制化調(diào)整。9.2 關(guān)鍵參數(shù)速查參數(shù)推薦范圍影響學(xué)習(xí)率1e-4 ~ 1e-3太大發(fā)散太小收斂慢Langevin 步長(zhǎng)0.01 ~ 0.1太大不穩(wěn)定太小采樣不足采樣步數(shù)10 ~ 100越多負(fù)相越準(zhǔn)但越慢噪聲系數(shù)與步長(zhǎng)相關(guān)約 sqrt(2ε)決定探索能力批大小64 ~ 256負(fù)相梯度方差與計(jì)算速度的權(quán)衡能量網(wǎng)絡(luò)寬度256 ~ 1024表達(dá)能力與過擬合風(fēng)險(xiǎn)的權(quán)衡需要注意這些參數(shù)不是獨(dú)立的。步長(zhǎng)、噪聲系數(shù)、溫度三者緊密耦合改動(dòng)一個(gè)通常需要同步調(diào)整其他幾個(gè)。9.3 面向工業(yè)應(yīng)用的建議如果你的目標(biāo)是工業(yè)落地我的建議是優(yōu)先試試 EBM 的異常檢測(cè)應(yīng)用而不是純生成任務(wù)。原因很簡(jiǎn)單異常檢測(cè)只需要能量值不需要在高維空間采樣繞過了最困難的配分函數(shù)問題訓(xùn)練穩(wěn)定性和產(chǎn)出價(jià)值都很可觀。另一個(gè)比較現(xiàn)實(shí)的落地方向是與現(xiàn)有模型做組合。比如把 EBM 的能量分?jǐn)?shù)作為其他系統(tǒng)的特征或約束項(xiàng)而不是獨(dú)立使用。我做過一個(gè)推薦系統(tǒng)的實(shí)驗(yàn)用 EBM 給用戶-物品對(duì)打分雖然單獨(dú)用效果一般但把能量分?jǐn)?shù)和協(xié)同過濾的特征拼接后效果有明顯的提升。9.4 最后一點(diǎn)經(jīng)驗(yàn)我個(gè)人在實(shí)際操作中的體會(huì)是EBM 的門檻不在理論上而在工程調(diào)試。你可能會(huì)花很多時(shí)間調(diào)整采樣步長(zhǎng)、噪聲系數(shù)、鏈的數(shù)量卻感覺效果始終差一口氣。這是正常的。EBM 對(duì)超參數(shù)敏感程度遠(yuǎn)高于 VAE 或 GAN但一旦你把整套調(diào)試流程捋順你會(huì)發(fā)現(xiàn)它在處理分布復(fù)雜、模態(tài)多樣的數(shù)據(jù)時(shí)有不可替代的優(yōu)勢(shì)。順著這個(gè)方向繼續(xù)擴(kuò)展后面可以寫概率模型的統(tǒng)計(jì)力學(xué)理論第二篇——比如基于流的模型與最優(yōu)傳輸?shù)囊暯腔蛘邚淖兎滞茢嗟狡骄鶊?chǎng)理論的對(duì)照。每一個(gè)主題挖下去都能挖出不少有意思的東西。