cs.LGSep 29, 2026

Counterfactual Probing for Parallel Unmasking with Hidden Forest Structure

Authors: Ryotaro Kawata, Satoshi Hayakawa, Taiji Suzuki

Organizations: The University of Tokyo · RIKEN AIP

Abstract

Masked generative models offer parallel token prediction, but accurate parallel sampling must account for dependencies among tokens. When dependencies are unknown, finding safe batches also costs model evaluations. We study whether total evaluations, including discovery, can be sublinear in sequence length NN; sublinear sequential depth then follows. We consider discrete distributions with hidden forest structure, accessed through a fixed approximate conditional oracle. Under explicit regularity conditions and uniform Hellinger error bounds, for any fixed target accuracy ε∈(0,1/8]\varepsilon\in(0,1/8] and sufficiently large NN, our sampler achieves seed-averaged total-variation error at most ε\varepsilon, with total masked-state submissions and sequential depth both bounded by O~(NCε−a)\widetilde{O}(N^C \varepsilon^{-a}) for constants 0<C<10<C<1 and a>0a>0. These guarantees use polynomial vocabulary size and an edge-response lower bound set by NN and ε\varepsilon. The sampler shares evaluations of hypothetical reveals across dependence tests to identify safe parallel batches without requiring full recovery of the hidden forest. A tunable parameter trades probing cost against irreversible commit rounds. In the same class, any admissible irreversible product-commit sampler attaining the same seed-averaged accuracy requires Ω(Ncεb)Ω(N^c \varepsilon^b) counterfactual submissions or commit rounds in the worst case, for constants c,b>0c,b>0.

Figures & tables

Appendix figures & tables7 assets

Supplementary material from the paper’s appendix.

Appendix

Explore similar work

Jun 22, 2026cs.LG

Walk fast but be careful: Understanding Parallel Sampling in Masked Diffusion

In this paper, we use random walks on graphs as a verifiable sandbox for studying parallel sampling strategies in masked diffusion models (MDMs). We train an MDM on random walk samples from a fixed graph. The graph and transition kernel are never shown to the model and serve as latent structure that is both controllable and enables evaluation. The framework provides a validity check for generated walks and a measure of distributional fidelity through the estimated transition kernel. Using simple graphs, we theoretically prove that parallel unmasking via widely used scores such as lowest entropy is not uniformly better than random parallel sampling; even with exact conditional probabilities, performance critically depends on the conditional dependence structure induced by the graph, a phenomenon difficult to isolate in benchmarks like Sudoku. We also develop training-free bisection samplers for MDMs, which take logarithmically many steps in the sequence length and are provably exact for random walks if the learned marginals are exact. Experiments on graph-walk tasks confirm that different parallel samplers perform better on different graph structures. Experiments on pretrained MDMs show that bisection-style samplers provide strong speed-quality tradeoffs on OpenWebText generation and reasoning benchmarks including GSM8K, MBPP, and HumanEval. Together, these results use graph walks to uncover conditional dependence as a key principle of parallel MDM sampling and translate this insight into efficient samplers that transfer to language generation and reasoning.
Jun 26, 2026cs.LG

VGB for Masked Diffusion Model: Efficient Test-time Scaling for Reward Satisfaction and Sample Editing

Inference-time scaling is a promising paradigm to improve generative models, especially when outputs must satisfy structural constraints or optimize downstream rewards. We consider Masked Diffusion Model (MDM) and introduce MDM-VGB, a discrete diffusion sampler that augments unmasking generation with theoretically principled reward-guided remasking. Inspired by the recent success of the classical Jerrum-Sinclair backtracking Markov chain in reward-tilted generation, MDM-VGB extends the backtracking random walk from a fixed prefix tree to a masked-state graph, allowing tokens to be unmasked and remasked at arbitrary positions. The resulting sampler favors unmasking and remasking moves that lead to higher-value partial configurations, enabling both effective high-reward generation and efficient repair of low-reward samples. We prove that MDM-VGB is robust to process-verifier noise and achieves quadratic complexity, while popular test-time heuristics such as best-of-NN can incur exponential complexity due to error accumulation. Our theoretical findings are corroborated by strong empirical performance, particularly on popular constraint-satisfaction and scientific benchmarks such as Sudoku and QM9.
Sep 29, 2026cs.CL

Reliable Parallel Decoding in Masked Diffusion Language Models

Masked diffusion language models (MDLMs) can generate text efficiently by predicting multiple masked tokens in parallel, but predictions from the same forward pass are not necessarily reliable when committed together. We study when parallel commitment is reliable. Our diagnostics show that confidence alone does not determine a reliable commitment order: confident predictions near the end of the sequence can fix an answer before its supporting computations are established, and downstream predictions become less reliable as the uncertainty of their upstream context grows. At the same time, a single forward pass can already resolve several masked tokens, and predictions that remain stable across the final layers are more likely to be correct. Based on these findings, we propose Reliable Parallel Decoding (RPD), a training-free method that selects candidates by layerwise prediction stability and final confidence, and commits them under a cumulative entropy budget over their preceding masked positions. RPD defers predictions with uncertain upstream context while committing the remaining candidates in parallel, without relying on a fixed block schedule. Across mathematical reasoning and code generation benchmarks on LLaDA and Dream, RPD achieves the highest decoding throughput among the evaluated methods while maintaining or improving accuracy.