Spaces:
Running on Zero
Running on Zero
Download decoding/common.py from wang2226/beyond-tokens-decoding: direct link, hf CLI and curl.
- Browser
- Download file 11.2 kB
-
https://huggingface.co/spaces/wang2226/beyond-tokens-decoding/resolve/main/decoding/common.py
- Command line
-
hf download hf://spaces/wang2226/beyond-tokens-decoding/decoding/common.py
-
curl -L -o common.py https://huggingface.co/spaces/wang2226/beyond-tokens-decoding/resolve/main/decoding/common.py
11.2 kB
| """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:] | |
| 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 | |
| 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 | |
| # --------------------------------------------------------------------------- | |
| 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 | |