"""Core block diffusion abstractions. The implementation is intentionally model-agnostic. A real dLLM adapter only needs to expose `forward()` that returns token logits and optionally a cache. """ from __future__ import annotations import time from dataclasses import asdict, dataclass, field from typing import Any, Protocol Token = int @dataclass(slots=True) class BlockDiffusionConfig: vocab_size: int = 128 mask_token_id: int = 0 eos_token_id: int = 2 block_size: int = 16 num_blocks: int = 1 steps: int = 8 remask_ratio: float = 0.5 use_cache: bool = False draft_width: int = 4 @dataclass(slots=True) class DecodeState: tokens: list[Token] mask: list[bool] confidences: list[float] cache: dict[str, Any] = field(default_factory=dict) @classmethod def masked(cls, length: int, mask_token_id: int) -> "DecodeState": return cls( tokens=[mask_token_id] * length, mask=[True] * length, confidences=[0.0] * length, cache={}, ) @dataclass(slots=True) class DecodeResult: tokens: list[Token] text: str nfe: int elapsed_s: float tokens_per_forward: float metadata: dict[str, Any] class MaskedLMAdapter(Protocol): vocab_size: int mask_token_id: int def forward(self, tokens: list[Token], cache: dict[str, Any] | None = None) -> tuple[list[list[float]], dict[str, Any]]: """Return per-position logits and an optional model cache.""" def decode(self, tokens: list[Token]) -> str: """Convert tokens to text for logging/evaluation.""" class ToyMaskedLMAdapter: """Deterministic toy adapter for smoke tests. It produces a simple repeating target sequence and increasing confidence for already stable positions. This validates sampler mechanics without needing a GPU or a downloaded checkpoint. """ def __init__(self, vocab_size: int = 128, mask_token_id: int = 0) -> None: self.vocab_size = vocab_size self.mask_token_id = mask_token_id def forward(self, tokens: list[Token], cache: dict[str, Any] | None = None) -> tuple[list[list[float]], dict[str, Any]]: logits: list[list[float]] = [] cache = dict(cache or {}) calls = int(cache.get("calls", 0)) + 1 for i, token in enumerate(tokens): target = 3 + (i % max(1, self.vocab_size - 3)) row = [-8.0] * self.vocab_size row[target] = 6.0 + min(calls, 8) * 0.25 if token != self.mask_token_id: row[token] = max(row[token], 5.5 + min(calls, 8) * 0.25) logits.append(row) cache["calls"] = calls return logits, cache def decode(self, tokens: list[Token]) -> str: return " ".join(str(t) for t in tokens) def argmax_with_confidence(logits: list[float]) -> tuple[int, float]: best_id = max(range(len(logits)), key=logits.__getitem__) best = logits[best_id] runner_up = max(v for i, v in enumerate(logits) if i != best_id) return best_id, best - runner_up def lowest_confidence_positions(confidences: list[float], candidates: list[int], count: int) -> set[int]: ordered = sorted(candidates, key=lambda i: confidences[i]) return set(ordered[: max(0, count)]) class BlockDiffusionSampler: method_name = "base" def __init__(self, adapter: MaskedLMAdapter, config: BlockDiffusionConfig) -> None: self.adapter = adapter self.config = config def decode(self, prompt_tokens: list[Token] | None = None) -> DecodeResult: prompt_tokens = prompt_tokens or [] generated_len = self.config.block_size * self.config.num_blocks state = DecodeState.masked(generated_len, self.config.mask_token_id) start = time.perf_counter() nfe = 0 for step in range(self.config.steps): logits, cache = self.adapter.forward(state.tokens, state.cache if self.config.use_cache else None) nfe += 1 state.cache = cache if self.config.use_cache else {} self.update_state(state, logits, step) elapsed = time.perf_counter() - start tokens = prompt_tokens + state.tokens return DecodeResult( tokens=tokens, text=self.adapter.decode(tokens), nfe=nfe, elapsed_s=elapsed, tokens_per_forward=len(state.tokens) / max(1, nfe), metadata={"method": self.method_name, "config": asdict(self.config)}, ) def update_state(self, state: DecodeState, logits: list[list[float]], step: int) -> None: raise NotImplementedError class ConfidenceRemaskSampler(BlockDiffusionSampler): """LLaDA/Dream-style fill then remask low-confidence positions.""" method_name = "confidence" def update_state(self, state: DecodeState, logits: list[list[float]], step: int) -> None: for i, row in enumerate(logits): token, conf = argmax_with_confidence(row) if state.mask[i] or conf >= state.confidences[i]: state.tokens[i] = token state.confidences[i] = conf state.mask[i] = False if step + 1 >= self.config.steps: return unmasked = [i for i, is_masked in enumerate(state.mask) if not is_masked] remask_count = int(len(unmasked) * self.config.remask_ratio * (1 - (step + 1) / self.config.steps)) for i in lowest_confidence_positions(state.confidences, unmasked, remask_count): state.tokens[i] = self.config.mask_token_id state.mask[i] = True class MultiBlockSampler(ConfidenceRemaskSampler): """Multi-block decoding with progressive block activation.""" method_name = "multiblock" def update_state(self, state: DecodeState, logits: list[list[float]], step: int) -> None: active_blocks = min(self.config.num_blocks, 1 + step * self.config.num_blocks // max(1, self.config.steps)) active_until = active_blocks * self.config.block_size inactive = range(active_until, len(state.tokens)) saved_tokens = {i: state.tokens[i] for i in inactive} saved_mask = {i: state.mask[i] for i in inactive} saved_conf = {i: state.confidences[i] for i in inactive} super().update_state(state, logits, step) for i in inactive: state.tokens[i] = saved_tokens[i] state.mask[i] = saved_mask[i] state.confidences[i] = saved_conf[i] class DMaxSampler(ConfidenceRemaskSampler): """DMax/TAD-style interface for distilled few-step block diffusion. The toy implementation changes only the step budget behavior. Real DMax/TAD reproduction should plug a trajectory-distilled adapter into this sampler. """ method_name = "dmax" def update_state(self, state: DecodeState, logits: list[list[float]], step: int) -> None: old_ratio = self.config.remask_ratio self.config.remask_ratio = old_ratio * 0.5 try: super().update_state(state, logits, step) finally: self.config.remask_ratio = old_ratio class SpeculativeSampler(ConfidenceRemaskSampler): """Draft/verify hook for DFlash/PRESTO/Fast-dLLM-style comparisons.""" method_name = "speculative" def update_state(self, state: DecodeState, logits: list[list[float]], step: int) -> None: super().update_state(state, logits, step) accepted = 0 for i in range(min(self.config.draft_width, len(state.tokens))): if state.confidences[i] > 8.0: accepted += 1 state.cache["accepted_draft_tokens"] = state.cache.get("accepted_draft_tokens", 0) + accepted SAMPLERS = { "confidence": ConfidenceRemaskSampler, "multiblock": MultiBlockSampler, "dmax": DMaxSampler, "speculative": SpeculativeSampler, }