0
| 本文作者: 叢末 百度智客聯(lián)盟 | 2026-07-23 12:03 | 專題:ICML 2019 |

作者丨幸麗娟
編輯丨岑 峰
當(dāng)大模型預(yù)訓(xùn)練還在靠一堆“煉丹術(shù)”般的啟發(fā)式規(guī)則苦苦調(diào)參時(shí),哈佛大學(xué)計(jì)算機(jī)科學(xué)系 Gordon McKay 教授 、Kempner 研究所聯(lián)合主任 Sham Kakade 拋出了一項(xiàng)反直覺的研究:一個(gè)簡單到離譜的二次模型,就能精準(zhǔn)預(yù)測大語言模型在動(dòng)態(tài)預(yù)訓(xùn)練中的絕大部分優(yōu)化效果。
他在 ICML 2026 的特邀演講《How Far Can Quadratics Take Us ? Lessons for LLM Pretraining》中,系統(tǒng)呈現(xiàn)并論證這個(gè)二次模型到底有多強(qiáng)大:
從推導(dǎo)任意時(shí)間學(xué)習(xí)的最優(yōu)策略、精確算出臨界批量大小與動(dòng)態(tài)最優(yōu)學(xué)習(xí)率,到揭示動(dòng)量的真實(shí)作用邊界……所有這些,都可以在這個(gè)樸素、清晰的二次框架下完成。
為什么一個(gè)如此簡單的模型能做到這一切?答案指向了同樣樸素的泰勒定理(Taylor's Theorem)。Sham Kakade 直接在真實(shí)大規(guī)模神經(jīng)網(wǎng)絡(luò)的不同 checkpoint 上做泰勒展開,構(gòu)建出對應(yīng)的二次模型,然后在這個(gè)二次模型上繼續(xù)訓(xùn)練,把損失軌跡與真實(shí)網(wǎng)絡(luò)逐點(diǎn)對比。
結(jié)果驚人:在 LLM 預(yù)訓(xùn)練的關(guān)鍵窗口內(nèi),兩條曲線近乎完全重合。也就是說,泰勒展開中那些被忽略的高階項(xiàng),在這個(gè)區(qū)間里幾乎沒有任何貢獻(xiàn)。這就意味著,真實(shí)網(wǎng)絡(luò)的預(yù)訓(xùn)練動(dòng)態(tài),在關(guān)鍵路徑上本身就是近似二次的。
正因如此,這個(gè)看似簡單的二次模型,才能在任意時(shí)間學(xué)習(xí)、批量大小縮放、動(dòng)量調(diào)優(yōu)等一系列核心問題上,給出既精準(zhǔn)又可驗(yàn)證的預(yù)測。它不是去擬合訓(xùn)練數(shù)據(jù),而是直接還原了預(yù)訓(xùn)練的本質(zhì)。
以下是 Sham Kakade 在 ICML 2026 大會(huì)上發(fā)表的演講精編稿,AI科技評論基于原英文演講內(nèi)容進(jìn)行了不改原意的翻譯編輯:
(注:本文包含大量基于真實(shí)高維張量、泰勒展開與動(dòng)力學(xué)方程的底層推演,字?jǐn)?shù)超過 10,000 字。建議預(yù)留 20 分鐘沉浸式閱讀,或先行收藏)

01
今天我要跟大家聊一聊基礎(chǔ)模型訓(xùn)練。現(xiàn)在,大語言模型需要大量的資源來訓(xùn)練,包括算力、資金和時(shí)間。坦率地說,我認(rèn)為大模型預(yù)訓(xùn)練(Pretraining)這一范式在短期內(nèi)不會(huì)消亡。

首先,很大程度上是因?yàn)槲覀儞碛械幕P停瑢λ邢掠稳蝿?wù)都極其重要。其次,持續(xù)學(xué)習(xí)(Continual Learning)極度困難,而且我確實(shí)認(rèn)為,目前阻力最小的路徑就是,當(dāng)出現(xiàn)新數(shù)據(jù)時(shí),直接對模型進(jìn)行重新訓(xùn)練。不過,這個(gè)內(nèi)部循環(huán)(Inner Loop)的形態(tài)可能會(huì)隨時(shí)間變化,有大量證據(jù)指向這一點(diǎn)。
“預(yù)訓(xùn)練”是我們構(gòu)建基礎(chǔ)模型的基本范式,我們的核心目標(biāo),便是在計(jì)算效率、數(shù)據(jù)利用率以及訓(xùn)練速度上將其推向極致,這就是本次演講的主題:理解我們用于基礎(chǔ)模型訓(xùn)練的一些優(yōu)化原理。

我們可能對其中一些非常基礎(chǔ)的問題尤其感興趣:比如如何判定模型訓(xùn)練的最優(yōu)停止時(shí)機(jī)?如何確定所需的數(shù)據(jù)總規(guī)模?能否設(shè)計(jì)出一種無需事先設(shè)定停止時(shí)間、能夠持續(xù)訓(xùn)練的方法?如何選取最優(yōu)的批量大小(Batch Size)?就串行運(yùn)行時(shí)間(Serial Runtime)而言,能將訓(xùn)練速度推向多快?動(dòng)量(Momentum)如何在大規(guī)模分布式訓(xùn)練中發(fā)揮作用?
這一系列問題都非常基礎(chǔ),而我將進(jìn)一步論證,一個(gè)極其簡單的模型,如何能為我們提供一些相當(dāng)深刻的洞見和思考。
我們首先著眼于一個(gè)特定的簡化模型,探索其能力極限,并依次探討任意時(shí)間學(xué)習(xí)(Anytime Learning)、批量大小縮放(Batch Size Scaling)、動(dòng)量(Momentum)等問題;隨后,我們將轉(zhuǎn)向一些新的研究工作,討論這些做法為何有效,以及用泰勒定理(Taylor's Theorem)進(jìn)行驗(yàn)證。
這里的方法論是:與其試圖去構(gòu)建一個(gè)復(fù)雜的模型,因?yàn)楫?dāng)我們分析它并將其與實(shí)踐對照時(shí),真實(shí)的動(dòng)力學(xué)遠(yuǎn)比這個(gè)復(fù)雜模型要強(qiáng)得多,那不如我們反過來,嘗試在一個(gè)可以精確分析的極簡模型探索能力極限,然后看看:這個(gè)帶有強(qiáng)分析的簡單模型,是否能給我們帶來任何洞見?

02
我們首先從一個(gè)特定的二次模型(Quadratic Model)開始聊。為什么?因?yàn)槲覀兲幚淼钠交瑩p失函數(shù)(Smooth Loss Functions)在局部都是二次的,而優(yōu)化本質(zhì)上就是在求解一系列局部二次問題。我們將嘗試精確求解一個(gè)全局二次模型,并看看這對訓(xùn)練實(shí)踐有什么啟示。
我們的目標(biāo)不是給出模型的上界,而是真正對正在發(fā)生的事情進(jìn)行近乎精確的分析。這個(gè)模型就是線性回歸(Linear Regression)——它看起來樸素簡單,卻具有豐富的學(xué)習(xí)動(dòng)態(tài)結(jié)構(gòu),能揭示很多對實(shí)踐有指導(dǎo)意義的信息。

模型設(shè)定:這是一個(gè)在線模型。每個(gè)時(shí)間步我們觀察一個(gè)樣本 x_t,服從均值為零、協(xié)方差矩陣為 H 的正態(tài)分布;目標(biāo) y_t 由 x_t 的線性函數(shù)加上獨(dú)立噪聲構(gòu)成。我們的目標(biāo)是最小化超額風(fēng)險(xiǎn)(Excess Risk)——R(w) 。
為此,我們采用平方損失(Square Loss),希望通過最小化 w 來降低損失,w 就是參數(shù)估計(jì)值與最優(yōu)估計(jì)值在 H 范數(shù)(H-norm)下的差值。
在線 SGD 更新規(guī)則,計(jì)算當(dāng)前樣本下平方損失關(guān)于參數(shù) w 的梯度,然后以學(xué)習(xí)率(Learning Rate)ηt 為步長,沿著負(fù)梯度方向更新參數(shù),這就是標(biāo)準(zhǔn)的隨機(jī)梯度下降(SGD)更新。
我們的目標(biāo)不僅是運(yùn)行算法,更是精確理解這個(gè)隨機(jī)過程的完整動(dòng)力學(xué)(Exact Dynamics)。為此,我們需要同時(shí)追蹤兩個(gè)關(guān)鍵統(tǒng)計(jì)量:均值(Mean)和協(xié)方差矩陣(Covariance Matrix)。

為了理解這里的動(dòng)態(tài),我們需要掌握該過程的均值以及協(xié)方差矩陣的演化。
均值動(dòng)態(tài):我們對 SGD 更新方程兩邊同時(shí)減去w?,將更新改寫為關(guān)于誤差的遞推形式。代入目標(biāo) y_t = x_t^T w + 噪聲 后,更新規(guī)則呈現(xiàn)為:
w_{t+1} - w* = (I - η x_t x_t^T)(w_t - w*) - η ε_(tái)t x_t
其中前面第一項(xiàng)是乘性噪聲(Multiplicative Noise),因?yàn)檎`差被乘以了一個(gè)隨機(jī)量 x_t;后面第二項(xiàng)是加性噪聲(Additive Noise)。
取期望后,加性噪聲消失,乘性噪聲簡化為 (I - ηH),均值遞推為:
E[w_{t+1} - w*] = (I - ηH) E[w_t - w*]
這意味著均值動(dòng)態(tài)與全批量梯度下降完全一致,只要學(xué)習(xí)率不太大,就會(huì)呈現(xiàn)收縮現(xiàn)象。
協(xié)方差動(dòng)態(tài):這里才是 SGD 與梯度下降真正不同的地方。為了完整理解這個(gè)過程,我們必須追蹤完整的協(xié)方差矩陣:
Σ_t = E[(w_t - w*)(w_t - w*)^T]

我們對誤差遞推方程取外積(即乘以自身的轉(zhuǎn)置),展開各項(xiàng),然后取期望,得到協(xié)方差矩陣的遞推表達(dá)式。其中出現(xiàn)了關(guān)鍵項(xiàng)——四階矩(Fourth Moment): E[x_t x_t^T Σ_t x_t x_t^T](詳見上圖標(biāo)紅公式)。
在高斯假設(shè)下,四階矩可用伊辛-維克定理(Isserlis' Theorem / Wick's Theorem) 化簡,得到 H Σ_t H 和 tr(HΣ_t) · H 兩項(xiàng),其中 tr(HΣ_t) 恰恰就是超額風(fēng)險(xiǎn)。代入?yún)f(xié)方差更新規(guī)則后,得到一個(gè)規(guī)模龐大但形式簡潔的遞推式。
有效噪聲結(jié)構(gòu):合并同類項(xiàng)后,有效噪聲由兩部分組成:
噪聲底限(Noise Floor):η2σ2H,源于標(biāo)簽噪聲,無法消除;
自生噪聲(Self-generated Noise):η2E[R(w_t)]H,基于當(dāng)前誤差,隨收斂而消失。
這意味著:噪聲沿著曲率方向分布,而非各向同性。在訓(xùn)練早期誤差大時(shí),自生噪聲也強(qiáng);隨著模型收斂,它逐漸減弱。
總結(jié)來說,這個(gè)看似簡單的動(dòng)態(tài)系統(tǒng)包含一個(gè)均值遞推(與梯度下降一致)和一個(gè)協(xié)方差遞推(龐大但形式簡潔)。其形式雖簡潔,耦合關(guān)系卻使得分析變得復(fù)雜。過去十年中,大量工作試圖精確理解這一過程,而我們正是站在這個(gè)基礎(chǔ)上向前推進(jìn)。

在某種意義上,我們先將這個(gè)二次模型視為一個(gè)"稻草人模型"(Straw Man)——看看它的預(yù)測是否經(jīng)得起實(shí)踐的檢驗(yàn),或者我們能否推翻它。接下來,我們將通過一系列簡短研究(Vignettes) 來探討它在若干實(shí)際問題上的表現(xiàn)。

03
▎任意時(shí)間學(xué)習(xí)問題:算出“不設(shè)停止時(shí)間”的最優(yōu)跑法
第一個(gè)問題,是關(guān)于任意時(shí)間學(xué)習(xí)(Anytime Learning)的一個(gè)非常基礎(chǔ)的問題。這是與一群非常出色的合作者共同完成的工作,其中 Alex Meterez 也在現(xiàn)場,大家可以向他提問。

現(xiàn)在我們來看大語言模型預(yù)訓(xùn)練(LLM Pretraining)場景下的任意時(shí)間學(xué)習(xí)問題。圖中展示了多條不同的學(xué)習(xí)曲線,橫軸為已處理的 Token 數(shù)量,縱軸為驗(yàn)證損失(Validation Loss)。每條曲線對應(yīng)不同的 Chinchilla 倍數(shù),例如 32 倍(32×)對應(yīng)的模型規(guī)模大約是 1.5 億參數(shù),在圖中以深黑色曲線表示,其損失沿訓(xùn)練過程的變化一目了然。我們可以看到每條曲線的形態(tài)都不同。

這里的關(guān)鍵在于計(jì)算效率(Compute Efficiency)——即訓(xùn)練模型所消耗的浮點(diǎn)運(yùn)算量(FLOPs)。如果我們在訓(xùn)練中途取一個(gè)中間檢查點(diǎn)(Intermediate Checkpoint),那么對于該時(shí)刻對應(yīng)的 Token 數(shù)而言,這個(gè)模型的性能是非常差的。舉個(gè)例子:假設(shè)我們在訓(xùn)練一個(gè) 2 倍 Chinchilla 規(guī)模的模型,如果在 1 倍 Chinchilla 的位置提前停止,并評估其損失,會(huì)發(fā)現(xiàn)它遠(yuǎn)不如從一開始就為 1 倍 Chinchilla 直接訓(xùn)練的模型。
造成這一現(xiàn)象的原因在于,我們采用了余弦衰減(Cosine Decay)作為學(xué)習(xí)率調(diào)度(Learning Rate Schedule),而該調(diào)度是基于預(yù)設(shè)的停止時(shí)間來設(shè)定的。這意味著,雖然模型在預(yù)設(shè)的停止點(diǎn)能達(dá)到一個(gè)較低的損失值,但在中途的任意時(shí)間點(diǎn),其損失表現(xiàn)都比較差。更糟糕的是,如果我們后續(xù)獲得了更多數(shù)據(jù),想要繼續(xù)訓(xùn)練,我們就卡住了,因?yàn)閷W(xué)習(xí)率已經(jīng)衰減到了零。那我們該如何解決這個(gè)問題?
因此,關(guān)于隨時(shí)學(xué)習(xí)(Anytime Learning),一個(gè)很自然的問題便是:如果我們能在訓(xùn)練過程中獲得更多數(shù)據(jù),是否可以在繼續(xù)訓(xùn)練的同時(shí)不損失計(jì)算效率?
同樣重要的是,當(dāng)我們花費(fèi)數(shù)月時(shí)間訓(xùn)練一個(gè)模型時(shí),我們希望能隨時(shí)取出一個(gè)中間檢查點(diǎn)(Intermediate Checkpoint),并準(zhǔn)確評估其在當(dāng)前時(shí)刻的真實(shí)表現(xiàn)。然而在現(xiàn)有框架下,如果我們中途取出這樣一個(gè)檢查點(diǎn),其損失往往很差,并不能反映我們在該時(shí)間點(diǎn)實(shí)際能達(dá)到的最佳性能。那么,我們該如何為一個(gè)未知的“停止時(shí)間”(Stopping Time)來設(shè)計(jì)訓(xùn)練策略呢?

我們的目標(biāo)是什么?我們設(shè)定了一個(gè)相當(dāng)嚴(yán)格的條件:能否在不預(yù)先知道停止時(shí)間的情況下,始終匹配一條特定的包絡(luò)線(Envelope)?所謂余弦包絡(luò)線,就是在余弦衰減調(diào)度下,對于任意給定的停止時(shí)間,我們所能達(dá)到的最佳損失軌跡。圖中我已將這些最佳點(diǎn)連成一條包絡(luò)線。我們希望找到一個(gè)單一的訓(xùn)練過程,能夠在每一個(gè)時(shí)間點(diǎn)上都與這條包絡(luò)線相匹配。
這是一個(gè)非常強(qiáng)的要求,但也具有明確的現(xiàn)實(shí)意義:如果我們獲得了更多數(shù)據(jù),就可以繼續(xù)訓(xùn)練而無需重新開始;同時(shí),在訓(xùn)練過程中,我們可以隨時(shí)評估模型,并確信在那一刻它的表現(xiàn)是接近最優(yōu)的。然而正如你所見,當(dāng)前的實(shí)際軌跡與包絡(luò)線之間存在巨大的差距。所以,問題已經(jīng)很清晰了:我們要去匹配那條包絡(luò)線。
顯然,現(xiàn)實(shí)中的模型非常復(fù)雜——它是一個(gè) Transformer,以某種復(fù)雜的方式被訓(xùn)練。但我們不妨回到我們的“玩具模型”,也就是二次模型,問一問:在這個(gè)模型里,我們對這個(gè)問題有什么認(rèn)識(shí)?這個(gè)問題在那里是否更容易解決?

第一個(gè)結(jié)論是:即使在簡單的二次模型中,任意時(shí)間學(xué)習(xí)本質(zhì)上也是困難的,即使在二維中也很難。 我們有一個(gè)非常簡單的下界(Lower Bound),它針對特定的學(xué)習(xí)率設(shè)定,但我認(rèn)為它能給出正確的直覺——對于任何多項(xiàng)式衰減的學(xué)習(xí)率調(diào)度,基本上都不存在一個(gè)真正的隨時(shí)學(xué)習(xí)方案。也就是說,如果你希望模型在任意時(shí)間點(diǎn)都達(dá)到最優(yōu),在很多時(shí)間點(diǎn)上,你總會(huì)與最優(yōu)結(jié)果差出一個(gè)條件數(shù)(Condition Number)的因子。
直觀上可以這樣理解:在一維情形下,正確的衰減方式是 1/√t;但在二維情形下,存在兩個(gè)不同的時(shí)間尺度,你希望學(xué)習(xí)率同時(shí)按這兩個(gè)尺度衰減,但這是不可能的。在高維中,這種復(fù)雜性只會(huì)進(jìn)一步加劇。
從大量已有工作中我們知道,在有限維度中,任意時(shí)間學(xué)習(xí)的最優(yōu)速率為 dσ2/n,而達(dá)到這一速率的最簡單方法可能是使用恒定學(xué)習(xí)率加平均,這是一個(gè)能在任意時(shí)間點(diǎn)都達(dá)到最優(yōu)速率的方案。如果我們不使用平均,我們也知道,在預(yù)先知道停止時(shí)間的情況下,可以得到接近最優(yōu)的結(jié)果。但關(guān)鍵在于,無論有限維還是無限維,我們在沒有平均的情況下要達(dá)到最優(yōu),都嚴(yán)重依賴于知道停止時(shí)間。
這就是表格中第一列所展示的內(nèi)容:當(dāng)知道停止時(shí)間時(shí),方案并不是任意時(shí)間方案;沒有平均時(shí)我們可以接近最優(yōu);但如果我們加入平均,恒定學(xué)習(xí)率加平均在無限維中同樣適用——只是情況更加微妙。我們可以進(jìn)一步擴(kuò)展這一視角,表明基于過程的特定衰減條件,1/√t 加上平均同樣可以實(shí)現(xiàn)隨時(shí)最優(yōu)。所以這個(gè)問題相當(dāng)微妙。
我不打算深入技術(shù)細(xì)節(jié),但核心結(jié)論是:即使在這樣一個(gè)簡單的“玩具模型”中,我們也能得到相當(dāng)豐富的答案。 我們看到,在沒有平均的情況下很難做到任意時(shí)間最優(yōu);在知道停止時(shí)間的情況下我們可以接近最優(yōu);同時(shí),我們也找到了一個(gè)候選的任意時(shí)間最優(yōu)過程。

當(dāng)然,這只是一個(gè)非常簡單的二次模型,那么平均(Averaging)的方法,為什么能在如此復(fù)雜的神經(jīng)網(wǎng)絡(luò)中起作用呢?我們能否用 1/t 這樣的衰減策略來替代余弦衰減?能否真正匹配那條包絡(luò)線?我們認(rèn)真對待了這些問題,并進(jìn)行了嘗試。
我認(rèn)為一個(gè)重要的背景是:人們在實(shí)踐中確實(shí)經(jīng)常使用平均策略,但他們通常是在已知停止時(shí)間的情況下進(jìn)行的。而我們進(jìn)一步深究的問題是:在不預(yù)先知道停止時(shí)間的情況下,平均策略能否真正起作用,并匹配包絡(luò)線?結(jié)果,它基本做到了。
圖中紅色星號標(biāo)出的是余弦衰減曲線,該模型在 32 倍 Chinchilla 規(guī)模下訓(xùn)練,我們繪制了其對應(yīng)的包絡(luò)線。而恒定學(xué)習(xí)率加平均,以及 1/√t 加平均,均為單次運(yùn)行(Single Run)的結(jié)果。這里我們評估的不是當(dāng)前迭代點(diǎn)(Iterate),而是某個(gè)窗口內(nèi)的移動(dòng)平均(Running Average)。
具體而言,我們使用的是指數(shù)移動(dòng)平均(EMA),本質(zhì)上是對一段窗口內(nèi)的迭代點(diǎn)進(jìn)行平均。這一操作易于實(shí)現(xiàn),且不帶來顯著的計(jì)算開銷。更重要的是,它確實(shí)直接貼合在包絡(luò)線上,而且這一貼合覆蓋了長達(dá) 30 倍 Chinchilla 的范圍。在論文中,我們將結(jié)果進(jìn)一步放大,并展示了相應(yīng)的后悔值(Regret)。
我認(rèn)為這相當(dāng)令人印象深刻:在 30 倍的范圍內(nèi),在這樣一個(gè)非凸問題上,它卻實(shí)實(shí)在在地直接落在包絡(luò)線上。值得注意的細(xì)節(jié)是,我們用于平均的窗口大約占整個(gè)訓(xùn)練運(yùn)行的 5%。隨著訓(xùn)練持續(xù),窗口也隨之拉長——如果窗口太短,效果會(huì)變差;如果太長,曲線則會(huì)向上彎曲。我們驚訝地發(fā)現(xiàn),即便是這樣一個(gè)較大的窗口,依然能產(chǎn)生穩(wěn)健的平均效果。
我了解到,一些基礎(chǔ)模型公司已經(jīng)在使用某種形式的平均策略,但也有些公司并未采用。我認(rèn)為,通過嚴(yán)格調(diào)優(yōu)所有相關(guān)參數(shù),我們實(shí)現(xiàn)的這個(gè)基準(zhǔn)是相當(dāng)扎實(shí)的,并且成功地匹配了包絡(luò)線。這再次與簡單的二次模型給出的預(yù)測一致。
論文中還提到了另一種相關(guān)方案,叫做“預(yù)熱穩(wěn)定衰減”(Warmup Stable Decay),它同樣能命中包絡(luò)線,但它并不算是真正的任意時(shí)間學(xué)習(xí)方案,因?yàn)樗枰褂孟喈?dāng)多的額外樣本來重新訓(xùn)練。而我們的方案則不同——只是一個(gè)單一過程、一次運(yùn)行,就貼合在包絡(luò)線上。

因此,實(shí)踐中的啟示是:通過某種形式的尾部平均(Tail Averaging),我們確實(shí)能夠在每一個(gè)Horizon上(至少在相當(dāng)寬的范圍上)匹配這條包絡(luò)線。我認(rèn)為這是一個(gè)相當(dāng)可靠的實(shí)踐結(jié)論,而且它與理論預(yù)測高度一致。它真正地告訴我們:我們不必在訓(xùn)練開始前就強(qiáng)行設(shè)定一個(gè)停止時(shí)間,因此,當(dāng)后續(xù)獲得更多數(shù)據(jù)時(shí),我們可以無縫地繼續(xù)訓(xùn)練。另一個(gè)有趣的發(fā)現(xiàn)是,這種平均策略似乎在約占 Token 總數(shù) 5% 的窗口上,也依舊有效。
▎批量大小與串行時(shí)間:精算“并行提速”的極限
下一個(gè)問題是:我們的訓(xùn)練任務(wù)什么時(shí)候停止?我們也有與許多才華橫溢的合作者完成的一系列工作,我想特別強(qiáng)了一些更為年輕的成員。

這個(gè)問題本質(zhì)上是:當(dāng)我們按模型規(guī)模進(jìn)行擴(kuò)展時(shí),批量大小(Batch Size)應(yīng)該如何隨之變化。在實(shí)踐中,我們通常依賴一些可預(yù)測的縮放規(guī)律(Scaling Laws),例如損失如何隨已處理的 Token數(shù)量或所投入的計(jì)算量而變化。那么,一個(gè)很自然的問題是:串行運(yùn)行時(shí)間(Serial Runtime)——即訓(xùn)練任務(wù)從開始到完成所需的實(shí)際時(shí)間——是如何隨模型規(guī)模和Token數(shù)量而擴(kuò)展的?因?yàn)槲覀冿@然不希望訓(xùn)練任務(wù)持續(xù)數(shù)個(gè)月之久,我們希望能夠盡快完成。因此,理解串行運(yùn)行時(shí)間的縮放規(guī)律至關(guān)重要。
而這個(gè)問題實(shí)際上可以歸結(jié)為另一個(gè)更基本的問題:臨界批量大小(Critical Batch Size) 是如何隨規(guī)模變化的。臨界批量大小是并行化中最自然的一個(gè)概念——它決定了我們能夠在多大程度上通過增加并行度來縮短串行運(yùn)行時(shí)間。接下來,我將帶領(lǐng)大家通過這張圖來深入理解這個(gè)問題。

那么,問題來了:在不損失計(jì)算效率的前提下,我們究竟能將串行運(yùn)行時(shí)間縮短到多短? 這里的“計(jì)算效率”指的是,我們選定一個(gè)目標(biāo)損失,達(dá)到該目標(biāo)損失所需的總浮點(diǎn)運(yùn)算量(Total FLOPs),而我們希望在任務(wù)完成得更快的同時(shí),不增加總浮點(diǎn)運(yùn)算量。換句話說,在給定目標(biāo)精度下,我們能多快完成任務(wù)?這本質(zhì)上就是批量大小的縮放問題。

先看線性縮放區(qū)域(Linear Scaling Regime)。圖中橫軸表示批量大小,縱軸表示在該批量大小下達(dá)到特定目標(biāo)損失所需的更新步數(shù)(Steps)。我們期望的理想情況是:批量大小翻倍時(shí),步數(shù)減半。如果這一關(guān)系成立,那么總浮點(diǎn)運(yùn)算量保持不變,而串行運(yùn)行時(shí)間減半——因?yàn)槲覀兛梢岳貌⑿谢瘉砑铀佟T诰€性縮放區(qū)域中,隨著我們持續(xù)翻倍批量大小并相應(yīng)減少步數(shù),這一線性關(guān)系能在一定范圍內(nèi)保持成立。
然而,最終這條藍(lán)色曲線會(huì)逐漸偏離完美的線性縮放區(qū)域。我們將臨界批量大小(Critical Batch Size) 定義為曲線開始顯著偏離(例如計(jì)算效率變差 20% 左右)時(shí)的批量大小。這是一個(gè)嚴(yán)格的標(biāo)準(zhǔn),因?yàn)槲覀兪冀K希望保持總浮點(diǎn)運(yùn)算量不變。臨界批量大小,本質(zhì)上就是完美縮放關(guān)系失效的臨界點(diǎn)。
那么,這里的關(guān)鍵問題是:臨界批量大小如何隨 Token 數(shù)量的增加而變化?又如何隨模型規(guī)模的增大而變化? 隨著我們不斷增大模型規(guī)模、使用越來越多的訓(xùn)練數(shù)據(jù),這些因素究竟會(huì)如何影響我們的串行運(yùn)行時(shí)間?
我們?nèi)绾位卮疬@些問題?讓我們回到我們的“玩具模型”——二次模型,看看它會(huì)給出什么樣的預(yù)測。

我們嘗試精確分析該模型的動(dòng)態(tài)過程,并理解其含義。當(dāng)我們將批量大小納入模型時(shí),修改精確動(dòng)態(tài)并不復(fù)雜——忽略不等式(實(shí)際上它是只差一個(gè)常數(shù)),均值動(dòng)態(tài)并不隨批量大小變化,但有效方差(Effective Variance)會(huì)按因子 B 縮小。現(xiàn)在我們需要理解的是,當(dāng)規(guī)模擴(kuò)大時(shí),如何刻畫這些動(dòng)態(tài)。需要特別強(qiáng)調(diào)的是,要真正精確理解這一過程,即使最終我們只關(guān)心損失值,我們也必須追蹤誤差在完整協(xié)方差矩陣中的傳遞——如果不追蹤整個(gè)系統(tǒng),就無法精確地得到結(jié)果。這就是我們的目標(biāo)。
理論會(huì)怎么說?這個(gè)問題既微妙又有趣。我們先固定模型規(guī)模,問:臨界批量大小如何隨 Token 數(shù)的變化而縮放?

我們的預(yù)測是:在有限維設(shè)定下,隨著 Token 數(shù)的增加,臨界批量大小最終將趨于一個(gè)與 Token 數(shù)無關(guān)的常數(shù)。直覺是這樣的:如果你在訓(xùn)練一開始就使用非常大的批量,實(shí)際上并沒有什么好處——因?yàn)楫?dāng)參數(shù)離最優(yōu)解還很遠(yuǎn)時(shí),我們?yōu)槭裁匆烟荻裙烙?jì)得那么精確呢?在遠(yuǎn)離最優(yōu)解的區(qū)域,使用過大的批量只是在浪費(fèi)樣本。當(dāng)然,如果完全不考慮計(jì)算效率,無限大的批量總是好的;但當(dāng)我們關(guān)心計(jì)算效率時(shí),如果隨著樣本數(shù)增加而不斷增大批量大小,本質(zhì)上就是在訓(xùn)練初期浪費(fèi)了大量樣本——因?yàn)槟菚r(shí)候根本不需要那么精確的梯度。
然而,微妙之處在于,在無限維設(shè)定下,上述結(jié)論不再成立。我認(rèn)為我們實(shí)際所處的實(shí)踐區(qū)間更接近無限維設(shè)定,因?yàn)槲覀兊哪P鸵?guī)模通常與 Token 數(shù)同量級,甚至更大——這正是 Chinchilla 縮放規(guī)律所表明的。因此,在無限維設(shè)定下,根據(jù)譜(Spectrum)的具體條件,臨界批量大小實(shí)際上會(huì)隨 Token 數(shù)按某個(gè)小于 1 的冪次縮放。
直覺上可以從偏差-方差權(quán)衡(Bias-Variance Tradeoff) 來理解:對于更長的訓(xùn)練運(yùn)行,實(shí)際上希望使用更大的批量大小。
在無限維情形下,當(dāng)我們采用平均策略并設(shè)定學(xué)習(xí)率時(shí),特定的偏差-方差權(quán)衡會(huì)使得:運(yùn)行越長,你越傾向于使用更小的學(xué)習(xí)率。這主要?dú)w因于過程中復(fù)雜的乘性噪聲結(jié)構(gòu)。進(jìn)一步說,由于這種偏差-方差權(quán)衡,在非常長的運(yùn)行中,你實(shí)際上會(huì)把初始學(xué)習(xí)率設(shè)為停止時(shí)間的函數(shù);類似地,你也希望把批量大小設(shè)為停止時(shí)間的函數(shù)。因此,這是一個(gè)具體的、可驗(yàn)證的預(yù)測:我們預(yù)期,至少在高維情形下,臨界批量大小會(huì)隨Token數(shù)的增加而縮放。
那么,我們對臨界批量大小如何隨模型規(guī)模縮放的預(yù)測是什么?這是模型規(guī)模縮放的問題。在這里,我們不完全依賴二次理論,而是借助一組關(guān)于平均場漸近(Mean Field Asymptotics) 的結(jié)果。這一系列工作表明,在特定的縮放條件下,當(dāng)取某個(gè)平均場極限時(shí),訓(xùn)練動(dòng)態(tài)會(huì)收斂到一個(gè)定義良好的漸近極限——而該極限不是模型規(guī)模 d 的函數(shù),因?yàn)殡S著我們縮放模型規(guī)模,極限行為本身并不依賴它。因此,平均場推理給出的結(jié)論是:臨界批量大小不應(yīng)該依賴于模型規(guī)模 d。
當(dāng)然,實(shí)際上我并不認(rèn)為我們真的處于那種所有訓(xùn)練動(dòng)態(tài)都收斂的平均場極限——但那也并非必要條件。看起來,某些標(biāo)量量(而非完整的高維學(xué)習(xí)動(dòng)態(tài))可能在整體極限還未達(dá)到之前就已經(jīng)進(jìn)入了平臺(tái)期,例如學(xué)習(xí)率遷移(Learning Rate Transfer) 現(xiàn)象似乎確實(shí)成立。從理論角度看,臨界批量大小很可能在訓(xùn)練動(dòng)態(tài)的整體極限到達(dá)之前,就已經(jīng)穩(wěn)定在某一個(gè)極限值附近,這似乎是合理的。
因此,從理論預(yù)測來看,我們的結(jié)論是:第一,臨界批量大小與Token數(shù)之間存在強(qiáng)縮放關(guān)系——運(yùn)行越長,你應(yīng)當(dāng)使用越大的批量;第二,臨界批量大小與模型規(guī)模之間的縮放關(guān)系較弱。

另一個(gè)問題是:這些理論預(yù)測能否得到驗(yàn)證?顯然,我們使用的是簡化模型,但它給出了很強(qiáng)的量化預(yù)測。我們對此進(jìn)行了實(shí)驗(yàn)。需要強(qiáng)調(diào)的是,這些實(shí)驗(yàn)都是經(jīng)過精心調(diào)優(yōu)的運(yùn)行——因?yàn)槊看芜\(yùn)行我們都在為給定的模型規(guī)模和Token預(yù)算尋找最優(yōu)的批量大小。
實(shí)踐中常見的做法是,在 Chinchilla 縮放(即模型規(guī)模增長時(shí),Token 數(shù)按約 20 倍模型規(guī)模同步增長)下,觀察臨界批量大小隨模型規(guī)模的縮放。但在這種設(shè)置下,兩個(gè)因素被混淆了:當(dāng)模型規(guī)模增長時(shí),Token數(shù)也在同時(shí)增長。因此,在 Chinchilla 縮放下,我們確實(shí)會(huì)看到臨界批量大小隨模型規(guī)模強(qiáng)烈增長,但這并不能區(qū)分出哪個(gè)因素才是真正的驅(qū)動(dòng)因素。
因此,我們需要解耦這兩個(gè)因素。首先,當(dāng)我們固定模型規(guī)模,單獨(dú)考察臨界批量大小隨 Token 數(shù) n 的變化時(shí),我們得到的指數(shù)與理論預(yù)測高度一致——即存在強(qiáng)依賴關(guān)系:固定模型規(guī)模,訓(xùn)練運(yùn)行越長,臨界批量大小隨 Token 數(shù)的增長越顯著。
接下來,我們做另一個(gè)方向的縮放:固定 Token 數(shù),調(diào)大模型規(guī)模。我們再次進(jìn)行細(xì)致的掃描實(shí)驗(yàn),試圖理解任務(wù)是否能更快完成——即考察臨界批量大小隨模型規(guī)模的縮放。結(jié)果再次與理論一致:臨界批量大小對模型規(guī)模的依賴非常弱。因此,通過真正解耦這兩個(gè)因素,實(shí)驗(yàn)結(jié)果與我們的預(yù)測吻合良好。

現(xiàn)在我們可以進(jìn)入一個(gè)更細(xì)微的問題。基于對縮放規(guī)律的理解,我們的第一步是將長串行過程通過臨界批量大小轉(zhuǎn)化為更多的并行性和更低的串行運(yùn)行時(shí)間。但如果我們正在采用學(xué)習(xí)率衰減調(diào)度,那么我們能否設(shè)計(jì)一種“批量提升(Batch Ramp)”策略——即不衰減學(xué)習(xí)率,而是通過某種方式在訓(xùn)練過程中逐步增大批量大小?我們能否用批量增大來替代學(xué)習(xí)率衰減,從而匹配原始衰減過程(例如余弦衰減)?再次,我們可以借助二次模型來理解其中的機(jī)制。

我們的目標(biāo)是以實(shí)際相關(guān)的方式實(shí)現(xiàn)這一策略,因此我們希望為 Adam 設(shè)計(jì)批量提升方案。對于 SGD,已有一些相關(guān)的工作,我們可以給出一個(gè)清晰的結(jié)論:如果我們連續(xù)執(zhí)行兩步更新,學(xué)習(xí)率分別為 η 和 β,在特定的方差異分母占優(yōu)的區(qū)域,我們可以通過將學(xué)習(xí)率加倍并將批量大小加倍,使得兩步更新等價(jià)于一步更新。
而對于我們更關(guān)心的 Adam 設(shè)置,我們可以在歸一化梯度(Normalized Gradient)設(shè)定下進(jìn)行分析。結(jié)果表明,在特定的方差主導(dǎo)區(qū)域,該等價(jià)關(guān)系同樣成立——我們可以在實(shí)踐中驗(yàn)證這一點(diǎn)。其含義是:在該區(qū)域,我們應(yīng)當(dāng)將學(xué)習(xí)率乘以 √2,并將批量大小加倍。換句話說,在實(shí)踐中,如果我們原本打算將 Adam 的學(xué)習(xí)率減半,那么我們應(yīng)該改為將學(xué)習(xí)率除以 √2,同時(shí)將批量大小加倍。
通過這樣做,我們可以在訓(xùn)練過程中更激進(jìn)地增大批量大小,從而有效減少串行運(yùn)行時(shí)間。我們的主張相當(dāng)強(qiáng)——這些隨機(jī)過程應(yīng)當(dāng)真正匹配,因此學(xué)習(xí)曲線應(yīng)當(dāng)是對齊的。我們在多種模型規(guī)模和不同批量大小的設(shè)置下都進(jìn)行了驗(yàn)證。
我們的關(guān)鍵主張是:這些隨機(jī)過程確實(shí)對齊了。我們按已見 token 數(shù)量來度量時(shí)間——我們關(guān)心的是加速墻鐘時(shí)間(Wall-clock Time),而不同過程之間的匹配方式,正是按已見 token 數(shù)量來計(jì)時(shí)的。
Seesaw算法的核心操作是:每當(dāng)原始調(diào)度將學(xué)習(xí)率乘以 √2 時(shí),我們改為將學(xué)習(xí)率乘以 2,并將批量大小加倍。我們按 token 數(shù)重新調(diào)整時(shí)間軸,理論上這一過程應(yīng)與原始余弦調(diào)度曲線對齊。如左圖所示,在大尺度上兩者幾乎完全重合,與理論預(yù)測一致;即便放大觀察,雖然能看到細(xì)微差距,但那是在極度放大的尺度下,兩條曲線實(shí)際上幾乎重疊在一起。
更有意思的是,如果看右圖(以步數(shù)為橫軸),我們通過這種批量提升(Batch Ramp),也就是我們提出的Seesaw Procedure,大約能提前 35% 完成訓(xùn)練。而這個(gè) 35% 實(shí)際上是理論上的最優(yōu)值。因此,我們可以在不損失計(jì)算效率的前提下,通過增大批量大小,實(shí)現(xiàn) 35% 的串行運(yùn)行時(shí)間節(jié)省——總浮點(diǎn)運(yùn)算量保持不變,因?yàn)閮蓷l曲線在達(dá)到相同驗(yàn)證損失時(shí),我們的方案加速了 35%。這與二次模型理論完全一致:當(dāng)我們按已見 token 數(shù)計(jì)時(shí)時(shí),兩條曲線直接對齊。

基于上述分析,可以得到三個(gè)結(jié)論:
實(shí)踐中我們?yōu)榇羞\(yùn)行時(shí)間建立的縮放規(guī)律,應(yīng)當(dāng)真正解耦 token 數(shù)量和模型規(guī)模——這與理論一致;
臨界批量大小強(qiáng)烈依賴于 token 預(yù)算,而對模型規(guī)模的依賴非常弱;
對于 Adam,我們提出了批量提升過程,可以在余弦衰減下匹配訓(xùn)練動(dòng)態(tài),并實(shí)現(xiàn) 35% 的串行運(yùn)行時(shí)間提升,做到曲線對齊。
到目前為止,我們看到的是理論曲線與實(shí)際曲線直接對齊,而非松散的上下界關(guān)系——這些學(xué)習(xí)曲線是實(shí)實(shí)在在地吻合在一起的。
▎動(dòng)量:必須同時(shí)調(diào)優(yōu)動(dòng)量與批量大小
最后一個(gè)討論,是關(guān)于動(dòng)量(Momentum)。

動(dòng)量在實(shí)踐中到底給我們帶來了什么?它是我們?yōu)閿?shù)不多的能真正加速優(yōu)化的工具之一,在各類場景下都廣泛有效。但在這種隨機(jī)設(shè)定下,它究竟帶來了什么好處?我們該如何隨批量大小調(diào)優(yōu)它?它如何幫助我們提高計(jì)算效率?我們可以再次求助于二次模型,看看它給出什么答案,然后到實(shí)踐中去驗(yàn)證。

我們想理解動(dòng)量如何影響計(jì)算效率(達(dá)到給定精度所需的總浮點(diǎn)運(yùn)算量)和串行運(yùn)行時(shí)間,在確定性情況下,我們有非常漂亮的理論結(jié)果:重球算法(Heavy Ball Algorithm,Polyak) 和 內(nèi)斯特羅夫加速算法(Nesterov's Acceleration) 能夠?qū)⒋羞\(yùn)行時(shí)間提升 √κ 倍,其中 κ 是條件數(shù)(Condition Number),即最大與最小特征值之比。這是一個(gè)非常優(yōu)美的結(jié)果,伴隨著許多優(yōu)雅的證明。
然而,對于 SGD,當(dāng)批量大小為 1 時(shí),早期結(jié)果表明這種加速實(shí)際上不再成立。你可以證明,重球和內(nèi)斯特羅夫在批量大小為 1 時(shí),在計(jì)算效率上沒有任何提升。這有點(diǎn)令人失望,因?yàn)閯?dòng)量是我們?yōu)閿?shù)不多的加速工具之一,而實(shí)踐中人們卻一直在使用它。為什么批量大小為 1 時(shí)如此微妙?
關(guān)鍵在于,這是一個(gè)不依賴于加性噪聲 σ2 的結(jié)果——實(shí)際上這里 σ2=0,它適用于一致線性系統(tǒng)(Consistent Linear System)。加速失效的根本原因,正是由于那種特殊的乘性噪聲結(jié)構(gòu)所導(dǎo)致的耦合動(dòng)態(tài)。這正是我們嚴(yán)重依賴“玩具模型”動(dòng)態(tài)的地方,它給出了一個(gè)強(qiáng)烈的預(yù)測:動(dòng)量至少在批量大小為 1 時(shí),不能提升計(jì)算效率。
不過,有一些工作(包括我自己參與的)研究過一種雙時(shí)間尺度算法(Two-timescale Algorithm),可以在小批量下獲得提升,這里我就不展開了。
這就引出了一個(gè)關(guān)于縮放的問題:顯然我們不想在實(shí)際中運(yùn)行批量大小為 1。我們知道,當(dāng)從批量大小 1 逐漸增大到全批量(Full Batch)時(shí),我們可以問:加速效應(yīng)何時(shí)會(huì)重新出現(xiàn)? 因?yàn)楫?dāng)批量大小趨于無窮時(shí),它確實(shí)減少了所需的步數(shù)。但我們也關(guān)心隨之而來的計(jì)算效率代價(jià)。因此,問題是:隨著批量大小的增加,動(dòng)量究竟會(huì)發(fā)生什么變化? 我們在“玩具模型”和實(shí)踐中都問了這個(gè)問題。
我們得到了一個(gè)相當(dāng)有趣的結(jié)論:
重球算法在任何批量大小下,計(jì)算效率都不優(yōu)于 SGD——也就是說,它從不減少達(dá)到目標(biāo)精度所需的總浮點(diǎn)運(yùn)算量。但它確實(shí)能通過 κ 因子改善串行運(yùn)行時(shí)間。這是正式的結(jié)論,詳見論文;其本質(zhì)是,動(dòng)量允許你將臨界批量大小(Critical Batch Size)增大 κ 倍,但總浮點(diǎn)運(yùn)算量仍與 SGD 相同。
稍微令人失望的是,那些優(yōu)雅的雙時(shí)間尺度算法在增大批量大小時(shí)似乎也沒有帶來太多額外好處。如果你愿意串行運(yùn)行很長時(shí)間,它們確實(shí)能在浮點(diǎn)運(yùn)算量上帶來一些顯著提升;但對于這些雙時(shí)間尺度算法,當(dāng)你試圖讓任務(wù)更快完成時(shí),最終又會(huì)回到?jīng)]有它們時(shí)的狀態(tài)——這有點(diǎn)令人失望。其中的細(xì)節(jié)較為微妙,具體可參考論文。

另一個(gè)問題是:理論雖然漂亮,但它能否經(jīng)得起實(shí)踐的驗(yàn)證?至少在合成實(shí)驗(yàn)中,結(jié)果是成立的。如上圖所示,橫軸是批量大小,縱軸是達(dá)到目標(biāo)損失所需的步數(shù),其中綠色曲線對應(yīng) SGD。我們發(fā)現(xiàn),當(dāng)調(diào)優(yōu)幾種不同的動(dòng)量算法時(shí),它們在給定的批量大小下所需的步數(shù)是相同的——至少在小批量區(qū)域如此。然而,動(dòng)量算法將完美縮放(Perfect Scaling)的窗口推得更遠(yuǎn)了。
這意味著什么?我們獲得了相同的浮點(diǎn)運(yùn)算量(因?yàn)槲覀冏裱昝揽s放律:批量翻倍,步數(shù)減半),但我們將縮放窗口向外擴(kuò)展了,從而改善了串行運(yùn)行時(shí)間,而計(jì)算效率并未提升。我們目前只在合成實(shí)驗(yàn)中驗(yàn)證了這一結(jié)論。
Shallow Way 等人一篇很出色的早期實(shí)證論文表明,這一結(jié)論在神經(jīng)網(wǎng)絡(luò)中同樣成立。他們確實(shí)運(yùn)行了這些實(shí)驗(yàn),并展示了與上述一致的趨勢:SGD 和動(dòng)量在計(jì)算效率上完全重合,但臨界批量大小被向外推移了。因此,這與理論在神經(jīng)網(wǎng)絡(luò)、復(fù)雜非凸系統(tǒng)上的預(yù)測直接吻合:浮點(diǎn)運(yùn)算量(FLOPs)沒有提升,但確實(shí)改善了臨界批量大小。

這里的關(guān)鍵結(jié)論是:當(dāng)我們考慮動(dòng)量并對其進(jìn)行調(diào)優(yōu)時(shí),不應(yīng)當(dāng)在固定批量大小下調(diào)優(yōu)動(dòng)量,因?yàn)槟菢涌床坏饺魏翁嵘N覀冃枰?/span>同時(shí)調(diào)優(yōu)動(dòng)量和批量大小,模型才能獲得真正的改善。如果不這樣做,你可能會(huì)覺得動(dòng)量沒什么用,但實(shí)際上,它對獲得串行運(yùn)行時(shí)間的提升至關(guān)重要。
因此,這是一個(gè)明確的結(jié)論,與簡單的高斯二次模型完全一致。總之,將批量大小與動(dòng)量一起調(diào)優(yōu),可以獲得更優(yōu)的串行運(yùn)行時(shí)間。

04
在剩下的時(shí)間里,我將討論一項(xiàng)非常新的工作,希望幾天后能發(fā)布在 arXiv 上。合作者還是之前那批,但這次新增了 Alex Damian,他在邊緣穩(wěn)定性(Edge Stability)方面做了很多優(yōu)秀的工作。

現(xiàn)在很多人會(huì)說:“我們不想再考慮這些簡單的二次模型了,它們對復(fù)雜模型能有什么啟發(fā)?”
但我認(rèn)為新一代研究者恰恰要認(rèn)真對待泰勒定理(Taylor's Theorem),并真正深入挖掘。當(dāng) Alex 加入后,我們團(tuán)隊(duì)的想法更傾向于:“讓我們直接展開泰勒定理,看看能否把常數(shù)都研究透,真正理解整個(gè)過程。”

我覺得這項(xiàng)工作相當(dāng)深入,因?yàn)槲覀冋嬲谧穯栆粋€(gè)核心問題:為什么二次模型在這里效果這么好? 如前所述,一種理解是,優(yōu)化本質(zhì)上只是一系列局部回歸(Local Regressions)。
但如果我們試圖認(rèn)真對待這個(gè)“局部回歸”的觀點(diǎn),并將其擴(kuò)展為全局模型呢?
我們的想法是:與其用 SGD 進(jìn)行一系列局部線性化訓(xùn)練,不如直接采用泰勒定理,在某個(gè) Checkpoint 展開,然后訓(xùn)練一個(gè)全局二次模型。
具體什么意思呢?如上圖所示,黑色曲線是一個(gè)模型( 1.5 億參數(shù))的訓(xùn)練運(yùn)行軌跡。我們在訓(xùn)練過程中取不同的檢查點(diǎn)(Checkpoints),然后在每個(gè) Checkpoint 處做泰勒展開,然后假設(shè)泰勒定理是精確成立的,并在該展開后的模型上進(jìn)行訓(xùn)練。
現(xiàn)在我們有了一個(gè)二次模型,我們在它上面訓(xùn)練,然后問:在這個(gè)二次模型上訓(xùn)練是否與原始過程匹配?有兩種方式可以做到這一點(diǎn)。
設(shè) f_θ 是神經(jīng)網(wǎng)絡(luò),它將輸入映射到目標(biāo)。我們對神經(jīng)網(wǎng)絡(luò)做線性化:
f_θ(x) = f_θ?(x) + ?f_θ?(x)?(θ ? θ?)
有兩種方式將其代入損失函數(shù):
方法一:GaussNewton - Prox-linear 方法。直接將線性化后的神經(jīng)網(wǎng)絡(luò)代入損失函數(shù)。這樣得到的模型在形式上類似于邏輯回歸。
方法二:直接對損失函數(shù)做泰勒展開。不對網(wǎng)絡(luò)本身做線性化,而是直接對損失函數(shù)L(θ) 在 θ?處展開到二階,得到包含梯度(一階項(xiàng))和 Hessian 矩陣(二階項(xiàng))的二次模型。
我們在不同 Checkpoint 進(jìn)行泰勒展開訓(xùn)練。這相當(dāng)于把泰勒模型這個(gè)“近似模型”當(dāng)作真實(shí)模型,在這個(gè)近似模型上繼續(xù)做優(yōu)化訓(xùn)練——而不是像 SGD 那樣,每步都回到原始網(wǎng)絡(luò)去做一系列局部更新。我們想知道:如果我們在全局上求解這個(gè)近似二次模型,它的軌跡能否追蹤原始訓(xùn)練過程?
結(jié)果相當(dāng)驚人。我們比較兩個(gè)過程:
過程一:取 θ?,忽略線性化,像往常一樣用余弦衰減訓(xùn)練——這就是黑色曲線;
過程二:在θ? 處,開始在線性化的模型上訓(xùn)練(兩種方式之一)。
現(xiàn)在我們可以直接比較這兩條損失曲線,看它們是否對齊。結(jié)果發(fā)現(xiàn),它們在 30%到50% 的訓(xùn)練窗口內(nèi)對齊得非常出色。具體來說,一旦超過大約 30% 的訓(xùn)練進(jìn)度(如放大圖所示),從每個(gè) Checkpoint 開始的線性化模型訓(xùn)練都能匹配黑色曲線;雖然最終性能會(huì)逐漸偏離實(shí)際過程,但在 30% 到 50% 左右,它們幾乎完全重合。

上圖左側(cè)所示:藍(lán)色曲線代表 Gaussian 方法,紅色曲線表示的是真實(shí) Hessian 二次模型。在大約 50% 的位置,它們幾乎重疊。在約占整個(gè)訓(xùn)練運(yùn)行約20%的區(qū)間內(nèi),在非凸優(yōu)化上訓(xùn)練的過程,其損失軌跡與在該區(qū)間內(nèi)的二次模型訓(xùn)練幾乎一致。
這個(gè)結(jié)果給了我們極大的支持,說明泰勒定理在訓(xùn)練動(dòng)態(tài)的相當(dāng)大范圍內(nèi)能提供有效的近似。這讓我們非常驚訝——這種近似一致性竟然能保持這么久。這篇論文很快就會(huì)發(fā)布。
最后,我再用一個(gè)”彩蛋“來結(jié)束這個(gè)部分。它與訓(xùn)練動(dòng)態(tài)不直接相關(guān),但結(jié)果相當(dāng)驚人。我們研究了大型神經(jīng)網(wǎng)絡(luò)的完整特征譜(Eigenspectrum)(上圖右側(cè)所示)。
以 1.5 億參數(shù)的神經(jīng)網(wǎng)絡(luò)為例,如果我們想做主成分分析(PCA)并觀察其完整的特征譜,面臨的挑戰(zhàn)是協(xié)方差矩陣規(guī)模高達(dá) 1.5 億 × 1.5 億,極其龐大。然而,數(shù)值線性代數(shù)方法,例如 Lanczos 求積法(Lanczos Quadrature),就能讓計(jì)算完整的特征譜成為可能。我們在 Kempner 上進(jìn)行了大規(guī)模計(jì)算,成功獲得了完整的譜密度。
其中有一些有趣的發(fā)現(xiàn):
譜密度大致呈現(xiàn)兩個(gè)冪律(Power Laws):第一個(gè)冪律延伸至詞匯表大小(Vocab Size)附近,在那里可以看到一個(gè)小的下降(圖中垂直虛線標(biāo)記處),之后則過渡到第二個(gè)冪律;
我們還比較了 Gauss-Newton 譜與 Hessian 譜,發(fā)現(xiàn)它們在詞匯表維度之前高度一致,之后才開始出現(xiàn)分叉;
進(jìn)一步分析特征向量的分布,發(fā)現(xiàn)它們呈現(xiàn)出明顯的結(jié)構(gòu)性分離。通過將特征向量投影到不同的網(wǎng)絡(luò)模塊——例如頂層線性層、注意力層等,我們發(fā)現(xiàn),前詞匯維度主要由網(wǎng)絡(luò)的頂層所捕獲,而超出該區(qū)域后,分布模式發(fā)生了顯著變化。
更多細(xì)節(jié)敬請期待我們的論文。我認(rèn)為這項(xiàng)研究令人印象深刻之處在于,我們直接借助數(shù)值量化的線性代數(shù)工具,去深入理解這些大規(guī)模網(wǎng)絡(luò)底層結(jié)構(gòu)的基本特征。

現(xiàn)在,我來做一下總結(jié)。我們始終堅(jiān)持一個(gè)核心視角:預(yù)訓(xùn)練本質(zhì)上是一系列短鏈?zhǔn)降木植炕貧w(Local Regressions),而我們的目標(biāo)是進(jìn)一步追問——如果我們把它當(dāng)作一個(gè)全局回歸模型來理解,它能為我們帶來什么樣的洞見?
事實(shí)證明,它帶來的并不僅僅是粗略的上界或松散的趨勢,而是具體、可驗(yàn)證、可操作的預(yù)測,這些能直接指導(dǎo)我們應(yīng)該如何設(shè)置縮放規(guī)律、如何設(shè)計(jì)訓(xùn)練策略。
我認(rèn)為這項(xiàng)工作中最精彩的部分在于:這種極為簡化的方法,到底能帶我們走多遠(yuǎn)? 而我們得到的答案是十分明確的。

我想強(qiáng)調(diào)的是,對二次模型能力的論證并非是依賴一堆松散的數(shù)學(xué)上界或漸近趨勢。我們所呈現(xiàn)的證據(jù)是直觀的、肉眼可見的:在多個(gè)截然不同的場景下,二次模型的理論預(yù)測與實(shí)際訓(xùn)練曲線都實(shí)現(xiàn)了直接重合。具體來說:
任意時(shí)間學(xué)習(xí)的結(jié)果直接貼合包絡(luò)線;
SeeSaw 批量提升調(diào)度中,二次模型算出了學(xué)習(xí)率衰減與批量大小增長之間的精確等價(jià)關(guān)系,且 SeeSaw 的學(xué)習(xí)曲線與基準(zhǔn)曲線精確對齊;
動(dòng)量的問題上,二次模型給出了一個(gè)相當(dāng)嚴(yán)格的結(jié)論:無論批量大小如何調(diào)整,重球算法在任何情況下都不能提升計(jì)算效率,這同樣與理論預(yù)測完全吻合;
最后,在泰勒定理的驗(yàn)證中,我們對真實(shí)的大規(guī)模神經(jīng)網(wǎng)絡(luò)在不同檢查點(diǎn)處做線性化處理,發(fā)現(xiàn)泰勒定理不是只在展開點(diǎn)附近有效,而是在相當(dāng)寬的窗口內(nèi)都能成立。
感謝各位的聆聽。本次演講的內(nèi)容,是多年來與許多杰出合作者共同努力的成果。Alex Meterez將在本周五的研討會(huì)上對更多技術(shù)細(xì)節(jié)進(jìn)行報(bào)告,我屆時(shí)也會(huì)進(jìn)一步討論關(guān)于完整特征譜的最新工作。如果想深入了解,歡迎參加研討會(huì)。 雷峰網(wǎng)雷峰網(wǎng)(公眾號:雷峰網(wǎng))雷峰網(wǎng)
一個(gè)人讀論文太孤單,一群人刷頂會(huì)才好玩。
ICML 2026 召開在即,我們正在召集一波含金量極高的 AI 研究者。群內(nèi)主打實(shí)時(shí)論文跟蹤與硬核技術(shù)探討,拒絕灌水。
? 進(jìn)群傳送門: 掃碼進(jìn)群或添加微信Vin_Vivid,備注:論文群 + 關(guān)注的 AI 方向。

搞科研/搞技術(shù),信息差很重要。
來,一起快人一步!
上車,帶你看遍全球 AI 頂會(huì)精華
可獨(dú)家暢覽:
專家演講PPT
大會(huì)報(bào)告全文
熱門論文解讀
學(xué)術(shù)新星訪談

掃描上方二維碼
或點(diǎn)擊「閱讀原文」關(guān)注專區(qū)。
雷峰網(wǎng)原創(chuàng)文章,未經(jīng)授權(quán)禁止轉(zhuǎn)載。詳情見轉(zhuǎn)載須知。
本專題其他文章