Download diffusion.py from Caffin/bert-dllm: direct link, hf CLI and curl.
- Browser
- Download file 7.21 kB
-
https://huggingface.co/spaces/Caffin/bert-dllm/resolve/main/diffusion.py
- Command line
-
hf download hf://spaces/Caffin/bert-dllm/diffusion.py
-
curl -L -o diffusion.py https://huggingface.co/spaces/Caffin/bert-dllm/resolve/main/diffusion.py
7.21 kB
| """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): ... | |
| 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") | |
| def generated_length(self) -> int: | |
| """Number of tokens denoised after the fixed conditioning prefix.""" | |
| return self.canvas_length - self.prefix_length | |
| 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) | |
| 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, | |
| ) | |