Simplex diffusion models improve discrete data generation accuracy
Simplex Diffusion Models
Machine Learning
Summary
Generating discrete data like text or code is tricky because early steps often lose important uncertainty information. The authors propose Simplex Diffusion Models that keep track of uncertainty during the generation process using a new mathematical framework. This method improves the quality and reliability of generated data, performing better than previous techniques on tasks like text modeling and code generation. It also allows faster and more effective sampling with fewer steps.
What this means in practice
- •For natural language processing teams: Generate higher quality text with improved uncertainty modeling and faster sampling.
- •For software engineering tool developers: Create more accurate code generation tools by using diffusion models that better preserve information across steps.
Authors
Justin Deschenaux, Alexandre Galashov, Andrew Campbell, Li Kevin Wenliang, James Thornton, Arnaud Doucet, Valentin De Bortoli
Abstract
Diffusion models have revolutionized generative modeling for continuous data through the gradual refinement of a belief state. This iterative refinement has not yet carried over to discrete diffusion models, which discard uncertainty at intermediate steps through categorical sampling (information collapse). We propose Simplex Diffusion Models (SDMs), a framework that lifts the diffusion process to the probability simplex to represent beliefs over categories. SDMs admit probability paths with closed-form reverse transitions and can be trained with a simple cross-entropy loss. Contrary to earlier proposals such as Dirichlet Flow Matching which requires integrating an ordinary differential equation, we introduce a DDIM-like sampler with a tunable level of stochasticity. Because SDMs operate on samples on the simplex, they can carry uncertainty across denoising steps, which mitigates information collapse. On OpenWebText, SDMs are competitive with strong Discrete Diffusion baselines, achieving $17.0$ GenPPL at $5.46$ unigram entropy in 64 sampling steps, close to real validation data. Even without Self-Conditioning (SC), SDMs outperform masked and uniform diffusion (with SC or predictor-corrector sampling) on code generation (TinyGSM, $T=0.1$; $49.0\%$ vs. $45.8\%$). Distilled down to 8 steps, SDMs solve $32.1\%$ of GSM8K problems, more than distilled Discrete Diffusion models with 128 steps ($21.4\%$).