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.
- llm2decision/__init__.py +20 -0
- llm2decision/__main__.py +63 -0
- llm2decision/_http.py +127 -0
- llm2decision/auth.py +144 -0
- llm2decision/bench.py +189 -0
- llm2decision/check.py +194 -0
- llm2decision/client.py +289 -0
- llm2decision/config.py +217 -0
- llm2decision/configs/forms.json +5 -0
- llm2decision/configs/models.json +250 -0
- llm2decision/configs/providers.json +42 -0
- llm2decision/errors.py +39 -0
- llm2decision/marks.py +237 -0
- llm2decision/prompt.py +60 -0
- llm2decision/py.typed +0 -0
- llm2decision/transports.py +316 -0
- llm2decision/types.py +205 -0
- llm2decision-0.1.0.dist-info/METADATA +343 -0
- llm2decision-0.1.0.dist-info/RECORD +22 -0
- llm2decision-0.1.0.dist-info/WHEEL +4 -0
- llm2decision-0.1.0.dist-info/entry_points.txt +2 -0
- llm2decision-0.1.0.dist-info/licenses/LICENSE +202 -0
llm2decision/__init__.py
ADDED
|
@@ -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"
|
llm2decision/__main__.py
ADDED
|
@@ -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)
|