File size: 7,213 Bytes
d8f20d1 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 | """Core routines for masked discrete language diffusion."""
from __future__ import annotations
from dataclasses import dataclass
from typing import Iterator, Literal, Protocol
import torch
SamplingStrategy = Literal["random", "confidence"]
class MaskedLanguageModel(Protocol):
"""Minimal ``transformers`` masked-LM protocol used by the sampler."""
def __call__(self, *, input_ids: torch.Tensor, attention_mask: torch.Tensor): ...
@dataclass(frozen=True)
class DiffusionConfig:
"""The fixed canvas and reverse process used by the article reproduction."""
canvas_length: int = 256
prefix_length: int = 16
denoising_steps: int = 10
def __post_init__(self) -> None:
if self.canvas_length <= self.prefix_length:
raise ValueError("canvas_length must be greater than prefix_length")
if self.denoising_steps < 1:
raise ValueError("denoising_steps must be positive")
@property
def generated_length(self) -> int:
"""Number of tokens denoised after the fixed conditioning prefix."""
return self.canvas_length - self.prefix_length
@property
def mask_probabilities(self) -> tuple[float, ...]:
"""Training mask rates, from fully masked to lightly masked."""
return tuple(
step / self.denoising_steps
for step in range(self.denoising_steps, 0, -1)
)
def masks_after_step(self, step: int) -> int:
"""Return the exact number of canvas tokens to re-mask after one pass."""
if not 1 <= step <= self.denoising_steps:
raise ValueError("step is outside the denoising schedule")
return round(self.generated_length * (self.denoising_steps - step) / self.denoising_steps)
@dataclass(frozen=True)
class DenoisingSnapshot:
"""A displayable state emitted after each denoising pass."""
step: int
input_ids: torch.Tensor
mask_positions: torch.Tensor
accepted_tokens: int
total_tokens: int
model_seconds: float
def prepare_conditioned_canvas(
tokenizer: object,
prompt: str,
config: DiffusionConfig,
device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor, int]:
"""Build an all-mask canvas while preserving the article's fixed prefix.
Short prompts are left-padded, exactly like the reference implementation.
Long prompts are clipped rather than silently changing the conditioning length.
"""
encoded = tokenizer(prompt, add_special_tokens=True, return_tensors="pt")
prompt_ids = encoded["input_ids"].squeeze(0).to(dtype=torch.long)
used_prompt_tokens = min(int(prompt_ids.numel()), config.prefix_length)
if prompt_ids.numel() >= config.prefix_length:
prefix = prompt_ids[: config.prefix_length]
else:
pad_id = getattr(tokenizer, "pad_token_id", None)
if pad_id is None:
raise ValueError("The tokenizer must define a pad_token_id")
padding = torch.full(
(config.prefix_length - prompt_ids.numel(),),
int(pad_id),
dtype=torch.long,
)
prefix = torch.cat((padding, prompt_ids))
mask_id = getattr(tokenizer, "mask_token_id", None)
if mask_id is None:
raise ValueError("The tokenizer must define a mask_token_id")
input_ids = torch.full(
(1, config.canvas_length), int(mask_id), dtype=torch.long, device=device
)
input_ids[0, : config.prefix_length] = prefix.to(device)
attention_mask = torch.ones_like(input_ids, device=device)
return input_ids, attention_mask, used_prompt_tokens
def _choose_remask_positions(
confidence: torch.Tensor,
config: DiffusionConfig,
target_count: int,
strategy: SamplingStrategy,
generator: torch.Generator,
) -> torch.Tensor:
"""Pick non-prefix positions to hide before the following reverse pass."""
if target_count == 0:
return torch.empty(0, dtype=torch.long, device=confidence.device)
positions = torch.arange(
config.prefix_length, config.canvas_length, device=confidence.device
)
if strategy == "confidence":
return positions[torch.topk(confidence[positions], target_count, largest=False).indices]
if strategy == "random":
permutation = torch.randperm(
positions.numel(), device=confidence.device, generator=generator
)
return positions[permutation[:target_count]]
raise ValueError(f"Unsupported sampling strategy: {strategy}")
def denoise_canvas(
model: MaskedLanguageModel,
tokenizer: object,
input_ids: torch.Tensor,
attention_mask: torch.Tensor,
config: DiffusionConfig,
*,
temperature: float,
strategy: SamplingStrategy,
generator: torch.Generator,
) -> Iterator[DenoisingSnapshot]:
"""Denoise a complete canvas in parallel and emit every reverse-process state.
``random`` reproduces the article's iterative re-masking. ``confidence`` retains
the highest-confidence predictions, which is a common improved dLLM decoder.
"""
if temperature <= 0:
raise ValueError("temperature must be positive")
mask_id = int(getattr(tokenizer, "mask_token_id"))
blocked_ids = list(dict.fromkeys(getattr(tokenizer, "all_special_ids", [])))
current = input_ids.clone()
mask_positions = current.eq(mask_id)
mask_positions[:, : config.prefix_length] = False
model_seconds = 0.0
for step in range(1, config.denoising_steps + 1):
started_at = torch.cuda.Event(enable_timing=True) if current.is_cuda else None
finished_at = torch.cuda.Event(enable_timing=True) if current.is_cuda else None
if started_at is not None:
started_at.record()
with torch.inference_mode():
logits = model(input_ids=current, attention_mask=attention_mask).logits
if finished_at is not None:
finished_at.record()
finished_at.synchronize()
model_seconds += started_at.elapsed_time(finished_at) / 1000
logits = logits / temperature
if blocked_ids:
logits[..., blocked_ids] = -torch.inf
probabilities = torch.softmax(logits, dim=-1)
confidence = probabilities.max(dim=-1).values[0]
active_positions = mask_positions[0]
active_probabilities = probabilities[0, active_positions]
sampled_tokens = torch.multinomial(
active_probabilities, 1, generator=generator
).squeeze(-1)
current[0, active_positions] = sampled_tokens
target_count = config.masks_after_step(step)
next_masked = _choose_remask_positions(
confidence,
config,
target_count,
strategy,
generator,
)
mask_positions = torch.zeros_like(current, dtype=torch.bool)
mask_positions[0, next_masked] = True
current[mask_positions] = mask_id
yield DenoisingSnapshot(
step=step,
input_ids=current.clone(),
mask_positions=mask_positions.clone(),
accepted_tokens=config.generated_length - target_count,
total_tokens=config.generated_length,
model_seconds=model_seconds,
)
|