File size: 7,213 Bytes
d8f20d1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Core routines for masked discrete language diffusion."""

from __future__ import annotations

from dataclasses import dataclass
from typing import Iterator, Literal, Protocol

import torch


SamplingStrategy = Literal["random", "confidence"]


class MaskedLanguageModel(Protocol):
    """Minimal ``transformers`` masked-LM protocol used by the sampler."""

    def __call__(self, *, input_ids: torch.Tensor, attention_mask: torch.Tensor): ...


@dataclass(frozen=True)
class DiffusionConfig:
    """The fixed canvas and reverse process used by the article reproduction."""

    canvas_length: int = 256
    prefix_length: int = 16
    denoising_steps: int = 10

    def __post_init__(self) -> None:
        if self.canvas_length <= self.prefix_length:
            raise ValueError("canvas_length must be greater than prefix_length")
        if self.denoising_steps < 1:
            raise ValueError("denoising_steps must be positive")

    @property
    def generated_length(self) -> int:
        """Number of tokens denoised after the fixed conditioning prefix."""
        return self.canvas_length - self.prefix_length

    @property
    def mask_probabilities(self) -> tuple[float, ...]:
        """Training mask rates, from fully masked to lightly masked."""
        return tuple(
            step / self.denoising_steps
            for step in range(self.denoising_steps, 0, -1)
        )

    def masks_after_step(self, step: int) -> int:
        """Return the exact number of canvas tokens to re-mask after one pass."""
        if not 1 <= step <= self.denoising_steps:
            raise ValueError("step is outside the denoising schedule")
        return round(self.generated_length * (self.denoising_steps - step) / self.denoising_steps)


@dataclass(frozen=True)
class DenoisingSnapshot:
    """A displayable state emitted after each denoising pass."""

    step: int
    input_ids: torch.Tensor
    mask_positions: torch.Tensor
    accepted_tokens: int
    total_tokens: int
    model_seconds: float


def prepare_conditioned_canvas(
    tokenizer: object,
    prompt: str,
    config: DiffusionConfig,
    device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor, int]:
    """Build an all-mask canvas while preserving the article's fixed prefix.

    Short prompts are left-padded, exactly like the reference implementation.
    Long prompts are clipped rather than silently changing the conditioning length.
    """
    encoded = tokenizer(prompt, add_special_tokens=True, return_tensors="pt")
    prompt_ids = encoded["input_ids"].squeeze(0).to(dtype=torch.long)
    used_prompt_tokens = min(int(prompt_ids.numel()), config.prefix_length)

    if prompt_ids.numel() >= config.prefix_length:
        prefix = prompt_ids[: config.prefix_length]
    else:
        pad_id = getattr(tokenizer, "pad_token_id", None)
        if pad_id is None:
            raise ValueError("The tokenizer must define a pad_token_id")
        padding = torch.full(
            (config.prefix_length - prompt_ids.numel(),),
            int(pad_id),
            dtype=torch.long,
        )
        prefix = torch.cat((padding, prompt_ids))

    mask_id = getattr(tokenizer, "mask_token_id", None)
    if mask_id is None:
        raise ValueError("The tokenizer must define a mask_token_id")

    input_ids = torch.full(
        (1, config.canvas_length), int(mask_id), dtype=torch.long, device=device
    )
    input_ids[0, : config.prefix_length] = prefix.to(device)
    attention_mask = torch.ones_like(input_ids, device=device)
    return input_ids, attention_mask, used_prompt_tokens


def _choose_remask_positions(
    confidence: torch.Tensor,
    config: DiffusionConfig,
    target_count: int,
    strategy: SamplingStrategy,
    generator: torch.Generator,
) -> torch.Tensor:
    """Pick non-prefix positions to hide before the following reverse pass."""
    if target_count == 0:
        return torch.empty(0, dtype=torch.long, device=confidence.device)

    positions = torch.arange(
        config.prefix_length, config.canvas_length, device=confidence.device
    )
    if strategy == "confidence":
        return positions[torch.topk(confidence[positions], target_count, largest=False).indices]
    if strategy == "random":
        permutation = torch.randperm(
            positions.numel(), device=confidence.device, generator=generator
        )
        return positions[permutation[:target_count]]
    raise ValueError(f"Unsupported sampling strategy: {strategy}")


def denoise_canvas(
    model: MaskedLanguageModel,
    tokenizer: object,
    input_ids: torch.Tensor,
    attention_mask: torch.Tensor,
    config: DiffusionConfig,
    *,
    temperature: float,
    strategy: SamplingStrategy,
    generator: torch.Generator,
) -> Iterator[DenoisingSnapshot]:
    """Denoise a complete canvas in parallel and emit every reverse-process state.

    ``random`` reproduces the article's iterative re-masking. ``confidence`` retains
    the highest-confidence predictions, which is a common improved dLLM decoder.
    """
    if temperature <= 0:
        raise ValueError("temperature must be positive")

    mask_id = int(getattr(tokenizer, "mask_token_id"))
    blocked_ids = list(dict.fromkeys(getattr(tokenizer, "all_special_ids", [])))
    current = input_ids.clone()
    mask_positions = current.eq(mask_id)
    mask_positions[:, : config.prefix_length] = False
    model_seconds = 0.0

    for step in range(1, config.denoising_steps + 1):
        started_at = torch.cuda.Event(enable_timing=True) if current.is_cuda else None
        finished_at = torch.cuda.Event(enable_timing=True) if current.is_cuda else None
        if started_at is not None:
            started_at.record()
        with torch.inference_mode():
            logits = model(input_ids=current, attention_mask=attention_mask).logits
        if finished_at is not None:
            finished_at.record()
            finished_at.synchronize()
            model_seconds += started_at.elapsed_time(finished_at) / 1000

        logits = logits / temperature
        if blocked_ids:
            logits[..., blocked_ids] = -torch.inf
        probabilities = torch.softmax(logits, dim=-1)
        confidence = probabilities.max(dim=-1).values[0]

        active_positions = mask_positions[0]
        active_probabilities = probabilities[0, active_positions]
        sampled_tokens = torch.multinomial(
            active_probabilities, 1, generator=generator
        ).squeeze(-1)
        current[0, active_positions] = sampled_tokens

        target_count = config.masks_after_step(step)
        next_masked = _choose_remask_positions(
            confidence,
            config,
            target_count,
            strategy,
            generator,
        )
        mask_positions = torch.zeros_like(current, dtype=torch.bool)
        mask_positions[0, next_masked] = True
        current[mask_positions] = mask_id
        yield DenoisingSnapshot(
            step=step,
            input_ids=current.clone(),
            mask_positions=mask_positions.clone(),
            accepted_tokens=config.generated_length - target_count,
            total_tokens=config.generated_length,
            model_seconds=model_seconds,
        )