File size: 7,868 Bytes
13c5606 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 | """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,
}
|