重參數化技巧
要訓練變分自編碼器,我們必須讓梯度反向傳播穿過一個步驟:從編碼器的分布抽出隨機樣本 z。但對梯度而言,隨機性是條死路——你無法問「如果我把平均值微調一下,那次隨機抽樣會怎麼變?」,因為那次抽樣,嗯,就是隨機的。重參數化技巧是一種巧妙的改寫,把隨機性挪出梯度的必經之路:我們不直接從一個其參數正在學習的分布抽 z,而是先抽出一個固定、不含參數的雜訊,再用那些參數對它做確定性的變換。
精確地說,假設編碼器為一個高斯潛在變數輸出平均值 μ 與標準差 σ,於是 z 服從 N(μ, σ²)。我們不從該分布抽 z,而是抽出標準常態雜訊 ε ∼ N(0, I)——它不依賴任何參數——並令 z = μ + σ ⊙ ε(逐元素相乘再相加)。這個 z 恰好具有正確的分布,但如今它是 μ 與 σ 的確定性、可微函數,而 ε 是外部輸入。損失對 μ 與 σ(進而對編碼器權重)的梯度,便直接穿過加法與乘法流回;唯一的隨機節點 ε 沒有任何參數需要微分。
這把一個高變異數的估計式(分數函數法,即 REINFORCE 梯度,可用但雜訊大)轉換成低變異數的路徑式梯度,這也是變分自編碼器能用一般隨機梯度下降穩定訓練的原因。此技巧適用於任何能寫成「不含參數的雜訊之確定性變換」的分布(位置—尺度族,以及更一般而言任何「可重參數化」的分布)。對離散潛在變數則不直接適用——那時要改用鬆弛(relaxation)方法,例如 Gumbel-Softmax/Concrete 分布,或退回使用分數函數估計式。
當 μ = 2、σ = 0.5,抽到 ε = 1.3 時得 z = 2 + 0.5×1.3 = 2.65;梯度 ∂z/∂μ = 1 與 ∂z/∂σ = ε = 1.3 都是明確定義的常數,因此 z 處的損失訊號能傳回產生 μ 與 σ 的編碼器——若我們是「憑空」抽出 z = 2.65,這就不可能辦到。