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