"""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, )