Walk Fast but Be Careful: Understanding Parallel Sampling in Masked Diffusion
Abstract
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.
Then back it, or bet against it.
Related papers
Open the market on this paper to see 7 more related papers.