Skip to content

Why Deterministic PRM Guidance Underperforms in Discrete Diffusion Reasoning

Conference: NeurIPS2026
arXiv: 2609.35472
Code: https://github.com/dLLM-PRM-Gap/
Area: LLM Reasoning
Keywords: discrete diffusion language models, process reward models, test-time compute, candidate diversity, outcome verifiers

TL;DR

This diagnostic study charges diffusion generation and reward scoring to a shared inference budget, finds that deterministic top-1 PRM guidance on Dream-7B trails independent sampling with a task-matched ORM, and locates the failures through candidate-pool, terminal-scoring, and readout controls.

Background & Motivation

A discrete diffusion language model does not progressively fix a left-to-right prefix: it repeatedly denoises a partially masked full sequence. An intermediate snapshot may already reveal a later answer fragment while still hiding an earlier condition needed for the calculation. This appears well suited to process reward models (PRMs): before the answer becomes fixed, a scorer could inspect intermediate states and allocate subsequent computation to promising branches. However, a prefix verifier effective in autoregressive reasoning is not automatically suited to evidence scattered across positions.

The problem extends beyond the reward model's classification ability. Guidance incurs additional scorer calls and deletes candidates early; a high-scoring state that eventually fails may permanently eliminate the ancestors of correct alternatives. Meanwhile, a scorer that distinguishes overall problem difficulty may not select correctly among multiple answers to the same problem. Reporting only PRM ROC-AUC or guided accuracy therefore cannot establish whether the additional computation is worthwhile.

The paper does not introduce a new generation network. It diagnoses an existing segmental branching and top-1 pruning recipe under controlled compute. Core idea: separately measure which correct candidates survive search and whether the scorer selects them, then locate the bottlenecks with final-state specialists, diversity-preserving sampling, and readout ablations.

Method

Overall Architecture

The input is a mathematics problem or programming task. The generator maintains a partially masked solution and produces a completed answer after 128 denoising steps. The diagnostic procedure first fits verifiers using outcome-labeled snapshots from training problems, then compares independent sampling and guided search under a shared forward-pass budget, and finally evaluates the candidate pool separately from the final choice.

The cross-mask PRM is an outcome-supervised intermediate-state value model, not a conventional process-supervised model trained on human judgments of individual reasoning steps. Across mask ratios, it predicts the final correctness of the trajectory associated with a state. The ORM and final-state PRM instead train only on fully decoded states. The former evaluates intermediate states, whereas the latter two specialize in completed answers; their training distributions and inference responsibilities must not be conflated.

The diagram depicts the diagnostic workflow, not a newly proposed network. Dashed edges represent training supervision or ablation relationships; solid edges represent evaluation data flow. PRM Guided outputs one answer after terminal pruning, whereas PRM Hybrid preserves all terminal candidates for pool diagnostics.

%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400}}}%%
flowchart TD
    A["Training trajectories<br/>final-correctness labels"] -.-> B["Outcome-supervised snapshots"]
    B -.->|Fit PRM and final-state verifiers| C["Unified cost protocol"]
    Q["Test problems"] --> C
    C -->|PRM guidance or independent sampling| D["Poolโ€“selection separation"]
    D -->|Hybrid or independent pool| P["Oracle ceiling and diversity"]
    D -->|Guided terminal PRM or ORM selection| O["Answer accuracy"]
    D -.-> E["Readout controls"]
    E -.-> O

Key Designs

1. Outcome-supervised snapshots: establish what the PRM actually learns

The authors run the generator on GSM8K training problems, save denoising snapshots, and assign the completed trajectory's binary correctness label to its intermediate states. This requires no human labeling of individual steps, but the target means โ€œwhether this sampled continuation eventually succeeded,โ€ not โ€œwhether every visible step is correct.โ€ Different stochastic continuations from a partial state may still yield correct or incorrect answers. Inherited labels are therefore noisy observations of state value, especially at high mask ratios.

The PRM combines a frozen diffusion language model backbone, trainable LoRA adapters, and a two-layer MLP reward head. It uses pooled solution hidden states together with a 256-dimensional denoising-step embedding. The main architectural comparison keeps data and the reward head fixed while changing bidirectional versus causal attention; readout receives a separate ablation. Final-state specialists see only completed solutions, and the final-state PRM also drops the step embedding. Consequently, the names โ€œPRMโ€ and โ€œORMโ€ cannot themselves explain the performance gap: the actual training protocol matters.

Highly masked snapshots contain little evidence for distinguishing final outcomes. Bidirectional PRM ROC-AUC falls from approximately 0.77 in the nearly decoded bucket to approximately 0.54 in the most-masked bucket. The information bound in Appendix F.1 states that, within a fixed mask bucket and with fixed class priors, if the mutual information between state and final correctness approaches zero, the best achievable AUC ceiling also approaches chance. This is a conditional statement: it does not prove that mutual information must decrease monotonically with mask ratio in the actual denoising distribution, nor that the learned PRM reaches the theoretical ceiling.

2. Unified cost protocol: charge intermediate and terminal scoring alike

PRM Guided starts from a fully masked state. Every \(b\) denoising steps, it replicates the current state into \(K\) branches, samples each branch independently, and retains only the state with the highest PRM score. All branches in the next segment share this retained ancestor rather than starting independent full trajectories. The final segment still scores every candidate and performs top-1 selection; if \(b\) does not divide the total step count, the final segment is shorter.

Let \(T\) denote the number of denoising steps and \(N\) the number of independent complete samples. Compute is measured as one forward pass through a dLLM-scale model for one candidate:

\[ C_{\mathrm{ORM}}(N)=NT+N,\qquad C_{\mathrm{PRM}}(K,b)=KT+K\left\lceil\frac{T}{b}\right\rceil. \]

ORM Rerank spends its budget on \(N\) independent complete trajectories and one terminal scoring pass per trajectory. At the headline setting \(T=128,b=64\), PRM Guided costs \(130K\) passes and ORM Rerank costs \(129N\). Setting \(K=N\) gives approximate matching within 0.8%, not exact equality. The small asymmetry favors PRM. Forward-pass matching must not be described as exact wall-clock matching either, since model calls, parallelism, and segment overhead differ. Verifier training compute is separate and is not included in these inference-cost formulas.

3. Poolโ€“selection separation: did the correct answer disappear, or was it not selected?

PRM Hybrid follows Guided through the preceding segments but preserves all \(K\) terminal candidates rather than pruning at the end. This enables Oracle@K: a problem counts as solvable whenever at least one pool candidate is correct. Oracle uses ground-truth correctness and is a perfect-selector ceiling, not a deployable algorithm. The Hybrid pool ceiling must not be reported as Guided's achieved accuracy.

The authors also count unique answers per problem and apply offline counterfactual pruning. At stored early, middle, and late states, a score-based top-M cut tests whether every lineage that eventually succeeds is removed. When two states are retained, each spawns \(K/2\) children in the next segment, keeping total branch width at \(K\). This controls the explanation that relaxed pruning helps merely by spending more compute.

A further control uses effective-sample-size (ESS)-tempered sequential Monte Carlo (SMC) at the same budget. Multiple particles receive tempered PRM weights and undergo systematic resampling when necessary. SMC substantially restores the candidate pool but still uses the cross-mask PRM for terminal weighted answer voting or highest-score particle selection. If Oracle recovers without an accuracy recovery, diversity is not the only bottleneck: terminal selection also matters. Applying the cross-mask PRM, ORM, and final-state PRM to the same independently sampled candidate pool further isolates terminal scoring ability.

Pooled ROC-AUC measures discrimination between correct and incorrect candidates across all problems; it is not within-problem ranking quality. A scorer that gives easy problems high scores and hard problems low scores may achieve high pooled AUC while assigning identical scores to all candidates of a given problem. The empirical PRM does carry within-problem signal, but moderate pairwise advantages do not guarantee a correct highest-scoring candidate. The appendix's conditionally independent comparison model is illustrative, not a general bound on actual PRM accuracy.

4. Readout controls: do not attribute pooling defects entirely to causal attention

Early token hidden states in a causal model cannot read the complete subsequent context. Mean-pooling visible tokens mixes many representations with only partial evidence. Keeping the remaining training setup fixed, the authors replace that readout with the hidden state at the last non-MASK, non-EOS position. On a fully decoded sequence, this position can aggregate the preceding complete text, directly testing whether the main causal PRM fails because of an unsuitable readout.

Last-token pooling raises final-state ROC-AUC from 0.6147 to 0.7274, recovering approximately 70% of the gap to bidirectional mean pooling at 0.7759. It also improves reranking, but reaches only 50.64% at \(N=32\), far below the ORM's 82.71%. Fixing readout does not eliminate the terminal-specialization gap associated with cross-mask training. Longer training and final-state training controls still leave a bidirectional advantage, but do not establish that every causal sequence-level scorer is theoretically unable to handle snapshots: Appendix F.4 covers only restricted prefix readouts.

A Worked Example

With \(K=8,b=64,T=128\), producing the first eight half-completed snapshots costs 512 generation passes and eight PRM scoring passes. If top-1 retains a snapshot containing an incorrect calculation, the other seven lineages cannot return even if their continuations could have succeeded. The second segment produces eight complete answers from the retained state for another 520 passes.

Guided then returns only the terminal PRM's highest-scoring answer; Hybrid exposes all eight completed answers to diagnostics. If all eight are wrong, no terminal selector can repair the pool. If a correct answer is present but not selected, terminal scoring is the problem. ORM Rerank at the corresponding scale instead spends 1,032 passes on eight independent complete solutions and terminal scoring, avoiding an intermediate decision that ties every candidate to one ancestor. These call counts illustrate the protocol, not a newly measured case result.

The paper also provides a terminal misranking example. A new device consumes 2 kWh per day at 1.50 dollars per kWh, so the additional weekly cost is 21 dollars. Among three independent candidates, the causal PRM scores the wrong 30-dollar answer at +1.001 and the wrong 15-dollar answer at โˆ’0.827, but the correct 21-dollar answer at โˆ’0.908. The ORM selects the correct solution from the full 32-candidate pool. This demonstrates the distinction between answer availability and correct selection, but does not reconstruct a measured Guided trajectory end to end. The source labels these candidates \(t=13,25,3\); they are treated here as candidate identifiers, not denoising-step indices.

Loss & Training

The main PRM trains with binary final-correctness labels and binary cross-entropy (BCE). LoRA targets q_proj and v_proj with rank 16 and alpha 32. Main-comparison models train for 2,000 steps at batch size 32, with a cosine learning-rate schedule from \(2\times10^{-5}\) to zero. The backbone is frozen. Training, tuning, and early stopping use training problems only; validation problems are separate from fitting problems, and test problems are not used for training.

The main GSM8K sampler uses temperature=0.5, alg_temp=0.5, and top_p=1.0. Training snapshots and testing share the sampler and snapshot schedule. Retrains of 15,000/31,000 steps and the 8,407-step final-state protocol are separate controls, not the main PRM training budget.

To test inherited-label noise, the authors also perform eight fresh continuations from each of 10,000 training states and use the success fraction as the label, comparing models under the same 500-step training control. Pooled AUC rises from 0.816 to 0.827, but inference accuracy moves only from 60.88% to 61.49%, with a difference interval including zero. More detailed outcome-value supervision does not remove the main gap in this control; that does not establish that human process labels, more rollouts, or other supervision would be ineffective.

Key Experimental Results

Main Results

The primary evidence comes from Dream-v0-Instruct-7B. GSM8K uses all 1,319 test problems with strict answer extraction; the 32-candidate pool contains 42,208 complete trajectories. The table selects headline results and task-matched controls. Accuracy is in %, and forward-pass costs are per problem.

Task and budget PRM Guided ORM Rerank ORM advantage (pp) Cost and reporting scope
GSM8K, K=N=8 65.18 ยฑ 0.75 75.13 9.95 1,040 vs 1,032; Guided mean over three runs
GSM8K, K=N=32 70.02 ยฑ 0.70 82.71 12.69 4,160 vs 4,128; Guided mean over two runs
MATH, K=N=8 20.80 30.65 9.85 Task-specific verifiers, 500 test problems
MBPP, task-matched control 50.88 ยฑ 3.57 63.04 ยฑ 0.84 12.16 257 held-out tasks; N/K not separately specified in this main-text row

Sources: Tables 1 and 2, Appendices B.2 and D.1. MATH verifier fitting and validation use 450/50 problems, separate from the 500 held-out test problems. The ORM result from rescoring shared stored candidates is 31.90%, not the table's 30.65%; the two must not be merged into one result. On MBPP, terminal PRM reranking reaches 65.47 ยฑ 0.74%, with its difference from the ORM including zero, yet substantially exceeds Guided at 50.88%. This isolates a setting where the terminal PRM is already competitive but guidance still incurs a loss.

The GSM8K single-sample baseline is 43.14%; majority voting reaches 60.05%/67.63% at N=8/32. Guided therefore carries useful signal and beats majority voting, but independent sampling with a terminal ORM makes better use of comparable compute. ORM@8 also exceeds every tested Guided budget, including K=32 at four times the scale and the best Guided sweep point, K=24 at 73.67%.

Ablation Study

Search on GSM8K, K=N=8 Oracle@8 (%) Unique answers per problem Cross-mask PRM final choice (%) ORM final choice (%)
Independent samples 81.05 4.31 42.84 75.13
Hybrid pool from top-1 guidance 67.30 ยฑ 1.24 1.75 65.18 (Guided output) Not run
Matched-budget SMC 77.89 ยฑ 0.57 3.95 ยฑ 0.02 65.48 ยฑ 0.12 (weighted voting) Not run

Sources: Table 3, Appendices B.1 and D.2. The Hybrid pool and Guided final output are complementary diagnostics, not identical output interfaces. SMC highest-score particle selection separately reaches 66.34 ยฑ 0.82%. Top-1 loses 13.75 pp of Oracle headroom, and SMC recovers 10.59 pp without bringing achieved selection to ORM performance. This supports separate pool-damage and terminal-selection bottlenecks. โ€œNot runโ€ must not be read as evidence that an ORM would fail on the restored pool.

Terminal scorer Pooled ROC-AUC on fully decoded states Reranking accuracy at N=8 (%) Reranking accuracy at N=32 (%)
Causal cross-mask PRM, mean pooling 0.6147 43.90 40.56
Causal cross-mask PRM, last-token pooling 0.7274 49.66 50.64
Bidirectional cross-mask PRM, mean pooling 0.7759 42.84 65.35
Bidirectional ORM, final-state-only training 0.9623 75.13 82.71
Final-state PRM, final-state-only training Not separately reported 75.40 ยฑ 0.05 82.79 ยฑ 0.11

Sources: Appendices C.3, C.4, and B.2, using the same independently sampled terminal candidate pool. Paired difference intervals between the final-state PRM and ORM include zero at both budgets. The table does not assign the ORM's AUC to the final-state PRM. The bidirectional cross-mask PRM even trails the causal last-token model at N=8, showing that relative pooled AUC does not directly establish within-problem top-1 performance at a particular budget.

Key Findings

  • Relaxing top-1 to top-2 at the same branch width raises Guided accuracy from 65.18% to 69.70%, still 5.43 pp below ORM@8. In the offline terminal-state cut, the risk of removing every correct lineage is 19.92%/5.93% for top-1/top-4.
  • Single evaluations at b=16/32/48/64 give 59.4%/62.8%/65.1%/66.5%. The 66.5% result is one run, whereas 65.18% in the headline table is a multi-run mean. Likewise, Appendix C.2's single-run gaps of 8.79/11.98 must not replace the headline 9.95/12.69.
  • AUC is 0.7702 for the nearly decoded snapshot bucket and 0.7759 for fully decoded states. The 0.816/0.827 values belong to the fresh-rollout supervision control. These are different slices, not points to merge into one mask-ratio curve.
  • Wall-clock validation reports approximately 156 seconds/problem for Guided@8 and 168 seconds/problem for ORM@8, not exact time matching. ORM@6 already achieves 72.40% in approximately 126 seconds, exceeding Guided@8 at 65.18%.

Highlights & Insights

  • Reporting Oracle alongside achieved accuracy distinguishes โ€œdeleting the answerโ€ from โ€œmisjudging the answer.โ€ This is more actionable than final accuracy alone: the former calls for preserving lineages, the latter for improving the terminal verifier.
  • Final-state PRM parity with the ORM directly counters misleading model labels. Calling a scorer a PRM does not imply that it must continuously guide search or cannot perform terminal verification; its actual training distribution matters more.
  • The SMC control makes diversity restoration a useful direction but not a sufficient solution in these experiments. Diversity-preserving search paired with a final-state specialist is a natural next test, not an already verified combination result.

Limitations & Future Work

  • The conclusion concerns deterministic segmental top-1 guidance and the tested controls on Dream-7B, not every PRM, diffusion language model, or stochastic search method. Tasks are limited to mathematics and code; open-ended generation is untested.
  • LLaDA cross-backbone evidence consists of configuration-averaged bidirectional/causal guided accuracies of 31.64%/22.25%, against 20.77% for a single sample. No LLaDA-specific ORM comparison was run, so it does not establish ORM superiority on a second backbone. Standard deviations are across configurations, not random seeds of one fixed configuration.
  • A GSM8K-trained ORM does not transfer to MATH500: it reaches only 6.10% at N=32. However, independent sampling and Guided use different temperatures in that OOD table, which establishes transfer failure rather than replacing the task-matched MATH control.
  • Removing right-half generated tokens changes AUC by only approximately 0.001, and reversed-input causal scoring is near chance, but neither excludes every use of right-side evidence. Reversal also changes the order familiar to the backbone and RoPE-relative positions; it is not proof of inherent causal-architecture incapacity.
  • Adaptive late guidance, explicit diversity kernels, appended CLS or learned-query readouts, and human step supervision merit testing. Main inference results exclude amortized verifier training: the study reports approximately 2,400 H20 GPU-hours overall, so near-matched inference calls do not establish equal combined training and inference cost.
  • vs autoregressive PRMs and Math-Shepherd: These primarily evaluate sequentially growing reasoning prefixes. This paper studies non-prefix masked snapshots and uses outcome-inherited labels rather than human step annotations. The lesson is to inspect supervision semantics and state distributions before transferring a reward-search recipe.
  • vs self-consistency and best-of-N verification: Majority voting relies primarily on repeated answers, whereas the ORM uses learned terminal-correctness signal. Strong task-matched verification gains cannot simply be attributed to drawing more samples.
  • vs dLLM particle search, remasking, and reward-free guidance: These change sampling or construct alternative rewards. This paper supplies cost accounting and poolโ€“selection decomposition as a common evaluation protocol, but does not empirically rank all these methods.

Rating

  • Novelty: 4/5 โ€” The contribution is controlled failure decomposition, not a new generator.
  • Experimental Thoroughness: 4/5 โ€” Same-pool, SMC, readout, supervision, and task controls are substantial; cross-backbone ORM and broader search controls remain missing.
  • Writing Quality: 4/5 โ€” The causal story is clear, but single/multi-run results, data slices, and protocols require careful separation.
  • Value: 4/5 โ€” Provides reusable compute-accounting and candidate-pool evaluation criteria for reward-guided reasoning.