Skip to content

Latest commit

 

History

2 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

PG-DLM: Inference-Time Scaling of Diffusion Language Models via Trajectory Refinement

Particle Gibbs for diffusion language models — a Markov chain over complete denoising trajectories that refines toward high reward without retraining.

arXiv

Overview

PG-DLM treats inference-time scaling of masked diffusion language models as sampling over trajectories. It provides one unified loop over three inference-time algorithms and a base model:

  • BON — Best-of-N: draw N samples, keep the highest-reward one.
  • SMC — Sequential Monte Carlo: propagate a particle set and resample by intermediate reward.
  • PG — Particle Gibbs: a Markov chain over full denoising trajectories that resamples them with conditional SMC (the pgadapt variant early-stops per example, adding an average number of particles scaling axis).

The same loop drives two base models:

--base_model model task reward
llada LLaDA-8B GSM8K reasoning Qwen2.5-Math-PRM-7B (process reward)
llada / mdlm LLaDA-8B / MDLM reward-guided steering toxicity / sentiment / CoLA classifiers

Installation

conda env create -f env.yml
conda activate pg-dlm

Run everything from the repository root (so pg_dlm is importable). env.yml includes a prebuilt flash-attn wheel (needed by the MDLM backbone) matching the pinned torch 2.6 / Python 3.11 / CUDA 12; if you change any of those, swap it for the matching wheel from the flash-attn releases.

Model weights (LLaDA-8B, kuleshov-group/mdlm-owt, the reward/eval classifiers) download from the Hugging Face Hub on first use — nothing is vendored. Point HF_HOME at a disk with room to spare.

Repository structure

pg_dlm/
  eval.py                  # single entry point: --base_model {llada,mdlm}, --algorithm, --dataset
  backends/                # base-model denoising, behind one interface
    llada.py               #   LLaDA block-wise diffusion (resample at block ends)
    mdlm.py                #   MDLM flat diffusion (resample every N steps); backbone from the Hub
  samplers/                # model-agnostic BON / SMC / PG (+ pgadapt)
  rewards/                 # math PRM (reward_hub) + toxicity / sentiment / cola classifiers
  tasks/                   # GSM8KDataset, SteeringDataset (PPLM prompts)
  parsing/                 # gsm8k_acc.py (accuracy) + steering_eval.py (toxicity/sentiment/cola rate)
scripts/                   # runnable experiment + parsing scripts

eval.py only saves generations; accuracy and steering metrics are computed afterward by the parsers. Run single-GPU with python -m pg_dlm.eval ... or multi-GPU with torchrun --nproc_per_node N -m pg_dlm.eval ... (the scripts do the latter).

Reproduce: GSM8K on LLaDA

bash scripts/run_llada_gsm8k_bon.sh 0            # Best-of-N   (GPU 0)
bash scripts/run_llada_gsm8k_smc.sh 0 1 2 3      # SMC         (GPUs 0-3)
bash scripts/run_llada_gsm8k_pg.sh  0 1 2 3      # Particle Gibbs

python -m pg_dlm.parsing.gsm8k_acc --dir results/llada_gsm8k_smc   # accuracy

Reproduce: reward-guided steering (LLaDA and MDLM)

# reward ∈ {toxicity, sentiment, cola}
bash scripts/run_mdlm_steering_bon.sh  toxicity 0      # MDLM   Best-of-N
bash scripts/run_mdlm_steering_smc.sh  toxicity 0      # MDLM   SMC
bash scripts/run_mdlm_steering_pg.sh   toxicity 0      # MDLM   Particle Gibbs

bash scripts/run_llada_steering_bon.sh toxicity 0      # LLaDA  Best-of-N
bash scripts/run_llada_steering_smc.sh toxicity 0      # LLaDA  SMC
bash scripts/run_llada_steering_pg.sh  toxicity 0      # LLaDA  Particle Gibbs

bash scripts/eval_steering.sh results/mdlm_steering_pg_toxicity toxicity   # score

steering_eval takes a *_samples.jsonl file or a whole run directory (one row per config) and reports the classifier rate for the steered reward.

Key options

flag meaning
--algorithm base / bestofn / smc / pg / pgadapt
--num_samples number of particles
--num_iters, --init_pg, --weight_threshold Particle-Gibbs refinement iterations / reference init / adaptive stop
--lmbda, --ess_ratio, --resample_freq potential temperature, ESS resample threshold, resample cadence
--reward_strategy full = exp(λ·rₜ) (GSM8K) · diff = exp(λ·(rₜ−rₜ₋₁)) (steering)
--compute_partial partial (score the denoised prefix — LLaDA) · x0 (score the sampled x0 — MDLM); defaults to the backend's natural choice

Citation

@inproceedings{dang2026inference,
  title = {Inference-Time Scaling of Diffusion Language Models via Trajectory Refinement},
  author = {Dang, Meihua and Han, Jiaqi and Xu, Minkai and Xu, Kai and Srivastava, Akash and Ermon, Stefano},
  booktitle = {Proceedings of the Third Conference on Language Modeling (COLM)},
  year = {2026},
}

Acknowledgements

Built on d1 (LLaDA generation + evaluation), MDLM (masked diffusion base model), and Feynman-Kac Diffusion Steering (steering rewards and prompts). See NOTICE for details. Licensed under Apache-2.0.

About

Inference-Time Scaling of Diffusion Language Models via Trajectory Refinement. COLM 2026.

Resources

Stars

0 stars

Watchers

1 watching

Forks

Releases

Packages

Contributors

Languages