echo-1 / ml /diffusion.py
lupodevelop's picture
echo-1 Stage 0 explainer: diffusion vs autoregressive, execution-verified
3afc977 verified
Raw
History Blame Contribute Delete
4 kB
"""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())