Particle Gibbs for diffusion language models — a Markov chain over complete denoising trajectories that refines toward high reward without retraining.
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
pgadaptvariant 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 |
conda env create -f env.yml
conda activate pg-dlmRun 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.
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).
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# 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 # scoresteering_eval takes a *_samples.jsonl file or a whole run directory (one row per config) and
reports the classifier rate for the steered reward.
| 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 |
@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},
}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.