Ouzhang's picture
Add files using upload-large-folder tool
13c5606 verified
Raw
History Blame Contribute Delete
7.87 kB
"""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,
}