1. ニュートン法(Newton's Method)
1.1 損失の二次近似
損失 $\mathcal L$ を最小化するパラメータ $\theta$ を求める。
\theta\in\mathbb R^P,\qquad g=\nabla_\theta\mathcal L(\theta),\qquad H=\nabla_\theta^2\mathcal L(\theta),\qquad\mathcal L\in C^2
\mathcal L(\theta+d)\approx\mathcal L(\theta)+g^\top d+\frac12d^\top Hd
\begin{aligned}
H\succ0\quad&\Longrightarrow\quad
\Delta\theta=\arg\min_d\left[g^\top d+\frac12d^\top Hd\right]\\
&\Longrightarrow\quad g+H\Delta\theta=0\\
&\Longrightarrow\quad\Delta\theta=-H^{-1}g
\end{aligned}
1.2 正定値性とステップ長
$H$ は正定値とは限らないので、$\lambda I$ を加える。
\begin{aligned}
B&=H+\lambda I\succ0,\qquad \lambda\ge0,\\
d&=-B^{-1}g,\\
g^\top d&=-g^\top B^{-1}g<0\qquad(g\ne0)
\end{aligned}
g^\top d<0\quad\Longrightarrow\quad\exists\,\bar\alpha>0:\ \mathcal L(\theta+\alpha d)<\mathcal L(\theta)\quad(0<\alpha<\bar\alpha)
\Delta\theta=\alpha d=-\alpha(H+\lambda I)^{-1}g
2. 自然勾配法:損失の曲率から分布の変化へ
自然勾配法は、更新の大きさを予測分布の変化で測る。
2.1 パラメータが表す分布
入力 $x$ は $\theta$ に依存しない分布 $\rho$ に従う。
\theta\longmapsto q_\theta(y\mid x),\qquad x\sim\rho,\qquad
\theta\longrightarrow\theta+d\ \Longrightarrow\ q_\theta\longrightarrow q_{\theta+d}
分布の変化をKLダイバージェンスで測る。
D_{\mathrm{KL}}(a\|b)=\mathbb E_{y\sim a}\left[\log\frac{a(y)}{b(y)}\right]\ge0,\qquad
D_{\mathrm{KL}}(a\|b)=0\iff a=b
\mathbb E_\theta[\cdot]:=\mathbb E_{x\sim\rho,\;y\sim q_\theta(\cdot\mid x)}[\cdot]
\begin{aligned}
\mathcal K_\theta(d)
&:=\mathbb E_{x\sim\rho}
D_{\mathrm{KL}}\!\left(q_\theta(\cdot\mid x)\middle\|q_{\theta+d}(\cdot\mid x)\right)\\
&=\mathbb E_\theta\left[\log\frac{q_\theta(y\mid x)}{q_{\theta+d}(y\mid x)}\right]
\end{aligned}
2.2 KLの二次近似とFisher情報行列
台が局所的に共通で、微分と期待値を交換できるとする。
s_\theta:=\nabla_\theta\log q_\theta(y\mid x),\qquad F:=\mathbb E_\theta[s_\theta s_\theta^\top]
v^\top Fv=\mathbb E_\theta[(v^\top s_\theta)^2]\ge0\quad\Longrightarrow\quad F\succeq0
\mathbb E_{y\sim q_\theta}[s_\theta]=\int\nabla_\theta q_\theta(y\mid x)\,\mathrm dy=\nabla_\theta\int q_\theta(y\mid x)\,\mathrm dy=\nabla_\theta1=0
\begin{aligned}
\mathcal K_\theta(0)&=0,\\
\left.\nabla_d\mathcal K_\theta(d)\right|_{d=0}&=-\mathbb E_\theta[s_\theta]=0,\\
\left.\nabla_d^2\mathcal K_\theta(d)\right|_{d=0}
&=-\mathbb E_\theta[\nabla_\theta^2\log q_\theta(y\mid x)]=F
\end{aligned}
\mathcal K_\theta(d)=\frac12d^\top Fd+o(\|d\|^2)
2.3 自然勾配の導出
$\lambda>0$ を加えると、$F$ が特異でも正定値になる。
v^\top(F+\lambda I)v=v^\top Fv+\lambda\|v\|^2\ge\lambda\|v\|^2>0\qquad(v\ne0)
$\eta>0$ は、損失の一次改善と分布変化の二次ペナルティの釣り合いを決める。
\begin{aligned}
\Delta\theta
&=\arg\min_d\left[g^\top d+\frac1{2\eta}d^\top Fd+\frac{\lambda}{2\eta}\|d\|^2\right],\\
0&=g+\frac1\eta(F+\lambda I)\Delta\theta,\\
\Delta\theta&=-\eta(F+\lambda I)^{-1}g
\end{aligned}
これをダンピング付き自然勾配更新と呼ぶ。
F\succ0\quad\Longrightarrow\quad\lim_{\lambda\to0}\Delta\theta=-\eta F^{-1}g
\begin{aligned}
H&=\nabla_\theta^2\mathcal L(\theta)&&\text{:損失の曲率}\\
F&=\left.\nabla_d^2\mathcal K_\theta(d)\right|_{d=0}&&\text{:分布の変化}
\end{aligned}
3. MLP:入力から予測分布を構成する
3.1 関数と確率モデル
MLP $\hat f$ の出力 $z$ を、固定した写像 $r$ で分布にする。
x\xrightarrow{\ \hat f(\cdot;\theta)\ }z\in\mathbb R^m\xrightarrow{\ r\ }q_\theta(\cdot\mid x),\qquad
q_\theta(y\mid x)=r\!\left(y;\hat f(x;\theta)\right)
r(y;z)=
\begin{cases}
\mathcal N(y;z,\tau^2I_m)&\text{回帰},\\
\operatorname{softmax}(z)_y&\text{分類}
\end{cases}
3.2 重み・バイアス・層
W_\ell\in\mathbb R^{n_\ell\times n_{\ell-1}},\qquad b_\ell\in\mathbb R^{n_\ell},\qquad \ell=1,\ldots,L,\qquad n_L=m
\begin{aligned}
h_0&=x,\\
a_\ell&=W_\ell h_{\ell-1}+b_\ell,&&1\le\ell\le L,\\
h_\ell&=\phi_\ell(a_\ell),&&1\le\ell\le L,\qquad\phi_L(u)=u,\\
z&=h_L=\hat f(x;\theta)
\end{aligned}
\theta=\operatorname{vec}(W_1,b_1,\ldots,W_L,b_L),\qquad P=\sum_{\ell=1}^L n_\ell(n_{\ell-1}+1)
バイアスは、入力に定数1を付けた重みと同じである。
Wh+b=\begin{bmatrix}W&b\end{bmatrix}\begin{bmatrix}h\\1\end{bmatrix}
ReLUの切り替え境界の位置は $b_j$ で決まる。
a_j=w_j^\top h+b_j=0\quad\Longleftrightarrow\quad w_j^\top h=-b_j
3.3 活性化関数:シグモイド・ReLU・SiLU
シグモイド関数(Sigmoid)
S(u):=\frac1{1+e^{-u}},\qquad 0<S(u)<1
\begin{aligned}
S'(u)&=S(u)(1-S(u)),\\
S''(u)&=S(u)(1-S(u))(1-2S(u))
\end{aligned}
\max_uS'(u)=S'(0)=\frac14,\qquad\lim_{|u|\to\infty}S'(u)=0,\qquad\operatorname{sign}S''(u)=-\operatorname{sign}u
ReLU(Rectified Linear Unit)
\operatorname{ReLU}(u)=\max(0,u)=\begin{cases}0&u\le0,\\u&u>0\end{cases}
\operatorname{ReLU}'(u)=\begin{cases}0&u<0,\\1&u>0\end{cases},\qquad
\operatorname{ReLU}''(u)=0\quad(u\ne0)
訓練入力 $x_i$ での前活性が0になるパラメータの集合を $\mathcal B$ とする。
\mathcal B:=\bigcup_{i,\ \ell<L,\ j}\left\{\theta:\ (a_{i,\ell}(\theta))_j=0\right\},\qquad\operatorname{vol}(\mathcal B)=0\ \ (\text{通常})
\theta\notin\mathcal B\quad\Longrightarrow\quad
\exists\,U\ni\theta:\ \operatorname{sign}a_{i,\ell}(\theta')=\operatorname{sign}a_{i,\ell}(\theta)\ \ (\theta'\in U)
\quad\Longrightarrow\quad\mathcal L\in C^2(U)
本記事の二階微分の式は $\theta\notin\mathcal B$ で成り立つ。
\partial_{\mathrm C}\operatorname{ReLU}(0)=[0,1],\qquad\text{自動微分の慣例}:\ \operatorname{ReLU}'(0):=0
SiLU(Sigmoid-weighted Linear Unit)
\begin{aligned}
\operatorname{SiLU}(u)&=uS(u),\\
\operatorname{SiLU}'(u)&=S(u)+uS(u)(1-S(u)),\\
\operatorname{SiLU}''(u)&=S(u)(1-S(u))\left[2+u(1-2S(u))\right]
\end{aligned}
\operatorname{SiLU}\in C^\infty(\mathbb R),\qquad
\lim_{u\to\infty}\operatorname{SiLU}'(u)=1,\qquad
\lim_{u\to-\infty}\operatorname{SiLU}'(u)=0
3.4 非線形性と深さ
T_\ell(h):=\phi_\ell(W_\ell h+b_\ell),\qquad
\hat f(\cdot;\theta)=T_L\circ\cdots\circ T_1,\qquad\theta\longmapsto[x\longmapsto\hat f(x;\theta)]
活性化が恒等写像なら、MLPはアフィン写像に退化する。
\begin{aligned}
\phi_\ell(u)=u\quad(\forall\ell)
&\Longrightarrow\quad\hat f(x;\theta)=Ax+c,\\
W_2(W_1x+b_1)+b_2
&=(W_2W_1)x+(W_2b_1+b_2)
\end{aligned}
4. KLダイバージェンスから損失へ
4.1 真の分布と負の対数尤度
\begin{aligned}
D_{\mathrm{KL}}(q_\theta\|q_{\theta+d})&:\quad\text{更新による分布の変化},\\
D_{\mathrm{KL}}(q_*\|q_\theta)&:\quad\text{真の分布とモデルのずれ}
\end{aligned}
\mathbb E_*[\cdot]:=\mathbb E_{x\sim\rho,\;y\sim q_*(\cdot\mid x)}[\cdot],\qquad
\ell(y,z):=-\log r(y;z)
各期待値は有限とする。
\begin{aligned}
\mathbb E_x D_{\mathrm{KL}}(q_*\|q_\theta)
&=\mathbb E_*[\log q_*(y\mid x)-\log q_\theta(y\mid x)]\\
&=C_*+\mathbb E_*[\ell(y,\hat f(x;\theta))],\\
C_*&:=\mathbb E_*[\log q_*(y\mid x)],\qquad \nabla_\theta C_*=0
\end{aligned}
\arg\min_\theta\mathbb E_x D_{\mathrm{KL}}(q_*\|q_\theta)
=\arg\min_\theta\mathbb E_*[\ell(y,\hat f(x;\theta))]
\mathcal D=\{(x_i,y_i)\}_{i=1}^N,\qquad z_i=\hat f(x_i;\theta),\qquad \ell_i=\ell(y_i,z_i)
\mathbb E_*[\ell(y,\hat f(x;\theta))]\approx\mathcal L(\theta):=\frac1N\sum_{i=1}^N\ell_i=-\frac1N\sum_{i=1}^N\log q_\theta(y_i\mid x_i)
4.2 回帰:Gaussian分布 → MSE
r(y;z)=\mathcal N(y;z,\tau^2I_m),\qquad y,z\in\mathbb R^m,\qquad\tau^2>0\ \text{は固定}
\begin{aligned}
\ell_i&=\frac1{2\tau^2}\|z_i-y_i\|^2+C_{\mathrm G},\\
C_{\mathrm G}&=\frac m2\log(2\pi\tau^2),\\
\mathcal L_{\mathrm{MSE}}(\theta)&:=\frac1{Nm}\sum_{i=1}^N\|z_i-y_i\|^2
\end{aligned}
\mathcal L(\theta)=\frac{m}{2\tau^2}\mathcal L_{\mathrm{MSE}}(\theta)+C_{\mathrm G}
\quad\Longrightarrow\quad
\arg\min_\theta\mathcal L=\arg\min_\theta\mathcal L_{\mathrm{MSE}}
4.3 分類:logits → softmax → クロスエントロピー
y_i\in\{1,\ldots,K\},\qquad z_i\in\mathbb R^K,\qquad m=K,\qquad t_{ik}=\mathbf1[y_i=k]
q_{ik}=r(k;z_i)=\frac{e^{z_{ik}}}{\sum_{j=1}^K e^{z_{ij}}},\qquad
q_i=(q_{i1},\ldots,q_{iK})^\top
\begin{aligned}
\ell_i
&=-\sum_{k=1}^K t_{ik}\log q_{ik}\\
&=-\log q_{i,y_i}\\
&=\log\sum_{k=1}^K e^{z_{ik}}-z_{i,y_i}
\end{aligned}
\mathcal L_{\mathrm{CE}}(\theta)=\frac1N\sum_{i=1}^N\left[\log\sum_{k=1}^K e^{z_{ik}}-z_{i,y_i}\right]
5. 自動微分:損失から重み・バイアスの勾配へ
5.1 順モードと逆モード
\mathcal L:\mathbb R^P\to\mathbb R,\qquad J_i:=\frac{\partial z_i}{\partial\theta}\in\mathbb R^{m\times P}
\begin{aligned}
\mathrm d\mathcal L
&=\frac1N\sum_{i=1}^N(\nabla_{z_i}\ell_i)^\top J_i\,\mathrm d\theta
=g^\top\mathrm d\theta,\\
J_{\mathcal L}&=\frac{\partial\mathcal L}{\partial\theta}=g^\top\in\mathbb R^{1\times P},\\
g&=\frac1N\sum_{i=1}^N J_i^\top\nabla_{z_i}\ell_i
\end{aligned}
順モードは $J_{\mathcal L}$ を右から、逆モードは左からベクトルに掛ける。
\begin{array}{lll}
\text{Forward mode}: & J_{\mathcal L}e_j=\dfrac{\partial\mathcal L}{\partial\theta_j}\quad(j=1,\ldots,P) & P\ \text{回}\\[3mm]
\text{Reverse mode}: & J_{\mathcal L}^\top 1=\nabla_\theta\mathcal L=g & 1\ \text{回}
\end{array}
5.2 MLPの逆伝播
\delta_{i,\ell}:=\nabla_{a_{i,\ell}}\ell_i,\qquad t_i=(t_{i1},\ldots,t_{iK})^\top
a_{i,L}=z_i\quad\Longrightarrow\quad
\delta_{i,L}=\nabla_{z_i}\ell_i=
\begin{cases}
\displaystyle\frac{z_i-y_i}{\tau^2}&\text{Gaussian},\\[2mm]
q_i-t_i&\text{softmax}
\end{cases}
\delta_{i,\ell}=\left(W_{\ell+1}^\top\delta_{i,\ell+1}\right)\odot\phi_\ell'(a_{i,\ell})
\quad(\ell=L-1,\ldots,1),\qquad(u\odot v)_j:=u_jv_j
隠れ層がすべてシグモイドのとき、勾配の減衰は $S'$ と重みの積で抑えられる。
D_{i,k}:=\operatorname{diag}\bigl(S'(a_{i,k})\bigr),\qquad\|D_{i,k}\|_2\le\frac14,\qquad
\delta_{i,\ell}=D_{i,\ell}W_{\ell+1}^\top\delta_{i,\ell+1}
\|\delta_{i,\ell}\|_2\le4^{-(L-\ell)}\left(\prod_{k=\ell+1}^{L}\|W_k\|_2\right)\|\delta_{i,L}\|_2
\begin{aligned}
\nabla_{W_\ell}\mathcal L&=\frac1N\sum_{i=1}^N\delta_{i,\ell}h_{i,\ell-1}^\top,\qquad
\nabla_{b_\ell}\mathcal L=\frac1N\sum_{i=1}^N\delta_{i,\ell},\\
g&=\operatorname{vec}\left(\nabla_{W_1}\mathcal L,\nabla_{b_1}\mathcal L,\ldots,\nabla_{W_L}\mathcal L,\nabla_{b_L}\mathcal L\right)
\end{aligned}
6. MLPからHessian・Fisherへ
6.1 損失の二階微分
u_i:=\nabla_{z_i}\ell_i=\delta_{i,L},\qquad C_i:=\nabla_{z_i}^2\ell_i,\qquad
g=\frac1N\sum_{i=1}^N J_i^\top u_i
$g$ をもう一度 $\theta$ で微分する。
H=\frac1N\sum_{i=1}^N\left[\underbrace{J_i^\top C_iJ_i}_{\text{出力での損失の曲率}}+\underbrace{\sum_{a=1}^m(u_i)_a\nabla_\theta^2z_{ia}}_{\text{出力自体の曲がり}}\right]
ReLUでも第2項は消えない。
z=w_2\operatorname{ReLU}(w_1x)\quad\Longrightarrow\quad
\frac{\partial^2z}{\partial w_1\partial w_2}=\operatorname{ReLU}'(w_1x)\,x=x\ne0\qquad(w_1x>0)
6.2 分布の微分
J(x):=\frac{\partial\hat f(x;\theta)}{\partial\theta},\qquad z=\hat f(x;\theta),\qquad
s_z(y):=\nabla_z\log r(y;z),\qquad F_z(z):=\mathbb E_{y\sim r(\cdot;z)}[s_zs_z^\top]
\begin{aligned}
s_\theta&=\nabla_\theta\log r\!\left(y;\hat f(x;\theta)\right)=J(x)^\top s_z,\\
F&=\mathbb E_{x\sim\rho}\,\mathbb E_{y\sim r(\cdot;z)}\left[J(x)^\top s_zs_z^\top J(x)\right]
=\mathbb E_{x\sim\rho}\left[J(x)^\top F_z(z)J(x)\right]
\end{aligned}
C_i=F_z(z_i)=
\begin{cases}
\tau^{-2}I_m&\text{Gaussian},\\
\operatorname{diag}(q_i)-q_iq_i^\top&\text{softmax}
\end{cases}
訓練入力上の一様分布 $\rho_N$ でのFisherを $F_N$ とする。
\rho_N:=\operatorname{Unif}\{x_1,\ldots,x_N\},\qquad
F_N:=F\big|_{\rho=\rho_N}=\frac1N\sum_{i=1}^N J_i^\top F_z(z_i)J_i
C_i=F_z(z_i)\quad\Longrightarrow\quad
H=F_N+\frac1N\sum_{i=1}^N\sum_{a=1}^m(u_i)_a\nabla_\theta^2z_{ia}
$H$ の第1項 $G$ を一般化ガウス・ニュートン行列(GGN)と呼ぶ。
G:=\frac1N\sum_{i=1}^N J_i^\top C_iJ_i,\qquad
\ell_i=\frac12\|z_i-y_i\|^2\ \Longrightarrow\ C_i=I_m\ \Longrightarrow\ G=\frac1N\sum_{i=1}^NJ_i^\top J_i
Gaussian回帰とsoftmax分類では、同じ $\eta,\lambda$ で、$F_N$ による自然勾配更新がGGNによる更新に一致する(Martens, 2020)。
G=F_N\quad\Longrightarrow\quad\Delta\theta_{\mathrm{Natural}}=-\eta(G+\lambda I)^{-1}g
\operatorname{rank}F_N\le\sum_{i=1}^N\operatorname{rank}\left(J_i^\top C_iJ_i\right)\le Nm,\qquad
P>Nm\ \Longrightarrow\ \det F_N=0
\text{softmax}:\quad C_i\mathbf1=0,\quad\operatorname{rank}C_i=K-1\quad\Longrightarrow\quad\operatorname{rank}F_N\le\min\{P,\,N(K-1)\}
ReLUの隠れユニット1個について、入る重みとバイアスを $c$ 倍し、出る重みを $1/c$ 倍する。
\begin{aligned}
&\theta_c:\ (w_{\mathrm{in}},b_{\mathrm{in}},w_{\mathrm{out}})\longmapsto(cw_{\mathrm{in}},cb_{\mathrm{in}},w_{\mathrm{out}}/c),\qquad
\hat f(x;\theta_c)=\hat f(x;\theta)\quad(c>0),\\
&v:=\left.\frac{\mathrm d\theta_c}{\mathrm dc}\right|_{c=1}=(w_{\mathrm{in}},b_{\mathrm{in}},-w_{\mathrm{out}})\ne0
\quad\Longrightarrow\quad J_iv=0\quad\Longrightarrow\quad F_Nv=0
\end{aligned}
softmaxでは、出力バイアスを全クラス同時にずらす方向も零方向である。
b_L\longmapsto b_L+c\mathbf1\ \Longrightarrow\ q_i\ \text{不変},\qquad
v_b:=\frac{\partial\theta}{\partial c}\ \Longrightarrow\ J_iv_b=\mathbf1\ \Longrightarrow\ F_Nv_b=\frac1N\sum_{i=1}^NJ_i^\top C_i\mathbf1=0
F_N\succeq0,\ \lambda>0\quad\Longrightarrow\quad F_N+\lambda I\succ0
6.3 経験Fisherとの違い
観測ラベルで勾配を掛け合わせた行列を、経験Fisher(empirical Fisher)と呼ぶ。
\tilde F_N:=\frac1N\sum_{i=1}^N\nabla_\theta\ell_i\,\nabla_\theta\ell_i^\top
=\frac1N\sum_{i=1}^N J_i^\top u_iu_i^\top J_i,\qquad u_i=-s_z(y_i)
\begin{aligned}
F_N&=\frac1N\sum_{i=1}^N J_i^\top\,\mathbb E_{y\sim r(\cdot;z_i)}\!\left[s_z(y)s_z(y)^\top\right]J_i,\\
\tilde F_N&=\frac1N\sum_{i=1}^N J_i^\top\,s_z(y_i)s_z(y_i)^\top J_i
\end{aligned}
(x_i,y_i)\overset{\text{i.i.d.}}{\sim}\rho(x)\,q_*(y\mid x),\quad
\theta\ \text{固定},\quad
q_\theta=q_*,\quad
\mathbb E_*\|s_\theta\|^2<\infty
\quad\Longrightarrow\quad\tilde F_N\xrightarrow{\ \text{a.s.}\ }F\quad(N\to\infty)
この条件の外では、線形回帰でも両者の前処理が大きく異なる(Kunstner et al., 2019)。
\text{二次の情報}:\ F_N=G\quad(\tilde F_N\ \text{は使わない})
6.4 負の曲率と鞍点
非退化な鞍点の近傍で、$H$ を固有分解する。
H=\sum_k\mu_kv_kv_k^\top,\quad\mu_k\ne0\ (\forall k),\qquad
-H^{-1}g=\sum_k\Delta\theta_k,\qquad
\Delta\theta_k:=-\frac{v_k^\top g}{\mu_k}v_k
各成分の方向微分 $g^\top\Delta\theta_k$ の符号は、固有値の符号で決まる。
v_k^\top g\ne0\quad\Longrightarrow\quad
g^\top\Delta\theta_k=-\frac{(v_k^\top g)^2}{\mu_k}
\begin{cases}
<0&\mu_k>0,\\
>0&\mu_k<0
\end{cases}
saddle-free Newton 法は、$\mu_k$ を $|\mu_k|$ に置き換えてこの符号反転を避ける(Dauphin et al., 2014)。
|H|:=\sum_k|\mu_k|v_kv_k^\top,\qquad
\Delta\theta_{\mathrm{SFN}}=-(|H|+\lambda_SI)^{-1}g,\qquad\lambda_S>0,\qquad
g=0\ \Longrightarrow\ \Delta\theta_{\mathrm{SFN}}=0
ダンピングはすべての固有値に同じ $\lambda_H$ を加える。
\Delta\theta=-\sum_k\frac{v_k^\top g}{\mu_k+\lambda_H}v_k,\qquad\lambda_H>-\min_k\mu_k
C_i\succeq0\ \Longrightarrow\ F_N=G\succeq0,\qquad
v^\top Hv<0\ \Longrightarrow\ v^\top(H-F_N)v<-v^\top F_Nv\le0
したがって自然勾配は常に降下方向である。
\lambda_F>0,\ g\ne0\quad\Longrightarrow\quad
g^\top\Delta\theta_{\mathrm{Natural}}=-\eta\,g^\top(F_N+\lambda_FI)^{-1}g<0
\quad\Longrightarrow\quad
\exists\,\bar\alpha>0:\ \mathcal L(\theta+\alpha\Delta\theta_{\mathrm{Natural}})<\mathcal L(\theta)\quad(0<\alpha<\bar\alpha)
6.5 更新式への接続
\begin{aligned}
H+\lambda_HI&\succ0,\qquad\lambda_H\ge0,\\
F_N+\lambda_FI&\succ0,\qquad\lambda_F>0,\\
\alpha,\eta&>0
\end{aligned}
\begin{aligned}
\Delta\theta_{\mathrm{Newton}}&=-\alpha(H+\lambda_HI)^{-1}g,\\
\Delta\theta_{\mathrm{Natural}}&=-\eta(F_N+\lambda_FI)^{-1}g,\\
\theta&\leftarrow\theta+\Delta\theta
\end{aligned}
6.6 大規模モデルでの計算
\text{メモリ}:\ O(P^2),\qquad\text{直接解法}:\ O(P^3),\qquad P\sim10^6\text{–}10^{11}
Hessian-free法(Krylov部分空間法)
共役勾配法で、次の対称正定値系を近似的に解く。
(G+\lambda I)d=-g,\qquad\lambda>0,\qquad G+\lambda I\succ0,\qquad
\text{必要な演算}:\ v\longmapsto Gv+\lambda v
$Gv$ は、行列を作らずに順モードと逆モードで求まる(Schraudolph, 2002)。
Gv=\frac1N\sum_{i=1}^NJ_i^\top\bigl(C_i(J_iv)\bigr),\qquad
J_iv:\ \text{順モード},\qquad J_i^\top w:\ \text{逆モード}
Hessian-free法は、残差や計算予算に応じて共役勾配法の反復を打ち切る(Martens, 2010)。
K-FAC(Kronecker-factored Approximate Curvature)
K-FACは、層ごとのブロックをクロネッカー積で近似する(Martens and Grosse, 2015)。
\begin{aligned}
\bar h_{i,\ell-1}&:=\begin{bmatrix}h_{i,\ell-1}\\1\end{bmatrix},\qquad
\operatorname{vec}\bigl(\nabla_{[W_\ell\ b_\ell]}\ell_i\bigr)=\bar h_{i,\ell-1}\otimes\delta_{i,\ell},\qquad y\sim r(\cdot;z_i)\ \text{で逆伝播},\\
F_{\ell\ell}&\approx\mathbb E[\bar h_{\ell-1}\bar h_{\ell-1}^\top]\otimes\mathbb E[\delta_\ell\delta_\ell^\top]=:A_{\ell-1}\otimes D_\ell,\\
A_{\ell-1},D_\ell\succ0\quad&\Longrightarrow\quad(A_{\ell-1}\otimes D_\ell)^{-1}=A_{\ell-1}^{-1}\otimes D_\ell^{-1}
\end{aligned}
softmaxの出力層では $D_L$ が特異になる。
\mathbf1^\top\delta_{i,L}=\mathbf1^\top(q_i-t_i)=0\quad\Longrightarrow\quad D_L\mathbf1=\mathbb E[\delta_L\delta_L^\top]\mathbf1=0
因子ごとのダンピングは、全体に $\lambda I$ を加える操作とは異なる。
(A+\gamma_AI)\otimes(D+\gamma_DI)=A\otimes D+\gamma_DA\otimes I+\gamma_AI\otimes D+\gamma_A\gamma_DI
Woodburyの公式による低ランク計算
F_N=M^\top M,\qquad
M:=\frac1{\sqrt N}\begin{bmatrix}C_1^{1/2}J_1\\\vdots\\C_N^{1/2}J_N\end{bmatrix}\in\mathbb R^{Nm\times P}
(M^\top M+\lambda I_P)^{-1}g=\frac1\lambda\left[g-M^\top\left(\lambda I_{Nm}+MM^\top\right)^{-1}Mg\right]
Nm\ll P\quad\Longrightarrow\quad P\times P\ \text{の求解}\ \longrightarrow\ Nm\times Nm\ \text{の求解},\qquad
M\ \text{の保持}:\ NmP\ \text{要素}
まとめ
\begin{aligned}
H&=G+\frac1N\sum_{i=1}^N\sum_{a=1}^m(u_i)_a\nabla_\theta^2z_{ia},&
G&=F_N\succeq0\ \ (\text{Gaussian, softmax}),\qquad\operatorname{rank}F_N\le Nm,\\
\Delta\theta_{\mathrm{Newton}}&=-\alpha(H+\lambda_HI)^{-1}g,&
\Delta\theta_{\mathrm{Natural}}&=-\eta(F_N+\lambda_FI)^{-1}g
\end{aligned}
参考文献
- Dauphin, Y. N., Pascanu, R., Gulcehre, C., Cho, K., Ganguli, S., & Bengio, Y. (2014). Identifying and attacking the saddle point problem in high-dimensional non-convex optimization. Advances in Neural Information Processing Systems 27. arXiv:1406.2572
- Kunstner, F., Balles, L., & Hennig, P. (2019). Limitations of the empirical Fisher approximation for natural gradient descent. Advances in Neural Information Processing Systems 32. arXiv:1905.12558
- Martens, J. (2010). Deep learning via Hessian-free optimization. Proceedings of the 27th International Conference on Machine Learning, 735–742. PDF
- Martens, J. (2020). New insights and perspectives on the natural gradient method. Journal of Machine Learning Research, 21(146), 1–76. JMLR
- Martens, J., & Grosse, R. (2015). Optimizing neural networks with Kronecker-factored approximate curvature. Proceedings of the 32nd International Conference on Machine Learning, PMLR 37, 2408–2417. arXiv:1503.05671
- Schraudolph, N. N. (2002). Fast curvature matrix-vector products for second-order gradient descent. Neural Computation, 14(7), 1723–1738. doi:10.1162/08997660260028683