分類問題では出力が 0 0 0 か 1 1 1 のラベルです。ここに最小二乗法をそのまま当てはめると、確率の範囲を外れた予測が出るうえ、判定に無関係なはずの遠くの点が決定境界を動かしてしまいます。
線形なスコア w T x \boldsymbol{w}^{\mathsf{T}}\boldsymbol{x} w T x を確率に変える写像としてシグモイド関数 σ ( z ) = 1 / ( 1 + e − z ) \sigma(z) = 1/(1+e^{-z}) σ ( z ) = 1/ ( 1 + e − z ) を使います。これは恣意的な選択ではなく、「対数オッズ(ロジット)が線形である」という仮定と同値です。
損失関数は手で選ぶものではありません。ベルヌーイ分布の最尤推定を書き下すと、負の対数尤度がそのまま交差エントロピー誤差になります。
勾配は ∇ L ( w ) = ∑ i ( μ i − y i ) x i \nabla L(\boldsymbol{w}) = \sum_i (\mu_i - y_i)\boldsymbol{x}_i ∇ L ( w ) = ∑ i ( μ i − y i ) x i という驚くほど簡単な形になります。シグモイドの微分が交差エントロピーの微分と約分するからで、この約分が学習の速さを支えています。
ヘッセ行列は X T S X ⪰ O X^{\mathsf{T}}SX \succeq O X T S X ⪰ O なので L L L は凸です。しかし停留条件は超越方程式で、線形回帰と違って閉じた解がありません。だから微分を使って数値的に探すしかない、というのが次章の勾配降下法につながります。
データが線形分離可能なとき最尤推定量は存在せず、重みは発散します。L 2 L^2 L 2 正則化を足すと最小点の存在と一意性が回復し、これはガウス事前分布による MAP 推定と一致します。
線形回帰と最小二乗法 で扱ったのは、身長から体重を予測するような、出力が実数の問題でした。しかし現場で解きたい問題の多くは、そうではありません。
このメールは迷惑メールか、そうでないか。
この検査値の組を持つ患者は、その病気に罹患しているか、していないか。
この画像に写っているのは猫か、猫でないか。
いずれも出力は「はい/いいえ」の 2 択です。こういう問題を 2 値分類問題 といいます。数学的には、特徴ベクトル x ∈ R d \boldsymbol{x} \in \mathbb{R}^{d} x ∈ R d に対してラベル y ∈ { 0 , 1 } y \in \{0, 1\} y ∈ { 0 , 1 } を予測する問題として定式化します。
ここで素朴な疑問が湧きます。ラベル y y y も所詮は数値なのだから、線形回帰をそのまま使えばよいのではないか。y ^ = w T x \hat{y} = \boldsymbol{w}^{\mathsf{T}}\boldsymbol{x} y ^ = w T x を最小二乗法で当てはめ、y ^ ≥ 0.5 \hat{y} \ge 0.5 y ^ ≥ 0.5 なら「はい」、そうでなければ「いいえ」と判定すればよさそうに見えます。実際にやってみると、何が起きるかがはっきりします。
勉強時間 t t t (時間)から試験の合否 y y y (1 1 1 が合格)を予測する、次の 4 点のデータを考えます。
y ^ = a t + b \hat{y} = at + b y ^ = a t + b を最小二乗法で当てはめます。t ˉ = 2.5 \bar{t} = 2.5 t ˉ = 2.5 、y ˉ = 0.5 \bar{y} = 0.5 y ˉ = 0.5 、∑ i ( t i − t ˉ ) 2 = 2.25 + 0.25 + 0.25 + 2.25 = 5 \sum_i (t_i - \bar{t})^2 = 2.25 + 0.25 + 0.25 + 2.25 = 5 ∑ i ( t i − t ˉ ) 2 = 2.25 + 0.25 + 0.25 + 2.25 = 5 、∑ i ( t i − t ˉ ) ( y i − y ˉ ) = 0.75 + 0.25 + 0.25 + 0.75 = 2 \sum_i (t_i - \bar{t})(y_i - \bar{y}) = 0.75 + 0.25 + 0.25 + 0.75 = 2 ∑ i ( t i − t ˉ ) ( y i − y ˉ ) = 0.75 + 0.25 + 0.25 + 0.75 = 2 ですから、
a = 2 5 = 0.4 , b = 0.5 − 0.4 × 2.5 = − 0.5. a = \frac{2}{5} = 0.4, \qquad b = 0.5 - 0.4 \times 2.5 = -0.5 . a = 5 2 = 0.4 , b = 0.5 − 0.4 × 2.5 = − 0.5.
判定は y ^ = 0.5 \hat{y} = 0.5 y ^ = 0.5 すなわち t = 2.5 t = 2.5 t = 2.5 を境界とし、4 点すべてを正しく分類します。一見うまくいっています。しかし予測値そのものを見ると y ^ ( 1 ) = − 0.1 \hat{y}(1) = -0.1 y ^ ( 1 ) = − 0.1 、y ^ ( 4 ) = 1.1 \hat{y}(4) = 1.1 y ^ ( 4 ) = 1.1 です。負の確率と、1 1 1 を超える確率 が出てきました。「合格する確率は − 10 % -10\% − 10% です」と言われても意味が取れません。
これは見た目の問題にとどまりません。ここに「20 時間勉強して合格した」という、常識的にはまったく問題のない 1 点 ( t , y ) = ( 20 , 1 ) (t, y) = (20, 1) ( t , y ) = ( 20 , 1 ) を追加します。t ˉ = 6 \bar{t} = 6 t ˉ = 6 、y ˉ = 0.6 \bar{y} = 0.6 y ˉ = 0.6 、∑ i ( t i − t ˉ ) 2 = 25 + 16 + 9 + 4 + 196 = 250 \sum_i (t_i - \bar{t})^2 = 25 + 16 + 9 + 4 + 196 = 250 ∑ i ( t i − t ˉ ) 2 = 25 + 16 + 9 + 4 + 196 = 250 、∑ i ( t i − t ˉ ) ( y i − y ˉ ) = 3 + 2.4 − 1.2 − 0.8 + 5.6 = 9 \sum_i (t_i - \bar{t})(y_i - \bar{y}) = 3 + 2.4 - 1.2 - 0.8 + 5.6 = 9 ∑ i ( t i − t ˉ ) ( y i − y ˉ ) = 3 + 2.4 − 1.2 − 0.8 + 5.6 = 9 なので
a = 9 250 = 0.036 , b = 0.6 − 0.036 × 6 = 0.384. a = \frac{9}{250} = 0.036, \qquad b = 0.6 - 0.036 \times 6 = 0.384 . a = 250 9 = 0.036 , b = 0.6 − 0.036 × 6 = 0.384.
境界は 0.036 t + 0.384 = 0.5 0.036t + 0.384 = 0.5 0.036 t + 0.384 = 0.5 すなわち t = 3.22 … t = 3.22\ldots t = 3.22 … に移動しました。その結果、t = 3 t = 3 t = 3 の点は y ^ ( 3 ) = 0.492 < 0.5 \hat{y}(3) = 0.492 < 0.5 y ^ ( 3 ) = 0.492 < 0.5 となり、もともと正しく分類できていた点が誤分類されます 。
なぜこうなるのでしょうか。二乗誤差 ( y ^ i − y i ) 2 (\hat{y}_i - y_i)^2 ( y ^ i − y i ) 2 は、y ^ i \hat{y}_i y ^ i が y i y_i y i から離れるほど罰を与えます。ところが分類の観点では、t = 20 t = 20 t = 20 の点について y ^ = 1.1 \hat{y} = 1.1 y ^ = 1.1 でも y ^ = 5 \hat{y} = 5 y ^ = 5 でも「合格側に十分入っている」という意味では同じく正解です。二乗誤差はこの「正解の側に深く入りすぎた」状態を誤差として数え、それを減らすために直線を寝かせてしまう。分類の目的関数として二乗誤差が不適切だ、ということです。
以上から、必要なものが 2 つはっきりします。
出力を ( 0 , 1 ) (0,1) ( 0 , 1 ) に押し込める仕組み 。予測値を確率として読めるようにしたい。
確率モデルから導かれる損失関数 。「0.9 0.9 0.9 と答えて正解だった」ことと「0.55 0.55 0.55 と答えて正解だった」ことの差を、恣意的でない基準で測りたい。
この 2 つを与えるのがロジスティック回帰です。1 番目にシグモイド関数(定義 3.1 )が、2 番目に最尤推定から出てくる交差エントロピー誤差(定義 4.1 )が対応します。そして最後に、その損失関数を最小化する段になって初めて「微分」が本質的に必要になります。線形回帰では正規方程式という連立一次方程式を解けば済んだのに対し、ロジスティック回帰では閉じた解が書けないからです(注意 5.4 )。
flowchart LR
A["特徴ベクトル x"] --> B["線形スコア z = w·x"]
B --> C["確率 p = シグモイド(z)"]
C --> D["負の対数尤度 = 交差エントロピー誤差 L"]
D --> E["勾配 grad L = 総和 (p - y) x"]
E --> F["重み w を更新"]
F -.-> B ロジスティック回帰の全体像。線形スコアを確率に変え、確率から損失を作り、損失の勾配で重みを直す。
以下、データは n n n 個の組 ( x 1 , y 1 ) , … , ( x n , y n ) (\boldsymbol{x}_1, y_1), \ldots, (\boldsymbol{x}_n, y_n) ( x 1 , y 1 ) , … , ( x n , y n ) で与えられ、x i ∈ R d \boldsymbol{x}_i \in \mathbb{R}^{d} x i ∈ R d 、y i ∈ { 0 , 1 } y_i \in \{0,1\} y i ∈ { 0 , 1 } とします。切片(バイアス)は特徴ベクトルに吸収 します。すなわち x i \boldsymbol{x}_i x i の第 1 成分は常に 1 1 1 であるとし、対応する重み w 1 w_1 w 1 が切片の役割を果たすものとします。こうしておくと式に切片が現れず、以後の計算がすべて w T x \boldsymbol{w}^{\mathsf{T}}\boldsymbol{x} w T x の形に統一されます。
計画行列 X X X を、第 i i i 行が x i T \boldsymbol{x}_i^{\mathsf{T}} x i T である n × d n \times d n × d 行列とします(線形回帰と最小二乗法 の 計画行列(定義 2.1)[線形回帰と最小二乗法] と同じ記法です)。ラベルを並べたベクトルを y = ( y 1 , … , y n ) T ∈ R n \boldsymbol{y} = (y_1, \ldots, y_n)^{\mathsf{T}} \in \mathbb{R}^{n} y = ( y 1 , … , y n ) T ∈ R n と書きます。
確率については、各 i i i について y i y_i y i が x i \boldsymbol{x}_i x i を与えたもとで条件付きに独立であることを仮定します。x i \boldsymbol{x}_i x i 自体の分布は一切モデル化しません(この点は 例 3.4 で再度触れます)。確率変数と期待値の基本については 確率変数と期待値 を参照してください。
ベクトルによる微分は成分ごとの偏微分を並べたもの、すなわち ∇ f ( w ) = ( ∂ f / ∂ w 1 , … , ∂ f / ∂ w d ) T \nabla f(\boldsymbol{w}) = (\partial f/\partial w_1, \ldots, \partial f/\partial w_d)^{\mathsf{T}} ∇ f ( w ) = ( ∂ f / ∂ w 1 , … , ∂ f / ∂ w d ) T とし、ヘッセ行列は ( ∇ 2 f ) j k = ∂ 2 f / ∂ w j ∂ w k (\nabla^2 f)_{jk} = \partial^2 f / \partial w_j \partial w_k ( ∇ 2 f ) j k = ∂ 2 f / ∂ w j ∂ w k とします。詳しくは 多変数関数の微分と偏微分 の ヘッセ行列の定義(定義 7.3)[多変数関数の微分と偏微分] を参照してください。対称行列 A A A について A ⪰ O A \succeq O A ⪰ O は半正定値、A ≻ O A \succ O A ≻ O は正定値を表します。
定義 3.1 (シグモイド関数(標準ロジスティック関数) )
関数 σ : R → R \sigma : \mathbb{R} \to \mathbb{R} σ : R → R を
σ ( z ) = 1 1 + e − z \sigma(z) = \frac{1}{1 + e^{-z}} σ ( z ) = 1 + e − z 1 で定める。これをシグモイド関数 、または標準ロジスティック関数 という。
名前の由来は、グラフが S S S 字(ギリシャ文字シグマの語幹 sigma + eides「〜のような形」)を描くことです。もともとは 1838 年に Verhulst が人口の増加を記述する微分方程式 d p d t = p ( 1 − p ) \frac{dp}{dt} = p(1-p) d t d p = p ( 1 − p ) の解として導入したものでした。この微分方程式自体が、次の性質 (3) と同じ形をしていることに注意してください。
シグモイド関数のグラフ。z = 0 で 0.5 を通り、両端で 0 と 1 に漸近します。
命題 3.2 (シグモイド関数の基本性質 )
定義 3.1 の σ \sigma σ について、次が成り立つ。
すべての z ∈ R z \in \mathbb{R} z ∈ R について 0 < σ ( z ) < 1 0 < \sigma(z) < 1 0 < σ ( z ) < 1 であり、σ \sigma σ は R \mathbb{R} R 上 C ∞ C^{\infty} C ∞ 級かつ狭義単調増加で、lim z → − ∞ σ ( z ) = 0 \lim_{z \to -\infty} \sigma(z) = 0 lim z → − ∞ σ ( z ) = 0 、lim z → + ∞ σ ( z ) = 1 \lim_{z \to +\infty} \sigma(z) = 1 lim z → + ∞ σ ( z ) = 1 。
すべての z z z について σ ( − z ) = 1 − σ ( z ) \sigma(-z) = 1 - \sigma(z) σ ( − z ) = 1 − σ ( z ) 。特に σ ( 0 ) = 1 / 2 \sigma(0) = 1/2 σ ( 0 ) = 1/2 。
すべての z z z について σ ′ ( z ) = σ ( z ) ( 1 − σ ( z ) ) = σ ( z ) σ ( − z ) > 0 \sigma'(z) = \sigma(z)\bigl(1 - \sigma(z)\bigr) = \sigma(z)\,\sigma(-z) > 0 σ ′ ( z ) = σ ( z ) ( 1 − σ ( z ) ) = σ ( z ) σ ( − z ) > 0 。
σ : R → ( 0 , 1 ) \sigma : \mathbb{R} \to (0,1) σ : R → ( 0 , 1 ) は全単射であり、その逆関数は σ − 1 ( p ) = log p 1 − p \sigma^{-1}(p) = \log \dfrac{p}{1-p} σ − 1 ( p ) = log 1 − p p (0 < p < 1 0 < p < 1 0 < p < 1 )で与えられる。
証明(命題 3.2) (1) すべての z z z について e − z > 0 e^{-z} > 0 e − z > 0 なので 1 + e − z > 1 > 0 1 + e^{-z} > 1 > 0 1 + e − z > 1 > 0 であり、したがって 0 < σ ( z ) = 1 / ( 1 + e − z ) < 1 0 < \sigma(z) = 1/(1+e^{-z}) < 1 0 < σ ( z ) = 1/ ( 1 + e − z ) < 1 です。z ↦ e − z z \mapsto e^{-z} z ↦ e − z は C ∞ C^{\infty} C ∞ 級で、分母 1 + e − z 1+e^{-z} 1 + e − z は決して 0 0 0 にならないので、商として σ \sigma σ も C ∞ C^{\infty} C ∞ 級です。狭義単調増加であることは (3) で示す σ ′ > 0 \sigma' > 0 σ ′ > 0 から従います。極限は、z → − ∞ z \to -\infty z → − ∞ のとき e − z → + ∞ e^{-z} \to +\infty e − z → + ∞ なので σ ( z ) → 0 \sigma(z) \to 0 σ ( z ) → 0 、z → + ∞ z \to +\infty z → + ∞ のとき e − z → 0 e^{-z} \to 0 e − z → 0 なので σ ( z ) → 1 \sigma(z) \to 1 σ ( z ) → 1 です。
(2) 定義から σ ( − z ) = 1 / ( 1 + e z ) \sigma(-z) = 1/(1+e^{z}) σ ( − z ) = 1/ ( 1 + e z ) です。一方
1 − σ ( z ) = 1 − 1 1 + e − z = ( 1 + e − z ) − 1 1 + e − z = e − z 1 + e − z 1 - \sigma(z) = 1 - \frac{1}{1+e^{-z}} = \frac{(1+e^{-z}) - 1}{1+e^{-z}} = \frac{e^{-z}}{1+e^{-z}} 1 − σ ( z ) = 1 − 1 + e − z 1 = 1 + e − z ( 1 + e − z ) − 1 = 1 + e − z e − z であり、分子分母に e z e^{z} e z を掛けると 1 e z + 1 \dfrac{1}{e^{z}+1} e z + 1 1 となって σ ( − z ) \sigma(-z) σ ( − z ) に一致します。z = 0 z = 0 z = 0 とすれば σ ( 0 ) = 1 − σ ( 0 ) \sigma(0) = 1 - \sigma(0) σ ( 0 ) = 1 − σ ( 0 ) 、すなわち σ ( 0 ) = 1 / 2 \sigma(0) = 1/2 σ ( 0 ) = 1/2 です。
(3) σ ( z ) = ( 1 + e − z ) − 1 \sigma(z) = (1+e^{-z})^{-1} σ ( z ) = ( 1 + e − z ) − 1 に合成関数の微分法を適用します。外側の微分が − ( 1 + e − z ) − 2 -(1+e^{-z})^{-2} − ( 1 + e − z ) − 2 、内側 1 + e − z 1+e^{-z} 1 + e − z の微分が − e − z -e^{-z} − e − z なので
σ ′ ( z ) = − ( 1 + e − z ) − 2 ⋅ ( − e − z ) = e − z ( 1 + e − z ) 2 . \sigma'(z) = -(1+e^{-z})^{-2} \cdot (-e^{-z}) = \frac{e^{-z}}{(1+e^{-z})^{2}} . σ ′ ( z ) = − ( 1 + e − z ) − 2 ⋅ ( − e − z ) = ( 1 + e − z ) 2 e − z . 他方、(2) の途中式より 1 − σ ( z ) = e − z 1 + e − z 1 - \sigma(z) = \dfrac{e^{-z}}{1+e^{-z}} 1 − σ ( z ) = 1 + e − z e − z なので
σ ( z ) ( 1 − σ ( z ) ) = 1 1 + e − z ⋅ e − z 1 + e − z = e − z ( 1 + e − z ) 2 \sigma(z)\bigl(1-\sigma(z)\bigr) = \frac{1}{1+e^{-z}} \cdot \frac{e^{-z}}{1+e^{-z}} = \frac{e^{-z}}{(1+e^{-z})^{2}} σ ( z ) ( 1 − σ ( z ) ) = 1 + e − z 1 ⋅ 1 + e − z e − z = ( 1 + e − z ) 2 e − z となり、両者は一致します。さらに (2) より 1 − σ ( z ) = σ ( − z ) 1 - \sigma(z) = \sigma(-z) 1 − σ ( z ) = σ ( − z ) なので σ ′ ( z ) = σ ( z ) σ ( − z ) \sigma'(z) = \sigma(z)\sigma(-z) σ ′ ( z ) = σ ( z ) σ ( − z ) です。(1) より σ ( z ) > 0 \sigma(z) > 0 σ ( z ) > 0 かつ σ ( − z ) > 0 \sigma(-z) > 0 σ ( − z ) > 0 なので σ ′ ( z ) > 0 \sigma'(z) > 0 σ ′ ( z ) > 0 です。
(4) (3) より σ \sigma σ は狭義単調増加なので単射です。(1) より σ \sigma σ は連続で、値域は ( 0 , 1 ) (0,1) ( 0 , 1 ) に含まれ、両端の極限が 0 0 0 と 1 1 1 なので、中間値の定理により ( 0 , 1 ) (0,1) ( 0 , 1 ) のすべての値を取ります。よって σ : R → ( 0 , 1 ) \sigma : \mathbb{R} \to (0,1) σ : R → ( 0 , 1 ) は全単射です。逆関数は p = 1 / ( 1 + e − z ) p = 1/(1+e^{-z}) p = 1/ ( 1 + e − z ) を z z z について解けば求まります。両辺の逆数を取って 1 + e − z = 1 / p 1 + e^{-z} = 1/p 1 + e − z = 1/ p 、すなわち e − z = ( 1 − p ) / p e^{-z} = (1-p)/p e − z = ( 1 − p ) / p 。両辺の対数を取って − z = log 1 − p p -z = \log\dfrac{1-p}{p} − z = log p 1 − p 、したがって z = log p 1 − p z = \log\dfrac{p}{1-p} z = log 1 − p p です。
∎
性質 (3) は本記事全体で最も使う式です。シグモイドの導関数がシグモイド自身の多項式で書ける という事実が、後で勾配の計算をきれいにします(定理 5.1 )。
( 0 , 1 ) (0,1) ( 0 , 1 ) に値を取る単調増加な滑らかな関数はいくらでもあります。たとえば標準正規分布の累積分布関数 Φ \Phi Φ でもよく、それを使ったモデルはプロビット回帰と呼ばれます。ではなぜシグモイドなのでしょうか。答えは 命題 3.2 の (4) にあります。
定義 3.3 (オッズとロジット )
0 < p < 1 0 < p < 1 0 < p < 1 に対し、p 1 − p \dfrac{p}{1-p} 1 − p p を確率 p p p のオッズ (odds)といい、その対数
logit ( p ) = log p 1 − p \operatorname{logit}(p) = \log \frac{p}{1-p} logit ( p ) = log 1 − p p を ロジット (logit)または対数オッズ という。命題 3.2 (4) より logit = σ − 1 \operatorname{logit} = \sigma^{-1} logit = σ − 1 である。
オッズは「起こる場合の数と起こらない場合の数の比」です。競馬の「3 倍」やスポーツの「2 対 1」と同じ言葉づかいで、p = 0.75 p = 0.75 p = 0.75 ならオッズは 3 3 3 、すなわち「3 対 1」です。確率は [ 0 , 1 ] [0,1] [ 0 , 1 ] という有界区間に閉じ込められていますが、オッズは ( 0 , ∞ ) (0, \infty) ( 0 , ∞ ) を、その対数であるロジットは R \mathbb{R} R 全体を動きます。確率を「線形に動かせる量」に変換する のがロジットの役割です。
したがって、μ = σ ( w T x ) \mu = \sigma(\boldsymbol{w}^{\mathsf{T}}\boldsymbol{x}) μ = σ ( w T x ) と置くことは、両辺にロジットを施した
log μ 1 − μ = w T x \log \frac{\mu}{1-\mu} = \boldsymbol{w}^{\mathsf{T}}\boldsymbol{x} log 1 − μ μ = w T x
と完全に同値 です。つまりロジスティック回帰の仮定は「確率が線形」でも「確率がシグモイド形」でもなく、対数オッズが特徴の線形結合である という一点に尽きます。シグモイドはその仮定を確率について解き直しただけのものです。
この仮定は、次の例が示すように、自然な状況で実際に成り立ちます。
例 3.4 (2 つの正規分布からシグモイドが出てくる )
クラス y ∈ { 0 , 1 } y \in \{0,1\} y ∈ { 0 , 1 } の事前確率を π 1 = P ( y = 1 ) \pi_1 = P(y=1) π 1 = P ( y = 1 ) 、π 0 = 1 − π 1 \pi_0 = 1-\pi_1 π 0 = 1 − π 1 とし、クラスごとの特徴の分布(クラス条件付き密度)を p ( x ∣ y = k ) p(\boldsymbol{x} \mid y=k) p ( x ∣ y = k ) とします。ベイズの定理(定理 2.2)[確率論とベイズ統計] より
P ( y = 1 ∣ x ) = p ( x ∣ y = 1 ) π 1 p ( x ∣ y = 1 ) π 1 + p ( x ∣ y = 0 ) π 0 = 1 1 + exp ( − a ) , a = log p ( x ∣ y = 1 ) π 1 p ( x ∣ y = 0 ) π 0 . P(y=1 \mid \boldsymbol{x}) = \frac{p(\boldsymbol{x}\mid y=1)\pi_1}{p(\boldsymbol{x}\mid y=1)\pi_1 + p(\boldsymbol{x}\mid y=0)\pi_0} = \frac{1}{1 + \exp(-a)},
\qquad a = \log \frac{p(\boldsymbol{x}\mid y=1)\pi_1}{p(\boldsymbol{x}\mid y=0)\pi_0}. P ( y = 1 ∣ x ) = p ( x ∣ y = 1 ) π 1 + p ( x ∣ y = 0 ) π 0 p ( x ∣ y = 1 ) π 1 = 1 + exp ( − a ) 1 , a = log p ( x ∣ y = 0 ) π 0 p ( x ∣ y = 1 ) π 1 . 真ん中の変形は、分子分母を p ( x ∣ y = 1 ) π 1 p(\boldsymbol{x}\mid y=1)\pi_1 p ( x ∣ y = 1 ) π 1 で割って p ( x ∣ y = 0 ) π 0 p ( x ∣ y = 1 ) π 1 = e − a \dfrac{p(\boldsymbol{x}\mid y=0)\pi_0}{p(\boldsymbol{x}\mid y=1)\pi_1} = e^{-a} p ( x ∣ y = 1 ) π 1 p ( x ∣ y = 0 ) π 0 = e − a を使っただけです。つまりシグモイドは何の仮定も置かずに現れます 。a a a が対数オッズそのものだからです。
残るのは「a a a が x \boldsymbol{x} x の 1 次式か」だけです。両クラスの分布が共通の共分散行列 Σ \Sigma Σ (正則)を持つ正規分布 N ( μ 1 , Σ ) N(\boldsymbol{\mu}_1, \Sigma) N ( μ 1 , Σ ) 、N ( μ 0 , Σ ) N(\boldsymbol{\mu}_0, \Sigma) N ( μ 0 , Σ ) のとき、正規化定数が約分して
a = log π 1 π 0 − 1 2 ( x − μ 1 ) T Σ − 1 ( x − μ 1 ) + 1 2 ( x − μ 0 ) T Σ − 1 ( x − μ 0 ) = ( μ 1 − μ 0 ) T Σ − 1 x − 1 2 μ 1 T Σ − 1 μ 1 + 1 2 μ 0 T Σ − 1 μ 0 + log π 1 π 0 \begin{aligned}
a &= \log\frac{\pi_1}{\pi_0} - \tfrac12 (\boldsymbol{x}-\boldsymbol{\mu}_1)^{\mathsf{T}}\Sigma^{-1}(\boldsymbol{x}-\boldsymbol{\mu}_1) + \tfrac12 (\boldsymbol{x}-\boldsymbol{\mu}_0)^{\mathsf{T}}\Sigma^{-1}(\boldsymbol{x}-\boldsymbol{\mu}_0) \\[2pt]
&= (\boldsymbol{\mu}_1 - \boldsymbol{\mu}_0)^{\mathsf{T}}\Sigma^{-1}\boldsymbol{x} \;-\; \tfrac12 \boldsymbol{\mu}_1^{\mathsf{T}}\Sigma^{-1}\boldsymbol{\mu}_1 + \tfrac12 \boldsymbol{\mu}_0^{\mathsf{T}}\Sigma^{-1}\boldsymbol{\mu}_0 + \log\frac{\pi_1}{\pi_0}
\end{aligned} a = log π 0 π 1 − 2 1 ( x − μ 1 ) T Σ − 1 ( x − μ 1 ) + 2 1 ( x − μ 0 ) T Σ − 1 ( x − μ 0 ) = ( μ 1 − μ 0 ) T Σ − 1 x − 2 1 μ 1 T Σ − 1 μ 1 + 2 1 μ 0 T Σ − 1 μ 0 + log π 0 π 1 となります。2 行目では 2 次の項 − 1 2 x T Σ − 1 x -\tfrac12\boldsymbol{x}^{\mathsf{T}}\Sigma^{-1}\boldsymbol{x} − 2 1 x T Σ − 1 x が両方の括弧から出て打ち消し合う ことを使いました(共分散行列が共通でなければ消えません)。残ったのは x \boldsymbol{x} x の 1 次式です。
1 次元で数値を入れてみます。μ 0 = 0 \mu_0 = 0 μ 0 = 0 、μ 1 = 2 \mu_1 = 2 μ 1 = 2 、分散 1 1 1 、π 1 = π 0 = 1 / 2 \pi_1 = \pi_0 = 1/2 π 1 = π 0 = 1/2 とすると
a = − 1 2 ( x − 2 ) 2 + 1 2 x 2 = 2 x − 2 , P ( y = 1 ∣ x ) = σ ( 2 x − 2 ) . a = -\tfrac12 (x-2)^2 + \tfrac12 x^2 = 2x - 2,
\qquad P(y=1\mid x) = \sigma(2x-2). a = − 2 1 ( x − 2 ) 2 + 2 1 x 2 = 2 x − 2 , P ( y = 1 ∣ x ) = σ ( 2 x − 2 ) . 境界 P = 1 / 2 P = 1/2 P = 1/2 は x = 1 x = 1 x = 1 、つまり 2 つの平均の中点です。ロジスティック回帰はこの a a a の係数 ( − 2 , 2 ) (-2, 2) ( − 2 , 2 ) を、μ k \boldsymbol{\mu}_k μ k や Σ \Sigma Σ を経由せずに直接推定するモデルだと読めます。
定義 3.6 (ロジスティック回帰モデル )
パラメータ w ∈ R d \boldsymbol{w} \in \mathbb{R}^{d} w ∈ R d に対し、特徴 x ∈ R d \boldsymbol{x} \in \mathbb{R}^{d} x ∈ R d を与えたときのラベル y ∈ { 0 , 1 } y \in \{0,1\} y ∈ { 0 , 1 } の条件付き分布を
P ( y = 1 ∣ x ; w ) = σ ( w T x ) , P ( y = 0 ∣ x ; w ) = 1 − σ ( w T x ) P(y = 1 \mid \boldsymbol{x};\boldsymbol{w}) = \sigma(\boldsymbol{w}^{\mathsf{T}}\boldsymbol{x}), \qquad
P(y = 0 \mid \boldsymbol{x};\boldsymbol{w}) = 1 - \sigma(\boldsymbol{w}^{\mathsf{T}}\boldsymbol{x}) P ( y = 1 ∣ x ; w ) = σ ( w T x ) , P ( y = 0 ∣ x ; w ) = 1 − σ ( w T x ) で定めるモデルをロジスティック回帰モデル という。z = w T x z = \boldsymbol{w}^{\mathsf{T}}\boldsymbol{x} z = w T x をロジット またはスコア 、μ = σ ( z ) \mu = \sigma(z) μ = σ ( z ) を予測確率 と呼ぶ。2 つの式はまとめて
P ( y ∣ x ; w ) = μ y ( 1 − μ ) 1 − y , μ = σ ( w T x ) P(y \mid \boldsymbol{x};\boldsymbol{w}) = \mu^{y}(1-\mu)^{1-y}, \qquad \mu = \sigma(\boldsymbol{w}^{\mathsf{T}}\boldsymbol{x}) P ( y ∣ x ; w ) = μ y ( 1 − μ ) 1 − y , μ = σ ( w T x ) と書ける(y = 1 y=1 y = 1 を代入すれば μ \mu μ 、y = 0 y=0 y = 0 を代入すれば 1 − μ 1-\mu 1 − μ になる)。すなわち y y y は成功確率 μ \mu μ のベルヌーイ分布に従う。
例 3.7 (係数はオッズの倍率である )
勉強時間 t t t から合格確率を予測するモデルが log μ 1 − μ = − 3 + 1.2 t \log\dfrac{\mu}{1-\mu} = -3 + 1.2\,t log 1 − μ μ = − 3 + 1.2 t と推定されたとします。係数 1.2 1.2 1.2 はどう読めばよいでしょうか。
t t t を 1 増やすと対数オッズが 1.2 1.2 1.2 増える、すなわちオッズが e 1.2 = 3.32 … e^{1.2} = 3.32\ldots e 1.2 = 3.32 … 倍になる 、というのが正確な読み方です。確認します。
t = 2 t = 2 t = 2 :z = − 0.6 z = -0.6 z = − 0.6 、μ = σ ( − 0.6 ) = 0.3543 \mu = \sigma(-0.6) = 0.3543 μ = σ ( − 0.6 ) = 0.3543 、オッズ = e − 0.6 = 0.5488 = e^{-0.6} = 0.5488 = e − 0.6 = 0.5488 。
t = 3 t = 3 t = 3 :z = 0.6 z = 0.6 z = 0.6 、μ = σ ( 0.6 ) = 0.6457 \mu = \sigma(0.6) = 0.6457 μ = σ ( 0.6 ) = 0.6457 、オッズ = e 0.6 = 1.8221 = e^{0.6} = 1.8221 = e 0.6 = 1.8221 。オッズ比は 1.8221 / 0.5488 = 3.320 = e 1.2 1.8221/0.5488 = 3.320 = e^{1.2} 1.8221/0.5488 = 3.320 = e 1.2 。
t = 5 t = 5 t = 5 :z = 3 z = 3 z = 3 、μ = 0.9526 \mu = 0.9526 μ = 0.9526 、オッズ = e 3 = 20.09 = e^{3} = 20.09 = e 3 = 20.09 。
t = 6 t = 6 t = 6 :z = 4.2 z = 4.2 z = 4.2 、μ = 0.9852 \mu = 0.9852 μ = 0.9852 、オッズ = e 4.2 = 66.69 = e^{4.2} = 66.69 = e 4.2 = 66.69 。オッズ比はやはり 66.69 / 20.09 = 3.320 = e 1.2 66.69/20.09 = 3.320 = e^{1.2} 66.69/20.09 = 3.320 = e 1.2 。
オッズ比はどこでも一定ですが、確率の増え方は一定ではありません。 t : 2 → 3 t : 2 \to 3 t : 2 → 3 では確率が 0.354 → 0.646 0.354 \to 0.646 0.354 → 0.646 (+ 0.29 +0.29 + 0.29 )と大きく動くのに対し、t : 5 → 6 t : 5 \to 6 t : 5 → 6 では 0.953 → 0.985 0.953 \to 0.985 0.953 → 0.985 (+ 0.03 +0.03 + 0.03 )しか動きません。すでに確率が 1 1 1 に近いところでは、オッズを 3 倍にしても確率はほとんど増えないからです。「係数 1.2 1.2 1.2 は確率を 1.2 1.2 1.2 増やす」という読み方は誤りです。
定義 3.6 はデータの生成規則を確率で書いたモデルです。こうしたモデルのパラメータを決める標準的な原理が最尤推定 、すなわち「手元のデータが最も起こりやすくなるパラメータを選ぶ」という方針です。
μ i ( w ) = σ ( w T x i ) \mu_i(\boldsymbol{w}) = \sigma(\boldsymbol{w}^{\mathsf{T}}\boldsymbol{x}_i) μ i ( w ) = σ ( w T x i ) と書きます。§2 の条件付き独立性の仮定から、観測されたラベル列 y 1 , … , y n y_1, \ldots, y_n y 1 , … , y n の同時確率、すなわち尤度 は積になります。
L ( w ) = ∏ i = 1 n P ( y i ∣ x i ; w ) = ∏ i = 1 n μ i y i ( 1 − μ i ) 1 − y i . \mathcal{L}(\boldsymbol{w}) = \prod_{i=1}^{n} P(y_i \mid \boldsymbol{x}_i;\boldsymbol{w}) = \prod_{i=1}^{n} \mu_i^{\,y_i}(1-\mu_i)^{1-y_i}. L ( w ) = i = 1 ∏ n P ( y i ∣ x i ; w ) = i = 1 ∏ n μ i y i ( 1 − μ i ) 1 − y i .
積のままでは扱いにくいので対数を取ります。log \log log は狭義単調増加なので、L \mathcal{L} L を最大にする w \boldsymbol{w} w と log L \log\mathcal{L} log L を最大にする w \boldsymbol{w} w は完全に一致します。さらに最適化の慣習に合わせて符号を反転させると、次の量が現れます。
定義 4.1 (交差エントロピー誤差(負の対数尤度) )
データ ( x i , y i ) i = 1 n (\boldsymbol{x}_i, y_i)_{i=1}^{n} ( x i , y i ) i = 1 n と 定義 3.6 に対し、μ i = σ ( w T x i ) \mu_i = \sigma(\boldsymbol{w}^{\mathsf{T}}\boldsymbol{x}_i) μ i = σ ( w T x i ) として
L ( w ) = − ∑ i = 1 n [ y i log μ i + ( 1 − y i ) log ( 1 − μ i ) ] L(\boldsymbol{w}) = -\sum_{i=1}^{n} \Bigl[\, y_i \log \mu_i + (1-y_i)\log(1-\mu_i) \,\Bigr] L ( w ) = − i = 1 ∑ n [ y i log μ i + ( 1 − y i ) log ( 1 − μ i ) ] を交差エントロピー誤差 、または負の対数尤度 という。0 < μ i < 1 0 < \mu_i < 1 0 < μ i < 1 (命題 3.2 (1))なので対数は常に定義され、L ( w ) > 0 L(\boldsymbol{w}) > 0 L ( w ) > 0 である。
「交差エントロピー」という呼び名は情報理論から来ています。有限集合上の 2 つの確率分布 p , q p, q p , q に対し H ( p , q ) = − ∑ k p k log q k H(p,q) = -\sum_{k} p_k \log q_k H ( p , q ) = − ∑ k p k log q k を交差エントロピーといいます。定義 4.1 の第 i i i 項は、ラベルが定める分布 p = ( 1 − y i , y i ) p = (1-y_i,\; y_i) p = ( 1 − y i , y i ) (y i y_i y i は 0 0 0 か 1 1 1 なので、これは一点に集中した分布です)とモデルの分布 q = ( 1 − μ i , μ i ) q = (1-\mu_i,\; \mu_i) q = ( 1 − μ i , μ i ) の交差エントロピーそのものです。
さらに H ( p , q ) = H ( p ) + D K L ( p ∥ q ) H(p,q) = H(p) + D_{\mathrm{KL}}(p \,\|\, q) H ( p , q ) = H ( p ) + D KL ( p ∥ q ) という分解があり、いまは p p p が一点分布なのでそのエントロピーは H ( p ) = 0 H(p) = 0 H ( p ) = 0 です。したがって
L ( w ) = ∑ i = 1 n D K L ( δ y i ∥ B e r ( μ i ) ) L(\boldsymbol{w}) = \sum_{i=1}^{n} D_{\mathrm{KL}}\bigl(\,\delta_{y_i} \,\big\|\, \mathrm{Ber}(\mu_i)\,\bigr) L ( w ) = i = 1 ∑ n D KL ( δ y i Ber ( μ i ) )
となります。交差エントロピー誤差を下げることは、モデルの予測分布を観測されたラベルの分布に近づけることと同じ だ、というのが情報理論側からの読み方です。
命題 4.2 (最尤推定と交差エントロピー最小化の同値性 )
L \mathcal{L} L を上の尤度、L L L を 定義 4.1 の交差エントロピー誤差とする。任意の w ∈ R d \boldsymbol{w} \in \mathbb{R}^{d} w ∈ R d について L ( w ) = − log L ( w ) L(\boldsymbol{w}) = -\log\mathcal{L}(\boldsymbol{w}) L ( w ) = − log L ( w ) が成り立ち、したがって集合として
arg max w ∈ R d L ( w ) = arg min w ∈ R d L ( w ) \operatorname*{arg\,max}_{\boldsymbol{w}\in\mathbb{R}^{d}} \mathcal{L}(\boldsymbol{w}) = \operatorname*{arg\,min}_{\boldsymbol{w}\in\mathbb{R}^{d}} L(\boldsymbol{w}) w ∈ R d arg max L ( w ) = w ∈ R d arg min L ( w ) である(両辺が空集合になることも含めて等号が成り立つ)。
証明(命題 4.2) L ( w ) = ∏ i μ i y i ( 1 − μ i ) 1 − y i \mathcal{L}(\boldsymbol{w}) = \prod_i \mu_i^{y_i}(1-\mu_i)^{1-y_i} L ( w ) = ∏ i μ i y i ( 1 − μ i ) 1 − y i の各因子は 命題 3.2 (1) より狭義正なので、L ( w ) > 0 \mathcal{L}(\boldsymbol{w}) > 0 L ( w ) > 0 であり対数が取れます。積の対数は対数の和なので
log L ( w ) = ∑ i = 1 n [ y i log μ i + ( 1 − y i ) log ( 1 − μ i ) ] = − L ( w ) . \log \mathcal{L}(\boldsymbol{w}) = \sum_{i=1}^{n}\Bigl[\, y_i \log\mu_i + (1-y_i)\log(1-\mu_i) \,\Bigr] = -L(\boldsymbol{w}). log L ( w ) = i = 1 ∑ n [ y i log μ i + ( 1 − y i ) log ( 1 − μ i ) ] = − L ( w ) . t ↦ − t t \mapsto -t t ↦ − t は R \mathbb{R} R 上の狭義単調減少な全単射なので、L ( w ) ≥ L ( w ′ ) \mathcal{L}(\boldsymbol{w}) \ge \mathcal{L}(\boldsymbol{w}') L ( w ) ≥ L ( w ′ ) と L ( w ) ≤ L ( w ′ ) L(\boldsymbol{w}) \le L(\boldsymbol{w}') L ( w ) ≤ L ( w ′ ) は同値です(log \log log の単調増加性と合わせて)。よって L \mathcal{L} L の最大点全体と L L L の最小点全体は集合として一致します。
∎
つまり 損失関数は設計するものではなく、確率モデルから導かれるもの です。§1.2 で二乗誤差がうまくいかなかったのは、二乗誤差が「出力が正規分布に従う」という別のモデルの負の対数尤度だったからだ、と言い換えることもできます(ガウス雑音の下での最尤推定(命題 3.3)[確率論とベイズ統計] )。0 / 1 0/1 0/1 のラベルは正規分布に従いません。
ヒント
実装では log ( 1 + e z ) \log(1+e^{z}) log ( 1 + e z ) を素直に計算すると z z z が大きいときに e z e^{z} e z が桁溢れします。定義 4.1 の第 i i i 項は log ( 1 + e z i ) − y i z i \log(1+e^{z_i}) - y_i z_i log ( 1 + e z i ) − y i z i と 1 本にまとめられ(log μ = − log ( 1 + e − z ) \log\mu = -\log(1+e^{-z}) log μ = − log ( 1 + e − z ) 、log ( 1 − μ ) = − log ( 1 + e z ) \log(1-\mu) = -\log(1+e^{z}) log ( 1 − μ ) = − log ( 1 + e z ) 、および log ( 1 + e − z ) = log ( 1 + e z ) − z \log(1+e^{-z}) = \log(1+e^{z}) - z log ( 1 + e − z ) = log ( 1 + e z ) − z から従います)、さらに log ( 1 + e z ) = max ( z , 0 ) + log ( 1 + e − ∣ z ∣ ) \log(1+e^{z}) = \max(z,0) + \log\bigl(1+e^{-|z|}\bigr) log ( 1 + e z ) = max ( z , 0 ) + log ( 1 + e − ∣ z ∣ ) と書き換えれば指数の引数が常に 0 0 0 以下になり、安全に計算できます。
例 4.3 (切片だけのモデルは陽に解ける )
特徴を使わず切片だけを持つモデル、すなわち d = 1 d = 1 d = 1 で x i = ( 1 ) \boldsymbol{x}_i = (1) x i = ( 1 ) の場合を考えます。μ i = σ ( b ) \mu_i = \sigma(b) μ i = σ ( b ) は i i i によらない定数です。n n n 個のうち k k k 個が y i = 1 y_i = 1 y i = 1 だとすると
L ( b ) = − [ k log σ ( b ) + ( n − k ) log ( 1 − σ ( b ) ) ] . L(b) = -\bigl[\, k \log\sigma(b) + (n-k)\log(1-\sigma(b)) \,\bigr]. L ( b ) = − [ k log σ ( b ) + ( n − k ) log ( 1 − σ ( b )) ] . 微分します。命題 3.2 (3) より d d b log σ ( b ) = σ ′ ( b ) σ ( b ) = 1 − σ ( b ) \dfrac{d}{db}\log\sigma(b) = \dfrac{\sigma'(b)}{\sigma(b)} = 1-\sigma(b) d b d log σ ( b ) = σ ( b ) σ ′ ( b ) = 1 − σ ( b ) であり、同じく d d b log ( 1 − σ ( b ) ) = − σ ′ ( b ) 1 − σ ( b ) = − σ ( b ) \dfrac{d}{db}\log(1-\sigma(b)) = \dfrac{-\sigma'(b)}{1-\sigma(b)} = -\sigma(b) d b d log ( 1 − σ ( b )) = 1 − σ ( b ) − σ ′ ( b ) = − σ ( b ) です。よって
L ′ ( b ) = − [ k ( 1 − σ ( b ) ) − ( n − k ) σ ( b ) ] = n σ ( b ) − k . L'(b) = -\bigl[\, k(1-\sigma(b)) - (n-k)\sigma(b) \,\bigr] = n\,\sigma(b) - k . L ′ ( b ) = − [ k ( 1 − σ ( b )) − ( n − k ) σ ( b ) ] = n σ ( b ) − k . L ′ ( b ) = 0 L'(b) = 0 L ′ ( b ) = 0 は σ ( b ) = k / n \sigma(b) = k/n σ ( b ) = k / n と同値です。0 < k < n 0 < k < n 0 < k < n なら k / n ∈ ( 0 , 1 ) k/n \in (0,1) k / n ∈ ( 0 , 1 ) なので 命題 3.2 (4) により解が一意に存在し
b ^ = logit ( k n ) = log k n − k . \hat{b} = \operatorname{logit}\!\left(\frac{k}{n}\right) = \log\frac{k}{n-k}. b ^ = logit ( n k ) = log n − k k . たとえば n = 100 n = 100 n = 100 、k = 30 k = 30 k = 30 なら b ^ = log ( 30 / 70 ) = log ( 3 / 7 ) = − 0.8473 \hat{b} = \log(30/70) = \log(3/7) = -0.8473 b ^ = log ( 30/70 ) = log ( 3/7 ) = − 0.8473 で、予測確率は σ ( − 0.8473 ) = 0.30 \sigma(-0.8473) = 0.30 σ ( − 0.8473 ) = 0.30 、すなわち経験的な正例率そのもの です。最尤推定が「素直な答え」を返していることが確認できます。
一方 k = 0 k = 0 k = 0 または k = n k = n k = n のときは k / n k/n k / n が ( 0 , 1 ) (0,1) ( 0 , 1 ) の外にあるので L ′ ( b ) = 0 L'(b) = 0 L ′ ( b ) = 0 に解はありません。k = n k = n k = n なら L ( b ) = n log ( 1 + e − b ) L(b) = n\log(1+e^{-b}) L ( b ) = n log ( 1 + e − b ) は b → + ∞ b \to +\infty b → + ∞ で 0 0 0 に近づきますが、決して 0 0 0 になりません。最尤推定量が存在しないのです。これは 定理 6.3 の最も簡単な場合です。
例 4.3 ではパラメータが 1 個だったので微分して解けました。一般の d d d ではどうなるでしょうか。まず勾配を計算します。
定理 5.1 (交差エントロピー誤差の勾配 )
x 1 , … , x n ∈ R d \boldsymbol{x}_1,\ldots,\boldsymbol{x}_n \in \mathbb{R}^{d} x 1 , … , x n ∈ R d 、y 1 , … , y n ∈ { 0 , 1 } y_1,\ldots,y_n \in \{0,1\} y 1 , … , y n ∈ { 0 , 1 } を任意に固定し、L L L を 定義 4.1 の交差エントロピー誤差、μ i ( w ) = σ ( w T x i ) \mu_i(\boldsymbol{w}) = \sigma(\boldsymbol{w}^{\mathsf{T}}\boldsymbol{x}_i) μ i ( w ) = σ ( w T x i ) とする。このとき L L L は R d \mathbb{R}^{d} R d 上 C ∞ C^{\infty} C ∞ 級であり、
∇ L ( w ) = ∑ i = 1 n ( μ i ( w ) − y i ) x i = X T ( μ ( w ) − y ) \nabla L(\boldsymbol{w}) = \sum_{i=1}^{n} \bigl(\mu_i(\boldsymbol{w}) - y_i\bigr)\,\boldsymbol{x}_i
= X^{\mathsf{T}}\bigl(\boldsymbol{\mu}(\boldsymbol{w}) - \boldsymbol{y}\bigr) ∇ L ( w ) = i = 1 ∑ n ( μ i ( w ) − y i ) x i = X T ( μ ( w ) − y ) が成り立つ。ここで X X X は第 i i i 行が x i T \boldsymbol{x}_i^{\mathsf{T}} x i T の n × d n\times d n × d 計画行列、μ ( w ) = ( μ 1 , … , μ n ) T \boldsymbol{\mu}(\boldsymbol{w}) = (\mu_1,\ldots,\mu_n)^{\mathsf{T}} μ ( w ) = ( μ 1 , … , μ n ) T である。
証明(定理 5.1) z i = w T x i = ∑ j = 1 d w j x i j z_i = \boldsymbol{w}^{\mathsf{T}}\boldsymbol{x}_i = \sum_{j=1}^{d} w_j x_{ij} z i = w T x i = ∑ j = 1 d w j x ij 、μ i = σ ( z i ) \mu_i = \sigma(z_i) μ i = σ ( z i ) と置きます。z i z_i z i は w \boldsymbol{w} w の 1 次式なので C ∞ C^{\infty} C ∞ 級、σ \sigma σ も C ∞ C^{\infty} C ∞ 級(命題 3.2 (1))、log \log log は ( 0 , ∞ ) (0,\infty) ( 0 , ∞ ) 上 C ∞ C^{\infty} C ∞ 級で μ i , 1 − μ i ∈ ( 0 , 1 ) \mu_i, 1-\mu_i \in (0,1) μ i , 1 − μ i ∈ ( 0 , 1 ) なので、合成と有限和として L L L は C ∞ C^{\infty} C ∞ 級です。
第 i i i 項を ℓ i = − [ y i log μ i + ( 1 − y i ) log ( 1 − μ i ) ] \ell_i = -\bigl[y_i\log\mu_i + (1-y_i)\log(1-\mu_i)\bigr] ℓ i = − [ y i log μ i + ( 1 − y i ) log ( 1 − μ i ) ] と書き、連鎖律を μ i → z i → w j \mu_i \to z_i \to w_j μ i → z i → w j の順に適用します。
第 1 段:μ i \mu_i μ i についての微分。
∂ ℓ i ∂ μ i = − y i μ i + 1 − y i 1 − μ i = − y i ( 1 − μ i ) + μ i ( 1 − y i ) μ i ( 1 − μ i ) = − y i + y i μ i + μ i − μ i y i μ i ( 1 − μ i ) = μ i − y i μ i ( 1 − μ i ) . \frac{\partial \ell_i}{\partial \mu_i} = -\frac{y_i}{\mu_i} + \frac{1-y_i}{1-\mu_i}
= \frac{-y_i(1-\mu_i) + \mu_i(1-y_i)}{\mu_i(1-\mu_i)}
= \frac{-y_i + y_i\mu_i + \mu_i - \mu_i y_i}{\mu_i(1-\mu_i)}
= \frac{\mu_i - y_i}{\mu_i(1-\mu_i)} . ∂ μ i ∂ ℓ i = − μ i y i + 1 − μ i 1 − y i = μ i ( 1 − μ i ) − y i ( 1 − μ i ) + μ i ( 1 − y i ) = μ i ( 1 − μ i ) − y i + y i μ i + μ i − μ i y i = μ i ( 1 − μ i ) μ i − y i . 途中で分母を μ i ( 1 − μ i ) \mu_i(1-\mu_i) μ i ( 1 − μ i ) に通分し、分子の y i μ i y_i\mu_i y i μ i と − μ i y i -\mu_i y_i − μ i y i が打ち消し合うことを使いました。
第 2 段:z i z_i z i についての微分。 命題 3.2 (3) より d μ i d z i = σ ′ ( z i ) = μ i ( 1 − μ i ) \dfrac{d\mu_i}{dz_i} = \sigma'(z_i) = \mu_i(1-\mu_i) d z i d μ i = σ ′ ( z i ) = μ i ( 1 − μ i ) 。
第 3 段:w j w_j w j についての微分。 z i = ∑ j w j x i j z_i = \sum_{j} w_j x_{ij} z i = ∑ j w j x ij より ∂ z i ∂ w j = x i j \dfrac{\partial z_i}{\partial w_j} = x_{ij} ∂ w j ∂ z i = x ij 。
3 つを掛け合わせると、第 1 段の分母 μ i ( 1 − μ i ) \mu_i(1-\mu_i) μ i ( 1 − μ i ) と第 2 段の因子 μ i ( 1 − μ i ) \mu_i(1-\mu_i) μ i ( 1 − μ i ) が約分 して
∂ ℓ i ∂ w j = μ i − y i μ i ( 1 − μ i ) ⋅ μ i ( 1 − μ i ) ⋅ x i j = ( μ i − y i ) x i j . \frac{\partial \ell_i}{\partial w_j} = \frac{\mu_i - y_i}{\mu_i(1-\mu_i)} \cdot \mu_i(1-\mu_i) \cdot x_{ij} = (\mu_i - y_i)\,x_{ij}. ∂ w j ∂ ℓ i = μ i ( 1 − μ i ) μ i − y i ⋅ μ i ( 1 − μ i ) ⋅ x ij = ( μ i − y i ) x ij . 約分が正当なのは μ i ( 1 − μ i ) ≠ 0 \mu_i(1-\mu_i) \ne 0 μ i ( 1 − μ i ) = 0 だからで、これは 命題 3.2 (1) の 0 < μ i < 1 0 < \mu_i < 1 0 < μ i < 1 から従います。i i i について和を取り、j = 1 , … , d j = 1,\ldots,d j = 1 , … , d を並べれば
∇ L ( w ) = ∑ i = 1 n ( μ i − y i ) x i . \nabla L(\boldsymbol{w}) = \sum_{i=1}^{n}(\mu_i - y_i)\boldsymbol{x}_i . ∇ L ( w ) = i = 1 ∑ n ( μ i − y i ) x i . 最後に、x i \boldsymbol{x}_i x i が X X X の第 i i i 行であることから ∑ i ( μ i − y i ) x i = X T ( μ − y ) \sum_i (\mu_i - y_i)\boldsymbol{x}_i = X^{\mathsf{T}}(\boldsymbol{\mu}-\boldsymbol{y}) ∑ i ( μ i − y i ) x i = X T ( μ − y ) です(X T X^{\mathsf{T}} X T の列が x i \boldsymbol{x}_i x i なので、X T X^{\mathsf{T}} X T とベクトルの積は列の線形結合になります)。
∎
この公式は形が線形回帰と瓜二つです。最小二乗法の勾配は X T ( X w − y ) X^{\mathsf{T}}(X\boldsymbol{w} - \boldsymbol{y}) X T ( X w − y ) でした。違いは予測が X w X\boldsymbol{w} X w から σ ( X w ) \sigma(X\boldsymbol{w}) σ ( X w ) に変わっただけです。「残差(予測 − - − 実測)を特徴で重み付けして足す 」という構造は共通しています。
系 5.2 (切片を含むモデルの平均較正 )
モデルが切片を含む、すなわちある j 0 j_0 j 0 についてすべての i i i で x i j 0 = 1 x_{i j_0} = 1 x i j 0 = 1 であるとする。このとき ∇ L ( w ∗ ) = 0 \nabla L(\boldsymbol{w}^{*}) = \boldsymbol{0} ∇ L ( w ∗ ) = 0 を満たす任意の w ∗ \boldsymbol{w}^{*} w ∗ について
1 n ∑ i = 1 n μ i ( w ∗ ) = 1 n ∑ i = 1 n y i \frac{1}{n}\sum_{i=1}^{n} \mu_i(\boldsymbol{w}^{*}) = \frac{1}{n}\sum_{i=1}^{n} y_i n 1 i = 1 ∑ n μ i ( w ∗ ) = n 1 i = 1 ∑ n y i が成り立つ。すなわち予測確率の平均は、データ中の正例の割合に一致する。
証明(系 5.2) 定理 5.1 より ∇ L \nabla L ∇ L の第 j 0 j_0 j 0 成分は ∑ i ( μ i − y i ) x i j 0 \sum_{i}(\mu_i - y_i)x_{ij_0} ∑ i ( μ i − y i ) x i j 0 です。仮定より x i j 0 = 1 x_{ij_0} = 1 x i j 0 = 1 なので、これは ∑ i ( μ i − y i ) \sum_i (\mu_i - y_i) ∑ i ( μ i − y i ) に等しくなります。∇ L ( w ∗ ) = 0 \nabla L(\boldsymbol{w}^{*}) = \boldsymbol{0} ∇ L ( w ∗ ) = 0 よりこの成分も 0 0 0 、すなわち ∑ i μ i = ∑ i y i \sum_i \mu_i = \sum_i y_i ∑ i μ i = ∑ i y i です。両辺を n n n で割れば主張を得ます。
∎
系 5.2 は、最尤推定されたロジスティック回帰が「全体としては当たっている」ことを保証します。100 人について平均 0.3 0.3 0.3 の確率を出したなら、実際に 30 人が正例だったということです。例 4.3 はこの系の d = 1 d=1 d = 1 の場合にほかなりません。
定理 5.1 の証明で起きた約分は、偶然ではありません。シグモイドと交差エントロピーは、そう組み合わせるために選ばれた対 です。二乗誤差と組み合わせるとどうなるかを見れば、その意味がはっきりします。
例 5.3 (二乗誤差だと勾配が消え、しかも凸でなくなる )
同じモデル μ i = σ ( w T x i ) \mu_i = \sigma(\boldsymbol{w}^{\mathsf{T}}\boldsymbol{x}_i) μ i = σ ( w T x i ) に二乗誤差 E ( w ) = 1 2 ∑ i ( μ i − y i ) 2 E(\boldsymbol{w}) = \frac12\sum_i (\mu_i - y_i)^2 E ( w ) = 2 1 ∑ i ( μ i − y i ) 2 を使うと、定理 5.1 の証明の第 1 段だけが変わって ∂ E / ∂ μ i = μ i − y i \partial E/\partial\mu_i = \mu_i - y_i ∂ E / ∂ μ i = μ i − y i となり、第 2 段の μ i ( 1 − μ i ) \mu_i(1-\mu_i) μ i ( 1 − μ i ) は約分されずに残ります。
∂ E ∂ w j = ∑ i = 1 n ( μ i − y i ) μ i ( 1 − μ i ) x i j . \frac{\partial E}{\partial w_j} = \sum_{i=1}^{n} (\mu_i - y_i)\,\mu_i(1-\mu_i)\,x_{ij}. ∂ w j ∂ E = i = 1 ∑ n ( μ i − y i ) μ i ( 1 − μ i ) x ij . 余分な因子 μ i ( 1 − μ i ) \mu_i(1-\mu_i) μ i ( 1 − μ i ) が何をするかを、d = 1 d = 1 d = 1 、x = 1 x = 1 x = 1 、y = 1 y = 1 y = 1 、w = − 10 w = -10 w = − 10 という「自信を持って間違えている」1 点で見ます。μ = σ ( − 10 ) = 4.5398 × 10 − 5 \mu = \sigma(-10) = 4.5398\times 10^{-5} μ = σ ( − 10 ) = 4.5398 × 1 0 − 5 なので
交差エントロピーの勾配:μ − y = − 0.99995 \mu - y = -0.99995 μ − y = − 0.99995 。
二乗誤差の勾配:( μ − y ) μ ( 1 − μ ) = − 4.5394 × 10 − 5 (\mu-y)\mu(1-\mu) = -4.5394\times 10^{-5} ( μ − y ) μ ( 1 − μ ) = − 4.5394 × 1 0 − 5 。
その比は約 22028 22028 22028 倍です。最も大きく間違えている点で、二乗誤差はほとんど何も学習しません。 シグモイドが飽和して σ ′ ≈ 0 \sigma' \approx 0 σ ′ ≈ 0 になるからで、これが「勾配消失」と呼ばれる現象の最も単純な形です。
さらに悪いことに、この E E E は凸でさえありません。同じ 1 点の設定で s = σ ( − w ) = 1 − μ s = \sigma(-w) = 1-\mu s = σ ( − w ) = 1 − μ と置くと E ( w ) = 1 2 s 2 E(w) = \frac12 s^2 E ( w ) = 2 1 s 2 であり、d s d w = − σ ′ ( − w ) = − s ( 1 − s ) \dfrac{ds}{dw} = -\sigma'(-w) = -s(1-s) d w d s = − σ ′ ( − w ) = − s ( 1 − s ) (命題 3.2 (2),(3))を使って
E ′ ( w ) = s ⋅ d s d w = − s 2 ( 1 − s ) , E ′ ′ ( w ) = ( − 2 s + 3 s 2 ) ⋅ d s d w = s 2 ( 1 − s ) ( 2 − 3 s ) . E'(w) = s\cdot\frac{ds}{dw} = -s^{2}(1-s), \qquad
E''(w) = \bigl(-2s + 3s^{2}\bigr)\cdot\frac{ds}{dw} = s^{2}(1-s)(2-3s). E ′ ( w ) = s ⋅ d w d s = − s 2 ( 1 − s ) , E ′′ ( w ) = ( − 2 s + 3 s 2 ) ⋅ d w d s = s 2 ( 1 − s ) ( 2 − 3 s ) . s ∈ ( 0 , 1 ) s \in (0,1) s ∈ ( 0 , 1 ) なので E ′ ′ E'' E ′′ の符号は 2 − 3 s 2-3s 2 − 3 s の符号、すなわち s < 2 / 3 s < 2/3 s < 2/3 かどうかで決まります。s = 1 / 2 s = 1/2 s = 1/2 (w = 0 w=0 w = 0 )では E ′ ′ = 0.0625 > 0 E'' = 0.0625 > 0 E ′′ = 0.0625 > 0 ですが、s = 0.9 s = 0.9 s = 0.9 (w = − log 9 = − 2.197 w = -\log 9 = -2.197 w = − log 9 = − 2.197 )では E ′ ′ = − 0.0567 < 0 E'' = -0.0567 < 0 E ′′ = − 0.0567 < 0 です。変曲点 w = − log 2 w = -\log 2 w = − log 2 をまたいで凸性が入れ替わります。
同じ 1 点で交差エントロピーは L ( w ) = − log σ ( w ) = log ( 1 + e − w ) L(w) = -\log\sigma(w) = \log(1+e^{-w}) L ( w ) = − log σ ( w ) = log ( 1 + e − w ) であり、L ′ ( w ) = − ( 1 − σ ( w ) ) L'(w) = -(1-\sigma(w)) L ′ ( w ) = − ( 1 − σ ( w )) 、L ′ ′ ( w ) = σ ( w ) ( 1 − σ ( w ) ) > 0 L''(w) = \sigma(w)(1-\sigma(w)) > 0 L ′′ ( w ) = σ ( w ) ( 1 − σ ( w )) > 0 なので狭義凸です。しかも w → − ∞ w \to -\infty w → − ∞ で L ′ ( w ) → − 1 L'(w) \to -1 L ′ ( w ) → − 1 と、勾配が消えません。
勾配が求まったので、最尤推定量は停留条件
X T ( σ ( X w ) − y ) = 0 X^{\mathsf{T}}\bigl(\sigma(X\boldsymbol{w}) - \boldsymbol{y}\bigr) = \boldsymbol{0} X T ( σ ( X w ) − y ) = 0
を満たすはずです(σ \sigma σ は成分ごとに作用させます)。ここが線形回帰との決定的な分かれ目です。
例 5.5 (勾配を 1 ステップ手で回す )
§1.2 のデータ(t = 1 , 2 , 3 , 4 t = 1,2,3,4 t = 1 , 2 , 3 , 4 、y = 0 , 0 , 1 , 1 y = 0,0,1,1 y = 0 , 0 , 1 , 1 )に切片付きで当てはめます。x i = ( 1 , t i ) T \boldsymbol{x}_i = (1, t_i)^{\mathsf{T}} x i = ( 1 , t i ) T 、w = ( b , a ) T \boldsymbol{w} = (b, a)^{\mathsf{T}} w = ( b , a ) T です。
初期点 w = ( 0 , 0 ) \boldsymbol{w} = (0,0) w = ( 0 , 0 ) 。 z i = 0 z_i = 0 z i = 0 なので μ i = σ ( 0 ) = 0.5 \mu_i = \sigma(0) = 0.5 μ i = σ ( 0 ) = 0.5 (命題 3.2 (2))。損失は
L ( 0 ) = − ∑ i = 1 4 log 0.5 = 4 log 2 = 2.7726. L(\boldsymbol{0}) = -\sum_{i=1}^{4}\log 0.5 = 4\log 2 = 2.7726 . L ( 0 ) = − i = 1 ∑ 4 log 0.5 = 4 log 2 = 2.7726. 残差は μ − y = ( 0.5 , 0.5 , − 0.5 , − 0.5 ) \boldsymbol{\mu}-\boldsymbol{y} = (0.5,\, 0.5,\, -0.5,\, -0.5) μ − y = ( 0.5 , 0.5 , − 0.5 , − 0.5 ) なので、定理 5.1 より
∇ L ( 0 ) = ( 0.5 + 0.5 − 0.5 − 0.5 0.5 ⋅ 1 + 0.5 ⋅ 2 − 0.5 ⋅ 3 − 0.5 ⋅ 4 ) = ( 0 − 2 ) . \nabla L(\boldsymbol{0}) = \begin{pmatrix} 0.5+0.5-0.5-0.5 \\ 0.5\cdot 1 + 0.5\cdot 2 - 0.5\cdot 3 - 0.5\cdot 4\end{pmatrix} = \begin{pmatrix} 0 \\ -2 \end{pmatrix}. ∇ L ( 0 ) = ( 0.5 + 0.5 − 0.5 − 0.5 0.5 ⋅ 1 + 0.5 ⋅ 2 − 0.5 ⋅ 3 − 0.5 ⋅ 4 ) = ( 0 − 2 ) . 切片方向の成分が 0 0 0 なのは 系 5.2 の通りで、μ ˉ = 0.5 = y ˉ \bar{\mu} = 0.5 = \bar{y} μ ˉ = 0.5 = y ˉ だからです。傾き方向は負なので、a a a を増やせば損失が減ります。
1 歩進める。 学習率 η = 0.1 \eta = 0.1 η = 0.1 で w ← w − η ∇ L ( w ) = ( 0 , 0.2 ) \boldsymbol{w} \leftarrow \boldsymbol{w} - \eta\nabla L(\boldsymbol{w}) = (0,\, 0.2) w ← w − η ∇ L ( w ) = ( 0 , 0.2 ) とします。z i = 0.2 , 0.4 , 0.6 , 0.8 z_i = 0.2, 0.4, 0.6, 0.8 z i = 0.2 , 0.4 , 0.6 , 0.8 、μ i = 0.5498 , 0.5987 , 0.6457 , 0.6900 \mu_i = 0.5498, 0.5987, 0.6457, 0.6900 μ i = 0.5498 , 0.5987 , 0.6457 , 0.6900 となり
L = − [ log 0.4502 + log 0.4013 + log 0.6457 + log 0.6900 ] = 0.7981 + 0.9130 + 0.4375 + 0.3711 = 2.5197. L = -\bigl[\log 0.4502 + \log 0.4013 + \log 0.6457 + \log 0.6900\bigr] = 0.7981+0.9130+0.4375+0.3711 = 2.5197 . L = − [ log 0.4502 + log 0.4013 + log 0.6457 + log 0.6900 ] = 0.7981 + 0.9130 + 0.4375 + 0.3711 = 2.5197. 確かに 2.7726 2.7726 2.7726 から減りました。新しい勾配は μ − y = ( 0.5498 , 0.5987 , − 0.3543 , − 0.3100 ) \boldsymbol{\mu}-\boldsymbol{y} = (0.5498,\,0.5987,\,-0.3543,\,-0.3100) μ − y = ( 0.5498 , 0.5987 , − 0.3543 , − 0.3100 ) から
∇ L = ( 0.5498 + 0.5987 − 0.3543 − 0.3100 0.5498 + 1.1974 − 1.0630 − 1.2401 ) = ( 0.4842 − 0.5559 ) \nabla L = \begin{pmatrix} 0.5498+0.5987-0.3543-0.3100 \\ 0.5498+1.1974-1.0630-1.2401 \end{pmatrix} = \begin{pmatrix} 0.4842 \\ -0.5559 \end{pmatrix} ∇ L = ( 0.5498 + 0.5987 − 0.3543 − 0.3100 0.5498 + 1.1974 − 1.0630 − 1.2401 ) = ( 0.4842 − 0.5559 ) です。今度は切片成分が正になりました。傾きだけを上げたせいで全体の予測が押し上げられ、平均較正が崩れたためです。次の歩は切片を下げつつ傾きを上げる方向に進みます。
上の計算を素直にコードにすると次のようになります。§4.2 の注意に従い、損失は log ( 1 + e z ) − y z \log(1+e^{z}) - yz log ( 1 + e z ) − y z の形にまとめ、さらに log ( 1 + e z ) = max ( z , 0 ) + log ( 1 + e − ∣ z ∣ ) \log(1+e^{z}) = \max(z,0)+\log(1+e^{-|z|}) log ( 1 + e z ) = max ( z , 0 ) + log ( 1 + e − ∣ z ∣ ) と書き換えて桁溢れを避けています。
def softplus ( z ) : # log(1 + exp(z)) を安全に計算
return np. maximum ( z , 0.0 ) + np. log1p ( np. exp ( - np. abs ( z )))
return float ( np. sum ( softplus ( z ) - y * z ))
def grad ( w , X , y ) : # 勾配の公式そのもの
mu = 1.0 / ( 1.0 + np. exp ( - (X @ w) ))
X = np. array ( [[ 1.0 , 1.0 ] , [ 1.0 , 2.0 ] , [ 1.0 , 3.0 ] , [ 1.0 , 4.0 ]] )
y = np. array ( [ 0.0 , 0.0 , 1.0 , 1.0 ] )
print ( loss ( w , X , y ) , grad ( w , X , y )) # 2.772588722239781 [ 0. -2.]
w = w - 0.1 * grad ( w , X , y )
出力される損失は 2.7726 → 2.5197 → 2.4706 → 2.4305 2.7726 \to 2.5197 \to 2.4706 \to 2.4305 2.7726 → 2.5197 → 2.4706 → 2.4305 と単調に減っていきます。
数値的に探すと決めたなら、次に確かめるべきは「探して見つかるのか」です。一般の関数では、勾配が 0 0 0 になる点が局所最小・局所最大・鞍点のどれかは分かりません。しかし交差エントロピー誤差にはよい性質があります。
定理 6.1 (交差エントロピー誤差の凸性 )
定理 5.1 と同じ設定のもとで、S ( w ) = diag ( μ 1 ( 1 − μ 1 ) , … , μ n ( 1 − μ n ) ) S(\boldsymbol{w}) = \operatorname{diag}\bigl(\mu_1(1-\mu_1), \ldots, \mu_n(1-\mu_n)\bigr) S ( w ) = diag ( μ 1 ( 1 − μ 1 ) , … , μ n ( 1 − μ n ) ) と置くと
∇ 2 L ( w ) = ∑ i = 1 n μ i ( 1 − μ i ) x i x i T = X T S ( w ) X \nabla^{2} L(\boldsymbol{w}) = \sum_{i=1}^{n} \mu_i(1-\mu_i)\,\boldsymbol{x}_i\boldsymbol{x}_i^{\mathsf{T}} = X^{\mathsf{T}}S(\boldsymbol{w})X ∇ 2 L ( w ) = i = 1 ∑ n μ i ( 1 − μ i ) x i x i T = X T S ( w ) X が成り立つ。この行列はすべての w \boldsymbol{w} w で半正定値であり、したがって L L L は R d \mathbb{R}^{d} R d 上の凸関数である。さらに rank X = d \operatorname{rank} X = d rank X = d (X X X の列が線形独立)ならば ∇ 2 L ( w ) ≻ O \nabla^{2}L(\boldsymbol{w}) \succ O ∇ 2 L ( w ) ≻ O がすべての w \boldsymbol{w} w で成り立ち、L L L は狭義凸である。
証明(定理 6.1) ヘッセ行列の計算。 定理 5.1 より ∂ L ∂ w j = ∑ i ( μ i − y i ) x i j \dfrac{\partial L}{\partial w_j} = \sum_i (\mu_i - y_i)x_{ij} ∂ w j ∂ L = ∑ i ( μ i − y i ) x ij です。y i y_i y i は定数なので、これをさらに w k w_k w k で微分すると μ i \mu_i μ i だけが効いて
∂ 2 L ∂ w j ∂ w k = ∑ i = 1 n ∂ μ i ∂ w k x i j = ∑ i = 1 n σ ′ ( z i ) ∂ z i ∂ w k x i j = ∑ i = 1 n μ i ( 1 − μ i ) x i k x i j \frac{\partial^{2} L}{\partial w_j \partial w_k} = \sum_{i=1}^{n} \frac{\partial \mu_i}{\partial w_k}\, x_{ij}
= \sum_{i=1}^{n} \sigma'(z_i)\,\frac{\partial z_i}{\partial w_k}\, x_{ij}
= \sum_{i=1}^{n} \mu_i(1-\mu_i)\, x_{ik} x_{ij} ∂ w j ∂ w k ∂ 2 L = i = 1 ∑ n ∂ w k ∂ μ i x ij = i = 1 ∑ n σ ′ ( z i ) ∂ w k ∂ z i x ij = i = 1 ∑ n μ i ( 1 − μ i ) x ik x ij となります。2 番目の等号で連鎖律、3 番目で 命題 3.2 (3) と ∂ z i / ∂ w k = x i k \partial z_i/\partial w_k = x_{ik} ∂ z i / ∂ w k = x ik を使いました。x i j x i k x_{ij}x_{ik} x ij x ik は行列 x i x i T \boldsymbol{x}_i\boldsymbol{x}_i^{\mathsf{T}} x i x i T の ( j , k ) (j,k) ( j , k ) 成分なので、行列としてまとめると ∇ 2 L = ∑ i μ i ( 1 − μ i ) x i x i T \nabla^2 L = \sum_i \mu_i(1-\mu_i)\boldsymbol{x}_i\boldsymbol{x}_i^{\mathsf{T}} ∇ 2 L = ∑ i μ i ( 1 − μ i ) x i x i T です。X X X の第 i i i 行が x i T \boldsymbol{x}_i^{\mathsf{T}} x i T であることから、これは X T S X X^{\mathsf{T}}S X X T S X に等しくなります。
半正定値性。 任意の v ∈ R d \boldsymbol{v}\in\mathbb{R}^{d} v ∈ R d について
v T ∇ 2 L ( w ) v = ∑ i = 1 n μ i ( 1 − μ i ) v T x i x i T v = ∑ i = 1 n μ i ( 1 − μ i ) ( x i T v ) 2 ≥ 0 \boldsymbol{v}^{\mathsf{T}}\nabla^{2}L(\boldsymbol{w})\boldsymbol{v}
= \sum_{i=1}^{n}\mu_i(1-\mu_i)\,\boldsymbol{v}^{\mathsf{T}}\boldsymbol{x}_i\boldsymbol{x}_i^{\mathsf{T}}\boldsymbol{v}
= \sum_{i=1}^{n}\mu_i(1-\mu_i)\,(\boldsymbol{x}_i^{\mathsf{T}}\boldsymbol{v})^{2} \;\ge\; 0 v T ∇ 2 L ( w ) v = i = 1 ∑ n μ i ( 1 − μ i ) v T x i x i T v = i = 1 ∑ n μ i ( 1 − μ i ) ( x i T v ) 2 ≥ 0 です。各項が非負なのは、命題 3.2 (1) より 0 < μ i < 1 0 < \mu_i < 1 0 < μ i < 1 すなわち μ i ( 1 − μ i ) > 0 \mu_i(1-\mu_i) > 0 μ i ( 1 − μ i ) > 0 であり、( x i T v ) 2 ≥ 0 (\boldsymbol{x}_i^{\mathsf{T}}\boldsymbol{v})^2 \ge 0 ( x i T v ) 2 ≥ 0 だからです。
凸性。 w 0 , w 1 ∈ R d \boldsymbol{w}_0, \boldsymbol{w}_1 \in \mathbb{R}^{d} w 0 , w 1 ∈ R d を任意に取り、h = w 1 − w 0 \boldsymbol{h} = \boldsymbol{w}_1 - \boldsymbol{w}_0 h = w 1 − w 0 、g ( t ) = L ( w 0 + t h ) g(t) = L(\boldsymbol{w}_0 + t\boldsymbol{h}) g ( t ) = L ( w 0 + t h ) と置きます。L L L は C ∞ C^{\infty} C ∞ 級(定理 5.1 )なので g g g は R \mathbb{R} R 上 C 2 C^{2} C 2 級で、連鎖律より g ′ ′ ( t ) = h T ∇ 2 L ( w 0 + t h ) h ≥ 0 g''(t) = \boldsymbol{h}^{\mathsf{T}}\nabla^{2}L(\boldsymbol{w}_0+t\boldsymbol{h})\boldsymbol{h} \ge 0 g ′′ ( t ) = h T ∇ 2 L ( w 0 + t h ) h ≥ 0 です。1 変数関数の 2 階導関数が非負なら凸なので g g g は [ 0 , 1 ] [0,1] [ 0 , 1 ] 上凸で、g ( t ) ≤ ( 1 − t ) g ( 0 ) + t g ( 1 ) g(t) \le (1-t)g(0) + t\,g(1) g ( t ) ≤ ( 1 − t ) g ( 0 ) + t g ( 1 ) 、すなわち
L ( ( 1 − t ) w 0 + t w 1 ) ≤ ( 1 − t ) L ( w 0 ) + t L ( w 1 ) ( 0 ≤ t ≤ 1 ) L\bigl((1-t)\boldsymbol{w}_0 + t\boldsymbol{w}_1\bigr) \le (1-t)L(\boldsymbol{w}_0) + t\,L(\boldsymbol{w}_1) \qquad (0\le t\le 1) L ( ( 1 − t ) w 0 + t w 1 ) ≤ ( 1 − t ) L ( w 0 ) + t L ( w 1 ) ( 0 ≤ t ≤ 1 ) が成り立ちます。w 0 , w 1 \boldsymbol{w}_0,\boldsymbol{w}_1 w 0 , w 1 は任意だったので L L L は凸です。
狭義凸性。 rank X = d \operatorname{rank}X = d rank X = d と仮定し、v ≠ 0 \boldsymbol{v}\ne\boldsymbol{0} v = 0 とします。上の等式で v T ∇ 2 L v = 0 \boldsymbol{v}^{\mathsf{T}}\nabla^{2}L\boldsymbol{v} = 0 v T ∇ 2 L v = 0 が起きるとすると、すべての項が非負なので各項が 0 0 0 、つまり μ i ( 1 − μ i ) ( x i T v ) 2 = 0 \mu_i(1-\mu_i)(\boldsymbol{x}_i^{\mathsf{T}}\boldsymbol{v})^2 = 0 μ i ( 1 − μ i ) ( x i T v ) 2 = 0 です。μ i ( 1 − μ i ) > 0 \mu_i(1-\mu_i) > 0 μ i ( 1 − μ i ) > 0 なので x i T v = 0 \boldsymbol{x}_i^{\mathsf{T}}\boldsymbol{v} = 0 x i T v = 0 がすべての i i i で成り立ち、これは X v = 0 X\boldsymbol{v} = \boldsymbol{0} X v = 0 を意味します。rank X = d \operatorname{rank}X = d rank X = d より X X X の核は { 0 } \{\boldsymbol{0}\} { 0 } なので v = 0 \boldsymbol{v} = \boldsymbol{0} v = 0 となり矛盾します。よって v ≠ 0 \boldsymbol{v}\ne\boldsymbol{0} v = 0 ならば v T ∇ 2 L v > 0 \boldsymbol{v}^{\mathsf{T}}\nabla^{2}L\boldsymbol{v} > 0 v T ∇ 2 L v > 0 、すなわち ∇ 2 L ≻ O \nabla^2 L \succ O ∇ 2 L ≻ O です。このとき上の g g g について g ′ ′ > 0 g'' > 0 g ′′ > 0 となり、g g g は狭義凸、したがって L L L も狭義凸です。
∎
系 6.2 (停留点は大域最小点 )
定理 6.1 の設定のもとで、w ∗ ∈ R d \boldsymbol{w}^{*}\in\mathbb{R}^{d} w ∗ ∈ R d が ∇ L ( w ∗ ) = 0 \nabla L(\boldsymbol{w}^{*}) = \boldsymbol{0} ∇ L ( w ∗ ) = 0 を満たすならば、w ∗ \boldsymbol{w}^{*} w ∗ は L L L の大域最小点である。逆に大域最小点は停留点である。
証明(系 6.2) 任意の w ∈ R d \boldsymbol{w}\in\mathbb{R}^{d} w ∈ R d を取り、h = w − w ∗ \boldsymbol{h} = \boldsymbol{w}-\boldsymbol{w}^{*} h = w − w ∗ 、g ( t ) = L ( w ∗ + t h ) g(t) = L(\boldsymbol{w}^{*}+t\boldsymbol{h}) g ( t ) = L ( w ∗ + t h ) と置きます。g g g は C 2 C^{2} C 2 級なので、テイラーの定理(1 変数、ラグランジュ剰余)より、ある θ ∈ ( 0 , 1 ) \theta\in(0,1) θ ∈ ( 0 , 1 ) が存在して
g ( 1 ) = g ( 0 ) + g ′ ( 0 ) + 1 2 g ′ ′ ( θ ) . g(1) = g(0) + g'(0) + \tfrac12 g''(\theta). g ( 1 ) = g ( 0 ) + g ′ ( 0 ) + 2 1 g ′′ ( θ ) . ここで g ( 1 ) = L ( w ) g(1) = L(\boldsymbol{w}) g ( 1 ) = L ( w ) 、g ( 0 ) = L ( w ∗ ) g(0) = L(\boldsymbol{w}^{*}) g ( 0 ) = L ( w ∗ ) 、g ′ ( 0 ) = ∇ L ( w ∗ ) T h = 0 g'(0) = \nabla L(\boldsymbol{w}^{*})^{\mathsf{T}}\boldsymbol{h} = 0 g ′ ( 0 ) = ∇ L ( w ∗ ) T h = 0 (仮定)、g ′ ′ ( θ ) = h T ∇ 2 L ( w ∗ + θ h ) h ≥ 0 g''(\theta) = \boldsymbol{h}^{\mathsf{T}}\nabla^{2}L(\boldsymbol{w}^{*}+\theta\boldsymbol{h})\boldsymbol{h} \ge 0 g ′′ ( θ ) = h T ∇ 2 L ( w ∗ + θ h ) h ≥ 0 (定理 6.1 の半正定値性)です。したがって L ( w ) ≥ L ( w ∗ ) L(\boldsymbol{w}) \ge L(\boldsymbol{w}^{*}) L ( w ) ≥ L ( w ∗ ) が任意の w \boldsymbol{w} w について成り立ちます。逆向きは、L L L が微分可能なので大域最小点で勾配が消える(フェルマーの定理)ことから従います。テイラーの定理については 平均値の定理とテイラーの定理 の 定理 5.3[平均値の定理とテイラーの定理] を参照してください。
∎
これが、勾配だけを頼りに探索してよい理由です。局所最小に捕まる心配がなく、勾配が消えた場所が答えです。深層学習の損失関数は一般に凸ではないので、この保証はロジスティック回帰の大きな利点です。
凸性は「見つかれば大域最適」を保証しますが、「見つかる」ことは保証しません。実際、次のよくある状況で最小点は存在しません。
定理 6.3 (線形分離可能なら最尤推定量は存在しない )
データ ( x i , y i ) i = 1 n (\boldsymbol{x}_i, y_i)_{i=1}^{n} ( x i , y i ) i = 1 n (n ≥ 1 n \ge 1 n ≥ 1 )が狭義に線形分離可能 であるとする。すなわち、あるベクトル v ∈ R d \boldsymbol{v}\in\mathbb{R}^{d} v ∈ R d が存在して
y i = 1 ⟹ x i T v > 0 , y i = 0 ⟹ x i T v < 0 y_i = 1 \implies \boldsymbol{x}_i^{\mathsf{T}}\boldsymbol{v} > 0, \qquad
y_i = 0 \implies \boldsymbol{x}_i^{\mathsf{T}}\boldsymbol{v} < 0 y i = 1 ⟹ x i T v > 0 , y i = 0 ⟹ x i T v < 0 がすべての i i i で成り立つとする。このとき 定義 4.1 の L L L について
inf w ∈ R d L ( w ) = 0 \inf_{\boldsymbol{w}\in\mathbb{R}^{d}} L(\boldsymbol{w}) = 0 w ∈ R d inf L ( w ) = 0 であるが、この下限は達成されない。さらに L ( w k ) → 0 L(\boldsymbol{w}_k)\to 0 L ( w k ) → 0 を満たす任意の点列 ( w k ) (\boldsymbol{w}_k) ( w k ) は ∥ w k ∥ → ∞ \|\boldsymbol{w}_k\| \to \infty ∥ w k ∥ → ∞ を満たす。
証明(定理 6.3) (a) L > 0 L > 0 L > 0 。 命題 3.2 (1) より 0 < μ i < 1 0 < \mu_i < 1 0 < μ i < 1 なので、y i = 1 y_i = 1 y i = 1 の項 − log μ i -\log\mu_i − log μ i は μ i < 1 \mu_i < 1 μ i < 1 より狭義正、y i = 0 y_i = 0 y i = 0 の項 − log ( 1 − μ i ) -\log(1-\mu_i) − log ( 1 − μ i ) は 1 − μ i < 1 1-\mu_i < 1 1 − μ i < 1 より狭義正です。n ≥ 1 n \ge 1 n ≥ 1 個の狭義正の数の和なので L ( w ) > 0 L(\boldsymbol{w}) > 0 L ( w ) > 0 がすべての w \boldsymbol{w} w で成り立ちます。
(b) L ( t v ) → 0 L(t\boldsymbol{v})\to 0 L ( t v ) → 0 。 t > 0 t > 0 t > 0 とし、c i = x i T v c_i = \boldsymbol{x}_i^{\mathsf{T}}\boldsymbol{v} c i = x i T v と置きます。w = t v \boldsymbol{w} = t\boldsymbol{v} w = t v のとき z i = t c i z_i = tc_i z i = t c i です。y i = 1 y_i = 1 y i = 1 の項は 定義 3.1 より log σ ( z ) = − log ( 1 + e − z ) \log\sigma(z) = -\log(1+e^{-z}) log σ ( z ) = − log ( 1 + e − z ) なので
− log σ ( t c i ) = log ( 1 + e − t c i ) → t → ∞ log 1 = 0 -\log\sigma(tc_i) = \log\bigl(1+e^{-tc_i}\bigr) \xrightarrow{\;t\to\infty\;} \log 1 = 0 − log σ ( t c i ) = log ( 1 + e − t c i ) t → ∞ log 1 = 0 です(仮定より c i > 0 c_i > 0 c i > 0 なので e − t c i → 0 e^{-tc_i}\to 0 e − t c i → 0 )。y i = 0 y_i = 0 y i = 0 の項は 命題 3.2 (2) より 1 − σ ( t c i ) = σ ( − t c i ) 1-\sigma(tc_i) = \sigma(-tc_i) 1 − σ ( t c i ) = σ ( − t c i ) なので
− log ( 1 − σ ( t c i ) ) = − log σ ( − t c i ) = log ( 1 + e t c i ) → t → ∞ 0 -\log\bigl(1-\sigma(tc_i)\bigr) = -\log\sigma(-tc_i) = \log\bigl(1+e^{tc_i}\bigr) \xrightarrow{\;t\to\infty\;} 0 − log ( 1 − σ ( t c i ) ) = − log σ ( − t c i ) = log ( 1 + e t c i ) t → ∞ 0 です(仮定より c i < 0 c_i < 0 c i < 0 なので e t c i → 0 e^{tc_i}\to 0 e t c i → 0 )。有限個の和なので L ( t v ) → 0 L(t\boldsymbol{v})\to 0 L ( t v ) → 0 です。
(c) 下限と非達成。 (a) より L > 0 L > 0 L > 0 、(b) より 0 0 0 にいくらでも近づけるので inf L = 0 \inf L = 0 inf L = 0 。しかし (a) より L ( w ) = 0 L(\boldsymbol{w}) = 0 L ( w ) = 0 となる w \boldsymbol{w} w は存在しないので、下限は達成されません。
(d) 発散。 L ( w k ) → 0 L(\boldsymbol{w}_k)\to 0 L ( w k ) → 0 かつ ( w k ) (\boldsymbol{w}_k) ( w k ) が有界だと仮定します。ボルツァーノ・ワイエルシュトラスの定理より収束する部分列 w k m → w ∞ \boldsymbol{w}_{k_m}\to\boldsymbol{w}_{\infty} w k m → w ∞ が取れます。L L L は連続(定理 5.1 より C ∞ C^{\infty} C ∞ 級)なので L ( w ∞ ) = lim m L ( w k m ) = 0 L(\boldsymbol{w}_{\infty}) = \lim_m L(\boldsymbol{w}_{k_m}) = 0 L ( w ∞ ) = lim m L ( w k m ) = 0 となり、(a) に矛盾します。よって ( w k ) (\boldsymbol{w}_k) ( w k ) は有界ではありません。さらに、もし有界な部分列があれば、その部分列も L → 0 L \to 0 L → 0 を満たすので同じ議論で矛盾します。有界な部分列が存在しないことは ∥ w k ∥ → ∞ \|\boldsymbol{w}_k\|\to\infty ∥ w k ∥ → ∞ にほかなりません。
∎
例 6.4 (重みが発散する様子 )
§1.2 のデータ(t = 1 , 2 t=1,2 t = 1 , 2 が y = 0 y=0 y = 0 、t = 3 , 4 t=3,4 t = 3 , 4 が y = 1 y=1 y = 1 )は t = 2.5 t = 2.5 t = 2.5 で分離できます。定理 6.3 の v \boldsymbol{v} v として v = ( − 2.5 , 1 ) T \boldsymbol{v} = (-2.5,\, 1)^{\mathsf{T}} v = ( − 2.5 , 1 ) T を取れば c i = t i − 2.5 = − 1.5 , − 0.5 , 0.5 , 1.5 c_i = t_i - 2.5 = -1.5, -0.5, 0.5, 1.5 c i = t i − 2.5 = − 1.5 , − 0.5 , 0.5 , 1.5 となり、符号条件が満たされます。w = α v \boldsymbol{w} = \alpha\boldsymbol{v} w = α v に沿って損失を計算すると次の通りです。
α \alpha α 1 2 5 10 20 ∥ w ∥ \|\boldsymbol{w}\| ∥ w ∥ 2.69 5.39 13.46 26.93 53.85 L ( w ) L(\boldsymbol{w}) L ( w ) 1.3510 0.7237 0.1589 0.01343 0.0000908
α = 1 \alpha = 1 α = 1 の値を手で確かめます。z i = − 1.5 , − 0.5 , 0.5 , 1.5 z_i = -1.5, -0.5, 0.5, 1.5 z i = − 1.5 , − 0.5 , 0.5 , 1.5 で、y = 0 y=0 y = 0 の 2 点の損失は log ( 1 + e z ) \log(1+e^{z}) log ( 1 + e z ) 、y = 1 y=1 y = 1 の 2 点の損失は log ( 1 + e − z ) \log(1+e^{-z}) log ( 1 + e − z ) ですから
L = log ( 1 + e − 1.5 ) + log ( 1 + e − 0.5 ) + log ( 1 + e − 0.5 ) + log ( 1 + e − 1.5 ) = 2 ( 0.2014 + 0.4741 ) = 1.3510. L = \log(1+e^{-1.5}) + \log(1+e^{-0.5}) + \log(1+e^{-0.5}) + \log(1+e^{-1.5}) = 2(0.2014 + 0.4741) = 1.3510 . L = log ( 1 + e − 1.5 ) + log ( 1 + e − 0.5 ) + log ( 1 + e − 0.5 ) + log ( 1 + e − 1.5 ) = 2 ( 0.2014 + 0.4741 ) = 1.3510. 損失は単調に 0 0 0 へ向かい、それに応じて ∥ w ∥ \|\boldsymbol{w}\| ∥ w ∥ は際限なく大きくなります。例 5.5 の勾配降下法をいつまでも回し続けると、重みが延々と大きくなり続けるのはこのためです。実用上の困り事は、分離できてしまうデータでは予測確率がすべて 0 0 0 か 1 1 1 に張り付き、「どのくらい確信があるか」という情報が失われることです。特徴の数 d d d がデータ数 n n n より多いときはほぼ必ず分離できてしまうので、これは例外的な事態ではありません。
定理 6.5 (L2 正則化つき最尤推定の存在と一意性 )
λ > 0 \lambda > 0 λ > 0 とし、定義 4.1 の L L L に対して
L λ ( w ) = L ( w ) + λ 2 ∥ w ∥ 2 L_{\lambda}(\boldsymbol{w}) = L(\boldsymbol{w}) + \frac{\lambda}{2}\|\boldsymbol{w}\|^{2} L λ ( w ) = L ( w ) + 2 λ ∥ w ∥ 2 と置く。データ ( x i , y i ) i = 1 n (\boldsymbol{x}_i,y_i)_{i=1}^n ( x i , y i ) i = 1 n には何の条件も課さない (分離可能でも、X X X が列フルランクでなくてもよい)。このとき L λ L_{\lambda} L λ は R d \mathbb{R}^{d} R d 上でただ一つの大域最小点 w ^ λ \hat{\boldsymbol{w}}_{\lambda} w ^ λ を持ち、それは方程式
X T ( σ ( X w ^ λ ) − y ) + λ w ^ λ = 0 X^{\mathsf{T}}\bigl(\sigma(X\hat{\boldsymbol{w}}_{\lambda}) - \boldsymbol{y}\bigr) + \lambda\,\hat{\boldsymbol{w}}_{\lambda} = \boldsymbol{0} X T ( σ ( X w ^ λ ) − y ) + λ w ^ λ = 0 の唯一の解である。
証明(定理 6.5) 存在。 定理 6.3 の (a) より L ≥ 0 L \ge 0 L ≥ 0 なので L λ ( w ) ≥ λ 2 ∥ w ∥ 2 L_{\lambda}(\boldsymbol{w}) \ge \frac{\lambda}{2}\|\boldsymbol{w}\|^{2} L λ ( w ) ≥ 2 λ ∥ w ∥ 2 です。一方 w = 0 \boldsymbol{w} = \boldsymbol{0} w = 0 では μ i = 1 / 2 \mu_i = 1/2 μ i = 1/2 なので L λ ( 0 ) = L ( 0 ) = n log 2 L_{\lambda}(\boldsymbol{0}) = L(\boldsymbol{0}) = n\log 2 L λ ( 0 ) = L ( 0 ) = n log 2 です。そこで R = 2 n log 2 / λ + 1 R = \sqrt{2n\log 2/\lambda} + 1 R = 2 n log 2/ λ + 1 と取れば、∥ w ∥ > R \|\boldsymbol{w}\| > R ∥ w ∥ > R のとき
L λ ( w ) ≥ λ 2 ∥ w ∥ 2 > λ 2 R 2 > n log 2 = L λ ( 0 ) L_{\lambda}(\boldsymbol{w}) \ge \frac{\lambda}{2}\|\boldsymbol{w}\|^{2} > \frac{\lambda}{2}R^{2} > n\log 2 = L_{\lambda}(\boldsymbol{0}) L λ ( w ) ≥ 2 λ ∥ w ∥ 2 > 2 λ R 2 > n log 2 = L λ ( 0 ) となります。よって L λ L_{\lambda} L λ の R d \mathbb{R}^{d} R d 上の下限は、閉球 B ˉ ( 0 , R ) = { w : ∥ w ∥ ≤ R } \bar{B}(\boldsymbol{0},R) = \{\boldsymbol{w} : \|\boldsymbol{w}\|\le R\} B ˉ ( 0 , R ) = { w : ∥ w ∥ ≤ R } 上の下限と一致します。B ˉ ( 0 , R ) \bar{B}(\boldsymbol{0},R) B ˉ ( 0 , R ) は R d \mathbb{R}^d R d の有界閉集合すなわちコンパクトで、L λ L_{\lambda} L λ は連続なのでワイエルシュトラスの最大値・最小値定理により最小値を取る点 w ^ λ \hat{\boldsymbol{w}}_{\lambda} w ^ λ が存在します。これは R d \mathbb{R}^{d} R d 全体での大域最小点です。
一意性。 ∥ w ∥ 2 \|\boldsymbol{w}\|^{2} ∥ w ∥ 2 のヘッセ行列は 2 I 2I 2 I なので ∇ 2 L λ ( w ) = X T S ( w ) X + λ I \nabla^{2}L_{\lambda}(\boldsymbol{w}) = X^{\mathsf{T}}S(\boldsymbol{w})X + \lambda I ∇ 2 L λ ( w ) = X T S ( w ) X + λ I です。任意の v ≠ 0 \boldsymbol{v}\ne\boldsymbol{0} v = 0 について 定理 6.1 の半正定値性より
v T ∇ 2 L λ v = v T X T S X v + λ ∥ v ∥ 2 ≥ λ ∥ v ∥ 2 > 0 \boldsymbol{v}^{\mathsf{T}}\nabla^{2}L_{\lambda}\boldsymbol{v} = \boldsymbol{v}^{\mathsf{T}}X^{\mathsf{T}}SX\boldsymbol{v} + \lambda\|\boldsymbol{v}\|^{2} \ge \lambda\|\boldsymbol{v}\|^{2} > 0 v T ∇ 2 L λ v = v T X T S X v + λ ∥ v ∥ 2 ≥ λ ∥ v ∥ 2 > 0 です。いま w 1 ≠ w 2 \boldsymbol{w}_1 \ne \boldsymbol{w}_2 w 1 = w 2 がともに大域最小点だとすると、どちらも停留点なので ∇ L λ ( w 1 ) = 0 \nabla L_{\lambda}(\boldsymbol{w}_1) = \boldsymbol{0} ∇ L λ ( w 1 ) = 0 です。系 6.2 の証明と同じテイラー展開を h = w 2 − w 1 ≠ 0 \boldsymbol{h} = \boldsymbol{w}_2-\boldsymbol{w}_1 \ne \boldsymbol{0} h = w 2 − w 1 = 0 に対して行うと、ある θ ∈ ( 0 , 1 ) \theta\in(0,1) θ ∈ ( 0 , 1 ) について
L λ ( w 2 ) = L λ ( w 1 ) + 0 + 1 2 h T ∇ 2 L λ ( w 1 + θ h ) h ≥ L λ ( w 1 ) + λ 2 ∥ h ∥ 2 > L λ ( w 1 ) L_{\lambda}(\boldsymbol{w}_2) = L_{\lambda}(\boldsymbol{w}_1) + 0 + \tfrac12\boldsymbol{h}^{\mathsf{T}}\nabla^{2}L_{\lambda}(\boldsymbol{w}_1+\theta\boldsymbol{h})\boldsymbol{h}
\ge L_{\lambda}(\boldsymbol{w}_1) + \frac{\lambda}{2}\|\boldsymbol{h}\|^{2} > L_{\lambda}(\boldsymbol{w}_1) L λ ( w 2 ) = L λ ( w 1 ) + 0 + 2 1 h T ∇ 2 L λ ( w 1 + θ h ) h ≥ L λ ( w 1 ) + 2 λ ∥ h ∥ 2 > L λ ( w 1 ) となり、w 2 \boldsymbol{w}_2 w 2 が最小点であることに矛盾します。よって最小点は一意です。
方程式。 定理 5.1 と ∇ ( λ 2 ∥ w ∥ 2 ) = λ w \nabla\bigl(\frac{\lambda}{2}\|\boldsymbol{w}\|^{2}\bigr) = \lambda\boldsymbol{w} ∇ ( 2 λ ∥ w ∥ 2 ) = λ w より ∇ L λ ( w ) = X T ( σ ( X w ) − y ) + λ w \nabla L_{\lambda}(\boldsymbol{w}) = X^{\mathsf{T}}(\sigma(X\boldsymbol{w})-\boldsymbol{y}) + \lambda\boldsymbol{w} ∇ L λ ( w ) = X T ( σ ( X w ) − y ) + λ w です。L λ L_{\lambda} L λ は凸(凸関数 L L L と凸関数 λ 2 ∥ w ∥ 2 \frac{\lambda}{2}\|\boldsymbol{w}\|^2 2 λ ∥ w ∥ 2 の和)なので 系 6.2 と同じ議論により、停留点であることと大域最小点であることは同値です。最小点が一意だったので、停留方程式の解も一意です。
∎
演習 7.1 易
定義 3.1 の σ \sigma σ について、σ ′ ′ ( z ) = σ ′ ( z ) ( 1 − 2 σ ( z ) ) \sigma''(z) = \sigma'(z)\bigl(1-2\sigma(z)\bigr) σ ′′ ( z ) = σ ′ ( z ) ( 1 − 2 σ ( z ) ) を示し、σ \sigma σ が z = 0 z = 0 z = 0 に変曲点を持つことを確かめてください。
解答 命題 3.2 (3) より σ ′ = σ ( 1 − σ ) = σ − σ 2 \sigma' = \sigma(1-\sigma) = \sigma - \sigma^{2} σ ′ = σ ( 1 − σ ) = σ − σ 2 です。これを z z z で微分すると、積の微分法(あるいは合成関数の微分法)により
σ ′ ′ = σ ′ − 2 σ σ ′ = σ ′ ( 1 − 2 σ ) \sigma'' = \sigma' - 2\sigma\sigma' = \sigma'(1 - 2\sigma) σ ′′ = σ ′ − 2 σ σ ′ = σ ′ ( 1 − 2 σ ) を得ます。命題 3.2 (3) より σ ′ > 0 \sigma' > 0 σ ′ > 0 なので、σ ′ ′ \sigma'' σ ′′ の符号は 1 − 2 σ ( z ) 1-2\sigma(z) 1 − 2 σ ( z ) の符号だけで決まります。σ \sigma σ は狭義単調増加で σ ( 0 ) = 1 / 2 \sigma(0) = 1/2 σ ( 0 ) = 1/2 (同 (2))ですから
z < 0 z < 0 z < 0 のとき σ ( z ) < 1 / 2 \sigma(z) < 1/2 σ ( z ) < 1/2 なので 1 − 2 σ ( z ) > 0 1-2\sigma(z) > 0 1 − 2 σ ( z ) > 0 、すなわち σ ′ ′ > 0 \sigma'' > 0 σ ′′ > 0 (下に凸)、
z = 0 z = 0 z = 0 のとき σ ′ ′ = 0 \sigma'' = 0 σ ′′ = 0 、
z > 0 z > 0 z > 0 のとき σ ( z ) > 1 / 2 \sigma(z) > 1/2 σ ( z ) > 1/2 なので σ ′ ′ < 0 \sigma'' < 0 σ ′′ < 0 (上に凸)
となります。z = 0 z = 0 z = 0 の前後で凹凸が入れ替わるので、z = 0 z = 0 z = 0 は変曲点です。σ ′ ( 0 ) = 1 2 ⋅ 1 2 = 1 4 \sigma'(0) = \frac12\cdot\frac12 = \frac14 σ ′ ( 0 ) = 2 1 ⋅ 2 1 = 4 1 なので、そこでの接線の傾きは 1 / 4 1/4 1/4 、これがシグモイドの最大傾斜です。
演習 7.2 標準
ある病気の罹患確率について log μ 1 − μ = − 4 + 0.8 x 1 + 1.5 x 2 \log\dfrac{\mu}{1-\mu} = -4 + 0.8\,x_1 + 1.5\,x_2 log 1 − μ μ = − 4 + 0.8 x 1 + 1.5 x 2 というモデルが推定されました。x 1 x_1 x 1 は年齢を 10 歳単位で測った値、x 2 x_2 x 2 は喫煙者なら 1 1 1 、非喫煙者なら 0 0 0 を取る変数です。
50 歳(x 1 = 5 x_1 = 5 x 1 = 5 )の非喫煙者の罹患確率を求めてください。
他の条件を固定したとき、喫煙者は非喫煙者に比べてオッズが何倍になりますか。
このモデルの上では、喫煙者であることは何歳分の加齢と同じだけオッズを押し上げますか。
解答 1. z = − 4 + 0.8 × 5 + 1.5 × 0 = − 4 + 4 = 0 z = -4 + 0.8\times 5 + 1.5\times 0 = -4 + 4 = 0 z = − 4 + 0.8 × 5 + 1.5 × 0 = − 4 + 4 = 0 なので、命題 3.2 (2) より μ = σ ( 0 ) = 0.5 \mu = \sigma(0) = 0.5 μ = σ ( 0 ) = 0.5 、すなわち 50 % 50\% 50% です。
2. x 2 x_2 x 2 を 0 0 0 から 1 1 1 に変えると対数オッズが 1.5 1.5 1.5 増えるので、オッズは e 1.5 = 4.4817 e^{1.5} = 4.4817 e 1.5 = 4.4817 倍になります。例 3.7 と同様に、これは x 1 x_1 x 1 の値によらず一定です。ただし確率が何倍になるかは x 1 x_1 x 1 に依存します。実際 x 1 = 5 x_1 = 5 x 1 = 5 では μ \mu μ が 0.5 → σ ( 1.5 ) = 0.8176 0.5 \to \sigma(1.5) = 0.8176 0.5 → σ ( 1.5 ) = 0.8176 と 1.635 1.635 1.635 倍にしかなりません。
3. 喫煙者 ( x 1 , 1 ) (x_1, 1) ( x 1 , 1 ) と、より年配の非喫煙者 ( x 1 + Δ , 0 ) (x_1 + \Delta, 0) ( x 1 + Δ , 0 ) の対数オッズが等しくなる Δ \Delta Δ を求めます。
− 4 + 0.8 x 1 + 1.5 = − 4 + 0.8 ( x 1 + Δ ) ⟺ 1.5 = 0.8 Δ ⟺ Δ = 1.875. -4 + 0.8x_1 + 1.5 = -4 + 0.8(x_1 + \Delta) \iff 1.5 = 0.8\,\Delta \iff \Delta = 1.875 . − 4 + 0.8 x 1 + 1.5 = − 4 + 0.8 ( x 1 + Δ ) ⟺ 1.5 = 0.8 Δ ⟺ Δ = 1.875. x 1 x_1 x 1 の単位が 10 歳なので 18.75 歳分 です。数値で確かめます。x 1 = 5 x_1 = 5 x 1 = 5 (50 歳)の喫煙者は z = − 4 + 4 + 1.5 = 1.5 z = -4 + 4 + 1.5 = 1.5 z = − 4 + 4 + 1.5 = 1.5 、x 1 = 6.875 x_1 = 6.875 x 1 = 6.875 (68.75 歳)の非喫煙者は z = − 4 + 0.8 × 6.875 = − 4 + 5.5 = 1.5 z = -4 + 0.8\times 6.875 = -4 + 5.5 = 1.5 z = − 4 + 0.8 × 6.875 = − 4 + 5.5 = 1.5 で一致します。両者の罹患確率はともに σ ( 1.5 ) = 0.8176 \sigma(1.5) = 0.8176 σ ( 1.5 ) = 0.8176 です。
この計算が x 1 x_1 x 1 の値に依らないのは、対数オッズが x 1 x_1 x 1 と x 2 x_2 x 2 の線形結合 だからです。交互作用項 x 1 x 2 x_1x_2 x 1 x 2 を入れたモデルではこの換算は年齢に依存し、一定の「歳数」では言い表せなくなります。
演習 7.3 標準
ラベルを y ~ i = 2 y i − 1 ∈ { − 1 , + 1 } \tilde{y}_i = 2y_i - 1 \in \{-1, +1\} y ~ i = 2 y i − 1 ∈ { − 1 , + 1 } と付け替えます。z i = w T x i z_i = \boldsymbol{w}^{\mathsf{T}}\boldsymbol{x}_i z i = w T x i として、定義 4.1 の L L L が
L ( w ) = ∑ i = 1 n log ( 1 + e − y ~ i z i ) L(\boldsymbol{w}) = \sum_{i=1}^{n} \log\bigl(1 + e^{-\tilde{y}_i z_i}\bigr) L ( w ) = i = 1 ∑ n log ( 1 + e − y ~ i z i ) と書けることを示し、この表示から勾配を計算して 定理 5.1 と一致することを確かめてください。
解答 表示。 第 i i i 項を場合分けします。y i = 1 y_i = 1 y i = 1 (y ~ i = + 1 \tilde{y}_i = +1 y ~ i = + 1 )のとき、定義 4.1 の項は − log μ i = − log σ ( z i ) -\log\mu_i = -\log\sigma(z_i) − log μ i = − log σ ( z i ) です。定義 3.1 の σ ( z ) = 1 / ( 1 + e − z ) \sigma(z) = 1/(1+e^{-z}) σ ( z ) = 1/ ( 1 + e − z ) より − log σ ( z i ) = log ( 1 + e − z i ) = log ( 1 + e − y ~ i z i ) -\log\sigma(z_i) = \log(1+e^{-z_i}) = \log(1+e^{-\tilde{y}_i z_i}) − log σ ( z i ) = log ( 1 + e − z i ) = log ( 1 + e − y ~ i z i ) です。
y i = 0 y_i = 0 y i = 0 (y ~ i = − 1 \tilde{y}_i = -1 y ~ i = − 1 )のとき、項は − log ( 1 − μ i ) -\log(1-\mu_i) − log ( 1 − μ i ) です。命題 3.2 (2) より 1 − σ ( z i ) = σ ( − z i ) 1-\sigma(z_i) = \sigma(-z_i) 1 − σ ( z i ) = σ ( − z i ) なので
− log ( 1 − μ i ) = − log σ ( − z i ) = log ( 1 + e z i ) = log ( 1 + e − ( − 1 ) z i ) = log ( 1 + e − y ~ i z i ) -\log(1-\mu_i) = -\log\sigma(-z_i) = \log\bigl(1+e^{z_i}\bigr) = \log\bigl(1+e^{-(-1)z_i}\bigr) = \log\bigl(1+e^{-\tilde{y}_i z_i}\bigr) − log ( 1 − μ i ) = − log σ ( − z i ) = log ( 1 + e z i ) = log ( 1 + e − ( − 1 ) z i ) = log ( 1 + e − y ~ i z i ) となり、どちらの場合も同じ式になります。
勾配。 u i = − y ~ i z i u_i = -\tilde{y}_i z_i u i = − y ~ i z i と置くと第 i i i 項は log ( 1 + e u i ) \log(1+e^{u_i}) log ( 1 + e u i ) で、d d u log ( 1 + e u ) = e u 1 + e u = σ ( u ) \dfrac{d}{du}\log(1+e^{u}) = \dfrac{e^{u}}{1+e^{u}} = \sigma(u) d u d log ( 1 + e u ) = 1 + e u e u = σ ( u ) です(分子分母を e u e^{u} e u で割れば σ \sigma σ の定義形になります)。連鎖律より ∂ u i ∂ w = − y ~ i x i \dfrac{\partial u_i}{\partial \boldsymbol{w}} = -\tilde{y}_i\boldsymbol{x}_i ∂ w ∂ u i = − y ~ i x i なので
∇ L ( w ) = − ∑ i = 1 n y ~ i σ ( − y ~ i z i ) x i . \nabla L(\boldsymbol{w}) = -\sum_{i=1}^{n} \tilde{y}_i\,\sigma(-\tilde{y}_i z_i)\,\boldsymbol{x}_i . ∇ L ( w ) = − i = 1 ∑ n y ~ i σ ( − y ~ i z i ) x i . 定理 5.1 と一致することを場合分けで確かめます。y i = 1 y_i = 1 y i = 1 のとき − y ~ i σ ( − y ~ i z i ) = − σ ( − z i ) = − ( 1 − μ i ) = μ i − 1 = μ i − y i -\tilde{y}_i\sigma(-\tilde{y}_iz_i) = -\sigma(-z_i) = -(1-\mu_i) = \mu_i - 1 = \mu_i - y_i − y ~ i σ ( − y ~ i z i ) = − σ ( − z i ) = − ( 1 − μ i ) = μ i − 1 = μ i − y i 。y i = 0 y_i = 0 y i = 0 のとき − y ~ i σ ( − y ~ i z i ) = + σ ( z i ) = μ i = μ i − y i -\tilde{y}_i\sigma(-\tilde{y}_iz_i) = +\sigma(z_i) = \mu_i = \mu_i - y_i − y ~ i σ ( − y ~ i z i ) = + σ ( z i ) = μ i = μ i − y i 。どちらも μ i − y i \mu_i - y_i μ i − y i に等しく、一致します。
この表示は「y ~ i z i \tilde{y}_i z_i y ~ i z i (マージン)が大きいほど損失が小さい」という構造を露わにしており、サポートベクターマシンのヒンジ損失 max ( 0 , 1 − y ~ i z i ) \max(0, 1-\tilde{y}_iz_i) max ( 0 , 1 − y ~ i z i ) と直接比較できる形になっています。
演習 7.4 難
λ > 0 \lambda > 0 λ > 0 とし、定理 6.5 の L λ L_{\lambda} L λ とその一意の最小点 w ^ λ \hat{\boldsymbol{w}}_{\lambda} w ^ λ を考えます。任意の w ∈ R d \boldsymbol{w}\in\mathbb{R}^{d} w ∈ R d について
L λ ( w ) ≥ L λ ( w ^ λ ) + λ 2 ∥ w − w ^ λ ∥ 2 L_{\lambda}(\boldsymbol{w}) \;\ge\; L_{\lambda}(\hat{\boldsymbol{w}}_{\lambda}) + \frac{\lambda}{2}\bigl\|\boldsymbol{w}-\hat{\boldsymbol{w}}_{\lambda}\bigr\|^{2} L λ ( w ) ≥ L λ ( w ^ λ ) + 2 λ w − w ^ λ 2 を示し、これを使って ∥ w ^ λ ∥ ≤ 2 n log 2 / λ \|\hat{\boldsymbol{w}}_{\lambda}\| \le \sqrt{2n\log 2/\lambda} ∥ w ^ λ ∥ ≤ 2 n log 2/ λ を導いてください。
解答 不等式。 h = w − w ^ λ \boldsymbol{h} = \boldsymbol{w} - \hat{\boldsymbol{w}}_{\lambda} h = w − w ^ λ 、g ( t ) = L λ ( w ^ λ + t h ) g(t) = L_{\lambda}(\hat{\boldsymbol{w}}_{\lambda} + t\boldsymbol{h}) g ( t ) = L λ ( w ^ λ + t h ) と置きます。L L L は C ∞ C^{\infty} C ∞ 級(定理 5.1 )、λ 2 ∥ w ∥ 2 \frac{\lambda}{2}\|\boldsymbol{w}\|^{2} 2 λ ∥ w ∥ 2 は多項式なので L λ L_{\lambda} L λ も C ∞ C^{\infty} C ∞ 級で、g g g は C 2 C^{2} C 2 級です。テイラーの定理より、ある θ ∈ ( 0 , 1 ) \theta\in(0,1) θ ∈ ( 0 , 1 ) について
L λ ( w ) = g ( 1 ) = g ( 0 ) + g ′ ( 0 ) + 1 2 g ′ ′ ( θ ) . L_{\lambda}(\boldsymbol{w}) = g(1) = g(0) + g'(0) + \tfrac12 g''(\theta). L λ ( w ) = g ( 1 ) = g ( 0 ) + g ′ ( 0 ) + 2 1 g ′′ ( θ ) . w ^ λ \hat{\boldsymbol{w}}_{\lambda} w ^ λ は最小点なので ∇ L λ ( w ^ λ ) = 0 \nabla L_{\lambda}(\hat{\boldsymbol{w}}_{\lambda}) = \boldsymbol{0} ∇ L λ ( w ^ λ ) = 0 、したがって g ′ ( 0 ) = ∇ L λ ( w ^ λ ) T h = 0 g'(0) = \nabla L_{\lambda}(\hat{\boldsymbol{w}}_{\lambda})^{\mathsf{T}}\boldsymbol{h} = 0 g ′ ( 0 ) = ∇ L λ ( w ^ λ ) T h = 0 です。また 定理 6.5 の証明で示した通り ∇ 2 L λ = X T S X + λ I \nabla^{2}L_{\lambda} = X^{\mathsf{T}}SX + \lambda I ∇ 2 L λ = X T S X + λ I で、定理 6.1 より X T S X ⪰ O X^{\mathsf{T}}SX \succeq O X T S X ⪰ O なので
g ′ ′ ( θ ) = h T ( X T S X + λ I ) h ≥ λ ∥ h ∥ 2 . g''(\theta) = \boldsymbol{h}^{\mathsf{T}}\bigl(X^{\mathsf{T}}S X + \lambda I\bigr)\boldsymbol{h} \ge \lambda\|\boldsymbol{h}\|^{2}. g ′′ ( θ ) = h T ( X T S X + λ I ) h ≥ λ ∥ h ∥ 2 . 以上を代入すれば L λ ( w ) ≥ L λ ( w ^ λ ) + λ 2 ∥ h ∥ 2 L_{\lambda}(\boldsymbol{w}) \ge L_{\lambda}(\hat{\boldsymbol{w}}_{\lambda}) + \frac{\lambda}{2}\|\boldsymbol{h}\|^{2} L λ ( w ) ≥ L λ ( w ^ λ ) + 2 λ ∥ h ∥ 2 を得ます。この性質を**λ \lambda λ -強凸性**といい、狭義凸性より強い「二次関数で下から押さえられる」という主張です。
上からの評価。 不等式で w = 0 \boldsymbol{w} = \boldsymbol{0} w = 0 と取ります。定理 6.5 の証明で見たように L λ ( 0 ) = L ( 0 ) = n log 2 L_{\lambda}(\boldsymbol{0}) = L(\boldsymbol{0}) = n\log 2 L λ ( 0 ) = L ( 0 ) = n log 2 (μ i = 1 / 2 \mu_i = 1/2 μ i = 1/2 が n n n 個)なので
n log 2 ≥ L λ ( w ^ λ ) + λ 2 ∥ w ^ λ ∥ 2 ≥ λ 2 ∥ w ^ λ ∥ 2 n\log 2 \;\ge\; L_{\lambda}(\hat{\boldsymbol{w}}_{\lambda}) + \frac{\lambda}{2}\|\hat{\boldsymbol{w}}_{\lambda}\|^{2} \;\ge\; \frac{\lambda}{2}\|\hat{\boldsymbol{w}}_{\lambda}\|^{2} n log 2 ≥ L λ ( w ^ λ ) + 2 λ ∥ w ^ λ ∥ 2 ≥ 2 λ ∥ w ^ λ ∥ 2 です。最後の不等号では 定理 6.3 の (a) と λ 2 ∥ w ^ λ ∥ 2 ≥ 0 \frac{\lambda}{2}\|\hat{\boldsymbol{w}}_\lambda\|^2 \ge 0 2 λ ∥ w ^ λ ∥ 2 ≥ 0 から L λ ( w ^ λ ) ≥ 0 L_{\lambda}(\hat{\boldsymbol{w}}_{\lambda}) \ge 0 L λ ( w ^ λ ) ≥ 0 であることを使いました。整理すると ∥ w ^ λ ∥ 2 ≤ 2 n log 2 / λ \|\hat{\boldsymbol{w}}_{\lambda}\|^{2} \le 2n\log 2/\lambda ∥ w ^ λ ∥ 2 ≤ 2 n log 2/ λ 、すなわち ∥ w ^ λ ∥ ≤ 2 n log 2 / λ \|\hat{\boldsymbol{w}}_{\lambda}\| \le \sqrt{2n\log 2/\lambda} ∥ w ^ λ ∥ ≤ 2 n log 2/ λ です。
データが線形分離可能でも重みがこの範囲に収まることが、定理 6.3 の発散を正則化が確かに止めていることの定量的な証拠になります。λ → 0 \lambda \to 0 λ → 0 で上界が ∞ \infty ∞ に発散するのも、正則化なしの状況と整合しています。
C. M. Bishop, Pattern Recognition and Machine Learning , Springer, 2006 — 第 4 章「Linear Models for Classification」。§4.2 に 例 3.4 の生成モデルからの導出、§4.3 にロジスティック回帰と IRLS がある。
T. Hastie, R. Tibshirani, J. Friedman, The Elements of Statistical Learning , 2nd ed., Springer, 2009 — 第 4 章「Linear Methods for Classification」。線形分離可能な場合の非有界性や正則化の扱いも含む。著者による公開版 。
S. Boyd, L. Vandenberghe, Convex Optimization , Cambridge University Press, 2004 — 第 3 章(凸関数と強凸性)、第 7 章(最尤推定を凸最適化として扱う例)。著者による公開版 。
I. Goodfellow, Y. Bengio, A. Courville, Deep Learning , MIT Press, 2016 — 第 6 章。出力ユニットの選択と交差エントロピー損失の対応、および 例 5.3 の勾配飽和の議論。公開版 。
J. Berkson, “Application of the Logistic Function to Bio-Assay”, Journal of the American Statistical Association 39 (1944), 357–365 — “logit” という語が導入された論文。
J. A. Nelder, R. W. M. Wedderburn, “Generalized Linear Models”, Journal of the Royal Statistical Society, Series A 135 (1972), 370–384 — ロジスティック回帰を一般化線形モデルの一例として位置づけた論文。
久保拓弥『データ解析のための統計モデリング入門』岩波書店、2012 — 一般化線形モデルの枠組みからロジスティック回帰を扱う章。統計モデリングの実践側からの入門として読みやすい。
ソフトマックス回帰。 クラスが K K K 個ある場合は、クラスごとに重み w k ∈ R d \boldsymbol{w}_k \in \mathbb{R}^{d} w k ∈ R d を用意し
P ( y = k ∣ x ) = exp ( w k T x ) ∑ l = 1 K exp ( w l T x ) P(y = k \mid \boldsymbol{x}) = \frac{\exp(\boldsymbol{w}_k^{\mathsf{T}}\boldsymbol{x})}{\sum_{l=1}^{K}\exp(\boldsymbol{w}_l^{\mathsf{T}}\boldsymbol{x})} P ( y = k ∣ x ) = ∑ l = 1 K exp ( w l T x ) exp ( w k T x )
とします。これをソフトマックス関数 といい、K = 2 K = 2 K = 2 のとき分子分母を exp ( w 0 T x ) \exp(\boldsymbol{w}_0^{\mathsf{T}}\boldsymbol{x}) exp ( w 0 T x ) で割れば P ( y = 1 ∣ x ) = σ ( ( w 1 − w 0 ) T x ) P(y=1\mid\boldsymbol{x}) = \sigma\bigl((\boldsymbol{w}_1-\boldsymbol{w}_0)^{\mathsf{T}}\boldsymbol{x}\bigr) P ( y = 1 ∣ x ) = σ ( ( w 1 − w 0 ) T x ) となり、シグモイドに戻ります(重みの差だけが意味を持つので、パラメータには K − 1 K-1 K − 1 個分の自由度しかありません)。損失はラベルを one-hot ベクトル t i \boldsymbol{t}_i t i (第 y i y_i y i 成分だけ 1 1 1 )として L = − ∑ i ∑ k t i k log μ i k L = -\sum_i \sum_k t_{ik}\log \mu_{ik} L = − ∑ i ∑ k t ik log μ ik で、これも交差エントロピーです。勾配は 定理 5.1 と同じ形の
∂ L ∂ w k = ∑ i = 1 n ( μ i k − t i k ) x i \frac{\partial L}{\partial \boldsymbol{w}_k} = \sum_{i=1}^{n} (\mu_{ik} - t_{ik})\,\boldsymbol{x}_i ∂ w k ∂ L = i = 1 ∑ n ( μ ik − t ik ) x i
になります(ソフトマックス+交差エントロピーの勾配(命題 7.1)[ニューラルネットワークと逆伝播] )。約分が起きる構造がそのまま保たれている、というのが要点です。
ニュートン法と IRLS。 定理 6.1 でヘッセ行列まで求めてあるので、勾配だけでなく 2 階情報も使えます。ニュートン法の更新
w ← w − ( X T S X ) − 1 X T ( μ − y ) \boldsymbol{w} \leftarrow \boldsymbol{w} - \bigl(X^{\mathsf{T}}SX\bigr)^{-1}X^{\mathsf{T}}(\boldsymbol{\mu}-\boldsymbol{y}) w ← w − ( X T S X ) − 1 X T ( μ − y )
は、右辺を整理すると w ← ( X T S X ) − 1 X T S z \boldsymbol{w} \leftarrow (X^{\mathsf{T}}SX)^{-1}X^{\mathsf{T}}S\boldsymbol{z} w ← ( X T S X ) − 1 X T S z (z = X w + S − 1 ( y − μ ) \boldsymbol{z} = X\boldsymbol{w} + S^{-1}(\boldsymbol{y}-\boldsymbol{\mu}) z = X w + S − 1 ( y − μ ) )という重み付き最小二乗法 の形に書き換えられます。重み S S S が反復ごとに更新されるので、これを IRLS (反復再重み付け最小二乗法)と呼びます。収束が速い一方、各反復で d × d d\times d d × d 行列の逆行列(実際には連立一次方程式)を扱うため d d d が大きいと重くなります。深層学習で 1 階の 勾配降下法 が使われるのは、この計算量の差が理由の一つです。