$K$個のクラスのあるカテゴリカル変数$X_0$の観測可能な条件$C$の下での分布を$P(X_0|C)$とする。これは生成対象としたいデータ分布である。
条件$C$に依存しない無情報のカテゴリカル分布$\boldsymbol{\pi}$があるとして、時刻$t: 0 \rightarrow 1$で分布$P(X_0|C)$から$\boldsymbol{\pi}$への離散的な拡散過程を以下のようにする。
P(X_t=j|X_0=i) = (1-\beta_t) \delta_{i,j} + \beta_t \pi_j
ただし、$\beta_t$は$\beta_0 = 0, \beta_1 = 1$で$t$によって単調増加の係数とする。
これは、$(1-\beta_t):\beta_t$で$\boldsymbol{\delta}_i$と$\boldsymbol{\pi}$を混合したカテゴリカル分布の確率である。ただし$\boldsymbol{\delta}_i$は$i$のみ$1$で、他の要素が$0$の確率ベクトル。
時刻$t=0$では$P(X_0|C)$、時刻$t=1$で$\boldsymbol{\pi}$となる。
これを$i,j$要素にする行列を$P_{0, t}$とすると、
P_{0, t} = (1 - \beta_t) I_K + \beta_t \boldsymbol{1} \boldsymbol{\pi}^T
である。
時刻$s < t$について同様に$P(X_t=j|X_s=i)$を$i,j$要素にする行列を$P_{s, t}$とする。この確率過程$X_t$がマルコフ過程であるとすれば、
P_{0, t} = P_{0, s} P_{s, t}
と遷移確率行列の積の関係が成り立つ。
$P_{0, s}$は逆行列を容易に計算できる。$\boldsymbol{\pi}$は確率ベクトルなので、$\boldsymbol{\pi}^T \boldsymbol{1} = 1$であり、
(\boldsymbol{1} \boldsymbol{\pi}^T) (\boldsymbol{1} \boldsymbol{\pi}^T) = \boldsymbol{1} \boldsymbol{\pi}^T
なので$\boldsymbol{1} \boldsymbol{\pi}^T$は冪等行列。したがって、
(\boldsymbol{1} \boldsymbol{\pi}^T) P_{0, s} = \boldsymbol{1} \boldsymbol{\pi}^T
であり、この両辺を$\beta_s$倍して$P_{0, s}$から引けば、
(I_K - \beta_s \boldsymbol{1} \boldsymbol{\pi}^T) P_{0, s} = (1-\beta_s) I_K
となる。つまり、
P_{0, s}^{-1} = \frac{1}{1 - \beta_s} \left(I_K - \beta_s \boldsymbol{1} \boldsymbol{\pi}^T \right)
である。
これにより、
\begin{eqnarray}
P_{s, t} &=& P_{0, s}^{-1} P_{0, t} \\
&=& \frac{1}{1 - \beta_s} \left(I_K - \beta_s \boldsymbol{1} \boldsymbol{\pi}^T \right) \left((1 - \beta_t) I_K + \beta_t \boldsymbol{1} \boldsymbol{\pi}^T \right) \\
&=& \left(1 - \frac{\beta_t - \beta_s}{1 - \beta_s} \right)I_K + \frac{\beta_t - \beta_s}{1 - \beta_s} \boldsymbol{1} \boldsymbol{\pi}^T
\end{eqnarray}
と計算できる。
意味はないが、この$P_{s,t}$を使えば$x_0 \sim P(X_0|C)$のサンプルから、
\begin{eqnarray}
x_{\frac{1}{N}} &\sim& P_{0, \frac{1}{N}}[x_0, \cdot] \\
x_{\frac{2}{N}} &\sim& P_{\frac{1}{N}, \frac{2}{N}}[x_{\frac{1}{N}}, \cdot] \\
\vdots \\
x_{1} &\sim& P_{\frac{N - 1}{N}, 1}[x_{\frac{N-1}{N}}, \cdot] \\
\end{eqnarray}
のようにステップバイステップで$P_{s, t}$が定めるカテゴリカル分布からサンプリングして構成した$x_{1}$が分布$\pi$に従う。
この逆方向の遷移を考えることで、$\pi$を起点に$x_0 \sim P(X_0|C)$からのサンプリングを構成することがDiscreate Diffusionの目的となる。
$X_0$の分布によらず、時刻$s < t$について、
\begin{eqnarray}
P(X_s = k| X_t=j, X_0=i) &=& \frac{P(X_s=k, X_t=j, X_0=i)}{P(X_t=j, X_0=i)} \\
&=& \frac{P(X_t=j|X_s=k) P(X_s=k| X_0=i)}{P(X_t=j | X_0=i)} \\
&=& \frac{P_{0, s}[i, k] P_{s, t}[k, j]}{P_{0,t}[i, j]}
\end{eqnarray}
である。
ここで、
F_{\theta}(i|j, C, t) = P(X_0=i | X_t=j, C)
となるような関数$F_{\theta}$が使用できれば、
\begin{eqnarray}
P(X_s = k| X_t = j) &=& \sum_{i=1}^{K} P(X_s = k| X_t=j, X_0=i)P(X_0=i | X_t=j, C) \\
&=& \sum_{i=1}^{K} F_{\theta}(i|j, C, t) \frac{P_{0, s}[i, k] P_{s, t}[k, j]}{P_{0,t}[i, j]}
\end{eqnarray}
で逆方向$t \rightarrow s$の遷移確率を計算できる。
この逆方向の遷移確率について、$P(X_s = k| X_t = j)$を$j, k$要素とする行列を$Q_{t, s}$とすれば、
\begin{eqnarray}
x_1 &\sim& \pi \\
x_{\frac{N-1}{N}} &\sim& Q_{1, \frac{N-1}{N}}[x_1, \cdot] \\
x_{\frac{N-2}{N}} &\sim& Q_{\frac{1}{N}, \frac{2}{N}}[x_{\frac{N-1}{N}}, \cdot] \\
\vdots \\
x_{0} &\sim& Q_{\frac{1}{N}, 0}[x_{\frac{1}{N}}, \cdot] \\
\end{eqnarray}
のようにステップバイステップにカテゴリカル分布で生成することで、$x_0$が元のデータ分布に従う。
この$X_t$から$X_0$の分布を計算する
F_{\theta}(i|j, C, t) = P(X_0=i | X_t=j, C)
は自明でなく、データ分布$P(X_0|C)$に依存する。データ分布$P(X_0|C)$は未知であるが、この分布からサンプリングされた観測データが存在するとする。
Discrete Diffusionのモデルでは$F_{\theta}$をニューラルネットワークで構成し、データから学習する。
F_{\theta}(i|j, C, t)
は$i=1,...,K$についての確率値であるので、クロスエントロピー損失で学習すればよい。
\begin{eqnarray}
\theta^* &=& \mathop{\rm arg~max}\limits_{\theta} \, \mathcal{L}(\theta)\\
\mathcal{L}(\theta) &=& E_{\substack{(x_0, C) \sim \mathrm{Data} \\ t \sim U(0, 1)\\ x_t \sim P(X_t|X_0=x_0)}} \left[-\log F_{\theta}(x_0|x_t, C, t) \right]
\end{eqnarray}
逆方向$t \rightarrow s$の遷移確率$P(X_s = k| X_t = j)$を求める際に
P(X_s=k|X_t=j, X_0=i) = \frac{P_{0, s}[i, k] P_{s, t}[k, j]}{P_{0,t}[i, j]}
を計算する必要がある。これは$X_0$の分布に依存せず、$\pi$と$\beta_s,\beta_t$から決まる。
具体的には、$\beta_{s,t} = \frac{\beta_t - \beta_s}{1 - \beta_s}$とすれば、
\frac{P_{0, s}[i, k] P_{s, t}[k, j]}{P_{0,t}[i, j]} = \frac{\left((1 - \beta_s) \delta_{i,k} + \beta_s \pi_k \right) \left((1 - \beta_{s,t}) \delta_{k,j} + \beta_{s,t} \pi_j \right)}{(1 - \beta_t) \delta_{i,j} + \beta_t \pi_j}
である。
P(X_s = k| X_t = j) = \sum_{i=1}^{K} F_{\theta}(i|j, C, t)\frac{P_{0, s}[i, k] P_{s, t}[k, j]}{P_{0,t}[i, j]}
なので、
\begin{eqnarray}
w_{i,j} &=& \frac{F_{\theta}(i|j, C, t)}{(1 - \beta_t) \delta_{i,j} + \beta_t \pi_j} \\
W_j &=& \sum_{i=1}^{K} w_{i, j}
\end{eqnarray}
とすれば、
\begin{eqnarray}
P(X_s = k| X_t = j) &=& \sum_{i=1}^{K} w_{i,j} \left((1 - \beta_s) \delta_{i,k} + \beta_s \pi_k \right) \left((1 - \beta_{s,t}) \delta_{k,j} + \beta_{s,t} \pi_j \right) \\
&=& \left((1 - \beta_{s,t}) \delta_{k,j} + \beta_{s,t} \pi_j \right) \sum_{i=1}^{K} w_{i,j} \left((1 - \beta_s) \delta_{i,k} + \beta_s \pi_k \right) \\
&=& \left((1 - \beta_{s,t}) \delta_{k,j} + \beta_{s,t} \pi_j \right) \left(w_{k, j} (1 - \beta_s) + \beta_s \pi_k W_j \right) \\
&=& \left((1 - \beta_t) \delta_{k, j} + (\beta_t - \beta_s)\pi_j\right)w_{k, j} \\
&& + \beta_s \pi_k \left((1 - \beta_{s,t}) \delta_{k,j} + \beta_{s,t} \pi_j \right) W_j \\
&=& (1 - \beta_t) \left(w_{j,j} + \frac{\beta_s}{1 - \beta_s} \pi_j W_j\right) \delta_{k, j} \\
&& + (\beta_t - \beta_s) \pi_j w_{k, j} \\
&& + \left((\beta_t - \beta_s) \frac{\beta_s}{1 - \beta_s} \pi_j W_j \right)\pi_k
\end{eqnarray}
となる。
\boldsymbol{\pi} = \begin{bmatrix} \pi_{1} \\ \vdots \\ \pi_{K}\end{bmatrix}, \boldsymbol{w}_j=\begin{bmatrix} w_{1, j} \\ \vdots \\ w_{K, j}\end{bmatrix}, \boldsymbol{q}_j =\begin{bmatrix} P(X_s = 1| X_t = j) \\ \vdots \\ P(X_s = K| X_t = j)\end{bmatrix}
とすれば、
\begin{eqnarray}
\boldsymbol{q}_j &=& (\beta_t - \beta_s) \pi_j \boldsymbol{w}_j + (\beta_t - \beta_s) \frac{\beta_s}{1 - \beta_s} \pi_j W_j \boldsymbol{\pi} \\
&& + (1 - \beta_t) \left(w_{j,j} + \frac{\beta_s}{1 - \beta_s} \pi_j W_j \right) \boldsymbol{\delta}_j
\end{eqnarray}
である。また、
\begin{eqnarray}
w_{j,j} + \frac{\beta_s}{1 - \beta_s} \pi_j W_j &=& \left( 1 + \frac{\beta_s \pi_j}{1 - \beta_s} \right) w_{j, j} + \frac{\beta_s \pi_j}{1 - \beta_s} \sum_{i \neq j}^{K} w_{i, j} \\
&=& \frac{1 - \beta_s + \beta_s \pi_j}{1 - \beta_s} \frac{F_{\theta}(j|j, C, t)}{1 - \beta_t + \beta_t \pi_j} + \frac{\beta_s \pi_j}{1 - \beta_s} \sum_{i \neq j}^{K} \frac{F_{\theta}(i|j, C, t)}{\beta_t \pi_j} \\
&=& \frac{1}{1 - \beta_s} \left( \frac{1 - \beta_s + \beta_s \pi_j}{1 - \beta_t + \beta_t \pi_j} F_{\theta}(j|j, C, t) + \frac{\beta_s}{\beta_t} \sum_{i \neq j}^{K} F_{\theta}(i|j, C, t) \right) \\
&=& \frac{1}{1 - \beta_s} \left( \frac{\beta_s}{\beta_t} + \left( \frac{1 - \beta_s + \beta_s \pi_j}{1 - \beta_t + \beta_t \pi_j} - \frac{\beta_s}{\beta_t} \right)F_{\theta}(j|j, C, t) \right) \\
\end{eqnarray}
より、
\begin{eqnarray}
\boldsymbol{q}_j &=& (\beta_t - \beta_s) \pi_j \boldsymbol{w}_j + (\beta_t - \beta_s) \frac{\beta_s}{1 - \beta_s} \pi_j W_j \boldsymbol{\pi} \\
&& + \frac{1 - \beta_t}{1 - \beta_s} \left( \frac{\beta_s}{\beta_t} + \left( \frac{1 - \beta_s + \beta_s \pi_j}{1 - \beta_t + \beta_t \pi_j} - \frac{\beta_s}{\beta_t} \right)F_{\theta}(j|j, C, t) \right) \boldsymbol{\delta}_j
\end{eqnarray}
とも表せる。
さらに、
\begin{eqnarray}
w_{k,j} &=& \frac{F_{\theta}(k|j, C, t)}{(1 - \beta_t) \delta_{k,j} + \beta_t \pi_j} \\
&=& \frac{F_{\theta}(k|j, C, t)}{\beta_t \pi_j} + \delta_{k, j}\left(\frac{F_{\theta}(k|j, C, t)}{1 - \beta_t + \beta_t \pi_j} - \frac{F_{\theta}(k|j, C, t)}{\beta_t \pi_j} \right) \\
&=& \frac{F_{\theta}(k|j, C, t)}{\beta_t \pi_j} - \delta_{k, j}\frac{1-\beta_t} {\beta_t \pi_j(1 - \beta_t + \beta_t \pi_j)} F_{\theta}(k|j, C, t) \\
\end{eqnarray}
なので、$\boldsymbol{F}_{\theta}(j, C, t)$
\boldsymbol{F}_{\theta}(j, C, t) = \begin{bmatrix} F_{\theta}(1|j, C, t) \\ \vdots \\ F_{\theta}(K|j, C, t)\end{bmatrix}
とすれば、
\boldsymbol{w}_j = \frac{1}{\beta_t \pi_j} \boldsymbol{F}_{\theta}(j, C, t) - \frac{1-\beta_t} {\beta_t \pi_j(1 - \beta_t + \beta_t \pi_j)} F_{\theta}(j|j,C,t)\boldsymbol{\delta}_j\
となり、$\boldsymbol{F}_{\theta}(j, C, t)$は確率ベクトルで和が$1$なので
\begin{eqnarray}
W_j &=& \frac{1}{\beta_t \pi_j} - \frac{1-\beta_t}{\beta_t \pi_j(1 - \beta_t + \beta_t \pi_j)} F_{\theta}(j|j,C,t) \\
&=& \frac{(1 - \beta_t)(1 - F_{\theta}(j|j,C,t)) + \beta_t \pi_j}{\beta_t \pi_j(1 - \beta_t + \beta_t \pi_j)}
\end{eqnarray}
となる。
したがって、係数$a_j$が存在して
\begin{eqnarray}
\boldsymbol{q}_j &=& \frac{\beta_t - \beta_s}{\beta_t}\boldsymbol{F}_{\theta}(j, C, t) + \frac{(\beta_t - \beta_s)\beta_s}{(1 - \beta_s) \beta_t}\left( 1 - \frac{(1 - \beta_t)F_{\theta}(j|j,C,t)}{1 - \beta_t + \beta_t \pi_j} \right) \boldsymbol{\pi} + a_j \boldsymbol{\delta}_j
\end{eqnarray}
となる。$\boldsymbol{q}_j, \boldsymbol{\delta}_j, \boldsymbol{F}_{\theta}(j, C, t), \boldsymbol{\pi}$は全て和が$1$のベクトルなので、
1 = \frac{\beta_t - \beta_s}{\beta_t} + \frac{(\beta_t - \beta_s)\beta_s}{(1 - \beta_s) \beta_t}\left( 1 - \frac{(1 - \beta_t)F_{\theta}(j|j,C,t)}{1 - \beta_t + \beta_t \pi_j} \right) + a_j
である。$a_j$について整理すると、
\begin{eqnarray}
a_j &=& 1 - \frac{\beta_t - \beta_s}{\beta_t} - \frac{(\beta_t - \beta_s)\beta_s}{(1 - \beta_s) \beta_t}\left( 1 - \frac{(1 - \beta_t)F_{\theta}(j|j,C,t)}{1 - \beta_t + \beta_t \pi_j} \right) \\
&=& \frac{\beta_s(1 - \beta_t)}{(1 - \beta_s)\beta_t} + \frac{(\beta_t - \beta_s)\beta_s(1 - \beta_t) F_{\theta}(j|j,C,t)}{(1 - \beta_s) \beta_t (1 - \beta_t + \beta_t \pi_j)} \\
&=& \frac{\beta_s(1 - \beta_t)}{(1 - \beta_s)\beta_t} \left( 1 + \frac{(\beta_t - \beta_s) F_{\theta}(j|j,C,t)}{1 - \beta_t + \beta_t \pi_j}\right)
\end{eqnarray}
となる。
したがって、
\begin{eqnarray}
\boldsymbol{q}_j &=& \frac{\beta_t - \beta_s}{\beta_t}\boldsymbol{F}_{\theta}(j, C, t) \\
&& + \frac{(\beta_t - \beta_s)\beta_s}{(1 - \beta_s) \beta_t} \frac{(1 - \beta_t)(1 - F_{\theta}(j|j,C,t)) + \beta_t \pi_j}{1 - \beta_t + \beta_t \pi_j} \boldsymbol{\pi} \\
&& + \frac{\beta_s(1 - \beta_t)}{(1 - \beta_s)\beta_t} \left( 1 + \frac{(\beta_t - \beta_s) F_{\theta}(j|j,C,t)}{1 - \beta_t + \beta_t \pi_j}\right) \boldsymbol{\delta}_j
\end{eqnarray}
となる。
これは、$\boldsymbol{F}_{\theta}(j, C, t)$による時刻$0$の予測確率の影響を$t \rightarrow s$でのノイズ減少の割合$\frac{\beta_t - \beta_s}{\beta_t}$分にスケールし、無情報の事前分布$\boldsymbol{\pi}$と混合し、クラスが変わらない確率$\boldsymbol{\delta}_j$で調整しているような構成になっている。