"""Autoregressive fill-in-the-middle baseline (EXPERIMENTS.md). Same backbone, causal mask. Trained left-to-right; loss only on the middle (hole + eos) span so the objective matches diffusion's hole-only loss. The FIM ordering ([pre] prefix [suf] suffix [mid] hole) lets the AR model condition on the suffix too — the fair counterpart to diffusion's bidirectional view. The decode is strictly sequential: one token per forward pass. That sequential cost is exactly what the iso-latency comparison weighs against diffusion's parallel refinement. """ from __future__ import annotations import torch import torch.nn.functional as F def loss(model, ids, loss_mask, attn_mask, tok): """Causal next-token CE on the mid span. ids,(B,T); loss_mask,(B,T) marks the mid tokens; attn_mask,(B,T) non-pad.""" logits = model(ids, attn_mask) # (B,T,V) # Predict token q from position q-1. pred = logits[:, :-1, :] target = ids[:, 1:] mask = loss_mask[:, 1:] ce = F.cross_entropy( pred.reshape(-1, pred.size(-1)), target.reshape(-1), reduction="none" ).view(target.shape) ce = ce * mask.float() return ce.sum() / mask.float().sum().clamp(min=1) @torch.no_grad() def sample(model, head_ids, tok, max_new: int, seq_len: int): """KV-cached greedy decode of the hole for one FIM prompt. head_ids: 1-D tensor [bos pre ...prefix... suf ...suffix... mid]. Returns the decoded hole string.""" out = model.generate(head_ids, max_new=max_new, eos_id=tok.eos_id) return tok.decode(out) def build_prompt(tok, prefix, suffix, taskcfg): """The FIM prompt head for inference: returns a list of ids, or None if the context doesn't fit. prefix/suffix are id lists (lua) or strings (char).""" pre = list(prefix) if tok.mode == "lua" else tok.encode(prefix) suf = list(suffix) if tok.mode == "lua" else tok.encode(suffix) head = [tok.bos_id, tok.pre_id] + pre + [tok.suf_id] + suf + [tok.mid_id] if len(head) + 1 >= taskcfg.seq_len: return None return head