Download python/tinydecide/engine.py from TheREZOR/TinyDecide: direct link, hf CLI and curl.
- Browser
- Download file 23.5 kB
-
https://huggingface.co/TheREZOR/TinyDecide/resolve/main/python/tinydecide/engine.py
- Command line
-
hf download hf://TheREZOR/TinyDecide/python/tinydecide/engine.py
-
curl -L -o engine.py https://huggingface.co/TheREZOR/TinyDecide/resolve/main/python/tinydecide/engine.py
23.5 kB
| """TinyDecide engine in Python + numpy. | |
| A line-by-line port of the JavaScript reference engine (tinydecide.js): the same tokenizer, request | |
| layout, block masks, ELECTRA encoder and typed heads. Weights are stored as float32; each operation | |
| accumulates in float64 and rounds its result to float32, as the JavaScript engine does. | |
| Offsets (`char` in span answers) are Python str indices (code points). The JavaScript engine uses | |
| UTF-16 indices, so the two differ only on text with characters outside the BMP; `text` is the same. | |
| """ | |
| from __future__ import annotations | |
| import json | |
| import math | |
| import time | |
| import unicodedata | |
| from pathlib import Path | |
| import numpy as np | |
| K_STATE, K_QTEXT, K_ANS, K_OPT, K_LV = range(5) | |
| TYPES = ("choice", "noul", "score", "span") | |
| K_BUCKETS = (1, 2, 4, 8) | |
| f32, f64 = np.float32, np.float64 | |
| # ------------------------------------------------------------------ tokenizer | |
| def _is_cjk(cp: int) -> bool: | |
| return ((0x4E00 <= cp <= 0x9FFF) or (0x3400 <= cp <= 0x4DBF) or (0x20000 <= cp <= 0x2A6DF) or | |
| (0x2A700 <= cp <= 0x2B73F) or (0x2B740 <= cp <= 0x2B81F) or (0x2B820 <= cp <= 0x2CEAF) or | |
| (0xF900 <= cp <= 0xFAFF) or (0x2F800 <= cp <= 0x2FA1F)) | |
| def _is_punct(ch: str) -> bool: | |
| c = ord(ch) | |
| if 33 <= c <= 47 or 58 <= c <= 64 or 91 <= c <= 96 or 123 <= c <= 126: | |
| return True | |
| return unicodedata.category(ch).startswith("P") | |
| # JavaScript's \s and String.prototype.trim(): WhiteSpace (incl. every Zs) and LineTerminator. | |
| _JS_SPACE = frozenset("\t\n\x0b\x0c\r " + | |
| "".join(chr(c) for c in range(0x2000, 0x200B))) | |
| def _is_js_space(ch: str) -> bool: | |
| return ch in _JS_SPACE or unicodedata.category(ch) == "Zs" | |
| def _js_trim(s: str) -> str: | |
| a, b = 0, len(s) | |
| while a < b and _is_js_space(s[a]): | |
| a += 1 | |
| while b > a and _is_js_space(s[b - 1]): | |
| b -= 1 | |
| return s[a:b] | |
| class WordPiece: | |
| """Lowercasing, accent-stripping WordPiece (BERT basic tokenizer + greedy longest match).""" | |
| def __init__(self, m: dict): | |
| self.vocab = {} | |
| for i, p in enumerate(m["vocab"]): | |
| if p is not None: | |
| self.vocab[p] = i | |
| self.unk = self.vocab[m["unk"]] | |
| self.prefix = m["prefix"] | |
| self.max_chars = m["max_chars"] | |
| def encode(self, text: str): | |
| """-> (ids, offsets); offsets are [start, end) str indices into `text`.""" | |
| chars = [] # normalised chars, each with its source range | |
| for i, ch in enumerate(text): | |
| cp, a, b = ord(ch), i, i + 1 | |
| cat = unicodedata.category(ch) | |
| if cp == 0 or cp == 0xFFFD or (cat in ("Cc", "Cf") and ch not in ("\t", "\n", "\r")): | |
| continue | |
| if ch in (" ", "\t", "\n", "\r") or cat == "Zs": | |
| chars.append((" ", a, b)) | |
| continue | |
| if _is_cjk(cp): | |
| chars.append((" ", a, a)) | |
| chars.append((ch, a, b)) | |
| chars.append((" ", b, b)) | |
| continue | |
| for d in unicodedata.normalize("NFD", ch): | |
| if unicodedata.category(d) == "Mn": | |
| continue | |
| for l in d.lower(): | |
| chars.append((l, a, b)) | |
| words, cur = [], [] | |
| for c in chars: | |
| if c[0] == " " or _is_js_space(c[0]): | |
| if cur: | |
| words.append(cur) | |
| cur = [] | |
| continue | |
| if _is_punct(c[0]): | |
| if cur: | |
| words.append(cur) | |
| cur = [] | |
| words.append([c]) | |
| continue | |
| cur.append(c) | |
| if cur: | |
| words.append(cur) | |
| ids, offsets = [], [] | |
| for w in words: | |
| if len(w) > self.max_chars: | |
| ids.append(self.unk) | |
| offsets.append([w[0][1], w[-1][2]]) | |
| continue | |
| pieces, start, bad = [], 0, False | |
| while start < len(w): | |
| end, got = len(w), -1 | |
| while start < end: | |
| s = "".join(c[0] for c in w[start:end]) | |
| if start > 0: | |
| s = self.prefix + s | |
| got = self.vocab.get(s, -1) | |
| if got >= 0: | |
| break | |
| end -= 1 | |
| if got < 0: | |
| bad = True | |
| break | |
| pieces.append((got, w[start][1], w[end - 1][2])) | |
| start = end | |
| if bad: | |
| ids.append(self.unk) | |
| offsets.append([w[0][1], w[-1][2]]) | |
| else: | |
| for pid, a, b in pieces: | |
| ids.append(pid) | |
| offsets.append([a, b]) | |
| return ids, offsets | |
| # ------------------------------------------------------------------ weights | |
| def load_tensors(meta: dict, buf) -> dict: | |
| """model.bin -> {name: float32 array}. Q4: 16 nibble bytes + one bf16 scale per 32 weights; | |
| int8: one float32 scale per row.""" | |
| t = {} | |
| for e in meta["tensors"]: | |
| shape = tuple(e["shape"]) | |
| n = int(np.prod(shape)) | |
| if e["dtype"] == "f32": | |
| arr = np.frombuffer(buf, dtype="<f4", count=n, offset=e["offset"]).astype(f32) | |
| elif e["dtype"] == "q4": | |
| rows, cols = shape | |
| nb = cols // 32 | |
| nib = np.frombuffer(buf, dtype=np.uint8, count=rows * nb * 16, offset=e["offset"]).reshape(rows, nb, 16) | |
| sc = np.frombuffer(buf, dtype="<u2", count=rows * nb, offset=e["scale_offset"]).astype(np.uint32) << 16 | |
| d = sc.view(np.float32).reshape(rows, nb, 1) | |
| lo = (nib & 15).astype(np.int8) - 8 | |
| hi = (nib >> 4).astype(np.int8) - 8 | |
| arr = (np.concatenate([lo, hi], axis=2).astype(f32) * d).reshape(rows, cols) | |
| else: | |
| rows = shape[0] | |
| q = np.frombuffer(buf, dtype=np.int8, count=n, offset=e["offset"]).reshape(rows, -1) | |
| s = np.frombuffer(buf, dtype="<f4", count=rows, offset=e["scale_offset"]) | |
| arr = q.astype(f32) * s[:, None].astype(f32) | |
| t[e["name"]] = np.ascontiguousarray(arr.reshape(shape)) | |
| return t | |
| # ------------------------------------------------------------------ ops (float64 math, float32 results) | |
| def _linear(x, W, b=None): | |
| y = x.astype(f64) @ W.astype(f64).T | |
| if b is not None: | |
| y += b | |
| return y.astype(f32) | |
| def _layernorm(x, w, b, eps): | |
| x = x.astype(f64) | |
| m = x.mean(axis=-1, keepdims=True) | |
| v = ((x - m) ** 2).mean(axis=-1, keepdims=True) | |
| r = 1.0 / np.sqrt(v + eps) | |
| y = (x - m) * r * w | |
| if b is not None: | |
| y = y + b | |
| return y.astype(f32) | |
| def _erf(x): | |
| """The JavaScript engine's erf: Taylor series below 0.5, Numerical Recipes erfc above.""" | |
| s, a = np.sign(x), np.abs(x) | |
| out = np.empty_like(a) | |
| small = a < 0.5 | |
| if small.any(): | |
| v = a[small] | |
| a2, term, acc = v * v, v.copy(), v.copy() | |
| for n in range(1, 12): | |
| term = term * (-a2 / n) | |
| acc = acc + term / (2 * n + 1) | |
| out[small] = 2 / math.sqrt(math.pi) * acc | |
| big = ~small | |
| if big.any(): | |
| v = a[big] | |
| t = 1 / (1 + 0.5 * v) | |
| y = t * np.exp(-v * v - 1.26551223 + t * (1.00002368 + t * (0.37409196 + t * (0.09678418 + t * (-0.18628806 + | |
| t * (0.27886807 + t * (-1.13520398 + t * (1.48851587 + t * (-0.82215223 + t * 0.17087277))))))))) | |
| out[big] = 1 - y | |
| return s * out | |
| def _gelu(x): | |
| x = x.astype(f64) | |
| return (0.5 * x * (1 + _erf(x / math.sqrt(2)))).astype(f32) | |
| def _lsm(a, T): | |
| a = np.asarray(a, dtype=f64) | |
| m = a.max() | |
| l = math.log(float(np.exp((a - m) / T).sum())) | |
| return (a - m) / T - l | |
| def _bucket(c): | |
| return sum(1 for e in K_BUCKETS if c > e) | |
| def _norm(a): | |
| a = np.asarray(a, dtype=f64) | |
| n = math.sqrt(float(a @ a)) or 1e-12 | |
| return a / n | |
| # ------------------------------------------------------------------ request layout | |
| def encode_request(tok: WordPiece, meta: dict, state: str, questions: list) -> dict: | |
| sp, F = meta["specials"], meta["format"] | |
| st_ids, st_off = tok.encode(state) | |
| n = min(len(st_ids), F["ts_max"] - 1) | |
| ids = [sp["<|state|>"]] + st_ids[:n] | |
| pos = list(range(len(ids))) | |
| blk = [0] * len(ids) | |
| kind = [K_STATE] * len(ids) | |
| st_idx = list(range(1, len(ids))) | |
| qs = [] | |
| for qi, q in enumerate(questions): | |
| k, typ = qi + 1, q["type"] | |
| if typ not in TYPES: | |
| raise ValueError(f"Question {k} has an unknown type {typ!r} (use one of {', '.join(TYPES)}).") | |
| n_opt = len(q.get("options") or []) | |
| if typ in ("choice", "score") and (n_opt < 2 or n_opt > F["k_max"]): | |
| raise ValueError(f"Question {k} needs 2 to {F['k_max']} options, not {n_opt}.") | |
| b_ids, b_kind = [sp[f"<|{typ}|>"]], [K_QTEXT] | |
| t = tok.encode(q["text"])[0] | |
| b_ids += t | |
| b_kind += [K_QTEXT] * len(t) | |
| opt_local = [] | |
| options = q.get("options") or [] | |
| if options: | |
| b_ids.append(sp["<|sep|>"]) | |
| b_kind.append(K_QTEXT) | |
| mk, mkind = ("<|o|>", K_OPT) if typ == "choice" else ("<|lv|>", K_LV) | |
| for o in options: | |
| oi = tok.encode(o)[0] | |
| b_ids += oi | |
| b_kind += [K_QTEXT] * len(oi) | |
| opt_local.append(len(b_ids)) | |
| b_ids.append(sp[mk]) | |
| b_kind.append(mkind) | |
| ans_local = len(b_ids) | |
| b_ids.append(sp["<|ans|>"]) | |
| b_kind.append(K_ANS) | |
| if len(b_ids) > F["q_max"]: | |
| raise ValueError(f"Question {k} is too long ({len(b_ids)} tokens, max {F['q_max']}).") | |
| base = len(ids) | |
| for j, (i_, kd) in enumerate(zip(b_ids, b_kind)): | |
| ids.append(i_) | |
| pos.append(F["p_q"] + j) | |
| blk.append(k) | |
| kind.append(kd) | |
| qs.append({"type": typ, "ans": base + ans_local, "opt": [base + j for j in opt_local]}) | |
| return {"ids": ids, "pos": pos, "blk": blk, "kind": kind, "st_idx": st_idx, "st_off": st_off[:n], | |
| "qs": qs, "truncated": len(st_ids) > n} | |
| def _continues(kind, fusion): | |
| if fusion == "all": | |
| return [True] * len(kind) | |
| if fusion == "markers": | |
| return [k in (K_STATE, K_ANS, K_OPT, K_LV) for k in kind] | |
| return [k in (K_STATE, K_ANS) for k in kind] | |
| def _allowed(blk, cont, fusion_layer): | |
| """Boolean [T, T] mask: row i may attend to column j (TinyDecide.allowedSets).""" | |
| T = len(blk) | |
| blk = np.asarray(blk) | |
| cont = np.asarray(cont, dtype=bool) | |
| live = cont if fusion_layer else np.ones(T, dtype=bool) | |
| same = (blk[:, None] == blk[None, :]) & live[None, :] | |
| if fusion_layer: | |
| state = (blk == 0) & live | |
| m = same | ((blk[:, None] != 0) & state[None, :]) | |
| dead = ~cont | |
| m[dead, :] = False | |
| m[dead, dead] = True | |
| else: | |
| m = same | |
| # a token whose block has no live member falls back to itself | |
| empty = ~m.any(axis=1) | |
| m[empty, empty] = True | |
| return m | |
| # ------------------------------------------------------------------ model | |
| class TinyDecide: | |
| def __init__(self, meta: dict, buf): | |
| self.meta = meta | |
| self.cfg = meta["cfg"] | |
| self.W = load_tensors(meta, buf) | |
| tk = meta["tokenizer"] | |
| if tk["kind"] != "wordpiece": | |
| raise ValueError("this engine build ships the WordPiece tokenizer only") | |
| self.tok = WordPiece(tk) | |
| # -------------------------------------------------------------- loading | |
| def load(cls, path=".") -> "TinyDecide": | |
| """Load meta.json + model.bin from a folder (a downloaded TheREZOR/TinyDecide repo, or its full/).""" | |
| p = Path(path) | |
| return cls._from_files(p / "meta.json", p / "model.bin") | |
| def from_bytes(cls, meta, bin_bytes) -> "TinyDecide": | |
| """meta: a dict or the JSON text/bytes of meta.json; bin_bytes: the contents of model.bin.""" | |
| if isinstance(meta, (bytes, bytearray, str)): | |
| meta = json.loads(meta) | |
| return cls(meta, bytes(bin_bytes)) | |
| def from_pretrained(cls, repo_id: str = "TheREZOR/TinyDecide", subfolder: str | None = None, | |
| revision: str | None = None, cache_dir: str | None = None) -> "TinyDecide": | |
| """Download (or reuse the cached) model from the Hugging Face Hub. subfolder="full" picks the 13.8M build.""" | |
| try: | |
| from huggingface_hub import hf_hub_download | |
| except ImportError: | |
| raise ImportError("from_pretrained needs huggingface_hub: pip install huggingface_hub " | |
| "(or download the repo and use TinyDecide.load(folder))") from None | |
| kw = dict(subfolder=subfolder, revision=revision, cache_dir=cache_dir) | |
| meta = hf_hub_download(repo_id, "meta.json", **kw) | |
| bin_ = hf_hub_download(repo_id, "model.bin", **kw) | |
| return cls._from_files(Path(meta), Path(bin_)) | |
| def _from_files(cls, meta_path: Path, bin_path: Path) -> "TinyDecide": | |
| meta = json.loads(Path(meta_path).read_text(encoding="utf-8")) | |
| return cls(meta, Path(bin_path).read_bytes()) | |
| def w(self, name): | |
| x = self.W.get(name) | |
| if x is None: | |
| raise KeyError("missing tensor " + name) | |
| return x | |
| def has(self, name): | |
| return name in self.W | |
| # -------------------------------------------------------------- encoder | |
| def _attend(self, q, k, v, mask, heads, scale): | |
| T, d = q.shape | |
| dh = d // heads | |
| q = q.astype(f64).reshape(T, heads, dh).transpose(1, 0, 2) | |
| k = k.astype(f64).reshape(T, heads, dh).transpose(1, 0, 2) | |
| v = v.astype(f64).reshape(T, heads, dh).transpose(1, 0, 2) | |
| s = (q @ k.transpose(0, 2, 1)) * scale # [H, T, T] | |
| s = np.where(mask[None], s, -np.inf) | |
| mx = s.max(axis=2, keepdims=True) | |
| e = np.exp(s.astype(f32).astype(f64) - mx).astype(f32).astype(f64) # scores kept as float32, as in JS | |
| p = e / e.sum(axis=2, keepdims=True) | |
| out = (p @ v).transpose(1, 0, 2).reshape(T, d) | |
| return out.astype(f32) | |
| def hidden(self, enc): | |
| cfg = self.cfg | |
| ids, pos = np.asarray(enc["ids"]), np.asarray(enc["pos"]) | |
| T, d, L_a = len(ids), cfg["d"], cfg.get("L_a") or 0 | |
| cont = _continues(enc["kind"], cfg["fusion"]) | |
| mask_c, mask_f = _allowed(enc["blk"], cont, False), _allowed(enc["blk"], cont, True) | |
| if cfg["arch"] != "electra": | |
| raise ValueError("unsupported architecture " + str(cfg["arch"])) | |
| P, t0 = self.w("m.posemb.weight"), self.w("m.type0") | |
| if self.has("m.word.a.weight"): # low-rank table: a[id] (rank r) times b (r -> e) | |
| A, B = self.w("m.word.a.weight"), self.w("m.word.b.weight") | |
| word = A[ids].astype(f64) @ B.astype(f64).T | |
| else: | |
| word = self.w("m.word.weight")[ids].astype(f64) | |
| raw = (word + P[pos].astype(f64) + t0.astype(f64)).astype(f32) | |
| n = _layernorm(raw, self.w("m.eln.weight"), self.w("m.eln.bias"), cfg["ln_eps"]) | |
| x = _linear(n, self.w("m.proj.weight"), self.w("m.proj.bias")) | |
| scale = 1 / math.sqrt(d / cfg["heads"]) | |
| for l in range(cfg["layers"]): | |
| p = f"m.blocks.{l}." | |
| mask = mask_c if l < L_a else mask_f | |
| q = _linear(x, self.w(p + "q.weight"), self.w(p + "q.bias")) | |
| k = _linear(x, self.w(p + "k.weight"), self.w(p + "k.bias")) | |
| v = _linear(x, self.w(p + "v.weight"), self.w(p + "v.bias")) | |
| y = self._attend(q, k, v, mask, cfg["heads"], scale) | |
| o = _linear(y, self.w(p + "o.weight"), self.w(p + "o.bias")) + x | |
| x = _layernorm(o, self.w(p + "ln1.weight"), self.w(p + "ln1.bias"), cfg["ln_eps"]) | |
| f = _gelu(_linear(x, self.w(p + "fc.weight"), self.w(p + "fc.bias"))) | |
| f2 = _linear(f, self.w(p + "fc2.weight"), self.w(p + "fc2.bias")) + x | |
| x = _layernorm(f2, self.w(p + "ln2.weight"), self.w(p + "ln2.bias"), cfg["ln_eps"]) | |
| return x | |
| # -------------------------------------------------------------- answer | |
| def answer(self, state: str, questions: list, protos: list | None = None) -> dict: | |
| """Answer every question about `state` in one pass. | |
| questions: [{"type": "choice"|"noul"|"score"|"span", "text": ..., "options": [...]}, ...] | |
| protos: optional, one entry per question (None or {"vec", "cnt", "center", "lam"}), | |
| see tinydecide.corrections.make_protos. | |
| """ | |
| t_start = time.perf_counter() | |
| enc = encode_request(self.tok, self.meta, state, questions) | |
| x = self.hidden(enc) | |
| cfg, meta = self.cfg, self.meta | |
| dh, T = cfg["dh_head"], len(enc["ids"]) | |
| hq = _layernorm(x, self.w("h.norm.weight"), self.w("h.norm.bias"), 1e-5) | |
| temp, beta, scale = meta["temp"], meta["beta"], self.w("h.scale") | |
| has_p = self.has("h.P.weight") | |
| sq_dh = math.sqrt(dh) | |
| def mv(name, v, bias=None): | |
| return _linear(v[None, :], self.w(name), self.w(bias) if bias else None)[0] | |
| def dot(a, b): | |
| return float(np.asarray(a, dtype=f64) @ np.asarray(b, dtype=f64)) | |
| out = [] | |
| for qi, q in enumerate(enc["qs"]): | |
| h_ans = hq[q["ans"]] | |
| P = _protos(protos[qi] if protos and qi < len(protos) else None) | |
| pv = mv("h.P.weight", h_ans) if has_p else None | |
| def proto_term(i): | |
| # centred-cosine prototype term: beta(k) * cos(p - c, mean_k - c) | |
| if P is None or not (P["cnt"][i] > 0): | |
| return 0.0 | |
| if has_p and P["center"] is not None: | |
| c = P["center"] | |
| a = (pv - c).astype(f32) | |
| b = (P["vec"][i * dh:(i + 1) * dh] - c).astype(f32) | |
| na = math.sqrt(dot(a, a)) or 1e-12 | |
| nb = math.sqrt(dot(b, b)) or 1e-12 | |
| return beta[_bucket(P["cnt"][i])] * dot(a, b) / (na * nb) | |
| return None # older model without a prototype projection | |
| def legacy(qv, i): | |
| return beta[_bucket(P["cnt"][i])] * dot(qv, _norm(P["vec"][i * dh:(i + 1) * dh])) | |
| if q["type"] in ("choice", "score"): | |
| sel = 1 if q["type"] == "score" else 0 | |
| qa = mv(f"h.A.{sel}.weight", h_ans) | |
| s = math.exp(float(scale[sel])) | |
| qv = _norm(qa) | |
| lam = P["lam"] if P is not None else 1 | |
| Tt = temp[2 if sel == 1 else 0] | |
| O = self.w(f"h.O.{sel}.weight") | |
| ovs = _linear(hq[q["opt"]], O) if q["opt"] else np.zeros((0, dh), f32) | |
| logits, zero_shot = [], [] | |
| for i in range(len(q["opt"])): | |
| z = s * dot(qa, ovs[i]) / sq_dh | |
| zero_shot.append(z / Tt) | |
| t = proto_term(i) | |
| if t is None: | |
| z += legacy(qv, i) if (P is not None and P["cnt"][i] > 0) else 0 | |
| else: | |
| z += lam * t | |
| logits.append(z) | |
| lg = np.asarray(logits, dtype=f64) | |
| ex = np.exp((lg - lg.max()) / Tt) | |
| probs = ex / ex.sum() | |
| k = len(probs) | |
| H = -sum(p * math.log(p) if p > 0 else 0.0 for p in probs) | |
| r = {"type": q["type"], "probs": probs.tolist(), "pick": int(np.argmax(probs)), | |
| "confidence": 1 - H / math.log(k), "qvec": (pv if has_p else qv).astype(f64).tolist(), | |
| "proj": has_p, "z0": zero_shot} | |
| if q["type"] == "score": | |
| r["score"] = sum(p * i / (k - 1) for i, p in enumerate(probs.tolist())) | |
| out.append(r) | |
| elif q["type"] == "noul": | |
| z = dot(self.w("h.noul.weight").reshape(-1), h_ans) + float(self.w("h.noul.bias")[0]) | |
| z0 = z / temp[1] | |
| nv = _norm(mv("h.noul_q.weight", h_ans)) | |
| if P is not None: | |
| lam = P["lam"] | |
| term = [] | |
| for i in (0, 1): | |
| t = proto_term(i) | |
| if t is not None: | |
| term.append(lam * t) | |
| else: | |
| term.append(legacy(nv, i) if P["cnt"][i] > 0 else 0.0) | |
| z += term[1] - term[0] | |
| out.append({"type": "noul", "p": 1 / (1 + math.exp(-z / temp[1])), | |
| "qvec": (pv if has_p else nv).astype(f64).tolist(), "proj": has_p, "z0": [0, z0]}) | |
| else: | |
| st_idx = enc["st_idx"] | |
| S = len(st_idx) | |
| qS = mv("h.sq.weight", h_ans) | |
| start = [dot(self.w("h.snull.weight").reshape(-1), h_ans) + float(self.w("h.snull.bias")[0])] | |
| if S: | |
| ks = _linear(hq[st_idx], self.w("h.sk.weight")) | |
| start += ((ks.astype(f64) @ qS.astype(f64)) / sq_dh).tolist() | |
| ls = _lsm(start, temp[3]) | |
| best = 0 | |
| for t in range(1, S): | |
| if ls[1 + t] > ls[1 + best]: | |
| best = t | |
| eq = mv("h.eq.weight", h_ans) | |
| es = mv("h.es.weight", hq[st_idx[best] if S else 0]) | |
| qe = (eq + es).astype(f32) | |
| span_max = meta["format"]["span_max"] | |
| e = [-1e4] * S | |
| win = list(range(best, min(S, best + span_max))) | |
| if win: | |
| ek = _linear(hq[[st_idx[t] for t in win]], self.w("h.ek.weight")) | |
| vals = (ek.astype(f64) @ qe.astype(f64)) / sq_dh | |
| for j, t in enumerate(win): | |
| e[t] = float(vals[j]) | |
| le = _lsm(e, temp[3]) if S else [] | |
| best_e = best | |
| for t in range(S): | |
| if le[t] > le[best_e]: | |
| best_e = t | |
| p_present = 1 - math.exp(ls[0]) | |
| text, char = "", None | |
| st_off = enc["st_off"] | |
| if S and best < len(st_off) and best_e < len(st_off): | |
| char = [st_off[best][0], st_off[best_e][1]] | |
| text = _js_trim(state[char[0]:char[1]]) | |
| out.append({"type": "span", "p_present": p_present, "tok": [best, best_e], | |
| "p_span": math.exp(ls[1 + best] + le[best_e]) if S else 0, "text": text, "char": char}) | |
| ms = (time.perf_counter() - t_start) * 1000 | |
| n_state = len(enc["st_idx"]) + 1 | |
| return {"answers": out, "tokens": {"state": n_state, "questions": T - n_state, "total": T}, | |
| "truncated": enc["truncated"], "ms": ms, "ids": enc["ids"]} | |
| def _protos(p): | |
| """Normalise one protos entry to float32 arrays (like the JS worker's Float32Array copies).""" | |
| if p is None: | |
| return None | |
| center = p.get("center") | |
| lam = p.get("lam") | |
| return {"vec": np.asarray(p["vec"], dtype=f32), "cnt": [int(c) for c in p["cnt"]], | |
| "center": None if center is None else np.asarray(center, dtype=f32), | |
| "lam": 1 if lam is None else lam} | |