TheREZOR's picture
Python, Rust and ESP32 engines, shared conformance set, promo video
f2878d0 verified
Raw History Blame Contribute Delete
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
@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}