Structured Denoising Diffusion Models in Discrete State-Spaces

Jacob Austin, Daniel D. Johnson, Jonathan Ho, Daniel Tarlow, Rianne van den Berg

Introduction

Generative modeling is a core problem in machine learning, useful both for benchmarking our ability to capture statistics of natural datasets and for downstream applications that require generating high-dimensional data like images, text, and speech waveforms. There has been a great deal of progress with the development of methods like GANs , VAEs , large autoregressive neural network models , normalizing flows , and others, each with their own tradeoffs in terms of sample quality, sampling speed, log-likelihoods, and training stability.

Recently, diffusion models have emerged as a compelling alternative for image and audio generation, achieving comparable sample quality to GANs and log-likelihoods comparable to autoregressive models with fewer inference steps. A diffusion model is a parameterized Markov chain trained to reverse a predefined forward process, which is a stochastic process constructed to gradually corrupt training data into pure noise. Diffusion models are trained using a stable objective closely related to both maximum likelihood and score matching , and they admit faster sampling than autoregressive models by using parallel iterative refinement .

Although diffusion models have been proposed in both discrete and continuous state spaces , most recent work has focused on Gaussian diffusion processes that operate in continuous state spaces (e.g. for real-valued image and waveform data). Diffusion models with discrete state spaces have been explored for text and image segmentation domains , but they have not yet been demonstrated as a competitive model class for large scale text or image generation.

Our aim in this work is to improve and extend discrete diffusion models by using a more structured categorical corruption process to shape data generation, as illustrated in Figure 1. Our models do not require relaxing or embedding discrete data (including images) into continuous spaces, and can embed structure or domain knowledge into the transition matrices used by the forward process. We achieve significantly improved results by taking advantage of this flexibility. We develop structured corruption processes appropriate for text data, using similarity between tokens to enable gradual corruption and denoising. Expanding further, we also explore corruption processes that insert [MASK] tokens, which let us draw parallels to autoregressive and mask-based generative models. Finally, we study discrete diffusion models for quantized images, taking inspiration from the locality exploited by continuous diffusion models. This leads to a particular choice of discrete corruption process that diffuses preferentially to more similar states and leads to much better results in the image domain.

Overall, we make a number of technical and conceptual contributions. Beyond designing several new structured diffusion models, we introduce a new auxiliary loss which stabilizes training of D3PMs and a family of noise schedules based on mutual information that lead to improved performance. We strongly outperform various non-autoregressive baselines for text generation on character-level text generation, and successfully scale discrete diffusion models to large vocabularies and long sequence lengths. We also achieve strong results on the image dataset CIFAR-10, approaching or exceeding the Gaussian diffusion model from Ho et al. on log-likelihoods and sample quality.

Background: diffusion models

Diffusion models are latent variable generative models characterized by a forward and a reverse Markov process. The forward process q(x1:T∣x0)=∏t=1Tq(xt∣xt−1)q(\bm{x}_{1:T}|\bm{x}_{0})=\prod_{t=1}^{T}q(\bm{x}_{t}|\bm{x}_{t-1}) corrupts the data x0∼q(x0)\bm{x}_{0}\sim q(\bm{x}_{0}) into a sequence of increasingly noisy latent variables x1:T=x1,x2,...,xT\bm{x}_{1:T}=\bm{x}_{1},\bm{x}_{2},...,\bm{x}_{T}. The learned reverse Markov process pθ(x0:T)=p(xT)∏t=1Tpθ(xt−1∣xt)p_{\theta}(\bm{x}_{0:T})=p(\bm{x}_{T})\prod_{t=1}^{T}p_{\theta}(\bm{x}_{t-1}|\bm{x}_{t}) gradually denoises the latent variables towards the data distribution. For example, for continuous data, the forward process typically adds Gaussian noise, which the reverse process learns to remove.

In order to optimize the generative model pθ(x0)p_{\theta}(\bm{x}_{0}) to fit the data distribution q(x0)q(\bm{x}_{0}), we typically optimize a variational upper bound on the negative log-likelihood:

When the number of time steps TT goes to infinity, both the forward process and the reverse process share the same functional form , allowing the use of a learned reverse process from the same class of distributions as that of the forward process. Furthermore, for several choices of the forward process the distribution q(xt∣x0)q(\bm{x}_{t}|\bm{x}_{0}) converges to a stationary distribution π(x)\pi(\bm{x}) in the limit t→∞t\rightarrow\infty independent of the value of x0\bm{x}_{0}. When the number of time steps TT is large enough and we choose π(x)\pi(\bm{x}) as the prior p(xT)p(\bm{x}_{T}), we can guarantee that the LTL_{T} term in (1) will approach zero regardless of the data distribution q(x0)q(\bm{x}_{0}). (Alternatively, one can use a learned prior pθ(xT)p_{\theta}(\bm{x}_{T}).)

While q(xt∣xt−1)q(\bm{x}_{t}|\bm{x}_{t-1}) can in theory be arbitrary, efficient training of pθp_{\theta} is possible when q(xt∣xt−1)q(\bm{x}_{t}|\bm{x}_{t-1}):

Permits efficient sampling of xt\bm{x}_{t} from q(xt∣x0)q(\bm{x}_{t}|\bm{x}_{0}) for an arbitrary time tt, allowing us to randomly sample timesteps and optimize each Lt−1L_{t-1} term individually with stochastic gradient descent,

Has a tractable expression for the forward process posterior q(xt−1∣xt,x0)q(\bm{x}_{t-1}|\bm{x}_{t},\bm{x}_{0}), which allows us to compute the KL divergences present in the Lt−1L_{t-1} term of (1).

The majority of recent work in continuous spaces defines the forward and reverse distributions as q(xt∣xt−1)=N(xt∣1−βtxt−1,βtI)q(\bm{x}_{t}|\bm{x}_{t-1})=\mathcal{N}\left(\bm{x}_{t}|\sqrt{1-\beta_{t}}\bm{x}_{t-1},\beta_{t}\bm{I}\right) and pθ(xt−1∣xt)=N(xt−1∣μθ(xt,t),Σθ(xt,t))p_{\theta}(\bm{x}_{t-1}|\bm{x}_{t})=\mathcal{N}\left(\bm{x}_{t-1}|\bm{\mu}_{\theta}(\bm{x}_{t},t),\bm{\Sigma}_{\theta}(\bm{x}_{t},t)\right), respectively. The aforementioned properties hold in the case of these Gaussian diffusion models: the forward process q(xt∣x0)q(\bm{x}_{t}|\bm{x}_{0}) converges to a stationary distribution, motivating the choice p(xT)=N(xT∣0,I)p(\bm{x}_{T})=\mathcal{N}\left(\bm{x}_{T}|\bm{0},\bm{I}\right), and both q(xt∣x0)q(\bm{x}_{t}|\bm{x}_{0}) and q(xt−1∣xt,x0)q(\bm{x}_{t-1}|\bm{x}_{t},\bm{x}_{0}) are tractable Gaussian distributions for which the KL divergence can be computed analytically.

Diffusion models for discrete state spaces

Diffusion models with discrete state spaces were first introduced by Sohl-Dickstein et al. , who considered a diffusion process over binary random variables. Hoogeboom et al. extended the model class to categorical random variables with transition matrices characterized by uniform transition probabilities. In their supplementary material, Song et al. also derived this extension, although no experiments were performed with this model class. Here, we briefly describe a more general framework for diffusion with categorical random variables which includes these models as special cases.

For scalar discrete random variables with KK categories xt,xt−1∈1,...,Kx_{t},x_{t-1}\in{1,...,K} the forward transition probabilities can be represented by matrices: [Qt]ij=q(xt=j∣xt−1=i)[\bm{Q}_{t}]_{ij}=q(x_{t}=j|x_{t-1}=i). Denoting the one-hot version of xx with the row vector x\bm{x}, we can write

Note that due to the Markov property of the forward process q(xt∣xt−1,x0)=q(xt∣xt−1)q(\bm{x}_{t}|\bm{x}_{t-1},\bm{x}_{0})=q(\bm{x}_{t}|\bm{x}_{t-1}). Assuming that the reverse process pθ(xt∣xt−1)p_{\theta}(\bm{x}_{t}|\bm{x}_{t-1}) is also factorized as conditionally independent over the image or sequence elements, the KL divergence between qq and pθp_{\theta} can be computed by simply summing over all possible values of each random variable; we thus satisfy criteria 1 and 2 discussed in Section 2. Depending on Qt\bm{Q}_{t}, the cumulative products Q‾t\overline{\bm{Q}}_{t} can often be computed in closed form, or simply precomputed for all tt. However, for large KK and large TT this may be prohibitive. In Appendix A.4 we discuss how to ensure Q‾t\overline{\bm{Q}}_{t} can still be computed efficiently in this case, allowing the framework to scale to a larger number of categories.

In the next section we discuss the choice of the Markov transition matrices Qt\bm{Q}_{t} and corresponding stationary distributions. From here on, we refer to the general class of diffusion models with discrete state spaces as Discrete Denoising Diffusion Probabilistic Models (D3PMs).

An advantage of the D3PM framework described above is the ability to control the data corruption and denoising process by choosing Qt\bm{Q}_{t}, in notable contrast to continuous diffusion, for which only additive Gaussian noise has received significant attention. Besides the constraint that the rows of Qt\bm{Q}_{t} must sum to one to conserve probability mass, the only other constraint in choosing Qt\bm{Q}_{t} is that the rows of Q‾t=Q1Q2…Qt\overline{\bm{Q}}_{t}=\bm{Q}_{1}\bm{Q}_{2}\ldots\bm{Q}_{t} must converge to a known stationary distributionIf a stationary distribution is not known, we can introduce a learned prior pθ(xT)p_{\theta}(\bm{x}_{T}); we note that this is equivalent to extending the forward process by appending a rank-one matrix QT+1\bm{Q}_{T+1} that ignores xT\bm{x}_{T} and produces a deterministic xT+1\bm{x}_{T+1}, then learning the reverse step pθ(xT∣xT+1)=pθ(xT)p_{\theta}(\bm{x}_{T}|\bm{x}_{T+1})=p_{\theta}(\bm{x}_{T}). when tt becomes large, which can be guaranteed while imposing minimal restrictions on Qt\bm{Q}_{t} (see Appendix A.1).

We argue that for most real-world discrete data, including images and text, it makes sense to add domain-dependent structure to the transition matrices Qt\bm{Q}_{t} as a way of controlling the forward corruption process and the learnable reverse denoising process. Below we briefly discuss the uniform transition matrices that have been studied in prior work , along with a set of structured transition matrices we have explored for our image and text dataset experiments; see Appendix A.2 for more details on each matrix type. We also note that this set is not exhaustive, and many other transition matrices could also be used within the D3PM framework.

Absorbing state (Appendix A.2.2). Motivated by the success of BERT and recent work on Conditional Masked Language Models (CMLMs) in text, we consider a transition matrix with an absorbing state (called [MASK]), such that each token either stays the same or transitions to [MASK] with some probability βt\beta_{t}. This does not impose particular relationships between categories, similar to uniform diffusion, but still allows corrupted tokens to be distinguished from original ones. Moreover, the stationary distribution is not uniform but has all the mass on the [MASK] token. For images, we reuse the grey pixel as the [MASK] absorbing token.

Discretized Gaussian (Appendix A.2.3). Instead of transitioning uniformly to any other state, for ordinal data we propose imitating a continuous space diffusion model by using a discretized, truncated Gaussian distribution. We choose a normalization such that the transition matrix is doubly stochastic, leading to a uniform stationary distribution. This transition matrix will transition between more similar states with higher probability, and is well suited for quantized ordinal data such as images.

Token embedding distance (Appendix A.2.4). Textual data does not have ordinal structure, but there may still be interesting semantic relationships. For instance, in a character level vocabulary vowels may be more similar to each other than they are to consonants. As a demonstration of the generality of the D3PM framework, we explore using similarity in an embedding space to guide the forward process, and construct a doubly-stochastic transition matrix that transitions more frequently between tokens that have similar embeddings while maintaining a uniform stationary distribution.

For uniform and absorbing-state diffusion, the cumulative products Q‾t\overline{\bm{Q}}_{t} can be computed in closed form (see Appendix A.4.1); the remainder can be precomputed.

2 Noise schedules

We consider several different options for the noise schedule of the forward process. For discretized Gaussian diffusion, we explore linearly increasing the variance of the Gaussian before discretizing it. (Note that a linear schedule for Qt\bm{Q}_{t} leads to a nonlinear amount of cumulative noise in Q‾t\overline{\bm{Q}}_{t}.) For uniform diffusion we use the cosine schedule which sets the cumulative probability of a transition to a cosine function, as introduced by Nichol and Dhariwal and adapted by Hoogeboom et al. . For a general set of transition matrices Qt\bm{Q}_{t} (such as the one based on token embeddings), previously proposed schedules may not be directly applicable. We consider linearly interpolating the mutual information between xt\bm{x}_{t} and x0\bm{x}_{0} to zero, i.e. I(xt;x0)≈(1−tT) H(x0)I(\bm{x}_{t};\bm{x}_{0})\approx(1-\frac{t}{T})\,H(\bm{x}_{0}). Interestingly, for the specific case of absorbing-state D3PMs, this schedule reduces to exactly the (T−t+1)−1(T-t+1)^{-1} schedule proposed by Sohl-Dickstein et al. for a Bernoulli diffusion process. See Appendix A.7 for more details.

3 Parameterization of the reverse process

Finally, when modeling ordinal discrete data, instead of predicting the logits of p~θ(x~0∣xt)\widetilde{p}_{\theta}(\widetilde{\bm{x}}_{0}|\bm{x}_{t}) directly with the output of a neural net, another option is to model the probabilities with a truncated discretized logistic distribution (see Appendix A.8). This provides an extra ordinal inductive bias to the reverse model and boosts FID and log-likelihood scores for images.

4 Loss function

Connection to existing probabilistic models for text

In this section we expand on interesting connections between the D3PM framework and several existing probabilistic and language modeling approaches.

Autoregressive models are (discrete) diffusion models: Consider a diffusion process that deterministically masks tokens one-by-one in a sequence of length N=TN=T: q([xt]i∣x0)=[x0]i if i<N−t else [MASK] q(\left[\bm{x}_{t}\right]_{i}\mid\bm{x}_{0})=[\bm{x}_{0}]_{i}\text{ if }i<N-t\text{ else [MASK] }. This is a deterministic forward process, so q(xt−1∣xt,x0)q(\bm{x}_{t-1}|\bm{x}_{t},\bm{x}_{0}) is a delta distribution on the xt\bm{x}_{t} sequence with one fewer mask: q([xt−1]i∣xt,x0)=δ[xt]i if i≠T−t else δ[x0]iq(\left[\bm{x}_{t-1}\right]_{i}|\bm{x}_{t},\bm{x}_{0})=\delta_{[\bm{x}_{t}]_{i}}\text{ if }i\neq T-t\text{ else }\delta_{[\bm{x}_{0}]_{i}}. While this process is not applied independently to each token, it can be recast as an independently-applied diffusion process on the product space [0...N]×V[0...N]\times\mathcal{V}, where each token is tagged with its position in the sequence, V\mathcal{V} is the vocabulary, and Q\bm{Q} is an N×∣V∣×N×∣V∣N\times|\mathcal{V}|\times N\times|\mathcal{V}| sparse matrix.

Because all tokens except the one at position i=T−ti=T-t have deterministic posteriors, the KL divergence DKL(q([xt−1]j∣xt,x0)∣∣pθ([xt−1]j∣xt))D_{KL}(q([\bm{x}_{t-1}]_{j}|\bm{x}_{t},\bm{x}_{0})\mid\mid p_{\theta}([\bm{x}_{t-1}]_{j}|\bm{x}_{t})) is zero for all other positions. The only token for which this is not true is the token at position ii, for which DKL(q([xt−1]i∣xt,x0)∣∣pθ([xt−1]i∣xt))=−log⁡pθ([x0]i∣xt)D_{KL}(q([\bm{x}_{t-1}]_{i}|\bm{x}_{t},\bm{x}_{0})\mid\mid p_{\theta}([\bm{x}_{t-1}]_{i}|\bm{x}_{t}))=-\log p_{\theta}([\bm{x}_{0}]_{i}|\bm{x}_{t}), the standard cross entropy loss for an autoregressive model.

(Generative) Masked Language-Models (MLMs) are diffusion models: Generative Masked Language Models (, ) are generative models that generate text from a sequence of [MASK] tokens. They are usually trained by sampling a sequence x0\bm{x}_{0}, masking kk tokens according to some schedule, and learning to predict the masked tokens given context. It turns out that a D3PM absorbing ([MASK]) model trained on the usual ELBO objective with the x0\bm{x}_{0}-parameterization from 3.3 reduces to a reweighted version of this MLM objective (see Appendix A.3 for a detailed derivation).

Text generation

For text, we experiment with generation on two datasets: text8 , a character-level dataset extracted from English-language Wikipedia, and the One Billion Word dataset (LM1B) , a large dataset of shuffled English-language sentences. For both, we train a D3PM uniform model based on the work by Hoogeboom et al. (D3PM uniform) and a model that masks tokens (D3PM absorbing). We also consider a model that transitions uniformly to nearest neighbors in a token embedding space (D3PM NN). We follow Hoogeboom et al. and use T=1000T=1000 timesteps, although we are also able to evaluate on fewer due to the parameterization in Section 3.3.

text8 is a character-level text dataset consisting of a small vocabulary of 27 tokens: the letters ‘a’-‘z’ and the ‘_’ whitespace token. We follow the convention of training and evaluating text8 in chunks of length 256 without any preprocessing . For nearest-neighbor D3PM, our nearest neighbor graph in character-space is shown in Appendix B.2.1. D3PM uniform models were trained with a cosine schedule from Hoogeboom et al. (ablations in Appendix B.2.1), while D3PM absorbing and D3PM NN models were trained with a mutual information schedule.

2 Text generation on LM1B

Text generation for large-scale text datasets and large vocabularies with discrete diffusion models has not been previously demonstrated. We include results from LM1B as a proof of concept, showing that these models can indeed scale (as discussed in Appendix A.4), and that the D3PM absorbing model continues to excel. All models were trained and evaluated on packed sequences of length 128128, using a sentencepiecehttps://github.com/google/sentencepiece vocabulary of size 81928192.

Table 2 contains results from experiments on LM1B. Overall, mask diffusion (D3PM absorbing) does relatively well, approaching the performance of a comparable autoregressive model of the same size, and scaling to far fewer steps, while uniform diffusion performs significantly worse. We find, surprisingly, that the D3PM NN model performs worse than the uniform model in terms of log likelihoods (although it demonstrates unique qualitative behavior). This suggests that word embedding similarity may not be a meaningful kind of locality in a diffusion process. We found the the Lλ=0.01L_{\lambda=0.01} loss worked best for the mask absorbing model, but reduced performance for the other models. We note the surprising scaling in perplexity in Figure 2, achieving strong results with as few as 10 inference steps. We also show samples from our model and completions from corrupted samples.

Image generation

We evaluate the performance of several D3PM models on the task of unconditional image generation with the dataset CIFAR-10 . We follow Ho et al. and use T=1000T=1000 timesteps for all models and verify that for all models the forward process converges to the stationary distribution within TT steps, yielding a value of at most LT≈10−5L_{T}\approx 10^{-5} bits per dimension. We train three versions of D3PM with different transition matrices: doubly stochastic matrices with uniform transition probabilities (D3PM uniform) , transition matrices with an absorbing state located at R, G and B values of 128 (D3PM absorbing) and doubly stochastic discretized Gaussian transition matrices (D3PM Gauss). For the D3PM uniform model we experimented with a linear βt\beta_{t} schedule as well as the cosine schedule as proposed in , with the cosine schedule producing the best results. For D3PM absorbing we use the schedule βt=(T−t+1)−1\beta_{t}=(T-t+1)^{-1} as also proposed in , which corresponds to increasing the probability of being in the absorbing state linearly over time. For D3PM Gauss we use the same linear schedule as in . See Appendix B.1 for more details on the experimental setup.

Related Work

Diffusion generative models were first proposed by Sohl-Dickstein et al. and have gained renewed attention recently due to strong results on image and waveform generation . Recent works have proposed improvements for diffusion model training, including importance sampling of the ELBO, better noise schedules and implicit diffusion models . Several works have also drawn connections to score matching , leading to improved sampling algorithms in the continuous-time limit .

While most works have considered continuous diffusion models, discrete diffusion-like models were described in and applied to text generation and image segmentation data in . Some works have dealt with discrete data by embedding it in continuous space and leveraging Gaussian diffusion, but have not applied this to text. Seff et al. also considered generation of discrete structured objects using a diffusion-like Markov corruption process.

For text, denoising autoencoders have a long history both in representation learning and more recently as generative models . These closely resemble our absorbing state diffusion variants for a particular schedule and transition matrix (see Section 4), although our framing allows us to compute log-likelihoods and experiment with alternative transition matrices. Other works have considered non-autoregressive translation and speech transcription via insertion and deletion , masking , and iteratively-refined sequence alignments .

Discussion

Acknowledgments and Disclosure of Funding

We would like to thank Hugo Larochelle for providing high-level feedback during the project, and Ben Poole for reviewing a draft version of this manuscript. We would also like to thank Julia Kreutzer and Xavier Garcia for helpful conversations about language experiments. We, the authors, declare to have no competing interests. The research conducted for this paper was entirely supported by Google.

References

Appendix A Additional details regarding D3PMs

As discussed in Section 3.1, there are two constraints on Qt\bm{Q}_{t} that allow it to be used within a D3PM: the rows of Qt\bm{Q}_{t} must sum to one to conserve probability mass, and the rows of Q‾t=Q1Q2…Qt\overline{\bm{Q}}_{t}=\bm{Q}_{1}\bm{Q}_{2}\ldots\bm{Q}_{t} must converge to a known stationary distribution as tt becomes large. Technically, it is also possible to use a learned prior pθ(xT)p_{\theta}(\bm{x}_{T}), but assuming this is still modeled under a conditional independence assumption, q(xT∣x0)q(\bm{x}_{T}|\bm{x}_{0}) must still be close to a stationary distribution for the LTL_{T} loss term to be small.

One way to ensure that this occurs is to chose Qt\bm{Q}_{t} as increasing powers of a doubly stochastic base matrix Q\bm{Q} (rows and columns sum to 1) with strictly positive entries. This is enough to ensure that Q\bm{Q} is is irreducible and aperiodic and that product Q‾t\overline{\bm{Q}}_{t} converges as t→∞t\rightarrow\infty to a uniform distribution over all states. To show this, consider πi=1/K\pi_{i}=1/K for i=1,...,Ki=1,...,K, and ∑i=1KQi,:=1\sum_{i=1}^{K}\bm{Q}_{i,:}=\bm{1} and ∑j=1KQ:,j=1\sum_{j=1}^{K}\bm{Q}_{:,j}=\bm{1}, then [Qπ]i=∑j=1KQi,jπj=1/K∑j=1KQi,j=1/K=πi[\bm{Q}\bm{\pi}]_{i}=\sum_{j=1}^{K}\bm{Q}_{i,j}\pi_{j}=1/K\sum_{j=1}^{K}\bm{Q}_{i,j}=1/K=\pi_{i}, thus the uniform distribution is an eigenvector of the transition matrix with eigenvalue 1. Convergence to this distribution follows from the Perron-Frobenius theorem for positive square matrices.

More generally, a similar argument shows that even for Qt\bm{Q}_{t} that are not powers of the same base matrix, as long as each Qt\bm{Q}_{t} is doubly stochastic, irreducible, and aperiodic, the uniform distribution is the only possible stationary distribution, and as long as the second largest eigenvalue of Qt\bm{Q}_{t} is bounded below, the cumulative product Q‾t\overline{\bm{Q}}_{t} will converge to the uniform distribution. In practice, we choose Qt\bm{Q}_{t} to add more noise as tt increases, which ensures that Q‾T\overline{\bm{Q}}_{T} is very close to reaching a uniform stationary distribution.

A.2 More details on possible choices of Markov transition matrices

The transition matrix described by Sohl-Dickstein et al. for the binary case, and extended by Hoogeboom et al. , to the categorical case, can be represented using the following K×KK\times K transition matrix

A.2.2 Diffusion with an absorbing state

For our diffusion models with an absorbing state mm, we use the following matrix:

For text generation, we let mm be the [MASK] token at index K−1K-1; this leads to a BERT-like training objective, which masks tokens according to some schedule and learns to denoise them iteratively (see Section 4). For image generation, we set mm to the gray RGB pixel (128,128,128)(128,128,128) at index K//2K//2.

A.2.3 Discretized Gaussian transition matrices

For our D3PM models applied to ordinal data, inspired by continuous-space diffusion models, we use the following K×KK\times K matrix:

Normalization is ensured by assigning the diagonal values to one minus the sum of each row (not including the diagonal entry). Note that due to the normalization of the off-diagonal values over the range {−K+1,...,K−1}\{-K+1,...,K-1\} the sum of each row excluding the diagonal entry is always smaller than 1. The result yields an irreducible doubly stochastic matrix and a forward process with a uniform stationary distribution. Similar to the continuous Gaussian diffusion model, the parameters βt\beta_{t} influence the variance of the forward process distributions.

A.2.4 Structured diffusion in text: using word-embedding distance to introduce locality

For text, we construct a kk-nearest neighbor adjacency matrix

constructed from a pre-trained embedding space over the vocabulary. Then we consider a symmetrized adjacency matrix of the form A=(G+GT)/(2k)\mathbf{A}=(\mathbf{G}+\mathbf{G}^{T})/(2k) where kk is the number of nearest neighbors of each node, and finally construct a doubly stochastic rate matrix with

Our final transition matrix is constructed as a matrix exponential of this rate matrix:

Since R\bm{R} is symmetric and sums to zero along each row, Qt\mathbf{Q}_{t} is doubly stochastic, which ensures we have a uniform stationary distribution (as long as GG is connected). Increasing αt\alpha_{t} over time allows us to add more noise for larger values of tt.

Assuming word embeddings are some metric for syntactic or semantic similarity, this results in a corruption process that gradually moves away from the ground-truth sentence, swapping words with nearest-neighbors in embedding space. For character level modeling, this is a graph over characters, which more often transitions for instance from vowels to other vowels than from vowels to consonants. For words, this could transition between semantically similar words.

For example, in Figure 4, we construct the forward process to diffuse from "dog" to "cat" or "cow", which are nearby in embedding space, but not to more distant words. We can either bootstrap this process by updating the transition matrix Q\bm{Q} dynamically during training, or use pretrained embeddings; we use pretrained embeddings for all of our experiments.

A.2.5 Band-diagonal transitions

A class of transition matrices that introduce local, ordinal inductive biases for structured data are band-diagonal transition matrices which only allow the corruption process to transition locally between states and biases the reverse process towards local iterative refinement. For example, in images, this can be used to allow transitions only between adjacent pixel values.

where vv is the number of nonzero off-diagonal elements of Q\bm{Q} above (and below) the main diagonal. Note that this is a doubly stochastic matrix, so the stationary distribution is uniform. We do not use these in our experiments.

A.2.6 Combinations of absorbing diffusion and other diffusion

A.3 Generative Masked Language Models are Diffusion Models

Generative Masked Language Models are generative models that generate text from a sequence of [MASK] tokens. These are usually trained by sampling a sequence x0\bm{x}_{0}, masking tokens according to some schedule, and learning to predict the masked tokens given context. The actual masking procedure can either be done independently, i.e. by masking each token with probability p=k/Tp=k/T, like Devlin et al. , or by sampling exactly kk tokens. The usual objective isSometimes the loss is un-normalized or normalized by the full sequence length.:

where we first sample a datapoint x0\bm{x}_{0}, sample a number of tokens to mask kk (either uniformly or according to some schedule), then mask that many tokens at random and compute a cross entropy loss over those masked tokens. We claim that this training objective is a (reweighted) absorbing-state D3PM objective with a particular noise schedule and the x0\bm{x}_{0}-parameterization from 3.3 (and indeed, that any absorbing-state D3PM model with [MASK] as the absorbing state will be a reweighted version of this loss with different weights assigned to different numbers of masked tokens kk).

Consider a D3PM with a schedule that masks tokens with probability βt\beta_{t}. The reverse process predicts p~θ(x0~∣xt)\widetilde{p}_{\theta}(\widetilde{\bm{x}_{0}}|\bm{x}_{t}), then uses the forward process to compute pθ(xt−1∣xt)∝∑q(xt−1,xt∣x0~)p~θ(x~0∣xt)p_{\theta}(\bm{x}_{t-1}|\bm{x}_{t})\propto\sum q(\bm{x}_{t-1},\bm{x}_{t}|\widetilde{\bm{x}_{0}})\widetilde{p}_{\theta}(\widetilde{\bm{x}}_{0}|\bm{x}_{t}). In the particular case of absorbing-state diffusion, for each masked token [xt]i=m[\bm{x}_{t}]_{i}=m in xt\bm{x}_{t}, we thus have

We note that for each unmasked token [xt]i=[x0]i[\bm{x}_{t}]_{i}=[\bm{x}_{0}]_{i}, the KL-divergence is zero since unmasked tokens cannot make any other type of transition other than becoming masked. Also, the term in the KL divergence due to the probability of mask transitions is a constant, since mask transitions are independent of the model parameters θ\theta. Our LtL_{t} term is then

where CC is independent of θ\theta and the sum is taken over the masked tokens in xt\bm{x}_{t}. For example, if we use β(t)=1/(T−t+1)\beta(t)=1/(T-t+1) from Sohl-Dickstein et al. , βt∏i=0t−1(1−βi)=1/T\beta_{t}\prod_{i=0}^{t-1}(1-\beta_{i})=1/T and 1−∏i=0t(1−βi)=(t−1)/T1-\prod_{i=0}^{t}(1-\beta_{i})=(t-1)/T, so q([xt−1]i=[x0]i∣[xt]i=m,x0)=1/tq([\bm{x}_{t-1}]_{i}=[\bm{x}_{0}]_{i}|[\bm{x}_{t}]_{i}=m,\bm{x}_{0})=1/t for non-mask tokens and we can simplify our LtL_{t} objective to

where xt\bm{x}_{t} masks tokens independently and uniformly with probability t/Tt/T. The LTL_{T} term in our ELBO is 0 for the 1/(T−t+1)1/(T-t+1) schedule, so the full objective (up to a constant) reduces to

Note that while this looks very similar to Equation 11 (with each term reweighted by 1/t1/t, the expected number of masked tokens) it is not exactly identical since masking is computed independently per-token position (instead of choosing exactly kk tokens to mask). This is an entirely practical way to do masking (and indeed some methods implement it this way).

Furthermore, since the masking probability varies linearly as 1−∏(1−βt)=t/T1-\prod(1-\beta_{t})=t/T, this is very close to uniformly sampling the number of masked tokens kk, but kk is actually drawn from a mixture of binomial distributions, i.e.

which is very close to uniform weight over terms, but slightly downweights terms near and TT. By upweighting terms near the boundary, you could in theory make this exactly uniform and thus exactly recover Equation 11. For instance, for 50 categories, absorbing-state diffusion produces the weighting shown in Figure 6.

A.4 Scaling to a large number of categories

When the number of categories KK is large, it can quickly become impractical to store all of the transition matrices Qt\bm{Q}_{t} in memory, as the memory usage grows like O(K2T)O(K^{2}T). And even if there is an algorithm to compute individual step matrices Qt\bm{Q}_{t} on demand, it may or may not be possible to do the same for the cumulative products Q‾t\overline{\bm{Q}}_{t}. We propose two approaches to scaling D3PMs to large numbers of categories that ensure cumulative products are efficient: using low-rank corruption and using matrix exponentials.

In the low-rank case, we consider structuring our transition matrices as

As an illustrative example, we describe in more detail how to efficiently represent uniform and absorbing-state transition matrices using the low-rank structure.

A.4.2 Matrix exponentials

In the matrix exponential case, we specify our transition matrices as

where R\bm{R} is a transition rate matrix and exp⁡\exp denotes the matrix exponential operation; the similar form for Qt\bm{Q}_{t} and Q‾t\overline{\bm{Q}}_{t} is a consequence of the “exponential of sums” property for commuting matrices. For efficiency, we further assume that each of the αt\alpha_{t} is an integer multiple ntα⋆n_{t}\alpha_{\star} of some common factor α⋆\alpha_{\star}, and precompute matrices exp⁡(2kα⋆R)\exp(2^{k}\alpha_{\star}\bm{R}) for 0≤k≤log⁡2(α‾T/α⋆)0\leq k\leq\log_{2}(\overline{\alpha}_{T}/\alpha_{\star}), where α‾T=∑t<Tαt\overline{\alpha}_{T}=\sum_{t<T}\alpha_{t}, taking space O(K2log⁡(α‾T/α⋆))O(K^{2}\log(\overline{\alpha}_{T}/\alpha_{\star})). Then, to compute matrix-vector products with Qt\bm{Q}_{t} or Q‾t\overline{\bm{Q}}_{t}, we can iteratively take products with a subset of these precomputed matrices based on the digits of a binary expansion of the desired multiple ntn_{t} in time O(K2log⁡(α‾T/α⋆))O(K^{2}\log(\overline{\alpha}_{T}/\alpha_{\star})).This is closely related to the well-known “exponentiation-by-squaring” technique.

As long as R\bm{R} has non-positive off-diagonal entries and sums to zero along each row, the matrix exponential produces a valid transition matrix Qt\bm{Q}_{t}; convergence to a specific stationary distribution can also be ensured by controlling the eigenvectors. In particular, if every column also sums to zero, the resulting Qt\bm{Q}_{t} will be doubly stochastic and will thus have a uniform stationary distribution.

We note that this parameterization can be viewed as a discretization of a continuous-time discrete-space Markov processes; we describe this connection in more detail in the following section.

A.5 Continuous-time Markov process transition rates

A conceptual way to understand these processes is to imagine a continuous Poisson process occurring in each state ii at rate γi(t)\bm{\gamma}_{i}(t) determining when a transition between states occurs. When a transition occurs (at time tt), a Markov transition occurs between states ii and jj with probability Πij(t)\Pi_{ij}(t). Many common stochastic processes fall into this family, including Poisson processes. Like in the case of stochastic differential equations (Song et al. ), we can derive a set of Kolomogorov equations (or Fokker-Planck equations in the continuous-state space case) that determine the marginal probability ∂qij(τ,t)\partial q_{ij}(\tau,t) of ending up in state jj at time tt having started in state ii at time ss. The general form of the Kolmogorov forward equations is

Now we can state and prove a theorem connecting continuous time Markov processes and matrix exponentials.

Let {xt}t≥0\{\bm{x}_{t}\}_{t\geq 0} be a discrete-space, continuous-time Markov process with (possibly time-dependent) transition probability matrix Π(t)\Pi(t) and transition rates γi(t)\bm{\gamma}_{i}(t). Then for a particle with an initial distribution q(xs)q(\bm{x}_{s}) at time ss, the probability of ending in state jj at time tt is

From the Kolmogorov equations for continuous-time Markov processes, we have the ODE

where Π(t)\Pi(t) is the transition probability matrix. Solving this as a first-order ODE using integrating factors yields the desired equation. ∎

In other words, the αt\alpha_{t} parameters in Equation 16 correspond to a discretization of the cumulative transition rate of a continuous-time process.

A.6 Continuous-limit of schedule from Sohl-Dickstein et al. [43]

Consider for example the schedule described by Sohl-Dickstein et al. for Bernoulli variables βt=1/(T−t+1)\beta_{t}=1/(T-t+1), i.e. the Bernoulli variable would stay the same with probability 1−βt=(T−t)/(T−t+1)1-\beta_{t}=(T-t)/(T-t+1) and transition with probability βt\beta_{t}. In this section, we show that a D3PM absorbing or D3PM uniform process with this schedule is exactly a discretization of a continuous-time jump process of the form described in Theorem 1.

We start by observing that both absorbing-state and uniform D3PM transition matrices can be expressed equivalently as matrix exponentials. In the uniform case, we have

In either case, by setting this equal to the explicit forms in Appendix A.2, we obtain the relationship

where βt\beta_{t} is defined as in Appendix A.2, and αt\alpha_{t} is the matrix exponential coefficient as used in the previous section. Using the correspondence discussed in the previous section, we also know

for the continuous-time transition rate function γ(s)\gamma(s). Defining βt=1/(T−t+1)\beta_{t}=1/(T-t+1), we have

Denoting the anti-derivative ∫γ(t)=F(t)\int\gamma(t)=F(t), we have log⁡(T−t)−log⁡(T−t+1)=−F(t)+F(t−1)\log(T-t)-\log(T-t+1)=-F(t)+F(t-1), so we can deduce F(t)=−log⁡(T−t)F(t)=-\log(T-t) (up to a constant offset). Taking a derivative then yields γ(t)=1/(T−t)\gamma(t)=1/(T-t), which has the same form as the original schedule but is now interpreted as a continuously-varying rate function instead of a probability (and is also shifted by 1 unit in time). Intuitively, we can interpret this as a schedule which assigns uniform probability of a transition occurring over the remaining time, but instead of dividing it between T−t+1T-t+1 discrete steps, we divide it across a continuous interval of size T−tT-t. We note that using larger values of TT is equivalent to performing a finer discretization on a scaled version of this continuous-time process.

A.7 Mutual-information-based noise schedule

An important part of designing the forward process for a diffusion process is to specify the noise schedule: how much noise is added at each step tt such that after TT steps the process has (approximately) reached the stationary distribution of the transition matrix. Previous work on continuous-state diffusion models has focused on controlling the variance of the continuous noise added at each step, but in a discrete state space it is less obvious how to measure or control the level of noise added.

For uniform or absorbing-state transition matrices, once a single transition occurs, all information about the original data point is lost. In this case, the schedule introduced by Sohl-Dickstein et al. is a natural choice, since it is designed to make this first transition for t/Tt/T of the elements by time tt. However, when the transition matrix imposes additional structure on the transitions, such as for our token-embedding based transition matrix, it is not sufficient to perturb t/Tt/T of the elements by time tt, since the value at time tt may be highly correlated with the value at time t−1t-1 even after a transition occurs; we thus explore using mutual information to quantify how much noise has been added. Here we describe the mutual-information-based schedules in more detail. We focus on transition matrices that are parameterized as matrix exponentials, i.e. they have the form

Inspired by the schedule introduced by Sohl-Dickstein et al. , we consider setting our αt\alpha_{t} such that tT\frac{t}{T} of the information about p(x0)p(\bm{x}_{0}) has been lost by time tt. Our goal is to find exponents such that

where HH denotes the entropy of a random variable, and p(x0)p(\bm{x}_{0}) denotes the distribution of a randomly chosen token in the data.

In practice, we estimate p(x0)p(\bm{x}_{0}) by computing empirical frequencies over the training set, and compute the value of the right-hand side of 17 for transition matrices exp⁡(αˉR)\exp(\bar{\alpha}\bm{R}) with 256 geometrically-spaced exponents αˉ\bar{\alpha} distributed in a large range (linear on a log scale between 1e-4 and 1e5). We then interpolate using a monotonic cubic spline to find the particular exponents αˉt\bar{\alpha}_{t} that ensure the above property holds approximately, and round them so that they are all multiples of a common factor α⋆\alpha_{\star} to ensure efficiency (as described in Appendix A.4). Finally, we set Qt=exp⁡((αˉt−αˉt−1)R)\bm{Q}_{t}=\exp((\bar{\alpha}_{t}-\bar{\alpha}_{t-1})\bm{R}).

It turns out that, for the specific case of absorbing-state diffusion with a [MASK] token, the mutual information schedule reduces to exactly the (T−t+1)−1(T-t+1)^{-1} schedule proposed by Sohl-Dickstein et al. . To see this, let mtm_{t} be the probability that a given value from time 0 has been replaced with [MASK] at time tt. We note then that

where we have used the fact that a mask token has zero probability under the data distribution. We also have the joint entropy

It follows that the mutual information schedule for masks is one that ensures mt=q(xt=[MASK]∣x0)=tTm_{t}=q(\bm{x}_{t}=\text{[MASK]}|\bm{x}_{0})=\frac{t}{T}. But this is exactly the (T−t+1)−1(T-t+1)^{-1} schedule. To see this, let βt\beta_{t} be the probability that a non-mask token becomes a mask token at time tt, and note that mt=1−∏s=1t(1−βs)m_{t}=1-\prod_{s=1}^{t}(1-\beta_{s}). Thus,

Interestingly, although the (T−t+1)−1(T-t+1)^{-1} schedule was designed for the case of a uniform transition matrix (an used for this purpose by Sohl-Dickstein et al. and Hoogeboom et al. ), the (T−t+1)−1(T-t+1)^{-1} schedule is NOT in general identical to the mutual information schedule in that setting. We leave further investigation of these schedules to future work.

A.8 Parameterizing the reverse process with a discretized truncated logistic distribution

Appendix B Experiments

We follow the same training and evaluation setup as used by Ho et al. . For completeness we repeat these settings here. The model architecture is based on the backbone of a PixelCNN++ architecture: a U-Net based on a Wide ResNet with weight normalization layers replaced by group normalization layers . The model has four feature map resolutions and two convolutional residual blocks for each resolution level. At the 16×1616\times 16 resolution level a self-attention block is placed between the convolutional blocks . The time step tt is included in the neural net through a Transformer sinusoidal position embedding in each residual block. Furthermore, we use the same hyperparameters and augmentation settings as in without tuning them: the dropout rate is set to 0.1; we use a learning rate of 2×10−42\times 10^{-4} with the Adam optimizer with standard settings, a batch size of 128; for evaluation we use an exponential moving average (EMA) for the model parameters with a decay factor of 0.99990.9999; and finally, we use random horizontal flips as augmentation during training.

We trained all our models for 1.5M steps on TPUv2 accelerators with a 4×44\times 4 topology. Our Inception and FID scores were computed on 50000 samples with the Inception-v3 model . We have included averages and standard deviations over models trained with 5 different seeds.

For the D3PM Gauss models with discretized Gaussian transition matrices as described in Appendix A.2.3, we use the same linear schedule for the βt\beta_{t}’s as in : βt\beta_{t} is linearly increased from 1×10−41\times 10^{-4} to 0.020.02. We did not explore any other noise schedules for D3PM Gauss models. For the D3PM uniform model (see Section A.2.1) we experimented with a linear schedule for βt\beta_{t} (linearly increasing from 0.020.02 to 11) and the cosine schedule as suggested by Hoogeboom et al. . Table 4 shows that the D3PM uniform model with a cosine schedule produces much better results than the same model with a linear βt\beta_{t} schedule. For the D3PM absorbing model (see Section A.2.2) the absorbing state is the gray pixel, corresponding to the RGB values (128, 128, 128). For these models we used a schedule that corresponds to increasing the probability of being in the absorbing state linearly over time: βt=(T−t+1)−1\beta_{t}=(T-t+1)^{-1}. This schedule was also proposed in Sohl-Dickstein et al. for diffusion with binary random variables, which has a uniform stationary distribution as opposed to the stationary distribution with all the mass on the absorbing state.

B.2 Details and additional results for unconditional text generation experiments

Our experiments using text8 and LM1B were performed with a standard transformer encoder following the T5 architecture with 12 layers and 70 million parameters (12 heads, mlp dim 3072, qkv dim 768). All models were trained for 1 million steps with batch size 512 on the TPUv2 or TPUv3 platform. Our code is implemented in JAX and Flax . For our experiments, we used learning rate 5×10−45\times 10^{-4} with a 10000 step learning rate warmup and inverse sqrt decay. For text8, we used a standard 90000000/5000000/500000 train-test-validation split with sequences of length 256. For LM1B, we used the standard test-train split from TFDS with 30,301,028 examples in the training set and 306,688 in the test set. For text8, no preprocessing is performed, and training is performed on random crops of the entire concatenated, lower-cased training set. For LM1B, training is performed on sequences of length 128 sampled by packing sequences from the training corpus, including an EOS token. Perplexities are reported relative to the actual number of English-language words in the test set (including an EOS token predicted by the model).

Our autoregressive transformer baseline was a standard transformer decoder with the same basic architecture (but including causal masking, as is standard for autoregressive models) with the same number of parameters.

Table 5 contains additional comparisons of hybrid losses. We found that the hybrid loss Lλ=0.01L_{\lambda=0.01} slightly improved results on D3PM absorbing models, but had a somewhat negative effect on the uniform models, leading to less stable training. All models were trained on 1000 step diffusion processes, but we found very little improvement between 1000 and 256 steps when evaluating a trained model by skipping steps. For all figures, steps were skipped evenly (except possibly for the last step if the number of evaluation steps did not divide 10001000). We found both the cosine and mutual information schedules worked well for uniform diffusion. We used the cosine variant introduced by Hoogeboom et al. , i.e.

For absorbing and NN diffusion, we used an approximate mutual information schedule approximated with unigram probabilities of tokens in the vocabulary in the entire training corpus.

Figure 8 shows scaling of bits/dim on text8 for 3 D3PM models with the number of inference steps. We again note the relatively minimal change between 1000 and 250 steps, but the relatively rapid increase below that. Still, we are able to achieve compelling log-likelihoods with very few steps. Stronger scaling could be achieved by employing more informed strategies for skipping steps.

B.2.2 Additional tables and figures for LM1B

B.3 Additional uncurated generation examples from various models