AI・機械学習
なぜ誤差逆伝播法は逆方向に進むのか
Why back propagation goes backward (gregorygundersen.com)
要約
誤差逆伝播法(バックプロパゲーション)がニューラルネットワークの学習において、なぜ「逆方向」に計算を進めるのかを、計算グラフと連鎖律の観点から解説しています。順方向の計算では計算量が二次的に増加してしまうため、効率的な逆方向の伝播が採用されていることを、第一原理から再構築する形で説明しています。
全文翻訳
Home Blog RSS なぜ誤差逆伝播法は逆方向に進むのか
誤差逆伝播法はニューラルネットワークの勾配を計算するアルゴリズムですが、なぜアルゴリズムが逆方向のパスを使用するのかは、必ずしも明白ではありません。その答えは、第一原理から誤差逆伝播法を再構築することを可能にします。
公開 2018年4月15日
ニューラルネットワークの学習に使用されるアルゴリズムである誤差逆伝播法(Rumelhart et al., 1986)の一般的な説明は、各ノードのエラーを逆方向に伝播させているというものです。しかし、私がこのアルゴリズムを最初に学んだとき、直接的な答えが見つからない疑問がありました。なぜ逆方向に進まなければならないのでしょうか?ニューラルネットワークは単なる合成関数であり、連鎖律を使用して合成関数の導関数を計算する方法はわかっています。なぜ順方向のパスで勾配を計算しないのでしょうか?この疑問に答えることが、誤差逆伝播法への理解を深めるのに役立つことがわかりました。読者はニューラルネットワークと勾配降下法について広く理解しており、誤差逆伝播法にもある程度精通していると仮定します。まず、誤差逆伝播法をいくつかの有用な概念と記法で設定し、次に順方向伝播アルゴリズムがなぜ最適ではないのかを説明します。
設定
誤差逆伝播法の目標は、ニューラルネットワーク f におけるすべての重み θi について、効率的に ∂f/∂θi を計算することであることを思い出してください。問題を枠組みするために、任意の重み θ1 と f のどこかにあるノード v について推論してみましょう。
明確にするために、ノード v は、入力の重み付き和を活性化関数 σ を通した後のノードの出力値を指します。つまり、
u = θ1t1 + θ2t2 + ⋯ + θntn
v = σ(u)
通常の図では、u、σ、v はすべて単一のノードとして、破線で示されます。
誤差逆伝播法を理解するために必要な最も重要な観察は次のとおりです。連鎖律のおかげで、∂f/∂θ1 のほとんどの計算は、各ノードで局所的に行うことができます。
∂f/∂θ1 = ∂f/∂v ⋅ ∂v/∂u ⋅ ∂u/∂θ1
∂v/∂u は解析的に計算できます。これは σ の定義に依存するだけです。そして、∂u/∂θ1 = t1 であることがわかっています。したがって、各ノード v では、∂f/∂v がわかっていれば、∂f/∂θ1 を計算できます。∂f/∂v を計算する上での課題は、下流のノードが v の値に依存することです。幸いなことに、多変数連鎖律が答えを提供します。
各 wi が単一変数関数 wi(v) である多変数関数 g(w1, w2, …, wm) が与えられた場合、多変数連鎖律は次のように述べています。
∂g/∂v = ∂/∂v g(w1(v), w2(v), …, wm(v)) = ∑j (∂g/∂wj ⋅ ∂wj/∂v)
したがって、任意の重み θi について ∂f/∂θi を計算でき、これは順方向パスで誤差逆伝播法を実装しようとするために必要な仕組みを備えていることを意味します。何が起こるか見てみましょう。
繰り返される項
任意の重み θi について偏微分 ∂f/∂θi を計算できる順方向伝播アルゴリズムを設計したいと考えています。上記の議論から、ノード v において、これは次と同等であることが示されました。
∂f/∂θi = ∂f/∂v ⋅ ∂v/∂θi
表記を簡単にするために、中間変数 u は省略しました。順方向伝播アルゴリズムを設計するために、重要な事実を形式化しましょう。ノード b がノード a に依存する有向計算グラフでは、ノード b で ∂b/∂a を計算することは不可能です。
この主張は明白であるはずです。計算グラフが関数 f(a) = b を表す場合、f としたがって b にアクセスせずに f'(a) を計算することは不可能です。私たちの設定では、ノード v に依存する各下流ノード wj について、wj/∂v をノード v で計算することは不可能です。したがって、∂f/∂v を計算するために、多変数連鎖律を使用して項を分解し、∂f/∂θi を計算するために必要な他の項を v に依存する各ノード wj に順方向に渡す必要があります。
∂f/∂θi = (∑j (∂f/∂wj ⋅ ∂wj/∂v)Compute on wj) ⋅ ∂v/∂θiPass forward
このアルゴリズムは、同じメッセージを何度も順方向に伝播させているため、計算量が爆発することがわかります。たとえば、同じ層にある異なる重み θi と θk について ∂f/∂θi と ∂f/∂θk を計算したい場合、∂v/∂θi と ∂v/∂θk を個別に計算する必要がありますが、他のすべての項は繰り返されます。
∂f/∂θi = (∑j (∑k ∂f/∂zk ⋅ ∂zk/∂wj) ⋅ ∂wj/∂v)Repeated terms ⋅ ∂v/∂θi
∂f/∂θk = (∑j (∑k ∂f/∂zk ⋅ ∂zk/∂wj) ⋅ ∂wj/∂v) ⋅ ∂v/∂θk
繰り返される項のメッセージパッシングの図を以下に示します。
この図は、なぜ誤差逆伝播法が逆方向に進むのかを理解するための鍵であると思います。これが重要な洞察です。下流の項、たとえば ∂wj/∂v にすでにアクセスできる場合、その項をノード v に逆方向にメッセージパッシングして ∂f/∂v を計算できます。各ノードは自身の局所的な項を渡すだけなので、逆方向のパスはノード数に対して線形時間で実行できます。
逆方向のパス
この説明が、有向非巡回グラフで導関数を計算しようとするときに、第一原理から誤差逆伝播法にどのように到達できるかを明確にすることを願っています。ノード b がノード a に依存する特定のノードでは、∂b/∂a を単純に a に逆方向にメッセージパッシングします。多変数連鎖律は、誤差逆伝播法の正当性を証明するのに役立ちます。下流の重み wj を持つ任意のノード v について、v が逆方向に伝播するメッセージを単純に合計すると、目的の導関数が計算されます。
∂f/∂v = ∑j (∂f/∂wj ⋅ ∂wj/∂v)
誤差逆伝播法が解決する主な計算問題を理解すれば、誤差を逆伝播させるという標準的な説明がより理にかなっていると思います。このプロセスは、一種のクレジット割り当て問題の解決策と見なすことができます。各ノードは、上流の隣接ノードに何が間違っていたかを伝えます。しかし、アルゴリズムがこのように機能する理由は、ナイーブな順方向伝播ソリューションではノード数に対して二次的な実行時間になるためです。
Rumelhart, D. E., Hinton, G. E., & Williams, R. J. (1986). Learning representations by back-propagating errors. Nat