llm2decision 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,20 @@
1
+ """Any hosted LLM as a decision model, used like the TypeSafe (Jev) SDK.
2
+
3
+ import llm2decision as l2d
4
+
5
+ with l2d.DecisionClient("qwen3.8-27b@nebius") as client:
6
+ r = client.system_one(state, {"team": l2d.Choice(criteria={"billing": "…", "other": None})})
7
+ r.choices["team"].choice
8
+ """
9
+ from .client import AsyncDecisionClient, DecisionClient
10
+ from .config import LLM2DecisionWarning, ModelInfo, bindings, list_models, providers
11
+ from .errors import (AuthError, ConfigError, LLM2DecisionError, MarksError, QuestionError, TransportError,
12
+ UnreadableAnswer)
13
+ from .marks import Marks, assign, read
14
+ from .types import (Answer, Choice, ChoiceAnswer, Meta, Noul, NoulAnswer, Score, ScoreAnswer, SystemOneResponse,
15
+ Tfu, TfuAnswer, Usage)
16
+
17
+ __all__ = ["DecisionClient", "AsyncDecisionClient", "list_models", "ModelInfo", "Noul", "Tfu", "Choice", "Score", "NoulAnswer", "TfuAnswer", "ChoiceAnswer",
18
+ "ScoreAnswer", "Answer", "SystemOneResponse", "Usage", "Meta", "LLM2DecisionError", "QuestionError",
19
+ "ConfigError", "AuthError", "TransportError", "UnreadableAnswer", "LLM2DecisionWarning", "bindings", "providers", "Marks", "MarksError", "assign", "read"]
20
+ __version__ = "0.1.0"
@@ -0,0 +1,63 @@
1
+ """`python -m llm2decision` (or `llm2decision`): list the known models, check one, or measure it.
2
+
3
+ llm2decision models [--provider nebius]
4
+ llm2decision check qwen3.8-27b@nebius [--save]
5
+ llm2decision check Qwen/Qwen3-8B --base-url http://localhost:8000/v1
6
+ llm2decision bench qwen3.8-27b@nebius [--data test.jsonl] [--limit 200] [--out answers.jsonl]
7
+ """
8
+ from __future__ import annotations
9
+
10
+ import argparse
11
+ import re
12
+ import sys
13
+
14
+ from . import bench
15
+ from .client import DecisionClient
16
+ from .config import list_models
17
+ from .errors import LLM2DecisionError
18
+
19
+
20
+ def main(argv=None) -> int:
21
+ ap = argparse.ArgumentParser(prog="llm2decision")
22
+ sub = ap.add_subparsers(dest="cmd", required=True)
23
+ m = sub.add_parser("models", help="the known bindings")
24
+ m.add_argument("--provider")
25
+ c = sub.add_parser("check", help="probe a model through its provider")
26
+ b = sub.add_parser("bench", help="measure a model on Decision Questions")
27
+ for p in (c, b):
28
+ p.add_argument("model")
29
+ p.add_argument("--provider")
30
+ p.add_argument("--base-url")
31
+ p.add_argument("--key-env", help="the environment variable holding the key")
32
+ p.add_argument("--env-file")
33
+ c.add_argument("--save", action="store_true", help="record a pass in the user's models.json")
34
+ b.add_argument("--data", help="a Decision Questions file (default: the published set, downloaded once)")
35
+ b.add_argument("--out", help="where answers are kept and resumed from (default: bench_<model>.jsonl)")
36
+ b.add_argument("--limit", type=int, help="the first N questions only")
37
+ b.add_argument("--workers", type=int, default=8)
38
+ a = ap.parse_args(argv)
39
+ if a.cmd == "models":
40
+ for x in list_models(a.provider):
41
+ cal = ",".join(x.calibrated) or "-"
42
+ print(f"{x.name:<34} {x.transport:<14} calibrated={cal:<34} checked={x.checked or '-'}")
43
+ return 0
44
+ try:
45
+ with DecisionClient(a.model, provider=a.provider, base_url=a.base_url, key_env=a.key_env,
46
+ env_file=a.env_file, **({"workers": 1} if a.cmd == "bench" else {})) as client:
47
+ if a.cmd == "check":
48
+ rep = client.check(save=a.save)
49
+ print(rep)
50
+ return 0 if rep.ok else 1
51
+ items = bench.load(a.data)[: a.limit]
52
+ out = a.out or f"bench_{re.sub(r'[^A-Za-z0-9._@-]+', '_', a.model)}.jsonl"
53
+ rows = bench.run(client, items, out, workers=a.workers,
54
+ progress=lambda s: print(s, file=sys.stderr, flush=True))
55
+ print(bench.report(rows))
56
+ return 0 if len(rows) == len(items) else 1
57
+ except LLM2DecisionError as e:
58
+ print(f"llm2decision: {e}", file=sys.stderr)
59
+ return 2
60
+
61
+
62
+ if __name__ == "__main__":
63
+ sys.exit(main())
llm2decision/_http.py ADDED
@@ -0,0 +1,127 @@
1
+ """JSON over HTTP with the standard library: one kept-alive connection per thread and host, proxies
2
+ from the environment (`HTTPS_PROXY`, `HTTP_PROXY`, `NO_PROXY`), no dependencies.
3
+
4
+ Network failures raise `NetworkError` (the transport retries them); an HTTP status is returned as is.
5
+ A kept-alive connection the server has closed in the meantime is reopened once, silently.
6
+ """
7
+ from __future__ import annotations
8
+
9
+ import base64
10
+ import http.client
11
+ import json
12
+ import ssl
13
+ import threading
14
+ import urllib.request
15
+ from dataclasses import dataclass
16
+ from email.message import Message
17
+ from typing import Any
18
+ from urllib.parse import unquote, urlsplit
19
+
20
+
21
+ class NetworkError(Exception):
22
+ """The request did not get an HTTP response: refused, reset, timed out, name not resolved."""
23
+
24
+
25
+ @dataclass
26
+ class Response:
27
+ status: int
28
+ headers: Message
29
+ text: str
30
+
31
+ def json(self) -> Any:
32
+ return json.loads(self.text)
33
+
34
+
35
+ def _proxy_for(scheme: str, host: str):
36
+ """The proxy URL the environment names for this scheme and host, unless `NO_PROXY` exempts the host."""
37
+ proxies = urllib.request.getproxies()
38
+ url = proxies.get(scheme)
39
+ if not url or urllib.request.proxy_bypass(host):
40
+ return None
41
+ return urlsplit(url if "://" in url else f"http://{url}")
42
+
43
+
44
+ def _proxy_auth(p) -> dict[str, str]:
45
+ if not p.username:
46
+ return {}
47
+ token = base64.b64encode(f"{unquote(p.username)}:{unquote(p.password or '')}".encode()).decode()
48
+ return {"Proxy-Authorization": f"Basic {token}"}
49
+
50
+
51
+ class Session:
52
+ """Default headers plus per-thread connections. Thread-safe: threads never share a connection."""
53
+
54
+ def __init__(self, headers: dict[str, str] | None = None):
55
+ self.headers = dict(headers or {})
56
+ self._local = threading.local()
57
+ self._lock = threading.Lock()
58
+ self._open: list[http.client.HTTPConnection] = []
59
+ self._ssl = ssl.create_default_context()
60
+
61
+ def _connection(self, scheme: str, host: str, port: int, timeout: float):
62
+ conns = self._local.__dict__.setdefault("conns", {})
63
+ key = (scheme, host, port)
64
+ conn = conns.get(key)
65
+ fresh = conn is None
66
+ if fresh:
67
+ proxy = _proxy_for(scheme, host)
68
+ if scheme == "https":
69
+ if proxy:
70
+ conn = http.client.HTTPSConnection(proxy.hostname, proxy.port or 8080, timeout=timeout,
71
+ context=self._ssl)
72
+ conn.set_tunnel(host, port, headers=_proxy_auth(proxy))
73
+ else:
74
+ conn = http.client.HTTPSConnection(host, port, timeout=timeout, context=self._ssl)
75
+ else:
76
+ target = proxy if proxy else None
77
+ conn = http.client.HTTPConnection(target.hostname if target else host,
78
+ (target.port or 8080) if target else port, timeout=timeout)
79
+ conn._llm2decision_proxy = proxy # an absolute URL in the request line
80
+ conns[key] = conn
81
+ with self._lock:
82
+ self._open.append(conn)
83
+ conn.timeout = timeout
84
+ if conn.sock is not None:
85
+ conn.sock.settimeout(timeout)
86
+ return conn, fresh
87
+
88
+ def _drop(self, scheme: str, host: str, port: int) -> None:
89
+ conn = self._local.__dict__.get("conns", {}).pop((scheme, host, port), None)
90
+ if conn is not None:
91
+ conn.close()
92
+
93
+ def post(self, url: str, body: Any, *, timeout: float, headers: dict[str, str] | None = None) -> Response:
94
+ u = urlsplit(url)
95
+ scheme, host = u.scheme, u.hostname or ""
96
+ port = u.port or (443 if scheme == "https" else 80)
97
+ path = (u.path or "/") + (f"?{u.query}" if u.query else "")
98
+ data = json.dumps(body).encode("utf-8")
99
+ hdrs = {**self.headers, **(headers or {}), "Content-Type": "application/json",
100
+ "Content-Length": str(len(data)), "Accept": "application/json"}
101
+ for _ in range(2):
102
+ conn, fresh = self._connection(scheme, host, port, timeout)
103
+ target = url if getattr(conn, "_llm2decision_proxy", None) else path
104
+ if getattr(conn, "_llm2decision_proxy", None):
105
+ hdrs.update(_proxy_auth(conn._llm2decision_proxy))
106
+ try:
107
+ conn.request("POST", target, body=data, headers=hdrs)
108
+ r = conn.getresponse()
109
+ raw = r.read()
110
+ except (http.client.RemoteDisconnected, BrokenPipeError, ConnectionResetError) as e:
111
+ self._drop(scheme, host, port)
112
+ if fresh: # a new connection failed: a real network error
113
+ raise NetworkError(type(e).__name__) from e
114
+ continue # a kept-alive one went stale: reopen once
115
+ except (OSError, http.client.HTTPException) as e:
116
+ self._drop(scheme, host, port)
117
+ raise NetworkError(type(e).__name__) from e
118
+ if (r.getheader("Connection") or "").lower() == "close":
119
+ self._drop(scheme, host, port)
120
+ return Response(r.status, r.headers, raw.decode("utf-8", "replace"))
121
+ raise NetworkError("connection closed") # pragma: no cover - two stale connections
122
+
123
+ def close(self) -> None:
124
+ with self._lock:
125
+ conns, self._open = self._open, []
126
+ for c in conns:
127
+ c.close()
llm2decision/auth.py ADDED
@@ -0,0 +1,144 @@
1
+ """Where a provider's key comes from and how it is sent — with nothing printed and nothing sent
2
+ where it does not belong.
3
+
4
+ Sources, first found wins:
5
+
6
+ 1. `key=` in code — a string, or a callable returning one (a keyring, a vault, a secrets manager)
7
+ 2. `key_env=` — the name of an environment variable chosen by the caller
8
+ 3. the provider's own variable from its configuration (`NEBIUS_API_KEY`, `OPENROUTER_API_KEY`, …)
9
+ 4. an env file — only when named: `env_file=` or `LLM2DECISION_ENV_FILE`; never picked up silently,
10
+ and it never overrides a variable already set in the environment
11
+
12
+ A provider configured with `auth: none` (a vLLM, llama-server or Ollama of your own) needs no key.
13
+
14
+ A key found in the environment or an env file (3, 4) is sent only to the provider's own host: with
15
+ `base_url` pointed elsewhere it is withheld, and only an explicit `key=` or `key_env=` is sent there.
16
+ A typo in an address must not hand a key to a stranger.
17
+
18
+ How it is sent comes from the provider: `Authorization: Bearer <key>` (OpenAI-style), `x-api-key`
19
+ (Anthropic), `api-key` (Azure) or any header with an optional scheme.
20
+ """
21
+ from __future__ import annotations
22
+
23
+ import os
24
+ from dataclasses import dataclass
25
+ from pathlib import Path
26
+ from typing import Callable, Mapping
27
+ from urllib.parse import urlparse
28
+
29
+ ENV_FILE_VAR = "LLM2DECISION_ENV_FILE"
30
+
31
+
32
+ from .errors import AuthError # noqa: E402 re-exported
33
+
34
+
35
+ @dataclass(frozen=True)
36
+ class AuthSpec:
37
+ """How a provider takes its key (from the provider's configuration)."""
38
+ mode: str = "header" # "header" or "none"
39
+ header: str = "Authorization"
40
+ scheme: str = "Bearer" # "" for a bare key, as `x-api-key` takes it
41
+ env: str = "" # the provider's own variable, e.g. NEBIUS_API_KEY
42
+ host: str = "" # the provider's own host: environment keys go only here
43
+
44
+ @classmethod
45
+ def from_config(cls, cfg: Mapping | None, base_url: str = "") -> "AuthSpec":
46
+ cfg = dict(cfg or {})
47
+ if cfg.get("mode") == "none" or cfg.get("type") == "none":
48
+ return cls(mode="none", host=_host(base_url))
49
+ return cls(mode="header", header=cfg.get("header", "Authorization"),
50
+ scheme=cfg.get("scheme", "Bearer"), env=cfg.get("env", ""),
51
+ host=cfg.get("host") or _host(base_url))
52
+
53
+
54
+ @dataclass(frozen=True)
55
+ class Credential:
56
+ """A resolved key and where it came from. `repr` and `str` never show the key."""
57
+ value: str | None
58
+ source: str # "code", "env:NAME", "env_file:PATH:NAME", "none"
59
+ explicit: bool # given by the caller (code, key_env) rather than found
60
+
61
+ def __repr__(self) -> str:
62
+ return f"Credential(source={self.source!r}, set={self.value is not None})"
63
+
64
+ __str__ = __repr__
65
+
66
+
67
+ def _host(url: str) -> str:
68
+ return (urlparse(url).hostname or "").lower()
69
+
70
+
71
+ def parse_env_file(path: str | Path) -> dict[str, str]:
72
+ """`NAME=value` lines: `export ` allowed, `#` comments and blank lines skipped, one level of
73
+ matching quotes removed, an inline ` #` comment cut from unquoted values."""
74
+ out: dict[str, str] = {}
75
+ for raw in Path(path).read_text(encoding="utf-8").splitlines():
76
+ line = raw.strip()
77
+ if not line or line.startswith("#"):
78
+ continue
79
+ if line.startswith("export "):
80
+ line = line[len("export "):].lstrip()
81
+ if "=" not in line:
82
+ continue
83
+ name, value = line.split("=", 1)
84
+ name, value = name.strip(), value.strip()
85
+ if len(value) >= 2 and value[0] == value[-1] and value[0] in "'\"":
86
+ value = value[1:-1]
87
+ elif " #" in value:
88
+ value = value.split(" #", 1)[0].rstrip()
89
+ if name:
90
+ out[name] = value
91
+ return out
92
+
93
+
94
+ def resolve(spec: AuthSpec, *, key: str | Callable[[], str] | None = None, key_env: str | None = None,
95
+ env_file: str | Path | None = None, base_url: str = "",
96
+ environ: Mapping[str, str] | None = None) -> Credential:
97
+ """The key for one provider, by the order in the module docstring."""
98
+ environ = os.environ if environ is None else environ
99
+ if spec.mode == "none" and key is None and key_env is None:
100
+ return Credential(None, "none", False)
101
+ if key is not None:
102
+ value = key() if callable(key) else key
103
+ if not isinstance(value, str) or not value.strip():
104
+ raise AuthError("the key given in code is empty")
105
+ return Credential(value.strip(), "code", True)
106
+ if key_env:
107
+ value = environ.get(key_env, "").strip()
108
+ if not value:
109
+ raise AuthError(f"key_env names {key_env!r}, which is not set or is empty")
110
+ return Credential(value, f"env:{key_env}", True)
111
+ looked = []
112
+ if spec.env:
113
+ looked.append(f"the environment variable {spec.env}")
114
+ value = environ.get(spec.env, "").strip()
115
+ if value:
116
+ return _for_host(Credential(value, f"env:{spec.env}", False), spec, base_url)
117
+ file = env_file or environ.get(ENV_FILE_VAR)
118
+ if file and spec.env:
119
+ looked.append(f"{spec.env} in the env file {file}")
120
+ p = Path(file)
121
+ if not p.is_file():
122
+ raise AuthError(f"env file {file} does not exist")
123
+ value = parse_env_file(p).get(spec.env, "").strip()
124
+ if value:
125
+ return _for_host(Credential(value, f"env_file:{file}:{spec.env}", False), spec, base_url)
126
+ where = ", ".join(looked) or "nowhere: the provider names no variable"
127
+ raise AuthError(f"no key for this provider; looked in {where}. Pass key=…, key_env=…, or "
128
+ f"env_file=… (or set {ENV_FILE_VAR})")
129
+
130
+
131
+ def _for_host(cred: Credential, spec: AuthSpec, base_url: str) -> Credential:
132
+ """A key found rather than given goes only to the provider's own host."""
133
+ target = _host(base_url) or spec.host
134
+ if spec.host and target and target != spec.host:
135
+ raise AuthError(f"a key from {cred.source.split(':')[0]} is meant for {spec.host}, not {target}; "
136
+ "pass key=… or key_env=… to send a key to another address")
137
+ return cred
138
+
139
+
140
+ def headers(spec: AuthSpec, cred: Credential) -> dict[str, str]:
141
+ """The request headers that carry the key."""
142
+ if cred.value is None:
143
+ return {}
144
+ return {spec.header: f"{spec.scheme} {cred.value}" if spec.scheme else cred.value}
llm2decision/bench.py ADDED
@@ -0,0 +1,189 @@
1
+ """Measure a model on Decision Questions, the set behind the README's model table.
2
+
3
+ llm2decision bench qwen3.8-27b@nebius # the published set, downloaded once
4
+ llm2decision bench my-qwen@local --data test.jsonl --limit 200
5
+
6
+ Every question goes out on its own, as in the table. Answers are appended to `--out` as they come, so an
7
+ interrupted run continues where it stopped when started again. The report gives accuracy over the
8
+ answered questions (all and per type), the share left undecided, and the calibration error when the
9
+ model returns probabilities.
10
+ """
11
+ from __future__ import annotations
12
+
13
+ import json
14
+ import os
15
+ import signal
16
+ import sys
17
+ import threading
18
+ import time
19
+ import urllib.request
20
+ from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait
21
+ from dataclasses import dataclass, field
22
+ from pathlib import Path
23
+
24
+ from .errors import LLM2DecisionError, TransportError
25
+
26
+ # the revision the README's model table was measured on; the cache is per revision, so a newer one is fetched anew
27
+ DATA_REVISION = "04da4c4131f23424b5d41e219459fa1ce5a8fa69"
28
+ DATA_URL = f"https://huggingface.co/datasets/mihailgribov/decision-questions/resolve/{DATA_REVISION}/data/test.jsonl"
29
+ TYPES = ("noul", "tfu", "choice", "score")
30
+ NAMES = {"noul": "yes/no", "tfu": "yes/no/unknown", "choice": "choice", "score": "score"}
31
+
32
+
33
+ def cache_dir() -> Path:
34
+ return Path(os.environ.get("XDG_CACHE_HOME") or Path.home() / ".cache") / "llm2decision"
35
+
36
+
37
+ def load(path: str | os.PathLike | None = None) -> list[dict]:
38
+ """The items of a Decision Questions file; without a path, the published set (downloaded once and cached)."""
39
+ if path is None:
40
+ path = cache_dir() / f"decision_questions_test_{DATA_REVISION[:12]}.jsonl"
41
+ if not path.exists():
42
+ path.parent.mkdir(parents=True, exist_ok=True)
43
+ print(f"downloading Decision Questions to {path}", file=sys.stderr, flush=True)
44
+ tmp = path.with_suffix(".part")
45
+ with urllib.request.urlopen(DATA_URL, timeout=120) as r, open(tmp, "wb") as f:
46
+ f.write(r.read())
47
+ tmp.replace(path)
48
+ with open(path, encoding="utf-8") as f:
49
+ return [json.loads(line) for line in f if line.strip()]
50
+
51
+
52
+ def question(item: dict) -> dict:
53
+ q = dict(item["question"])
54
+ if isinstance(q.get("criteria"), str):
55
+ q["criteria"] = json.loads(q["criteria"])
56
+ return q
57
+
58
+
59
+ def state(item: dict):
60
+ return json.loads(item["state"]) if item.get("state_format") == "json" else item["state"]
61
+
62
+
63
+ def ask(client, item: dict) -> dict:
64
+ """One item, one call; the answer as probabilities over the item's answer labels."""
65
+ a = client.system_one(state(item), {"q": question(item)}).answers["q"]
66
+ if item["type"] == "noul":
67
+ probs = {"yes": a.noul, "no": 1 - a.noul}
68
+ else:
69
+ probs = {str(k): v for k, v in a.probabilities.items()}
70
+ return {"id": item["id"], "type": item["type"], "answer": str(item["answer"]),
71
+ "probs": {k: round(v, 6) for k, v in probs.items()},
72
+ "answered": a.meta.answered, "logprobs": a.meta.logprobs}
73
+
74
+
75
+ def _done(out: Path) -> dict[str, dict]:
76
+ if not out.exists():
77
+ return {}
78
+ rows = {}
79
+ for line in out.read_text(encoding="utf-8").splitlines():
80
+ if line.strip():
81
+ r = json.loads(line)
82
+ rows[r["id"]] = r
83
+ return rows
84
+
85
+
86
+ def run(client, items: list[dict], out: str | os.PathLike, *, workers: int = 8, progress=print) -> dict[str, dict]:
87
+ """Ask what `out` does not hold yet, appending each answer as it comes; returns every row of `out`
88
+ for the given items. Ctrl-C stops after the requests in flight; a refused request stops the run."""
89
+ out = Path(out)
90
+ rows = _done(out)
91
+ todo = [x for x in items if x["id"] not in rows]
92
+ progress(f"{len(items) - len(todo)}/{len(items)} already in {out}, {len(todo)} to ask")
93
+ stop = threading.Event()
94
+ previous = None
95
+ if threading.current_thread() is threading.main_thread():
96
+ def on_signal(sig, _):
97
+ stop.set()
98
+ progress("stopping: finishing the requests in flight…")
99
+ previous = signal.signal(signal.SIGINT, on_signal)
100
+ errors, t0, n0 = 0, time.time(), len(rows)
101
+ try:
102
+ with open(out, "a", encoding="utf-8") as f, ThreadPoolExecutor(workers) as ex:
103
+ it, running = iter(todo), {}
104
+ while True:
105
+ while not stop.is_set() and len(running) < workers:
106
+ x = next(it, None)
107
+ if x is None:
108
+ break
109
+ running[ex.submit(ask, client, x)] = x
110
+ if not running:
111
+ break
112
+ for fut in wait(running, timeout=1, return_when=FIRST_COMPLETED)[0]:
113
+ x = running.pop(fut)
114
+ try:
115
+ r = fut.result()
116
+ except LLM2DecisionError as e:
117
+ errors += 1
118
+ progress(f" {x['id']}: {str(e)[:160]}")
119
+ if isinstance(e, TransportError) and "HTTP 4" in str(e) and "429" not in str(e):
120
+ stop.set() # a refusal (key, balance, model): the rest would fail too
121
+ continue
122
+ f.write(json.dumps(r, ensure_ascii=False) + "\n")
123
+ f.flush()
124
+ rows[r["id"]] = r
125
+ done = len(rows) - n0
126
+ if done % 50 == 0:
127
+ rate = done / max(time.time() - t0, 1e-9)
128
+ progress(f" {len(rows)}/{len(items)} {rate:.1f}/s errors {errors}")
129
+ finally:
130
+ if previous is not None:
131
+ signal.signal(signal.SIGINT, previous)
132
+ wanted = {x["id"] for x in items}
133
+ rows = {k: v for k, v in rows.items() if k in wanted}
134
+ progress(f"{len(rows)}/{len(items)} answered and saved, errors {errors}")
135
+ if len(rows) < len(items):
136
+ progress("to continue: run the same command again")
137
+ return rows
138
+
139
+
140
+ def _right(r: dict) -> bool:
141
+ return max(r["probs"], key=r["probs"].get) == r["answer"]
142
+
143
+
144
+ def _ece(rows: list[dict], bins: int) -> float:
145
+ e = 0.0
146
+ for b in range(bins):
147
+ inside = [r for r in rows if b / bins < max(r["probs"].values()) <= (b + 1) / bins]
148
+ if inside:
149
+ conf = sum(max(r["probs"].values()) for r in inside) / len(inside)
150
+ acc = sum(map(_right, inside)) / len(inside)
151
+ e += len(inside) / len(rows) * abs(acc - conf)
152
+ return e
153
+
154
+
155
+ def ece(rows: list[dict], bins: int = 10) -> float:
156
+ """Expected calibration error of the top answer's probability over the answered rows: per question
157
+ type, then weighted by the type's share, so that easy types do not hide a badly calibrated one."""
158
+ rows = [r for r in rows if r["answered"]]
159
+ by = [[r for r in rows if r["type"] == t] for t in TYPES]
160
+ return sum(len(g) * _ece(g, bins) for g in by if g) / len(rows) if rows else 0.0
161
+
162
+
163
+ @dataclass
164
+ class Report:
165
+ n: int
166
+ accuracy: float | None
167
+ by_type: dict[str, float | None] = field(default_factory=dict)
168
+ undecided: float = 0.0
169
+ calibration_error: float | None = None # None: the model answers in text
170
+
171
+ def __str__(self) -> str:
172
+ f = lambda v: "-" if v is None else f"{v:.3f}"
173
+ lines = [f"questions {self.n}", f"accuracy {f(self.accuracy)} (over answered questions)"]
174
+ lines += [f" {NAMES[t]:<15}{f(v)}" for t, v in self.by_type.items()]
175
+ lines.append(f"undecided {100 * self.undecided:.1f}%")
176
+ lines.append("calibration err " + ("no probabilities (text answers)" if self.calibration_error is None
177
+ else f(self.calibration_error)))
178
+ return "\n".join(lines)
179
+
180
+
181
+ def report(rows) -> Report:
182
+ rows = list(rows.values() if isinstance(rows, dict) else rows)
183
+ answered = [r for r in rows if r["answered"]]
184
+ acc = lambda rs: sum(map(_right, rs)) / len(rs) if rs else None
185
+ by = {t: acc([r for r in answered if r["type"] == t]) for t in TYPES if any(r["type"] == t for r in rows)}
186
+ probs = answered and all(r["logprobs"] for r in answered)
187
+ return Report(n=len(rows), accuracy=acc(answered), by_type=by,
188
+ undecided=(1 - len(answered) / len(rows)) if rows else 0.0,
189
+ calibration_error=ece(rows) if probs else None)