進階最佳化
銳度感知最小化(sharpness-aware minimization,SAM)
兩個網路可以達到完全相同的訓練損失,泛化卻天差地遠。一個由來已久的直覺是:銳利的極小值——稍微推動權重損失就陡升——泛化得比平坦的極小值差,後者在整個鄰域內損失都維持得很低。銳度感知最小化把這個直覺化為明確的目標:不只是找損失低的權重,而是找出「整個周圍區域損失都低」的權重。
形式上,SAM 對 theta 最小化「以 theta 為中心、半徑 rho 的球內最壞情況損失」,這是一個極小極大問題。內層的最大化用單一步梯度上升來近似,得到一個擾動 epsilon-hat,等於 rho 乘上正規化後的梯度;SAM 接著在擾動點 theta 加 epsilon-hat 處求梯度,並用它來更新 theta。因此每一步要花兩次前向加反向傳播,而非一次。
經驗上,SAM 在視覺與語言基準上改善泛化,並增加對標籤雜訊的穩健性,這與「平坦極小值」的說法一致。缺點是梯度成本加倍——部分由高效與自適應變體緩解——以及多出一個超參數 rho,它設定鄰域半徑,而且要讓效果出現就真的得調它。
\min_\theta \max_{\|\epsilon\|\le \rho} L(\theta+\epsilon),\qquad \hat\epsilon = \rho\,\frac{\nabla L(\theta)}{\|\nabla L(\theta)\|}
最小化半徑 rho 球內的最壞損失;用一步上升近似內層最大化。
SAM 在原始權重座標下衡量銳度,因此不具重新參數化不變性;自適應版的 SAM(adaptive SAM)會逐參數重新縮放鄰域以處理這點。
又稱
另見