| name | unmasking-diffusion-policies |
| title | Learning Unmasking Policies for Diffusion Language Models |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2512.09106 |
| keywords | ["diffusion language models","unmasking policies","reinforcement learning","semi-autoregressive","block-based generation"] |
| description | Learn which tokens to unmask during diffusion sampling via reinforcement learning instead of heuristics. Policies eliminate manual tuning and scale across block sizes—crucial when semi-autoregressive generation needs dynamic, learned unmasking strategies. |
Overview
Rather than relying on manual confidence thresholds for token unmasking, this approach casts diffusion language model sampling as an MDP where a lightweight transformer policy learns optimal unmasking decisions based on token confidences.
When to Use
- Diffusion language model inference optimization
- Semi-autoregressive generation with block-based sampling
- Need for dynamic unmasking beyond confidence thresholds
- Scaling across different block sizes
- Learning unmasking strategies instead of manual tuning
When NOT to Use
- Scenarios where heuristic thresholds work adequately
- Autoregressive decoding without masking
- Tasks not benefiting from RL optimization
Core Technique
RL-based policy for unmasking token selection:
class UnmaskingPolicy:
def __init__(self):
self.policy = nn.Sequential(
nn.Linear(vocab_size, 256),
nn.ReLU(),
nn.Linear(256, vocab_size)
)
def learn_unmasking_policy(self, dllm, dataset):
"""Train policy via RL on dLLM sampling."""
for batch in dataset:
confidences = dllm.get_token_confidences(batch)
unmasking_logits = self.policy(confidences)
unmasking_probs = torch.softmax(unmasking_logits, dim=-1)
decisions = torch.multinomial(unmasking_probs, num_samples=1)
sampled_tokens = dllm.sample_with_unmasking(decisions)
full_output = dllm.full_diffusion_sampling(batch)
reward = compute_similarity(sampled_tokens, full_output)
loss = -reward * torch.log(unmasking_probs[decisions])
loss.backward()
.optimizer.step()
():
unmasking_logits = .policy(confidences)
decisions = torch.argmax(unmasking_logits, dim=-)
decisions