JOVANA
Explore Library Glossary Getting Started Three Levels Fields How it works Mission
Join the mission
All guides

排程、權重平均與平坦極小值

迭代點停在哪裡,和它如何前進一樣重要。把餘弦排程、權重平均、Lookahead 與銳度感知最小化學成一組工具,瞄準平坦、能泛化的極小值。

排程是演算法的一部分

初學者把學習率排程當成附帶設定。到研究所層級,它是最佳化器身分的一部分:同一條更新規則,搭配固定學習率、階梯衰減或餘弦曲線,會產生性質迥異的軌跡與不同的最終解。訓練早期需要大學習率以快速橫越地景;訓練後期需要小學習率以沉入某個盆地、而不在谷壁間嘎嘎作響。排程就是你在這兩個區制之間插值的方式。

這裡還有一個微妙的動力學故事。現代分析顯示,大批次的全網路訓練常常徘徊在穩定性邊緣(edge of stability)——學習率正好落在最陡曲率方向僅僅勉強穩定之處,因此損失一邊振盪、一邊平均而言仍在下降。完整圖像請見穩定性邊緣;實務上的結論是,一個在末段降低學習率的排程,同時也在降低迭代點被允許佔據的曲率門檻,溫和地把它導向更平坦的區域。

\lambda_{\max}\!\left(\nabla^2 L(\theta)\right) \;\approx\; \frac{2}{\eta}

在稳定性边缘,最大的海森矩阵特征值(即锐度)会徘徊在学习率倒数的两倍附近。

餘弦退火與熱重啟

餘弦退火讓學習率沿著半條餘弦曲線從峰值衰減到接近零。相較於階梯衰減,它在中間學習率上停留更久、並柔和地著陸,經驗上能讓深度網路收斂得更乾淨。它幾乎總是搭配開頭一段短暫的線性暖身(warmup)——數百到數千步——把學習率從接近零逐步拉高,使早期高變異的更新不至於炸掉剛初始化的權重。

\eta_t = \eta_{\min} + \tfrac{1}{2}\left(\eta_{\max}-\eta_{\min}\right)\!\left(1+\cos\!\frac{T_{\mathrm{cur}}}{T_{\max}}\pi\right)

余弦退火沿半余弦曲线将学习率从峰值衰减到接近零。

def lr(t, base, warmup, total):
    if t < warmup:
        return base * t / warmup            # linear warmup
    p = (t - warmup) / (total - warmup)     # 0 -> 1
    return 0.5 * base * (1 + cos(pi * p))    # cosine decay to ~0
先暖身再餘弦衰減——多數 transformer 預訓練的預設排程。

熱重啟(warm restart)變體會週期性地把學習率重置回峰值。每次重啟都把迭代點踢出當前盆地,隨後的衰減讓它沉入一個可能不同、也可能更好的極小值。額外好處:在每次重啟前取的快照構成一個免費的多樣模型集成——這是從單次訓練取得類似集成穩健性的便宜方法。

平均的是權重,而非只是預測

隨機權重平均(SWA)簡單到幾乎難以置信。照常訓練進入一個良好解的區域,接著切換到固定或循環學習率,並平均沿途造訪過的權重。因為 SGD 迭代點在寬盆地的邊緣彈跳、而非坐在谷底,這些點的平均會落在更靠近中心之處——一個更平坦、更寬的解,往往比任何單一快照泛化得更好。

同樣的本能也活在 Lookahead 的內層迴圈裡,它維護兩組權重:一組快速權重由任意基礎最佳化器更新 `k` 步,一組慢速權重則每 `k` 步朝快速權重最後落點走一小步。慢速權重實際上是快速軌跡的滑動平均,能抑制內層最佳化器的雜訊,並且幾乎不需調參就常常改善穩定性。

銳度感知最小化:對鄰域做最佳化

上述三項工具都只是間接瞄準平坦極小值。銳度感知最小化(SAM)則直接對平坦性做最佳化。它最小化的不是某一點的損失,而是當前權重周圍一個小球內的最壞情況損失。每一步包含兩次傳遞:先在梯度方向走一小步上升,找出鄰近損失最高的點,再用那個擾動點上的梯度走下降步。

\min_{\theta}\; \max_{\lVert\epsilon\rVert_2 \le \rho}\; L(\theta + \epsilon)

锐度感知最小化在半径为 ρ 的邻域内最小化最坏情况损失,从而直接偏向平坦极小值。

g  = grad(loss, w)
e  = rho * g / (g.norm() + eps)   # ascent to the worst nearby point
g2 = grad(loss, w + e)            # gradient at the perturbed point
w  = w - lr * g2                  # descend using the sharpness-aware gradient
SAM 每步花兩次前向-反向傳遞,但明確懲罰陡峭的方向。

SAM 直接連結到關於平坦極小值與泛化的文獻:位於寬而平坦盆地的解對權重擾動穩健,往往能更好地遷移到留出資料上。SAM 本質上是把這個先驗折進損失裡。代價是梯度計算翻倍,因此各種變體會把擾動分散到微批次上、或只週期性地施加,以一小部分開銷取回大半好處。

一個框架:對盆地做最佳化,而非對一點

把這一節收攏起來。單純訓練最小化的是最終迭代點的損失。排程、SWA、Lookahead 與 SAM 全都把這個目標往側邊推向周圍區域的幾何。它們有一部分是在兌現梯度下降自身就已展現的同一個現象——見隱式正則化——但它們把偏好平坦、穩健解的傾向從偶然變成明確且可控。

迭代点最终停在哪里才是关键——调度、SWA 与 SAM 共同的目标,是引导它落入宽而平坦的盆地,而非陡峭的极小值。

交互式迭代点沿损失曲面滚动并停入盆地。

  1. 永遠先暖身,再退火——餘弦是安全的預設;若想要免費的快照集成就加上熱重啟。
  2. 在最終階段加上 SWA,幾乎免費地提升泛化、且無推論成本。
  3. 當泛化是瓶頸、且你負擔得起梯度翻倍時,才動用 SAM。