bert-dllm / diffusion.py
chenhanqi
Build BERT masked diffusion Space
d8f20d1
Raw History Blame Contribute Delete
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): ...
@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,
)