"""Shared decoding primitives. Everything that touches the KV cache lives here, so that transformers API changes stay in one place. All contrast/guidance math runs in float32 log-space. """ from __future__ import annotations import math import os import time from typing import Iterable import torch import torch.nn.functional as F from transformers import DynamicCache TOPK = 8 # entries kept per distribution in traces ROUND = 4 # decimals kept for probabilities in traces # --------------------------------------------------------------------------- # Timing # --------------------------------------------------------------------------- def sync(device: torch.device) -> None: if device.type == "cuda": torch.cuda.synchronize(device) elif device.type == "mps": torch.mps.synchronize() def now(device: torch.device) -> float: sync(device) return time.perf_counter() # --------------------------------------------------------------------------- # Token-tracked KV cache # --------------------------------------------------------------------------- def lcp(a: list[int], b: list[int]) -> int: """Length of the longest common prefix of two token lists.""" n = min(len(a), len(b)) i = 0 while i < n and a[i] == b[i]: i += 1 return i class KV: """A model plus a DynamicCache that remembers which tokens it holds. Invariant: ``cache.get_seq_length() == len(self.toks)``. The cache is created without a config, so every layer is a plain, croppable ``DynamicLayer``. For Gemma 3 this means sliding-window layers keep full KV; the window itself is still enforced by the attention mask. """ def __init__(self, model): self.model = model self.device = model.device self.cache = DynamicCache() self.toks: list[int] = [] # (tokens fed, seconds) per forward pass self.log: list[tuple[int, float]] = [] def crop_to(self, n: int) -> None: extra = len(self.toks) - n if extra > 0: # Negative counts remove tokens; the positive form is deprecated in 5.17. self.cache.crop(-extra) del self.toks[n:] @torch.no_grad() def run(self, seq: list[int], keep: int = 1, hidden: bool = False): """Make the cache cover ``seq``; return the output for its last ``keep`` positions. Crops back to the longest common prefix of the cached tokens and ``seq``, re-feeding at least ``keep`` tokens so their logits exist, then feeds the rest. This single rule covers speculative rollback, prompt lookup, Jacobi blocks and caches kept in lockstep. """ if keep < 1 or keep > len(seq): raise ValueError(f"keep={keep} must be in [1, {len(seq)}]") n = min(lcp(self.toks, seq), len(seq) - keep) self.crop_to(n) new = seq[n:] x = torch.tensor([new], device=self.device) t0 = now(self.device) out = self.model( input_ids=x, past_key_values=self.cache, use_cache=True, logits_to_keep=keep, output_hidden_states=hidden, ) t1 = now(self.device) self.toks.extend(new) self.log.append((len(new), t1 - t0)) return out @property def n_forward(self) -> int: return len(self.log) def decode_forward_mean(self) -> float: """Mean seconds of the non-prefill forwards (the first forward is the prefill).""" rest = [s for _, s in self.log[1:]] return sum(rest) / len(rest) if rest else float("nan") def total_time(self) -> float: return sum(s for _, s in self.log) # --------------------------------------------------------------------------- # Logit utilities # --------------------------------------------------------------------------- def log_probs(logits: torch.Tensor, temperature: float = 1.0) -> torch.Tensor: return F.log_softmax(logits.float() / max(float(temperature), 1e-5), dim=-1) def apc_mask(logp: torch.Tensor, beta: float) -> torch.Tensor: """Adaptive plausibility constraint (Li et al. 2023): keep tokens with p >= beta * max p.""" if beta <= 0: return torch.ones_like(logp, dtype=torch.bool) return logp >= logp.max(dim=-1, keepdim=True).values + math.log(beta) def jsd(p_log: torch.Tensor, q_log: torch.Tensor) -> torch.Tensor: """Jensen-Shannon divergence between one distribution [V] and J distributions [J, V].""" p_log = p_log.unsqueeze(0).expand_as(q_log) p, q = p_log.exp(), q_log.exp() m_log = (0.5 * (p + q)).clamp_min(1e-30).log() kl_pm = torch.where(p > 0, p * (p_log - m_log), torch.zeros_like(p)).sum(-1) kl_qm = torch.where(q > 0, q * (q_log - m_log), torch.zeros_like(q)).sum(-1) return (0.5 * (kl_pm + kl_qm)).clamp_min(0.0) def rep_penalty(scores: torch.Tensor, seen: Iterable[int], penalty: float) -> torch.Tensor: """HF-style repetition penalty over every token already in the context.""" if penalty == 1.0: return scores ids = torch.tensor(sorted(set(seen)), device=scores.device, dtype=torch.long) if ids.numel() == 0: return scores scores = scores.clone() s = scores.index_select(-1, ids) s = torch.where(s < 0, s * penalty, s / penalty) scores.index_copy_(-1, ids, s) return scores def make_generator(seed: int) -> torch.Generator: """CPU generator, so sampled runs are reproducible across devices.""" return torch.Generator(device="cpu").manual_seed(int(seed)) def sample(probs: torch.Tensor, gen: torch.Generator) -> int: return int(torch.multinomial(probs.float().cpu(), 1, generator=gen)) def pick(scores: torch.Tensor, greedy: bool, temperature: float, gen: torch.Generator) -> int: """Argmax, or sample from softmax(scores / T). Scores may contain -inf.""" if greedy: return int(scores.argmax()) probs = torch.softmax(scores.float() / max(float(temperature), 1e-5), dim=-1) return sample(probs, gen) def rank_of(scores: torch.Tensor, token: int) -> int: """0-based rank of ``token`` under ``scores`` (0 = top-1).""" return int((scores > scores[token]).sum()) # --------------------------------------------------------------------------- # Early exit (logit lens) # --------------------------------------------------------------------------- def final_norm(model): base = getattr(model, "model", None) norm = getattr(base, "norm", None) if base is not None else None if norm is None: norm = model.get_decoder().norm return norm def early_exit_logits(model, h: torch.Tensor, apply_norm: bool = True) -> torch.Tensor: """Project pre-norm residual states [..., d] to vocabulary logits (float32).""" if apply_norm: h = final_norm(model)(h) z = model.get_output_embeddings()(h).float() cfg = model.config.get_text_config() cap = getattr(cfg, "final_logit_softcapping", None) if cap: z = cap * torch.tanh(z / cap) scale = getattr(cfg, "logits_scaling", None) if scale: z = z / scale return z # --------------------------------------------------------------------------- # Trace helpers # --------------------------------------------------------------------------- def topk_entries(probs: torch.Tensor, disp, k: int = TOPK) -> list[list]: """[[id, display string, p], ...] for the k most likely tokens of a probability vector.""" k = min(k, probs.shape[-1]) v, i = torch.topk(probs.float(), k) # masked tokens (probability exactly 0) carry no information; always keep the top entry pairs = [(t, x) for j, (t, x) in enumerate(zip(i.tolist(), v.tolist())) if x > 0 or j == 0] return [[t, disp(t), round(x, ROUND)] for t, x in pairs] def token_pieces(tok, ids: list[int]) -> list[str]: """Display text for each token, via incremental decoding (handles split UTF-8).""" pieces: list[str] = [] prev = "" for i in range(len(ids)): cur = tok.decode(ids[: i + 1], skip_special_tokens=False) cp = len(os.path.commonprefix([prev, cur])) if cp < len(prev): # A replacement char from an incomplete byte sequence was resolved; trim it # from earlier pieces so the concatenation stays equal to the decoded text. drop = len(prev) - cp j = len(pieces) - 1 while drop > 0 and j >= 0: cut = min(drop, len(pieces[j])) pieces[j] = pieces[j][: len(pieces[j]) - cut] drop -= cut j -= 1 pieces.append(cur[cp:]) prev = cur return pieces def distinct_n(ids: list[int], n: int = 2) -> float: grams = [tuple(ids[i : i + n]) for i in range(len(ids) - n + 1)] return round(len(set(grams)) / len(grams), ROUND) if grams else 1.0 def first_divergence(a: list[int], b: list[int]) -> int | None: n = lcp(a, b) return None if n == min(len(a), len(b)) else n def make_run(label: str, tok, ids: list[int], finish: str, **extra) -> dict: """A Run record shared by every renderer.""" return { "label": label, "ids": list(ids), "toks": token_pieces(tok, list(ids)), "text": tok.decode(ids, skip_special_tokens=True), "finish": finish, **extra, } def forward_stats(**kvs: KV) -> dict: """Forward counts and timings for named KV caches.""" return { "n_forward": {name: kv.n_forward for name, kv in kvs.items()}, "time_ms": { **{f"{name}_total": round(1000 * kv.total_time(), 2) for name, kv in kvs.items()}, **{f"{name}_fwd_mean": round(1000 * kv.decode_forward_mean(), 3) for name, kv in kvs.items()}, }, } # --------------------------------------------------------------------------- # Autoregressive baseline # --------------------------------------------------------------------------- @torch.no_grad() def ar_decode( model, tok, prompt_ids: list[int], max_new: int, stop_ids: set[int], *, disp, greedy: bool = True, temperature: float = 1.0, rep: float = 1.0, gen: torch.Generator | None = None, label: str = "Autoregressive", record: bool = True, ) -> dict: """Plain token-by-token decoding with a KV cache; the baseline for every method.""" kv = KV(model) seq = list(prompt_ids) out: list[int] = [] steps: list[dict] = [] finish = "length" t0 = now(kv.device) for i in range(max_new): z = kv.run(seq).logits[0, -1].float() z = rep_penalty(z, seq, rep) lp = log_probs(z, 1.0 if greedy else temperature) t = pick(lp, greedy, 1.0, gen) if record: probs = lp.exp() steps.append({"i": i, "id": t, "p": round(float(probs[t]), ROUND), "top": {"base": topk_entries(probs, disp)}}) seq.append(t) out.append(t) if t in stop_ids: finish = "eos" break wall = now(kv.device) - t0 run = make_run(label, tok, out, finish, steps=steps, **forward_stats(main=kv)) run["time_ms"]["wall"] = round(1000 * wall, 2) return run