深層学習

拡散モデル

データに段階的にノイズを加える過程を定め、その逆向きに各段階のノイズを予測して画像を生成する確率モデル。

  • A|中核
  • 深層学習

数式の記号で止まったら 記号の読み方 (∂・⊙・転置・上付き添字を、読み方から)

ひとことで言うと

拡散モデルは、データを少しずつ壊して最終的にガウスノイズへ近づける前向き過程と、その逆向きにノイズを取り除いてデータへ戻す逆過程を組み合わせた生成モデルです。DDPM では、逆過程を直接「きれいな画像」へ写すのではなく、各時刻で加わったノイズ ϵ\boldsymbol{\epsilon} をニューラルネットワーク ϵθ(xt,t)\boldsymbol{\epsilon}_{\theta}(\mathbf{x}_t,t) に予測させます。

写真を何段階も薄い霧で覆い、最後には何も見えなくする工程を考えます。学習では、霧の濃さを指定して「今回加えた霧は何だったか」を当てさせます。生成時は完全な霧から始め、当てた霧を一段ずつ取り除きます。1回で復元しないことが、拡散モデルの見た目と計算量を決めます。

なぜ必要か

画像を一発で生成する写像を学習させると、複雑なデータ分布を一度の予測に押し込む必要があります。拡散モデルは、ノイズ量の異なる多数の復元問題へ分解します。入力画像 x0\mathbf{x}_0 に時刻 tt のノイズを加えた xt\mathbf{x}_t を作り、ネットワークは「この段階で混ざったノイズ」を予測します。教師信号は自分で加えたノイズなので、画像どうしの対応を用意する必要はありません。

これは GAN のように生成器と識別器を敵対的に最適化する枠組みではありません。変分推論に基づく確率モデルとして逆過程を学習し、サンプリングはその逆過程を順に実行します。VAE のような潜在変数モデルとの接点はありますが、ここで扱う DDPM の xt\mathbf{x}_t はデータと同じ次元の中間状態です。潜在空間で同じ考え方を行う潜在拡散や、条件を入力に加える条件付き生成は別の拡張です。

観点DDPM
学習の教師自分で加えたノイズ
生成の開始点標準正規ノイズ
生成の進み方時刻を逆向きに一段ずつ更新

仕組み

前向き過程は、データ x0\mathbf{x}_0 から時刻 TT のノイズ xT\mathbf{x}_T までをマルコフ連鎖で作ります。βt\beta_t は時刻 tt に加えるノイズの分散、I\mathbf{I} は単位行列です。

q(xt∣xt−1)=N(1−βt xt−1, βtI)q(\mathbf{x}_t\mid\mathbf{x}_{t-1}) = \mathcal{N}\left(\sqrt{1-\beta_t}\,\mathbf{x}_{t-1},\,\beta_t\mathbf{I}\right)

αt=1−βt\alpha_t=1-\beta_t、αˉt=∏s=1tαs\bar{\alpha}_t=\prod_{s=1}^{t}\alpha_s と置くと、途中の状態は一段ずつ生成せず直接サンプルできます。

xt=αˉt x0+1−αˉt ϵ,ϵ∼N(0,I)\mathbf{x}_t=\sqrt{\bar{\alpha}_t}\,\mathbf{x}_0+\sqrt{1-\bar{\alpha}_t}\,\boldsymbol{\epsilon},\qquad \boldsymbol{\epsilon}\sim\mathcal{N}(\mathbf{0},\mathbf{I})

x0\mathbf{x}_0 は元データ、xt\mathbf{x}_t は時刻 tt のノイズ画像、ϵ\boldsymbol{\epsilon} は標準正規ノイズです。tt が進むほど信号の係数 αˉt\sqrt{\bar{\alpha}_t} は小さくなり、最終状態を標準正規分布に近づけます。学習時は画像と時刻 tt を選び、この式で xt\mathbf{x}_t を作ります。

逆過程は、ノイズから始めて t=T,T−1,…,1t=T,T-1,\ldots,1 の順に戻ります。モデルが表す一段の遷移を pθ(xt−1∣xt)p_\theta(\mathbf{x}_{t-1}\mid\mathbf{x}_t) とし、平均をニューラルネットワークで決めます。DDPM の実装で中心になるのは、平均をノイズ予測でパラメータ化することです。

μθ(xt,t)=1αt(xt−βt1−αˉt ϵθ(xt,t))\boldsymbol{\mu}_\theta(\mathbf{x}_t,t)=\frac{1}{\sqrt{\alpha_t}}\left(\mathbf{x}_t-\frac{\beta_t}{\sqrt{1-\bar{\alpha}_t}}\,\boldsymbol{\epsilon}_\theta(\mathbf{x}_t,t)\right)

μθ\boldsymbol{\mu}_\theta は逆過程の平均、θ\theta はネットワークの学習パラメータです。学習の重い変分下限を、実装では次の単純なノイズ予測損失として扱います。

Lsimple=Et,x0,ϵ[∥ϵ−ϵθ(xt,t)∥2]L_{\mathrm{simple}}=\mathbb{E}_{t,\mathbf{x}_0,\boldsymbol{\epsilon}}\left[\left\|\boldsymbol{\epsilon}-\boldsymbol{\epsilon}_\theta(\mathbf{x}_t,t)\right\|^2\right]

つまり、学習の1回は「時刻を選ぶ→ノイズを加える→加えたノイズを予測する→二乗誤差を最小化する」です。生成時には予測したノイズから逆過程の平均を計算し、必要な分散のランダムノイズを加えて次の状態を作ります。これを何段も繰り返すため、生成は1回の順伝播で終わらず遅くなります。ここで tt をネットワークに渡し忘れると、ノイズの強さが異なる復元問題を区別できません。

試験でどう問われるか

問われ方正解に寄る条件引っかけ
前向き過程の説明データに小さいガウスノイズを段階的に加え、最終的にノイズへ近づける前向きが画像を生成する過程だとする
逆過程の説明ノイズから始め、時刻を TT から0へ戻す時刻を0から TT へ進める
学習対象の説明ϵ\boldsymbol{\epsilon} そのものではなく、入力 (xt,t)(\mathbf{x}_t,t) から加えたノイズを予測するネットワーク画像を直接回帰する、とだけ説明する
生成が遅い理由多数の逆拡散ステップで順伝播を繰り返す学習データの読み込みが主因とする
GAN との違い識別器との敵対的学習ではなく、ノイズ予測の損失で逆過程を学習する必ず識別器を同時に更新するとする

実装で確かめる

前向き過程の閉形式を NumPy で確認します。x0 に対して時刻を変えると、同じ乱数の形でも信号係数とノイズ係数の組が変わります。

import numpy as np

rng = np.random.default_rng(0)
x0 = np.array([1.0, -0.5, 0.25])
betas = np.array([0.01, 0.05, 0.10])
alphas = 1.0 - betas
alpha_bar = np.cumprod(alphas)
epsilon = rng.normal(size=x0.shape)

for t in range(len(betas)):
    xt = np.sqrt(alpha_bar[t]) * x0 + np.sqrt(1 - alpha_bar[t]) * epsilon
    print(t + 1, np.round(xt, 4))

このコードで作られる xt\mathbf{x}_t は、各時刻のノイズを独立に何度も足した結果と同じ分布です。実装では alpha_bar の添字をデータの時刻と取り違えやすく、t=1t=1 を配列の0番目に対応させるかを最初に固定します。

前向き過程の βtβ_t と αˉt\bar{\alpha}_t は逆過程の式にも現れます。学習時と生成時でスケジュール、時刻の範囲、画像のスケーリングを変えると、ノイズ予測が合っていても逆過程の入力分布がずれます。

取り違えやすいもの

用語拡散モデルとの切り分け
前向き過程学習用にデータを壊す固定の確率過程。通常、ここをニューラルネットワークが学習するわけではない
逆過程ノイズからデータへ戻す学習対象の確率過程。ネットワークは時刻ごとのノイズや平均を予測する
ノイズ予測逆過程を実装しやすくするパラメータ化。ネットワークの出力をそのまま画像と解釈しない
VAEエンコーダ・デコーダで潜在表現を一度に扱う。DDPM は多数の時刻の中間状態を使う
GAN生成器と識別器を敵対的に学習する。DDPM の基本学習はノイズ予測の回帰損失である
自己回帰モデル画素やトークンなどを順序づけて生成する。DDPM も多段だが、時刻ごとのノイズ除去を行う

想起チェック

拡散モデルの前向き過程と逆過程は、それぞれ何をするか

前向き過程はデータに段階的にガウスノイズを加え、逆過程はノイズから始めてその手順を逆向きにたどります。

DDPM のネットワークは各時刻で何を予測するか

入力 xt\mathbf{x}_t と時刻 tt から、前向き過程で加えたノイズ ϵ\boldsymbol{\epsilon} を予測します。その二乗誤差が単純化された学習目的です。

生成が遅くなりやすい理由は何か

完全なノイズから画像まで、逆過程の遷移を時刻 TT から0へ多数回実行するためです。各段階でネットワークの予測が必要になります。

GAN と比べたときの基本的な学習の違いは何か

GAN のような生成器と識別器の敵対的最適化ではなく、DDPM は加えたノイズを予測する損失で逆過程を学習します。

出典