This paper post-trains Masked diffusion language model (dLM)s to match reference text at the feature level by minimizing MMD (Maximum Mean Discrepancy) between hidden states of a frozen pretrained diffusion model, letting a generator produce usable text in far fewer sampling steps without a teacher trajectory or a jointly trained critic.
Diffusion language models generate text by repeatedly denoising a noisy sequence. The problem: standard training optimizes per-token cross-entropy or per-latent squared error, which fits local marginals but does not guarantee good samples when you only run a handful of denoising steps. In practice, you either pay for many sampling steps or accept worse text.
Existing speed-ups fall into two camps. Distillation methods compress many teacher steps into fewer student steps, so you need a well-tuned teacher and paired trajectories. Distribution-matching methods train an auxiliary network alongside the generator to supply gradients, which adds its own moving target. Both work, but both require extra machinery. The authors ask whether you can get a useful distributional signal directly from samples, with no teacher trajectory and no co-trained auxiliary model.
The core idea: judge generated text by comparing its internal representations to those of real text, inside a frozen pretrained diffusion model. If two batches of sequences produce similar hidden-state distributions under the same extractor, they are probably similar as distributions over text.
The comparison tool is MMD (Maximum Mean Discrepancy) with a Gaussian Gaussian RBF kernel. Given a batch of real sequences and a batch of generated sequences, MMD computes three averages of kernel similarities: real-to-real, generated-to-generated, and real-to-generated. Minimizing the combination pulls generated features toward real ones (the cross term) while pushing generated features apart from each other (the within-generated term), which prevents collapse onto a few favored outputs.
A key trick: instead of pooling each sequence into one embedding, they keep the hidden state at every scored token position. One forward pass through the frozen extractor yields many feature samples per sequence, which matters when you only have one reference per prompt. To keep the MMD estimator unbiased despite within-sequence correlation, same-sequence token pairs are excluded from the within-distribution terms.
•
Discrete models (masked and hybrid masked-uniform diffusion). Sampling tokens is non-differentiable, so they use REINFORCE with a Leave-one-out baseline: draw G independent batches of generated sequences for the same context, use each batch’s negative MMD as its reward, and let the other batches serve as the baseline. They instantiate this as MDLM-MMD on top of MDLM and as DMax-MMD on top of DMax.
•
Continuous models (Embedded Language Flows (ELF) latents). Latents are differentiable, so MMD gradients flow straight through the generator and the frozen extractor. The generator is trained to map Gaussian noise to clean latents in one step. For multi-step inference, they use Self-conditioning: feed the previous prediction back in, refine it a few times, then decode. An optional distillation stage (Iterative Refinement Distillation (IRD)) compresses that refinement loop back into one step.
# ELF-MMD training step (continuous case)
x_real = encode(ref_tokens)
z = randn_like(x_real)
sc = zeros_like(x_real)
for _ in range(randrange(n)): # bootstrap self-conditioning
sc = stopgrad(net(z, t=0, sc=sc)) # no gradients through warmup
x_gen = net(z, t=0, sc=sc) # differentiable final pass
f_real = stopgrad(extractor(x_real, t=1, sc=x_real))
f_gen = extractor(x_gen, t=1, sc=x_gen)
loss = mmd_rbf(f_gen, f_real) # token-level RBF MMD
The method is evaluated on OpenWebText (unconditional), GSM8K via training on TinyGSM (math reasoning), and 16B DMax-Math and DMax-Coder checkpoints on math and code benchmarks.
•
Unconditional text, discrete. At matched unigram entropy near the OpenWebText reference, MDLM-MMD gets roughly 17–21% lower generative perplexity than the IDLM baseline across 8, 16, and 32 sampling steps. Generative perplexity here is measured by scoring samples under a pretrained GPT-2 Large.
•
Unconditional text, continuous. ELF-MMD improves perplexity-entropy trade-offs over the ELF baselines; adding IRD cuts generative perplexity by ~30 points (T5 latents) and ~10 points (GPT-2 latents) at 4 sampling steps.
•
GSM8K. The continuous ELF-MMD+IRD lifts accuracy at 4 steps from 14.2% to 20.8% and at 8 steps from 27.5% to 32.5%, versus 15.9% and 23.5% for the progressive-distillation baseline. Peak accuracy at 64 steps is 36.3%.
•
16B scale. With about 1.7 to 2.5 GPU-hours of post-training on 8 H100s, DMax-Math-MMD raises tokens per forward pass by 10.3 to 16.5% across four math benchmarks at similar or higher accuracy. DMax-Coder-MMD gains +2.4 points on HumanEval-Instruct and +3.8 points on MBPP-Instruct, with higher throughput too.
•
Ablations. Removing the generated-generated repulsive term causes continuous models to collapse quickly and hurts discrete models. Token-level RBF beats sequence-level RBF, a linear kernel (mean matching), and per-position feature regression. Features from a pretrained DLM beat raw token embeddings or features from a separately pretrained autoregressive model.
A caveat on causality: the paper shows token-level RBF outperforms the alternatives it compares, but the choice of extractor layer, kernel bandwidth, and representation space all matter and are not jointly optimized.
•
If you are working on a diffusion language model and want fewer sampling steps without building a distillation pipeline or co-training a critic, this gives a lightweight post-training recipe. The 16B result suggests a few GPU-hours can shift the accuracy-vs-throughput curve of an existing checkpoint. Worth testing on your own checkpoint if a frozen copy of the same model is a plausible extractor.
•
If you are picking between distribution-matching objectives, the ablation is informative: keep the repulsive within-generated term, keep token-level features rather than pooled ones, and prefer an RBF kernel over mean matching. These are concrete defaults to try before investing in heavier methods.
•
The method does not apply directly to autoregressive LLMs. It is designed for diffusion generators where you can either sample multiple completions per context (for the REINFORCE variant) or differentiate through generated latents (for the continuous variant).
•
Code is released at yandex-research/dlm-mmd.
•
Performance depends on the representation space, the RBF bandwidth, and which layer of the extractor you tap. The authors flag representation design as open work.
•
The theoretical identifiability argument (characteristic kernel plus injective features implies unique distribution matching) is cited for token-level features but not proven to uniquely identify the full sequence distribution. Matching feature marginals is a useful signal, not a guarantee of distributional equality.
•
Gains on benchmarks like GSM8K and HumanEval do not establish general reliability. The 16B results come from existing DMax checkpoints and the released DMax training data; behavior on other base models or domains is untested here.
•
REINFORCE needs at least B=2 generated sequences per context and benefits from G=4 batches for the leave-one-out baseline, which raises per-step cost even if total training time stays short.