"""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)}