Research
Representation-Space MMD for Diffusion Language Models
Overview Research area: Generative modeling for language — specifically post-training / fine-tuning methods for diffusion language models (DLMs), sitting at the intersection of diffusion distillation,

- arXiv
- 2610.06648
- Published
- 2026-10-05
- Authors
- Ilya Drobyshevskiy, Ilia Sudakov, Maksim Semenov, Denis Kuznedelev, Maksim Ignatov, Pavel Temirchev, Nikita Balagansky, Viacheslav Meshchaninov, Nikita Gushchin, Dmitry Baranchuk
AI summary
Overview
Research area: Generative modeling for language — specifically post-training / fine-tuning methods for diffusion language models (DLMs), sitting at the intersection of diffusion distillation, distribution matching, and representation learning.
Technical level: Advanced. The paper assumes familiarity with diffusion language modeling, flow matching, kernel methods (MMD, RBF kernels), policy-gradient / REINFORCE optimization, and self-conditioning.
Scope: The paper proposes a post-training objective that minimizes Maximum Mean Discrepancy between generated and reference text in the feature space of a frozen pretrained DLM, and validates it on discrete (MDLM, DMax) and continuous (ELF) diffusion language models from small scale up to 16B parameters.
What This Paper Is About
Diffusion language models generate text by iteratively denoising, but their standard training objectives (token-wise cross-entropy for discrete models, squared-error latent prediction for continuous models) do not necessarily produce high-quality samples in few steps. Existing acceleration methods either need teacher-derived targets or require jointly trained auxiliary models. This paper asks whether a useful distributional training signal can be computed directly from generated and reference samples, and answers by matching distributions in the feature space of a frozen pretrained DLM using Maximum Mean Discrepancy.
Key Contributions
-
Representation-space MMD objective for DLMs. The authors use a frozen pretrained DLM to extract contextual hidden states at individual token positions,
φ(x) = (φ₁(x), …, φ_L(x)) ∈ ℝ^{L×D}, obtaining multiple feature observations per sequence from a single extractor pass. They match the distributions of these features between generated and reference sequences with a Gaussian RBF kernel. -
A discrete realization optimized with policy gradients. Sequences are sampled from the token distributions produced by a single denoiser forward pass, and the MMD reward is optimized with REINFORCE using a leave-one-out baseline across
Ggroups. This is instantiated as MDLM-MMD (masked-token prediction) and DMax-MMD (hybrid masked–uniform diffusion), where MMD trains both masked prediction and refinement of the model's own token predictions. -
A continuous realization with direct differentiation. ELF-MMD is initialized from a pretrained Embedded Language Flows (ELF) model and trained to map Gaussian noise to clean latent sequences in a single step, with MMD gradients propagating through the frozen ELF feature extractor and the generated latents. Self-conditioning extends the one-step generator to multi-step sampling via iterative refinement, and iterative refinement distillation (IRD) is optionally applied afterward.
-
Evaluation across scales, including 16B. Experiments cover OpenWebText and TinyGSM, plus scaling to 16B DMax-Math and DMax-Coder checkpoints, where MMD post-training improves the accuracy–efficiency trade-off on math and code benchmarks.
Main Findings
-
Lower generative perplexity at matched entropy on OpenWebText. MDLM-MMD (initialized from pretrained MDLM, MMD rewards from frozen MDLM features, policy-gradient group size
G = 4) achieves lower generative perplexity than DiDi-Instruct, IDLM, and IDLM-REINFORCE at matched entropy across 8, 16, and 32 sampling steps. Interpolating the curves at the reference-data entropy gives approximately 17–21% lower generative perplexity than IDLM across the three sampling budgets. Evaluation uses 1,000 generated samples with sequences packed toL = 1024, scored with generative perplexity under pretrained GPT-2 Large and average unigram entropy, with data entropyH ≈ 5.43as reference. -
Continuous models improve at most sampling budgets. ELF-MMD (ELF-B architecture with either T5-small or GPT-2 Large final-layer hidden representations) improves over ELF and ELF-PD across most budgets. At 8 steps it achieves lower generative perplexity and entropy closer to the reference value than ELF-PD; at 32 steps it outperforms ELF on both metrics. Adding IRD at 4 steps reduces generative perplexity relative to ELF-MMD by approximately 30 points (T5 encoder) and approximately 10 points (GPT-2 encoder), while bringing entropy closer to the reference value for both encoders.
-
Better accuracy–computation trade-offs on GSM8K. Training on TinyGSM with sequences packed to
L = 512and evaluating final-answer accuracy on the GSM8K test set: MDLM-MMD achieves higher accuracy than MDLM, IDLM, IDLM-REINFORCE, and DiDi-Instruct at moderate and high decoding budgets, reaching approximately 54% accuracy at approximately 49 steps. -
Large accuracy gains at few steps for continuous conditional generation. With ELF-B using GPT-2 Small final-layer hidden representations and shifted timestep schedules (shift 32 for ELF, 128 for ELF-PD), accuracy at 4 and 8 steps rises from 14.2% to 20.8% and from 27.5% to 32.5%, compared with 15.9% and 23.5% for ELF-PD. At 64 steps, ELF-MMD+IRD reaches its highest accuracy of 36.3%, compared with 35.2% for ELF-MMD and 31.6% for ELF.
-
Pass@k improvements at small sample budgets. For MDLM-based models, MDLM-MMD outperforms baselines at 4–8 sampling steps for
k ≤ 8, and at 16–32 steps it achieves higher pass@k for most evaluated values ofk, with smaller differences at larger sample budgets. ELF-MMD also improves pass@k at smaller sample budgets, with gaps narrowing askincreases. -
Improved decoding parallelism at 16B scale. On 16B DMax-Math-MMD (threshold 0.85) and DMax-Coder-MMD (threshold 0.9), trained for 400 steps taking approximately 13 minutes (DMax-Math) and approximately 19 minutes (DMax-Coder) on eight NVIDIA H100 GPUs — approximately 1.7 and approximately 2.5 GPU-hours per run. DMax-Math-MMD increases tokens generated per forward (TPF) by 10.3–16.5% over the reported DMax-Math operating points while achieving similar or higher accuracy across GSM8K, MATH500, Minerva-Algebra, and ASDIV. On code, DMax-Coder-MMD improves accuracy by 2.4 percentage points on HumanEval-Instruct and 3.8 points on MBPP-Instruct while increasing TPF on both benchmarks. Reported DMax-Math-MMD results are GSM8K 92.1 accuracy / 6.15 TPF, MATH500 76.0 / 6.84, Minerva-Algebra 92.1 / 8.19, ASDIV 92.9 / 6.20; DMax-Coder-MMD gives HumanEval-Instruct 85.9 / 8.07 and MBPP-Instruct 83.0 / 6.10.
-
Token-level RBF is the strongest loss formulation. In the ablation on GSM8K, token-level RBF gives the strongest accuracy–computation trade-off, with larger gains for discrete models and more modest improvements for continuous models. Attraction-only and feature-regression losses yield lower accuracy, particularly for continuous models. A linear-kernel MMD baseline (MSE between generated and reference feature means) and sequence-level RBF on mean-pooled embeddings are also compared. With
B = 1, the attraction-only objective leads to worse performance on discrete models and rapid collapse in continuous models. OpenWebText ablations similarly favor token-level RBF in most setups.
Methodology in Plain English
The core idea is to compare the distributions of generated and real text not by the text itself, but by the internal features a pretrained diffusion language model computes for it. The authors freeze such a model and use its hidden states as a scoring space.
Instead of pooling a whole sequence into one vector, they keep the representations at individual token positions. This yields many comparison points per sequence from a single forward pass — especially useful when only one reference example is available per conditioning context.
The comparison metric is Maximum Mean Discrepancy with a Gaussian RBF kernel. Squared MMD has three terms: how similar real features are to each other, how similar generated features are to each other, and how similar generated features are to real ones. The cross term pulls generated features toward the reference; the generated–generated term pushes generated samples apart to avoid collapse. Within-sequence comparisons are excluded from the within-distribution terms to preserve unbiasedness, which means at least two generated sequences (B ≥ 2) are needed. For conditional settings the real–real term is dropped because it does not depend on the model parameters, so training can use a single reference per condition.
For discrete models, samples are drawn from the factorized token distributions of one denoiser pass, so sampling is not differentiable. The authors therefore use REINFORCE: the negative MMD estimate becomes a reward, with a leave-one-out baseline computed across G independent batches to reduce gradient variance. The resulting surrogate multiplies each batch's advantage by the summed log-probabilities of the sequences in that batch. Masked models (MDLM-MMD) score only masked positions; the hybrid DMax variant applies the objective to both masked prediction and refinement of the model's own predictions. After MMD post-training, each discrete model's original sampling procedure is used unchanged.
For continuous models, the generator is differentiable, so MMD gradients flow directly through the frozen extractor and the generated latents. ELF-MMD is trained as a one-step latent generator from Gaussian noise, with features extracted at t = 1 with self-conditioning set to the clean latents. Bootstrapping addresses the training–inference mismatch: the number of preliminary refinement steps is sampled from {0, …, n−1}, run without gradients at fixed noise, and MMD is applied only to the final differentiable prediction. At inference, self-conditioning feeds each prediction back to refine the output, giving K + 1 total network evaluations. Optionally, iterative refinement distillation freezes the trained generator as a teacher and distills a K-step trajectory into a single student prediction.
Why This Matters
Impact on research. The paper shows that a post-training signal for diffusion language models can be derived purely from comparing generated and reference samples in a frozen model's representation space — without teacher-forced denoising targets, without a jointly trained discriminator or auxiliary denoiser, and without full sampling trajectories. This simplifies the training pipeline relative to distillation and auxiliary-model approaches (IDLM, D-MMD, DiDi-Instruct), and connects language generation to representation-space distribution matching ideas already used in visual generation. It also demonstrates the approach scales to a 16B hybrid diffusion model.
Real-world applications (plausible given the results):
- Faster on-device or edge text generation, where few-step diffusion sampling is needed and the whole model must be quantized or small.
- Code assistants that generate at higher decoding parallelism (the paper measures tokens generated per forward on HumanEval-Instruct and MBPP-Instruct).
- Mathematical reasoning assistants, where the paper evaluates GSM8K, MATH500, Minerva-Algebra, and ASDIV.
- Model providers with existing pretrained diffusion checkpoints who want a cheaper post-training step — the 16B runs took approximately 1.7 and approximately 2.5 GPU-hours on eight NVIDIA H100 GPUs.
Industry relevance. The reported training cost is small relative to pretraining, and the method applies to released checkpoints (the 16B experiments start from public DMax-Math and DMax-Coder weights). That makes it attractive as a low-cost efficiency pass over existing deployment models. The paper reports that code is available at yandex-research/dlm-mmd.
Future Directions
- Representation design. The authors note that performance depends on choices such as the RBF bandwidth and the representation space used for distribution matching, and identify representation design as a natural direction. They used features from a single layer on clean inputs, and suggest combining features across layers or noise levels, or using an ensemble of extractors.
- Combining MMD with other objectives. The paper points to combining MMD with other training objectives for discrete DLMs, citing the gains from applying IRD after ELF-MMD post-training in the continuous setting as evidence that combinations help.
- Alternative multi-step sampling schemes. The authors note that MMD is also applicable to continuous DLM variants without self-conditioning, which could instead use alternative multi-step sampling schemes such as consistency sampling.
- Theoretical grounding. The paper observes that characteristic-kernel results and recent injectivity results for causal transformers motivate the choice of DLM token features, but explicitly state that these do not establish that matching token-feature distributions uniquely identifies the sequence distribution. Tightening this link remains open.
Target Audience
Researchers and practitioners working on diffusion language models, diffusion distillation, and few-step text generation; machine learning engineers who want to speed up sampling of existing DLM checkpoints without retraining from scratch; and researchers interested in representation-space distribution matching objectives, kernel methods for generative modeling, or policy-gradient fine-tuning of discrete generative models. Readers need a strong background in diffusion models and stochastic optimization to follow the derivations.
Authors’ abstract
We introduce a post-training method for diffusion language models (DLMs) that minimizes Maximum Mean Discrepancy (MMD) between generated and reference distributions in the feature space of a frozen pretrained DLM. To estimate MMD, we retain contextual features at individual token positions, obtaining multiple observations per sequence from a single extractor pass. We optimize this objective using policy gradients for discrete models and direct differentiation through generated latents for continuous models. In both cases, computing the loss directly from these features enables efficient post-training without full sampling trajectories or jointly trained auxiliary models. Experiments show lower generative perplexity at comparable entropy on OpenWebText and better accuracy-computation trade-offs on GSM8K. On 16B DMax-LLaDA2.0 models with hybrid masked-uniform diffusion, we increase decoding parallelism with similar or higher accuracy on math and code benchmarks.