| name | prism-diffusion-scaling |
| title | Prism: Efficient Test-Time Scaling via Hierarchical Search for Diffusion Language Models |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2602.01842 |
| keywords | ["Test-Time Scaling","Discrete Diffusion","Inference Optimization","Adaptive Search","Self-Verification"] |
| description | Scale inference efficiency for discrete diffusion language models through hierarchical trajectory search with adaptive pruning and self-verified feedback. Achieve 3-4× speedup versus best-of-N with equal quality. |
Prism: Efficient Test-Time Scaling for Discrete Diffusion Language Models
Most inference-scaling methods optimize for autoregressive decoding, but discrete diffusion models generate sequences through parallel iterative denoising—fundamentally different. Naive best-of-N approaches require O(NT) function evaluations (N trajectories, T denoising steps) which is computationally prohibitive. Prism introduces three innovations: hierarchical trajectory search that progressively prunes low-quality trajectories during mid-denoising, self-verified feedback reusing the model as verifier, and partial remasking preserving high-confidence tokens while exploring alternatives.
The key insight is that diffusion's bidirectional context enables effective mid-generation pruning unavailable to autoregressive models.
Core Concept
Prism uses three complementary mechanisms:
-
Hierarchical Trajectory Search (HTS): Divide inference into stages with geometric decay of active trajectories from N → K during early denoising when "logic skeletons" stabilize, then final refinement with K survivors
-
Self-Verified Feedback (SVF): Reuse the dLLM itself as verifier through dedicated Yes/No prompts on intermediate completions, eliminating overhead of separate models
-
Local Branching via Partial Remasking: Preserve high-confidence tokens as fixed "logic skeleton" while selectively re-masking low-confidence positions, enabling diverse exploration within fixed budget
These reduce complexity to approximately O(N + KT), achieving significant speedups.
Architecture Overview
- Trajectory Initializer: Launch N parallel trajectories
- Denoising Step Executor: Standard diffusion denoising step on all active trajectories
- Confidence Scorer: Estimate token confidence from model logits
- Pruning Engine: Geometric decay schedule reducing trajectories from N → K
- Self-Verification Module: Query dLLM for Yes/No feedback on intermediates
- Partial Remasking: Selectively remask low-confidence positions
- Skeleton Preservation: Keep high-confidence tokens fixed across variants
Implementation
The method involves trajectory management, confidence scoring, and adaptive pruning.
Initialize and manage parallel diffusion trajectories:
import torch
from typing import ,
:
():
.trajectories = [initial_noise.clone() _ (num_trajectories)]
.active_mask = torch.ones(num_trajectories, dtype=torch.)
.confidence_scores = []
.step_count =
():
.active_mask = torch.zeros((.trajectories), dtype=torch.)
.active_mask[new_active_indices] =
.trajectories = [.trajectories[i] i new_active_indices]
():
[traj traj, active (.trajectories, .active_mask) active]
():
.step_count = step_idx
batch = TrajectoryBatch(initial_noise, num_trajectories=)