Simplex Relaxation for Discrete Diffusion
Nanyang Technological University Singapore · The University of Tokyo · Purdue University · Institute for Advanced Intelligence and Computing (IAIC), A*STAR · Centre for Frontier AI Research (CFAR), A*STAR
cs.CL
Submitted: 2026-08-11
Updated: 2026-09-30
Code: https://github.com/alicommit-malp/sudoku
License: http://arxiv.org/licenses/nonexclusive-distrib/1.0/
Importance score: 95/100
The gist: The paper introduces Simplax, "an exact Dirichlet–categorical augmentation of uniform discrete diffusion" that "preserves the original uniform diffusion process as its categorical marginal" while
Terminology
Summary
The paper introduces Simplax, an exact Dirichlet–categorical augmentation of uniform discrete diffusion
that preserves the original uniform diffusion process as its categorical marginal
while introducing an auxiliary simplex-valued variable
to enrich the training objective and reverse transitions. The key idea is that "the simplex variable is not introduced as a replacement for the categorical state, but as an auxiliary random variable that is probabilistically coupled to it and can be used to construct the training objective and reverse transition."
The method augments the standard uniform discrete diffusion process by introducing, for each time t, a simplex-valued variable wt ∈ ΔK−1 through the conditional:
q(wt zt, x) = Dir(wt; ηt pt + zt),
where "ηt > 0 is a concentration parameter and
the additive one-hot count zt anchors the relaxed state to the sampled discrete token."
The construction yields an exact Dirichlet–categorical hierarchy with four key properties (Proposition 1):
-
Marginal: q(wt x) = Dir(wt; ηt pt)
-
Exact decoder: q(zt wt, x) = q(zt wt) = Cat(zt; wt)
-
Reverse categorical posterior: q(zs wt, x) = Cat(zs; ρst(x, wt)) with ρst(x, wt):= ps ⊙ [αts wt ⊘ pt + (1 − αts)⟨wt, π ⊘ pt⟩1]
-
Reverse Dirichlet mixture: q(ws wt, x) = Σk ρst,k(x, wt) Dir(ws; ηs ps + ek)
The paper derives a tractable Rao–Blackwellized reverse-bridge objective. The direct simplex bridge KL divergence is generally intractable because q(ws wt, x) is a Dirichlet mixture.
Instead, the authors define:
L̄zszt,wt(wt, zt, x; s, t):= Eq(ẽztwt)[KL(q(zs ẽzt, x) ∥ q(zs ẽzt, x̂θ))],
where ẽzt denotes a second categorical variable satisfying ẽzt ∼ q(ẽzt wt) = Cat(ẽzt; wt), conditionally independently of the network input zt given wt.
Proposition 2 gives the exact closed form:
L̄zszt,wt(wt, x̂θ, x; s, t) = ⟨wt, log p̂t − log pt⟩ + ⟨ρst(x, wt), log ps − log p̂s⟩.
This is fully tractable and eliminates sampling noise associated with the auxiliary decoder sample ẽzt.
Proposition 3 establishes that the objective has a non-degenerate infinitesimal limit:
L̄zszt,wt(wt, x̂θ, x; t − ∆, t) = ∆ lct(wt, x̂θ, x, t) + o(∆),
where lct(wt, x̂θ, x, t) = λ(t)[⟨wt, π ⊘ p̂t⟩ − ⟨wt, π ⊘ pt⟩⟨pt, log p̂t⟩ + ⟨π ⊙ (wt ⊘ pt), log p̂t⟩].
The paper notes this can be understood as a simplex-relaxed continuous-time analogue of the UDLM objective.
The default sampler is the stochastic ancestral sampler implied by the Dirichlet–categorical hierarchy.
Generation starts from:
wtN ∼ Dir(ηtN π), ztN ∼ Cat(wtN).
Each reverse step draws:
zs ∼ Cat(ρst(x̂θ, wt)), ws ∼ Dir(ηs p̂s + zs).
The paper emphasizes that the categorical input also has a computational advantage
since zt is stored as an integer token index and its embedding is obtained by lookup,
whereas Feeding the dense simplex vector wt instead requires computing wtT E for the vocabulary embedding matrix E at every sequence position.
The paper evaluates on OpenWebText with GPT-2 BPE tokenizer (V = 50,257), sequence length 1,024, and a 179M-parameter diffusion transformer. Key findings:
Design diagnostics:
-
Self-conditioning:
the preferred setting depends on the inference budget: omitting self-conditioning is better near the data-entropy operating point at NFE = 16, whereas using it is better at NFE = 128
-
Denoiser input:
the zt-input model attains lower Gen. PPL at comparable Gen. ENT at both NFE values
compared to wt-input -
UDLM initialization:
UDLM initialization improves the Gen. PPL–Gen. ENT frontier at both NFE values
Main results: Simplax has the lowest Gen. PPL under all three evaluators at NFE = 16 and 1,024. At NFE = 128, it is best under GPT-2 Large and GPT-2 XL, while LangFlow is best under Llama-2 7B.
At NFE = 16: Simplax achieves Gen. PPL of 90.5 (GPT-2 L), 93.1 (GPT-2 XL), 49.3 (Llama-2) with entropy 5.45.
At NFE = 128: Simplax achieves Gen. PPL of 56.9 (GPT-2 L), 58.9 (GPT-2 XL), 31.4 (Llama-2) with entropy 5.45.
At NFE = 1,024: Simplax achieves Gen. PPL of 45.1 (GPT-2 L), 46.8 (GPT-2 XL), 25.5 (Llama-2) with entropy 5.44.
All models are trained exclusively on puzzles with 30 clues
and evaluated across clue densities from 40 down to 17 clues, plus unconditional generation. The paper states: Simplax achieves the highest performance among the compared methods across all conditional and unconditional settings in Table 2.
Key results:
-
40 clues: Simplax 98.55% (best; next best Duo 97.00%)
-
35 clues: Simplax 91.05% (best; next best MDLM 85.15%)
-
30 clues: Simplax 61.75% (best; next best Duo 48.85%)
-
25 clues: Simplax 25.90% (best; next best Duo 16.00%)
-
20 clues: Simplax 8.80% (best; next best Duo 4.80%)
-
17 clues: Simplax 1.20% (best; next best FLM 0.45%)
-
Unconditional validity: Simplax 95.85% (best; next best Duo 80.95%)
The paper highlights: Its advantage extends beyond the 30-clue training distribution to both more and less conditioned inputs, including the challenging low-clue regimes.
The paper distinguishes Simplax from:
-
Standard discrete diffusion (D3PM, etc.):
we keep the original categorical forward process and reverse posterior, and do not replace the primary generative state
-
Auxiliary-variable methods (Di4C, VADD, CoDD, Duo, FLM, CADD, CANDI): "Unlike methods that use the auxiliary variable as the denoiser input or primary generative state, the zt-input Simplax formulation retains the categorical state as the network input and uses the simplex variable to define the reverse-bridge objective and sampler"
-
Simplex diffusion methods (DDSM, Dirichlet Flow Matching, etc.): "these methods are close to ours in geometry... but differ in role: in our formulation, the simplex variable is not the primary generative state, but an exact auxiliary bridge attached to a standard discrete diffusion process"
The paper acknowledges: "The present formulation is specialized to uniform categorical corruption and introduces an auxiliary simplex-valued state whose computational overhead relative to standard discrete diffusion has not been fully characterized. Moreover, the concentration schedule remains an additional design choice rather than being determined by the theory."
Improvements for AI systems
Based on the paper, here are the specific improvements I can make to AI systems:
- Exact Dirichlet–Categorical Augmentation for Discrete Diffusion Models
-
I can implement the auxiliary simplex variable
w tcoupled to the categorical statez tviaq(w t z t, x) = Dir(w t; η t p t + z t), preserving the original uniform diffusion as a marginal while enriching the training signal. -
This yields a tractable Rao–Blackwellized objective (Proposition 2) with closed-form loss
⟨w t, log p̂ t − log p t⟩ + ⟨ρ st(x, w t), log p s − log p̂ s⟩, eliminating sampling noise from auxiliary decoder samples.
- Improved Reverse Transition Accuracy via Exact Posterior Mixtures
-
I can use the exact reverse categorical posterior
q(z s w t, x) = Cat(z s; ρ st(x, w t))withρ st = p s ⊙ [α ts w t ⊘ p t + (1 − α ts)⟨w t, π ⊘ p t⟩1], improving token-level reconstruction over standard discrete diffusion samplers. -
The reverse Dirichlet mixture
q(w s w t, x) = Σ k ρ st,k Dir(w s; η s p s + e k)provides a more faithful generative path, especially under low inference budgets (NFE = 16).
- Computational Efficiency via Categorical Network Input
-
I can keep the integer token index
z tas the network input (stored as integer, embedding via lookup) rather than dense simplex vectors, reducing memory and compute per step—critical for large vocabularies (e.g., 50K tokens) and long sequences (1,024). -
This enables faster training and sampling without sacrificing quality, as demonstrated by lower Gen. PPL at both NFE = 16 and NFE = 128 compared to simplex-input variants.
- Robustness Across Inference Budgets via Adaptive Self-Conditioning
-
I can dynamically choose whether to use self-conditioning based on the inference budget: omit it near data-entropy operating points (NFE = 16) for better perplexity, and enable it at higher budgets (NFE = 128) for improved generation quality.
-
This adaptive strategy optimizes the Gen. PPL–Gen. ENT frontier, as shown in the paper’s design diagnostics.
- Superior Constrained Generation for Structured Outputs
-
I can apply this method to constrained categorical generation tasks (e.g., Sudoku with variable clue densities) achieving 98.55% validity at 40 clues and 95.85% unconditional validity—outperforming all baselines (Duo, MDLM, FLM) by 1.5–15 percentage points.
-
The method generalizes beyond training distribution (30 clues) to both more (40) and less (17) conditioned inputs, indicating robustness to distribution shift.
- Continuous-Time Training Objective for Stable Optimization
-
I can use the non-degenerate infinitesimal limit (Proposition 3)
l ct(w t, x̂ θ, x, t) = λ(t)[⟨w t, π ⊘ p̂ t⟩ − ⟨w t, π ⊘ p t⟩⟨p t, log p̂ t⟩ + ⟨π ⊙ (w t ⊘ p t), log p̂ t⟩]as a continuous-time loss, enabling stable gradient flow and better convergence than discrete-step objectives. -
This also provides a principled way to schedule noise levels without ad-hoc tuning.
-
Generate higher-quality text with lower perplexity (e.g., 45.1 Gen. PPL under GPT-2 Large at NFE=1,024) and better entropy calibration than existing discrete diffusion models (MDLM, UDLM, Duo, FLM).
-
Handle constrained generation tasks (e.g., puzzle solving, structured data) with high validity even under sparse conditioning (e.g., 1.2% validity at 17 clues, still best-in-class).
-
Scale efficiently to large vocabularies and long sequences due to integer-token inputs, making it practical for production text generation.
-
Adapt inference compute flexibly—from fast generation (NFE=16) to high-quality generation (NFE=1,024)—without retraining, by tuning self-conditioning and step count.
-
Maintain exactness of the underlying discrete diffusion process while leveraging continuous relaxations, avoiding approximation errors common in simplex-only or auxiliary-variable-only methods.
Sources
Related papers
- Exploring Solution Divergence and Its Effect on Large Language Model Problem Solving
- Ishigaki-IDS-Bench: A Benchmark for Generating Information Delivery Specification from BIM Information Requirements
- Subliminal Steering: Stronger Encoding of Hidden Signals
- MedStruct-S: A Benchmark for Key Discovery, Key-Conditioned QA and Semi-Structured Extraction from OCR Clinical Reports
- The End of Transformers? On Challenging Attention and the Rise of Sub-Quadratic Architectures
- Untangling the Mechanisms of Misleading Context in Medical Question Answering