"""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="> 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=" 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 @classmethod 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") @classmethod 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)) @classmethod 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_)) @classmethod 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}