TinyDecide / python /tinydecide /corrections.py
TheREZOR's picture
Python, Rust and ESP32 engines, shared conformance set, promo video
f2878d0 verified
Raw History Blame Contribute Delete
3.17 kB
"""Corrections: turn stored examples into the `protos` argument of TinyDecide.answer().
Same recipe as the playground (corrections.js): one prototype per option (the mean of its examples),
a centre (the mean of every vector seen for this question, or of the examples when none is given),
and a trust weight `lam` chosen by leave-one-out over the examples themselves.
An example is what answer() returned for a message, plus the option a person said was right:
{"v": answer["qvec"], "z": answer["z0"]} stored in lists[k] for the correct option k.
For a noul question, k = 1 means "true" and k = 0 means "false".
"""
from __future__ import annotations
import math
import numpy as np
K_BUCKETS = (1, 2, 4, 8)
LAMBDAS = (0, 0.25, 0.5, 1, 2)
def bucket_k(c: int) -> int:
return sum(1 for e in K_BUCKETS if c > e)
def mean_vec(lst, dh: int) -> list:
m = [0.0] * dh
n = len(lst)
for e in lst:
v = e["v"]
for i in range(dh):
m[i] += v[i] / n
return m
def _cos_c(a, b, c) -> float:
a, b, c = (np.asarray(x, dtype=np.float64) for x in (a, b, c))
x, y = a - c, b - c
d, na, nb = float(x @ y), float(x @ x), float(y @ y)
return d / (math.sqrt(na * nb) or 1e-12)
def _term_for(v, lists, c, beta):
return [beta[bucket_k(len(l))] * _cos_c(v, mean_vec(l, len(v)), c) if l else 0.0 for l in lists]
def lambda_for(type_: str, lists, c, beta) -> float:
"""How much to trust this question's corrections: leave-one-out log-likelihood over the examples."""
all_ = [(k, j, e) for k, l in enumerate(lists) for j, e in enumerate(l)]
if len(all_) < 2:
return 0.25
best, best_ll = 0, -math.inf
for lam in LAMBDAS:
ll = 0.0
for k, j, e in all_:
rest = [[x for jj, x in enumerate(l) if jj != j] if kk == k else l for kk, l in enumerate(lists)]
t = _term_for(e["v"], rest, c, beta)
if type_ == "noul":
z = [lam * t[0], e["z"][1] + lam * t[1]]
else:
z = [x + lam * t[i] for i, x in enumerate(e["z"])]
m = max(z)
lse = m + math.log(sum(math.exp(x - m) for x in z))
ll += z[k] - lse
if ll > best_ll + 1e-9:
best, best_ll = lam, ll
return best
def make_protos(type_: str, lists, beta, center=None):
"""lists: one list of examples per option (2 for noul). beta: model.meta["beta"].
center: optional mean qvec over messages asked with this question.
Returns {"vec", "cnt", "center", "lam"} or None when there are no examples (span takes none)."""
if type_ == "span" or not any(len(l) for l in lists):
return None
dh = len(next(l for l in lists if l)[0]["v"])
K = len(lists)
c = list(center) if center is not None else mean_vec([e for l in lists for e in l], dh)
vec = np.zeros(K * dh, dtype=np.float32)
for k, l in enumerate(lists):
if l:
vec[k * dh:(k + 1) * dh] = np.asarray(mean_vec(l, dh), dtype=np.float64)
return {"vec": vec, "cnt": [len(l) for l in lists], "center": np.asarray(c, dtype=np.float32),
"lam": lambda_for(type_, lists, c, beta)}