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