Skip to content
AI.info

Research

Learn from Your Mistakes: Self-Correcting Masked Diffusion Models

Overview Research area: Discrete generative modeling, specifically masked diffusion models (MDMs) for text, code, math, and molecular string generation. Technical level: Advanced. The paper builds on

Learn from Your Mistakes: Self-Correcting Masked Diffusion Models
arXiv
2602.11590
Published
2026-02-12
Authors
Yair Schiff, Omer Belhasin, Roy Uziel, Guanghan Wang, Marianne Arriola, Gilad Turok, Ran Zilberstein, Michael Elad, Volodymyr Kuleshov

AI summary

Overview

Research area: Discrete generative modeling, specifically masked diffusion models (MDMs) for text, code, math, and molecular string generation.

Technical level: Advanced. The paper builds on continuous-time variational objectives for discrete diffusion and introduces an augmented training loss with a predictor-corrector derivation, though the resulting algorithm changes are small.

Scope: The paper proposes ProSeCo (Progressive Self-Correction), a training and sampling framework that teaches a single masked diffusion model to both unmask tokens and revise tokens it has already decoded, and validates it on fine-tuning of an 8B MDM, guided molecule design, and unconditional text generation.

What This Paper Is About

Masked diffusion models generate text by unmasking many tokens in parallel, but once a token is unmasked it can never change, so early mistakes accumulate and degrade the final sample. This paper trains the same model to also act as a corrector that can revise already-decoded positions, and then interleaves correction steps with unmasking steps during sampling. The goal is to improve both the quality of generated outputs and the speed-quality trade-off relative to standard MDMs.

Key Contributions

  1. A joint decoding-and-correction training framework. The model is trained with an added cross-entropy self-correction loss (the authors' "CMDM"/"SCMDM" objective) in which the model's own argmax-decoded outputs are treated as corrupted inputs that must be mapped back to the clean data, alongside the standard MDM denoising loss.

  2. Simple training and sampling algorithms. Training requires one additional forward pass and one extra loss term on top of standard MDM training; sampling interleaves corrector loops with unmasking steps, controlled by two user-facing hyperparameters: corrector frequency (omega) and number of corrector steps per loop (S).

  3. Weight tying and a principled weighting scheme. Corrector and denoiser weights are tied (phi = theta) to avoid extra memory, a stop-gradient is applied to the corrector input for stability, and the correction loss reuses the MDM weight alpha-dot-t / (1 - alpha-t).

  4. Broad empirical validation. Experiments across math and code benchmarks, guided molecule design, and unconditional text generation show better quality-efficiency trade-offs (reported up to about 4x faster sampling) and inference-time compute scaling (reported up to about 1.2x improvement on benchmarks), plus released code, model weights, and a project page.

Main Findings

  • Fine-tuning the 8B LLaDA-Base model with ProSeCo beats standard MDM fine-tuning on every reported benchmark. Pass@1 results: HumanEval 69.51, MBPP 57.41, GSM8K 91.36, MATH 51.98 for ProSeCo SFT, versus 58.54, 56.88, 88.86, and 46.60 for the vanilla SFT baseline.

  • Adding ProSeCo Max Sampling raises scores further. With corrector sampling enabled, ProSeCo reaches HumanEval 72.56, MBPP 69.31, GSM8K 92.19, and MATH 55.06. The paper reports that ProSeCo beats the comparably sized instruction-tuned autoregressive model Llama3.1-Instruct (63.41, 70.90, 81.05, 47.38) on 3 out of 4 tasks.

  • ProSeCo outperforms other diffusion correctors in the same table. Reported numbers for LLaDA-Instruct + ReMDM are 43.90, 45.50, 83.93, 43.76; for LLaDA1.5, 45.12, 46.83, 84.00, 42.54; and for ReMeDi-Instruct, 71.30, 57.80, 86.30, 51.40. ProSeCo's Max Sampling setting exceeds all of these.

  • Quality-efficiency frontier improves. The paper reports about 2-4x speed-ups relative to LLaDA decoding without sacrificing accuracy, achieved by decoding 4-8 tokens per unmasking step, applying corrector loops every 2nd decoding iteration, and using up to 4 NFEs per corrector loop. A "Balanced" configuration gives a moderate compute increase with significant accuracy gains, and a "Max" configuration scales test-time compute for the highest reported results.

  • Parallel decoding degrades standard MDMs but not ProSeCo. Increasing parallel decoding hurts sample quality for standard MDMs, while ProSeCo recovers from the introduced mistakes and extends the parallel-decoding/quality Pareto frontier. A throughput analysis is reported in Appendix D.1.

  • Corrector hyperparameters are robust. Varying frequency (omega) and steps per loop (S) still beats baseline accuracy at each token-parallelism level, and for fast regimes (4 or 8 tokens per step), more frequent corrector loops can match or beat the baseline tokens/step = 1 result with significant speed-up (Appendix D.2).

  • Correction sampling alone does not help an untrained baseline. Applying the corrector sampling procedure to a vanilla SFT model fails to correct errors (Appendix D.4), because that model never learned to change already-unmasked tokens.

  • Guided molecule design improves the property-diversity Pareto frontier. On QM9 SMILES with classifier-free guidance, ProSeCo pushes the frontier in the desired direction for both ring count and drug-likeness (QED), most starkly for ring count. No numeric values for this experiment appear in the provided content.

  • Unconditional text generation on OpenWebText. With 5000 samples of L = 1024 tokens, ProSeCo reports MAUVE of 0.295, 0.557, 0.597, and 0.604 at T = 128, 256, 512, 1024; GPT2-Large perplexity of 23.1, 16.5, 13.2, 10.9; and entropy of 5.45, 5.39, 5.29, 5.22. For comparison, MDLM reports MAUVE 0.015/0.023/0.031/0.042 and perplexity 61.5/55.8/53.0/51.3; ReMDM reports 0.057/0.216/0.350/0.403 and 42.5/30.5/21.1/28.6; PRISM reports 0.118/0.294/0.423/0.527 and 21.5/18.0/16.4/15.3; GIDD reports 0.268/0.284/0.334/0.356 and 95.1/80.5/76.9/76.1. The AR baseline (T = 1024) reports MAUVE 0.76 and perplexity 12.1, and the data reference row is MAUVE 1.00, perplexity 14.8, entropy 5.44. Even at 256 steps, ProSeCo is reported to match PRISM at 2x or ReMDM at 4x the inference budget.

  • The time-varying correction weight matters. Ablating fixed lambda in {0.1, 1, 10} against time-varying lambda-t in {0.1, 1, 10} times alpha-dot-t/(1 - alpha-t) for ring count guidance shows that including alpha-dot-t/(1 - alpha-t) improves performance and that the model is robust to the scaling factor (Appendix D.5).

Methodology in Plain English

Standard MDMs are trained only to fill in masked positions. The denoising network is never asked to change a token that is already visible, so at generation time its predictions for unmasked positions carry no useful information and errors compound.

ProSeCo reframes the model's own output as a corruption. During training, the model first produces predictions for a masked sequence; those predictions are converted into hard tokens by taking the argmax, and that all-unmasked sequence is fed back into the same network as an input. The network is then trained with an extra cross-entropy term to reconstruct the original clean data from this self-generated sequence. A stop-gradient is placed on the argmax step for stability, and the corrector shares weights with the denoiser so no separate model is needed. The correction term reuses the same weighting the MDM loss uses, so heavily masked (harder) examples are down-weighted in the correction loss as well.

At sampling time, the model alternates between two roles. In unmasking mode it behaves like a normal MDM, taking a partially masked sequence and committing some high-confidence tokens. In corrector mode the sequence is passed through the network repeatedly (up to S times), and after each pass the already-unmasked positions in the latent sequence are overwritten with the corrector's predictions. Corrective loops are triggered every omega unmasking iterations, and the corrector's final logits are used for the posterior sampling step instead of the denoiser's original logits. The user trades compute for quality by choosing how often and how many times to run the corrector.

The paper also derives the correction loss from a predictor-corrector argument for discrete diffusion: the corrector acts as an MCMC step intended to close the gap between the model's unconditional marginals and the true marginals, and optimizing the corrector to satisfy a proportionality condition yields a loss equivalent to the correction term.

Why This Matters

Impact on research. The work challenges the assumption that decoded tokens must be locked in, and shows that a corrector can be obtained essentially for free by tying weights and adding one loss term and one forward pass. This makes self-correction compatible with existing pre-trained MDM backbones such as LLaDA, unlike prior corrector work that requires a distinct architecture. It also connects MDM correction to the learned predictor-corrector literature for discrete diffusion and to remasking/error-identification methods, positioning self-generated errors as a more "informed" noise than uniform categorical noise.

Real-world applications.

  • Code generation and program synthesis assistants, where pass@1 on benchmarks such as HumanEval and MBPP directly translates into developer productivity.
  • Mathematical reasoning and tutoring tools, where GSM8K and MATH accuracy determines whether step-by-step solutions can be trusted.
  • Molecule and drug discovery pipelines, where guided generation must maximize a target property such as ring count or QED without collapsing sample diversity or producing invalid SMILES strings.
  • Latency-sensitive text generation, where the reported 2-4x reduction in function evaluations matters for serving cost.

Industry relevance. The largest experiment is supervised fine-tuning of an 8B parameter model over roughly 400B tokens on a modified Llama-Nemotron-Post-Training dataset, which is a realistic industrial fine-tuning scale. The method's low implementation cost (one extra forward pass, a stop-gradient, one hyperparameter pair for sampling) makes it practical to bolt onto existing MDM deployment stacks, and the reported ability to trade NFEs for quality gives operators a dial for balancing serving cost against output quality.

Future Directions

  • Extending the corrector to larger pre-trained MDMs. The paper notes MDMs have reached the 8B scale and up to 100B parameters, and that LLaDA 2.1 scaled the GIDD mixed-noise framework to 16B-100B parameters, leaving open whether self-correction scales the same way.
  • Comparing against, or combining with, other correction families. The paper compares to remasking strategies (ReMDM, ReMeDi, GStar) and to a distinct Hollow-Transformer-based corrector (Zhao et al.); whether self-correction can be layered on top of these is not established in the provided content.
  • Reducing the inference overhead of corrector loops. The sampling procedure adds forward passes, and while the reported trade-offs favor ProSeCo, choosing omega and S automatically for a target compute budget is left as a guidance question rather than a solved one.
  • Exploring correction for settings beyond the mask corruption process, since the authors argue model-generated errors are a more faithful noise distribution than uniform categorical noise.

Target Audience

Researchers and engineers working on discrete diffusion and non-autoregressive text generation, particularly those fine-tuning or deploying masked diffusion models such as LLaDA. It is also relevant to practitioners in molecular generation using classifier-free guidance, and to readers interested in learned predictor-corrector sampling and inference-time compute scaling. Readers without a background in diffusion objectives or discrete latent variable models will find the theory sections challenging; the algorithm listings and benchmark tables are the most accessible entry points.

Authors’ abstract

Masked diffusion models (MDMs) have emerged as a promising alternative to autoregressive models, enabling parallel token generation while achieving competitive performance. Despite these advantages, MDMs face a fundamental limitation: once tokens are unmasked, they remain fixed, leading to error accumulation and ultimately degrading sample quality. We address this by proposing a framework that trains a model to perform both unmasking and correction. By reusing outputs from the MDM denoising network as inputs for corrector training, we train a model to recover from potential mistakes. During generation we apply additional corrective refinement steps between unmasking ones in order to change decoded tokens and improve outputs. We name our training and sampling method Progressive Self-Correction (ProSeCo) for its unique ability to iteratively refine an entire sequence, including already generated tokens. We conduct extensive experimental validation across multiple conditional and unconditional tasks, demonstrating that \method~yields better quality-efficiency trade-offs (up to ~4x faster sampling) and enables inference-time compute scaling to further increase sample quality beyond standard MDMs (up to ~1.2x improvement on benchmarks).

Read the original paper