TinyDecide / tinydecide.js
TheREZOR's picture
Python, Rust and ESP32 engines, shared conformance set, promo video
f2878d0 verified
Raw History Blame Contribute Delete
19.6 kB
// TinyDecide decision model in plain JavaScript (browser worker or Node). Mirrors the Python
// reference exactly: reflex/format.py (layout + masks), reflex/backbones.py (ELECTRA path),
// reflex/model.py (native path), reflex/heads.py (typed heads, temperatures, prototypes).
// verify.mjs checks it against torch outputs computed from the same int8 weights.
const K_STATE = 0, K_QTEXT = 1, K_ANS = 2, K_OPT = 3, K_LV = 4;
const TYPES = ["choice", "noul", "score", "span"];
const K_BUCKETS = [1, 2, 4, 8];
// ------------------------------------------------------------------ tokenizers
const RE_CTRL = /[\p{Cc}\p{Cf}]/u, RE_ZS = /\p{Zs}/u, RE_P = /\p{P}/u, RE_MN = /\p{Mn}/u;
function isCJK(cp) {
return (cp >= 0x4E00 && cp <= 0x9FFF) || (cp >= 0x3400 && cp <= 0x4DBF) || (cp >= 0x20000 && cp <= 0x2A6DF) ||
(cp >= 0x2A700 && cp <= 0x2B73F) || (cp >= 0x2B740 && cp <= 0x2B81F) || (cp >= 0x2B820 && cp <= 0x2CEAF) ||
(cp >= 0xF900 && cp <= 0xFAFF) || (cp >= 0x2F800 && cp <= 0x2FA1F);
}
function isPunct(ch) {
const c = ch.codePointAt(0);
if ((c >= 33 && c <= 47) || (c >= 58 && c <= 64) || (c >= 91 && c <= 96) || (c >= 123 && c <= 126)) return true;
return RE_P.test(ch);
}
export class WordPiece {
constructor(m) {
this.vocab = new Map();
m.vocab.forEach((p, i) => { if (p !== null) this.vocab.set(p, i); });
this.unk = this.vocab.get(m.unk);
this.prefix = m.prefix; this.maxChars = m.max_chars;
}
// -> {ids, offsets} with offsets in JS string (UTF-16) indices of the original text
encode(text) {
const chars = []; // normalised chars, each tagged with its source range
let u = 0;
for (const ch of text) {
const cp = ch.codePointAt(0), a = u, b = u + ch.length;
u = b;
if (cp === 0 || cp === 0xFFFD || (RE_CTRL.test(ch) && ch !== "\t" && ch !== "\n" && ch !== "\r")) continue;
if (ch === " " || ch === "\t" || ch === "\n" || ch === "\r" || RE_ZS.test(ch)) { chars.push([" ", a, b]); continue; }
if (isCJK(cp)) { chars.push([" ", a, a]); chars.push([ch, a, b]); chars.push([" ", b, b]); continue; }
for (const d of ch.normalize("NFD")) {
if (RE_MN.test(d)) continue;
for (const l of d.toLowerCase()) chars.push([l, a, b]);
}
}
const words = []; // arrays of [ch, a, b]
let cur = [];
const flush = () => { if (cur.length) words.push(cur); cur = []; };
for (const c of chars) {
if (c[0] === " " || /\s/u.test(c[0])) { flush(); continue; }
if (isPunct(c[0])) { flush(); words.push([c]); continue; }
cur.push(c);
}
flush();
const ids = [], offsets = [];
for (const w of words) {
if (w.length > this.maxChars) { ids.push(this.unk); offsets.push([w[0][1], w[w.length - 1][2]]); continue; }
const pieces = [];
let start = 0, bad = false;
while (start < w.length) {
let end = w.length, got = -1;
while (start < end) {
let s = "";
for (let i = start; i < end; i++) s += w[i][0];
if (start > 0) s = this.prefix + s;
const id = this.vocab.get(s);
if (id !== undefined) { got = id; break; }
end--;
}
if (got < 0) { bad = true; break; }
pieces.push([got, w[start][1], w[end - 1][2]]);
start = end;
}
if (bad) { ids.push(this.unk); offsets.push([w[0][1], w[w.length - 1][2]]); }
else for (const [id, a, b] of pieces) { ids.push(id); offsets.push([a, b]); }
}
return { ids, offsets };
}
}
// ------------------------------------------------------------------ weights
export class Weights {
constructor(meta, buf) {
this.t = {};
for (const e of meta.tensors) {
const n = e.shape.reduce((a, b) => a * b, 1);
let arr;
if (e.dtype === "f32") arr = new Float32Array(buf, e.offset, n).slice();
else if (e.dtype === "q4") { // device Q4_0: 16 nibble bytes + one bf16 scale per 32
const rows = e.shape[0], cols = e.shape[1], nb = cols / 32;
const nib = new Uint8Array(buf, e.offset, rows * nb * 16), sc = new Uint16Array(buf, e.scale_offset, rows * nb);
const f = new Float32Array(1), u = new Uint32Array(f.buffer);
arr = new Float32Array(n);
for (let r = 0; r < rows; r++) for (let b = 0; b < nb; b++) {
u[0] = sc[r * nb + b] << 16; const d = f[0];
const no = (r * nb + b) * 16, wo = r * cols + b * 32;
for (let k = 0; k < 16; k++) { const byte = nib[no + k]; arr[wo + k] = ((byte & 15) - 8) * d; arr[wo + k + 16] = ((byte >> 4) - 8) * d; }
}
} else {
const q = new Int8Array(buf, e.offset, n), s = new Float32Array(buf, e.scale_offset, e.shape[0]);
const cols = e.shape[1];
arr = new Float32Array(n);
for (let r = 0; r < e.shape[0]; r++) { const sc = s[r], o = r * cols; for (let c = 0; c < cols; c++) arr[o + c] = q[o + c] * sc; }
}
this.t[e.name] = { data: arr, shape: e.shape };
}
}
get(n) { const x = this.t[n]; if (!x) throw new Error("missing tensor " + n); return x; }
has(n) { return n in this.t; }
}
// y[t] = W x[t] + b ; W [out, in] row-major, x [T, in]
function linear(x, T, din, W, b, dout) {
const y = new Float32Array(T * dout);
for (let t = 0; t < T; t++) {
const xo = t * din, yo = t * dout;
for (let j = 0; j < dout; j++) {
const wo = j * din;
let s = b ? b[j] : 0;
for (let i = 0; i < din; i++) s += W[wo + i] * x[xo + i];
y[yo + j] = s;
}
}
return y;
}
function layernorm(x, T, d, w, b, eps) {
const y = new Float32Array(T * d);
for (let t = 0; t < T; t++) {
const o = t * d;
let m = 0; for (let i = 0; i < d; i++) m += x[o + i]; m /= d;
let v = 0; for (let i = 0; i < d; i++) { const z = x[o + i] - m; v += z * z; } v /= d;
const r = 1 / Math.sqrt(v + eps);
for (let i = 0; i < d; i++) y[o + i] = (x[o + i] - m) * r * w[i] + (b ? b[i] : 0);
}
return y;
}
function rmsnorm(x, T, d, w, eps = 1e-6) {
const y = new Float32Array(T * d);
for (let t = 0; t < T; t++) {
const o = t * d;
let s = 0; for (let i = 0; i < d; i++) s += x[o + i] * x[o + i];
const r = 1 / Math.sqrt(s / d + eps);
for (let i = 0; i < d; i++) y[o + i] = x[o + i] * r * w[i];
}
return y;
}
// erf, Abramowitz-Stegun 7.1.26 is 1.5e-7; use a higher-precision rational (W. J. Cody style via erfc series)
function erf(x) {
const s = Math.sign(x), a = Math.abs(x);
if (a < 0.5) { // Taylor series, fast convergence near 0
let sum = a, term = a, a2 = a * a;
for (let n = 1; n < 12; n++) { term *= -a2 / n; sum += term / (2 * n + 1); }
return s * 2 / Math.sqrt(Math.PI) * sum;
}
// continued fraction for erfc (Lentz), accurate to ~1e-12 for a >= 0.5
const t = 1 / (1 + 0.5 * a);
const y = t * Math.exp(-a * a - 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)))))))));
return s * (1 - y);
}
const gelu = (x) => 0.5 * x * (1 + erf(x / Math.SQRT2));
const geluTanh = (x) => 0.5 * x * (1 + Math.tanh(0.7978845608028654 * (x + 0.044715 * x * x * x)));
// ------------------------------------------------------------------ request layout
export function encodeRequest(tok, meta, req) {
const sp = meta.specials, F = meta.format;
const st = tok.encode(req.state);
const n = Math.min(st.ids.length, F.ts_max - 1);
const ids = [sp["<|state|>"], ...st.ids.slice(0, n)];
const pos = ids.map((_, i) => i), blk = ids.map(() => 0), kind = ids.map(() => K_STATE);
const stIdx = []; for (let i = 1; i < ids.length; i++) stIdx.push(i);
const stOff = st.offsets.slice(0, n);
const qs = [];
req.questions.forEach((q, qi) => {
const k = qi + 1, nOpt = (q.options || []).length;
if ((q.type === "choice" || q.type === "score") && (nOpt < 2 || nOpt > F.k_max))
throw new Error(`Question ${k} needs 2 to ${F.k_max} options, not ${nOpt}.`);
const bIds = [sp["<|" + q.type + "|>"]], bKind = [K_QTEXT];
const t = tok.encode(q.text).ids; bIds.push(...t); t.forEach(() => bKind.push(K_QTEXT));
const optLocal = [];
if (q.options && q.options.length) {
bIds.push(sp["<|sep|>"]); bKind.push(K_QTEXT);
const [mk, mkind] = q.type === "choice" ? ["<|o|>", K_OPT] : ["<|lv|>", K_LV];
for (const o of q.options) {
const oi = tok.encode(o).ids; bIds.push(...oi); oi.forEach(() => bKind.push(K_QTEXT));
optLocal.push(bIds.length); bIds.push(sp[mk]); bKind.push(mkind);
}
}
const ansLocal = bIds.length; bIds.push(sp["<|ans|>"]); bKind.push(K_ANS);
if (bIds.length > F.q_max) throw new Error(`Question ${k} is too long (${bIds.length} tokens, max ${F.q_max}).`);
const base = ids.length;
bIds.forEach((id, j) => { ids.push(id); pos.push(F.p_q + j); blk.push(k); kind.push(bKind[j]); });
qs.push({ type: q.type, ans: base + ansLocal, opt: optLocal.map((j) => base + j) });
});
return { ids, pos, blk, kind, stIdx, stOff, qs, truncated: st.ids.length > n };
}
function continues(kind, fusion) {
if (fusion === "all") return kind.map(() => true);
if (fusion === "markers") return kind.map((k) => k === K_STATE || k === K_ANS || k === K_OPT || k === K_LV);
return kind.map((k) => k === K_STATE || k === K_ANS);
}
// ------------------------------------------------------------------ model
export class TinyDecide {
constructor(meta, buf) {
this.meta = meta; this.cfg = meta.cfg; this.W = new Weights(meta, buf);
const tk = meta.tokenizer;
if (tk.kind !== "wordpiece") throw new Error("this engine build ships the WordPiece tokenizer only");
this.tok = new WordPiece(tk);
}
w(n) { return this.W.get(n).data; }
// attention over allowed keys; q,k,v [T, d]; returns [T, d]
attend(q, k, v, T, allowed, heads, scale) {
const d = this.cfg.d, dh = d / heads, out = new Float32Array(T * d);
const sc = new Float32Array(T);
for (let i = 0; i < T; i++) {
const keys = allowed[i];
if (!keys) continue;
for (let h = 0; h < heads; h++) {
const qo = i * d + h * dh;
let mx = -Infinity;
for (let a = 0; a < keys.length; a++) {
const ko = keys[a] * d + h * dh;
let s = 0; for (let c = 0; c < dh; c++) s += q[qo + c] * k[ko + c];
s *= scale; sc[a] = s; if (s > mx) mx = s;
}
let z = 0; for (let a = 0; a < keys.length; a++) { sc[a] = Math.exp(sc[a] - mx); z += sc[a]; }
const oo = i * d + h * dh;
for (let a = 0; a < keys.length; a++) {
const p = sc[a] / z, vo = keys[a] * d + h * dh;
for (let c = 0; c < dh; c++) out[oo + c] += p * v[vo + c];
}
}
}
return out;
}
allowedSets(enc, cont, fusionLayer) {
const T = enc.ids.length, blk = enc.blk;
const byBlk = new Map();
for (let j = 0; j < T; j++) {
if (fusionLayer && !cont[j]) continue;
if (!byBlk.has(blk[j])) byBlk.set(blk[j], []);
byBlk.get(blk[j]).push(j);
}
const state = byBlk.get(0) || [];
const sets = new Array(T);
for (let i = 0; i < T; i++) {
if (fusionLayer && !cont[i]) { sets[i] = [i]; continue; }
const own = byBlk.get(blk[i]) || [i];
sets[i] = (!fusionLayer || blk[i] === 0) ? own : state.concat(own);
}
return sets;
}
hidden(enc) {
const cfg = this.cfg, T = enc.ids.length, d = cfg.d, L_a = cfg.L_a || 0;
const cont = continues(enc.kind, cfg.fusion);
const setsC = this.allowedSets(enc, cont, false), setsF = this.allowedSets(enc, cont, true);
let x;
if (cfg.arch === "electra") {
const e = cfg.emb_dim, P = this.w("m.posemb.weight"), t0 = this.w("m.type0");
const raw = new Float32Array(T * e);
if (this.W.has("m.word.a.weight")) { // low-rank table: a[id] (rank r) times b (r -> e)
const A = this.w("m.word.a.weight"), B = this.w("m.word.b.weight"), r = this.W.get("m.word.a.weight").shape[1];
for (let t = 0; t < T; t++) for (let c = 0; c < e; c++) {
let s = 0; const ao = enc.ids[t] * r, bo = c * r;
for (let j = 0; j < r; j++) s += A[ao + j] * B[bo + j];
raw[t * e + c] = s + P[enc.pos[t] * e + c] + t0[c];
}
} else {
const W = this.w("m.word.weight");
for (let t = 0; t < T; t++) for (let c = 0; c < e; c++) raw[t * e + c] = W[enc.ids[t] * e + c] + P[enc.pos[t] * e + c] + t0[c];
}
const n = layernorm(raw, T, e, this.w("m.eln.weight"), this.w("m.eln.bias"), cfg.ln_eps);
x = linear(n, T, e, this.w("m.proj.weight"), this.w("m.proj.bias"), d);
for (let l = 0; l < cfg.layers; l++) {
const p = `m.blocks.${l}.`, sets = l < L_a ? setsC : setsF;
const q = linear(x, T, d, this.w(p + "q.weight"), this.w(p + "q.bias"), d);
const k = linear(x, T, d, this.w(p + "k.weight"), this.w(p + "k.bias"), d);
const v = linear(x, T, d, this.w(p + "v.weight"), this.w(p + "v.bias"), d);
const y = this.attend(q, k, v, T, sets, cfg.heads, 1 / Math.sqrt(d / cfg.heads));
const o = linear(y, T, d, this.w(p + "o.weight"), this.w(p + "o.bias"), d);
for (let i = 0; i < T * d; i++) o[i] += x[i];
x = layernorm(o, T, d, this.w(p + "ln1.weight"), this.w(p + "ln1.bias"), cfg.ln_eps);
const ffDim = this.W.get(p + "fc.weight").shape[0];
const f = linear(x, T, d, this.w(p + "fc.weight"), this.w(p + "fc.bias"), ffDim);
for (let i = 0; i < f.length; i++) f[i] = gelu(f[i]);
const f2 = linear(f, T, ffDim, this.w(p + "fc2.weight"), this.w(p + "fc2.bias"), d);
for (let i = 0; i < T * d; i++) f2[i] += x[i];
x = layernorm(f2, T, d, this.w(p + "ln2.weight"), this.w(p + "ln2.bias"), cfg.ln_eps);
}
return x;
}
throw new Error("unsupported architecture " + cfg.arch);
}
// protos: per question {vec: Float32Array[K*dh] | null, cnt: Int32Array[K]} or undefined
answer(req, protos) {
const t0 = (typeof performance !== "undefined" ? performance.now() : Date.now());
const enc = encodeRequest(this.tok, this.meta, req);
const x = this.hidden(enc);
const cfg = this.cfg, d = cfg.d, dh = cfg.dh_head, T = enc.ids.length;
const hq = layernorm(x, T, d, this.w("h.norm.weight"), this.w("h.norm.bias"), 1e-5);
const row = (i) => hq.subarray(i * d, (i + 1) * d);
const mv = (name, v, bias) => linear(v, 1, d, this.w(name), bias ? this.w(bias) : null, this.W.get(name).shape[0]);
const temp = this.meta.temp, beta = this.meta.beta, scale = this.w("h.scale");
const dot = (a, b) => { let s = 0; for (let i = 0; i < a.length; i++) s += a[i] * b[i]; return s; };
const norm = (a) => { const n = Math.sqrt(dot(a, a)) || 1e-12; return a.map((z) => z / n); };
const bucket = (c) => K_BUCKETS.reduce((b, e) => (c > e ? b + 1 : b), 0);
const out = [];
enc.qs.forEach((q, qi) => {
const hAns = row(q.ans), P = protos && protos[qi];
const hasP = this.W.has("h.P.weight");
const pv = hasP ? mv("h.P.weight", hAns) : null;
// centred-cosine prototype term: beta(k) * cos(p - c, mean_k - c)
const protoTerm = (i) => {
if (!P || !(P.cnt[i] > 0)) return 0;
if (hasP && P.center) {
const c = P.center, a = new Float32Array(dh), b = new Float32Array(dh);
for (let j = 0; j < dh; j++) { a[j] = pv[j] - c[j]; b[j] = P.vec[i * dh + j] - c[j]; }
const na = Math.sqrt(dot(a, a)) || 1e-12, nb = Math.sqrt(dot(b, b)) || 1e-12;
return beta[bucket(P.cnt[i])] * dot(a, b) / (na * nb);
}
return null; // older model without a prototype projection
};
if (q.type === "choice" || q.type === "score") {
const sel = q.type === "score" ? 1 : 0;
const qa = mv(`h.A.${sel}.weight`, hAns);
const s = Math.exp(scale[sel]);
const qv = norm(Array.from(qa));
const lam = P && P.lam != null ? P.lam : 1;
const zeroShot = [];
const logits = q.opt.map((oi, i) => {
const ov = mv(`h.O.${sel}.weight`, row(oi));
let z = s * dot(qa, ov) / Math.sqrt(dh);
zeroShot.push(z / temp[sel === 1 ? 2 : 0]);
const t = protoTerm(i);
if (t === null) z += (P && P.cnt[i] > 0) ? beta[bucket(P.cnt[i])] * dot(qv, norm(Array.from(P.vec.subarray(i * dh, (i + 1) * dh)))) : 0;
else z += lam * t;
return z;
});
const Tt = temp[sel === 1 ? 2 : 0];
const m = Math.max(...logits), ex = logits.map((z) => Math.exp((z - m) / Tt)), Z = ex.reduce((a, b) => a + b, 0);
const probs = ex.map((e) => e / Z);
const k = probs.length, H = -probs.reduce((a, p) => a + (p > 0 ? p * Math.log(p) : 0), 0);
const r = { type: q.type, probs, pick: probs.indexOf(Math.max(...probs)), confidence: 1 - H / Math.log(k),
qvec: hasP ? Array.from(pv) : qv, proj: hasP, z0: zeroShot };
if (q.type === "score") r.score = probs.reduce((a, p, i) => a + p * i / (k - 1), 0);
out.push(r);
} else if (q.type === "noul") {
let z = dot(this.w("h.noul.weight"), hAns) + this.w("h.noul.bias")[0];
const z0 = z / temp[1];
const nv = norm(Array.from(mv("h.noul_q.weight", hAns)));
if (P) {
const lam = P.lam != null ? P.lam : 1;
const term = [0, 1].map((i) => {
const t = protoTerm(i);
if (t !== null) return lam * t;
return P.cnt[i] > 0 ? beta[bucket(P.cnt[i])] * dot(nv, norm(Array.from(P.vec.subarray(i * dh, (i + 1) * dh)))) : 0;
});
z += term[1] - term[0];
}
out.push({ type: "noul", p: 1 / (1 + Math.exp(-z / temp[1])), qvec: hasP ? Array.from(pv) : nv, proj: hasP,
z0: [0, z0] });
} else {
const S = enc.stIdx.length;
const qS = mv("h.sq.weight", hAns);
const ks = enc.stIdx.map((i) => mv("h.sk.weight", row(i)));
const start = [dot(this.w("h.snull.weight"), hAns) + this.w("h.snull.bias")[0], ...ks.map((kk) => dot(qS, kk) / Math.sqrt(dh))];
const lsm = (a, Tt) => { const m = Math.max(...a); const l = Math.log(a.reduce((s, z) => s + Math.exp((z - m) / Tt), 0)); return a.map((z) => (z - m) / Tt - l); };
const ls = lsm(start, temp[3]);
let best = 0; for (let t = 1; t < S; t++) if (ls[1 + t] > ls[1 + best]) best = t;
const eq = mv("h.eq.weight", hAns), es = mv("h.es.weight", row(enc.stIdx[best] ?? 0));
const qe = eq.map((z, i) => z + es[i]);
const e = enc.stIdx.map((i, t) => (t >= best && t < best + this.meta.format.span_max) ? dot(qe, mv("h.ek.weight", row(i))) / Math.sqrt(dh) : -1e4);
const le = S ? lsm(e, temp[3]) : [];
let bestE = best; for (let t = 0; t < S; t++) if (le[t] > le[bestE]) bestE = t;
const pPresent = 1 - Math.exp(ls[0]);
let text = "", char = null;
if (S && enc.stOff[best] && enc.stOff[bestE]) { char = [enc.stOff[best][0], enc.stOff[bestE][1]]; text = req.state.slice(char[0], char[1]).trim(); }
out.push({ type: "span", p_present: pPresent, tok: [best, bestE], p_span: S ? Math.exp(ls[1 + best] + le[bestE]) : 0, text, char });
}
});
const t1 = (typeof performance !== "undefined" ? performance.now() : Date.now());
return { answers: out, tokens: { state: enc.stIdx.length + 1, questions: T - enc.stIdx.length - 1, total: T },
truncated: enc.truncated, ms: t1 - t0, ids: enc.ids };
}
}