File size: 3,997 Bytes
3afc977
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Masked-diffusion objective and sampling (LLaDA-style, TRAINING.md).

Training: per example sample t~U(0,1); mask each hole position with prob t;
predict masked tokens with bidirectional attention; cross-entropy on masked
positions weighted by 1/t (a principled upper bound on the NLL). Context
(prefix/suffix) is never masked — conditioning is built in (DESIGN.md §2).

Sampling: start with the whole hole region masked, run N denoising steps,
unmasking the most confident positions on a linear schedule. N is the
test-time-compute knob that spends the iso-latency budget (EVALUATION.md).
"""

from __future__ import annotations

import torch
import torch.nn.functional as F


def loss(model, ids, region, block_id, attn_mask, tok):
    """Block-causal masked-diffusion loss with GLOBAL averaging (Nemotron): every
    masked token across the batch is weighted 1/t and the sum is divided by the
    total masked-token count, so examples contribute in proportion to how many
    noisy tokens they have (lower gradient variance than per-sequence averaging).
    ids,(B,T) clean; region,(B,T) generated positions; block_id,(B,T); attn,(B,T)."""
    B, T = ids.shape
    device = ids.device
    t = torch.rand(B, device=device).clamp_(0.02, 1.0)  # avoid 1/t blow-up

    r = torch.rand(B, T, device=device)
    selected = region & (r < t.unsqueeze(1))

    inp = torch.where(selected, torch.full_like(ids, tok.mask_id), ids)
    logits = model.forward_blocks(inp, attn_mask, region, block_id)

    ce = F.cross_entropy(
        logits.view(-1, logits.size(-1)), ids.view(-1), reduction="none"
    ).view(B, T)
    w = (1.0 / t).unsqueeze(1)  # per-sample weight, broadcast over T
    ce = ce * selected.float() * w
    denom = selected.sum().clamp(min=1)  # global masked-token count
    return ce.sum() / denom


@torch.no_grad()
def sample(model, ids, region, attn_mask, tok, steps: int, tile_size: int = 16):
    """Tile-based denoising WITH revision (ReMDM-style).

    Each global pass re-processes the whole region (bidirectional), proposes new
    values everywhere (revision, not monotonic fill), groups the region into
    tiles, and RE-MASKS the lowest-confidence tiles to reprocess them next pass.
    Coherence between tiles emerges over passes; more steps help instead of
    freezing early mistakes. Returns filled ids (B,T)."""
    ids = ids.clone()
    ids[region] = tok.mask_id

    for step in range(1, steps + 1):
        logits = model(ids, attn_mask)
        probs = F.softmax(logits, dim=-1)
        conf, pred = probs.max(dim=-1)  # (B,T)

        # Fraction of tiles to leave masked after this step (linear to 0).
        frac_masked = 1.0 - step / steps

        for b in range(ids.size(0)):
            idxr = region[b].nonzero(as_tuple=True)[0]  # contiguous region positions
            R = idxr.numel()
            if R == 0:
                continue
            # Revision: commit proposed values across the WHOLE region.
            ids[b, idxr] = pred[b, idxr]
            if step == steps or frac_masked <= 0:
                continue

            n_tiles = (R + tile_size - 1) // tile_size
            keep_masked = round(n_tiles * frac_masked)
            if keep_masked <= 0:
                continue
            # Aggregate confidence per tile; re-mask the weakest tiles.
            c = conf[b, idxr]
            tile_conf = []
            for t in range(n_tiles):
                seg = c[t * tile_size : (t + 1) * tile_size]
                tile_conf.append(seg.mean())
            tile_conf = torch.stack(tile_conf)
            weak = torch.argsort(tile_conf)[:keep_masked]  # lowest confidence
            for t in weak.tolist():
                lo = t * tile_size
                hi = min((t + 1) * tile_size, R)
                ids[b, idxr[lo:hi]] = tok.mask_id

    return ids


def decode_hole(tok, row_ids, region_row) -> str:
    """The predicted block string: the region positions in order."""
    return tok.decode(row_ids[region_row].tolist())