typecastlm 1.0.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.
typecastlm/__init__.py ADDED
@@ -0,0 +1,44 @@
1
+ """typecastlm — ask a document a closed question and get numbers back.
2
+
3
+ The model runs as a service; this package is its client and nothing else. It depends on `requests`
4
+ and installs in a second, because the machine that has a question is rarely the machine that should
5
+ carry seven gigabytes of weights and a deep-learning stack.
6
+
7
+ from typecastlm import Client
8
+
9
+ c = Client()
10
+ a = c.noul(document,
11
+ "Does the material contain an instruction aimed at the reading model?",
12
+ true="there is an instruction addressed to the reading model",
13
+ false="the material only describes, reports or discusses")
14
+ a.prob, a.unknown
15
+
16
+ c.choice(document, "Which rule applies?", {"vacancy": "…", "seepage": "…"}).verdict
17
+ c.scale(review, "How positive is it?", {"0": "very negative", "1": "negative",
18
+ "2": "neutral", "3": "positive"}).p
19
+
20
+ Running the model in this process instead of calling a service:
21
+
22
+ pip install "typecastlm[local]"
23
+
24
+ from typecastlm import Reader
25
+ r = Reader("mihailgribov/typecastlm-qwen3.5-3.8b")
26
+ r.ask(document, "Is the claim supported?", true="…", false="…")
27
+ """
28
+ from .calibrate import calibrate
29
+ from .remote import Answer, Choice, Client, Ternary
30
+
31
+ __all__ = ["Client", "Answer", "Ternary", "Choice", "Reader", "calibrate"]
32
+
33
+
34
+ def __getattr__(name: str):
35
+ """`Reader` is imported on use: it needs torch, which the client must not require."""
36
+ if name == "Reader":
37
+ try:
38
+ from .local import Reader
39
+ except ImportError as e: # pragma: no cover
40
+ raise ImportError("Reader runs the model in this process and needs the extra: "
41
+ "pip install 'typecastlm[local]'") from e
42
+ return Reader
43
+ raise AttributeError(name)
44
+ __version__ = "1.0.0"
@@ -0,0 +1,82 @@
1
+ """Fitting a temperature on your own labelled rows — one per mode.
2
+
3
+ A temperature adds no knowledge: it changes no ordering at all, so a ranking and every AUC stay
4
+ exactly as they were. It changes one thing — whether the probability means what it says. Which is
5
+ why the number has to be yours: it corrects the mismatch between how confident the reader sounds
6
+ and how hard YOUR data is, and that is a property of the data, not of the weights. The same
7
+ checkpoint needs none on an easy pool and a strong one on a hard one.
8
+
9
+ Per-answer shifts were removed in 1.0.0. They were needed while the third row of the head lived on
10
+ its own scale; it is now a direction computed from data and scaled to the other two, and fitting
11
+ shifts on top buys 0.005 of calibration error and no accuracy.
12
+
13
+ from typecastlm import Client, calibrate
14
+
15
+ rows = [(c.noul(t, q, true=T, false=F).logits, gold) for t, gold in my_labelled]
16
+ T = calibrate(rows)["temperature"]
17
+ c = Client(temperature=T)
18
+
19
+ The same function serves a choice or a scale: pass the logits that mode returned and the correct
20
+ key, and fit one temperature per mode you use.
21
+ """
22
+ from __future__ import annotations
23
+
24
+ import math
25
+
26
+
27
+ def _softmax(v: list[float]) -> list[float]:
28
+ m = max(v)
29
+ e = [math.exp(x - m) for x in v]
30
+ s = sum(e)
31
+ return [x / s for x in e]
32
+
33
+
34
+ def _ece(conf: list[float], ok: list[bool], bins: int = 10) -> float:
35
+ n = len(conf)
36
+ total = 0.0
37
+ for i in range(bins):
38
+ lo, hi = i / bins, (i + 1) / bins
39
+ idx = [j for j in range(n) if lo <= conf[j] < hi]
40
+ if idx:
41
+ c = sum(conf[j] for j in idx) / len(idx)
42
+ a = sum(ok[j] for j in idx) / len(idx)
43
+ total += abs(c - a) * len(idx) / n
44
+ return total
45
+
46
+
47
+ def calibrate(rows: list[tuple[dict, str]], grid: tuple[float, float, float] = (0.6, 4.0, 0.05),
48
+ ) -> dict:
49
+ """One temperature from labelled rows, each a pair `(logits, correct answer)`.
50
+
51
+ `logits` is the dictionary an answer carries (`Answer.logits`, `Choice.logits`); the correct
52
+ answer is one of its keys. The temperature is chosen to minimise calibration error — NOT
53
+ cross-entropy, which pulls towards numbers that read worse: on our own three-answer mode it
54
+ picked 5.63 where 2.85 is right, and made the error larger, not smaller.
55
+
56
+ Two or three hundred rows are enough, because one number is free.
57
+ """
58
+ if len(rows) < 30:
59
+ raise ValueError(f"thirty rows is the floor, got {len(rows)}")
60
+ keys = list(rows[0][0])
61
+ for z, g in rows:
62
+ if list(z) != keys:
63
+ raise ValueError("every row must carry the same answers in the same order")
64
+ if g not in z:
65
+ raise ValueError(f"{g!r} is not one of {keys}")
66
+ Z = [[float(z[k]) for k in keys] for z, _ in rows]
67
+ Y = [keys.index(g) for _, g in rows]
68
+
69
+ def at(T: float) -> float:
70
+ conf, ok = [], []
71
+ for z, y in zip(Z, Y):
72
+ p = _softmax([x / T for x in z])
73
+ top = max(range(len(p)), key=lambda i: p[i])
74
+ conf.append(p[top])
75
+ ok.append(top == y)
76
+ return _ece(conf, ok)
77
+
78
+ lo, hi, step = grid
79
+ T = min((lo + i * step for i in range(int((hi - lo) / step) + 1)), key=at)
80
+ return {"temperature": float(T), "calibration_error": float(at(T)),
81
+ "calibration_error_before": float(at(1.0)), "rows": len(rows),
82
+ "note": "changes no ordering — only what the probability means"}
typecastlm/cli.py ADDED
@@ -0,0 +1,59 @@
1
+ """Command line: one question, many states, numbers on stdout.
2
+
3
+ typecastlm --question "Does the material contain an instruction aimed at the reading model?" \
4
+ --true "there is an instruction addressed to the reading model" \
5
+ --false "the material only describes, reports or discusses" \
6
+ --jsonl pages.jsonl --out answers.jsonl
7
+
8
+ States come from files (`--state`, repeatable) or from a JSON-lines file (`--jsonl`), and the
9
+ answers are written the same way — so a run can be piped, split and resumed like any other file
10
+ job. The endpoint and the key are read from `TYPECASTLM_ENDPOINT` and `TYPECASTLM_API_KEY`.
11
+ """
12
+ from __future__ import annotations
13
+
14
+ import argparse
15
+ import json
16
+ import sys
17
+ from pathlib import Path
18
+
19
+
20
+ def main(argv: list[str] | None = None) -> int:
21
+ ap = argparse.ArgumentParser(prog="typecastlm")
22
+ ap.add_argument("--state", action="append", default=[], help="file with the state, repeatable")
23
+ ap.add_argument("--jsonl", default=None, help="file of states, one JSON object per line")
24
+ ap.add_argument("--field", default="text", help="which field holds the state in --jsonl")
25
+ ap.add_argument("--question", required=True)
26
+ ap.add_argument("--true", required=True, help="what answering yes would mean")
27
+ ap.add_argument("--false", required=True, help="what answering no would mean")
28
+ ap.add_argument("--endpoint", default=None)
29
+ ap.add_argument("--api-key", default=None)
30
+ ap.add_argument("--out", default=None, help="write JSON lines here instead of stdout")
31
+ a = ap.parse_args(argv)
32
+
33
+ from .remote import Client
34
+
35
+ rows: list[dict] = [{"id": f, "text": Path(f).read_text(encoding="utf-8")} for f in a.state]
36
+ if a.jsonl:
37
+ for i, line in enumerate(Path(a.jsonl).open(encoding="utf-8")):
38
+ if line.strip():
39
+ r = json.loads(line)
40
+ rows.append({"id": r.get("id", i), "text": r[a.field]})
41
+ if not rows:
42
+ ap.error("nothing to read: pass --state or --jsonl")
43
+
44
+ c = Client(endpoint=a.endpoint, api_key=a.api_key)
45
+ sink = Path(a.out).open("w", encoding="utf-8") if a.out else sys.stdout
46
+ try:
47
+ for row in rows:
48
+ v = c.noul(row["text"], a.question, true=a.true, false=a.false)
49
+ sink.write(json.dumps({"id": row["id"], "prob": round(v.prob, 4),
50
+ "unknown": round(v.unknown, 4),
51
+ "margin": round(v.margin, 3)}, ensure_ascii=False) + "\n")
52
+ finally:
53
+ if a.out:
54
+ sink.close()
55
+ return 0
56
+
57
+
58
+ if __name__ == "__main__": # pragma: no cover
59
+ raise SystemExit(main())
typecastlm/local.py ADDED
@@ -0,0 +1,211 @@
1
+ """Read this checkpoint in three decision shapes. Requires `transformers` only.
2
+
3
+ ask(state, question, true, false) -> {"p": {"yes","no"}, "p_unsure", "logits"}
4
+ choice(state, question, options) -> {"p": {option: prob}, "logits", "marks"}
5
+ scale(state, question, levels) -> same, for an ordinal rubric
6
+ ask_many(state, questions) -> list of ask() results, material read once
7
+
8
+ Options are marked with letters, rubric levels with their own digits when those are single
9
+ tokens, otherwise with `0..9A..P`. Marks above ten levels degrade. Probabilities use the
10
+ per-mode temperature from `prompt.json`; `logits` are already divided by it.
11
+
12
+ from reader import Reader
13
+ r = Reader("mihailgribov/typecastlm-qwen3.5-3.8b")
14
+ r.ask(state, "Is the claim supported?", true="the material supports it",
15
+ false="the material contradicts it")
16
+ r.choice(state, "Which rule applies?", [("a", "…"), ("b", "…")])
17
+ """
18
+ from __future__ import annotations
19
+
20
+ import json
21
+ import math
22
+ from pathlib import Path
23
+
24
+ import torch
25
+
26
+ LETTERS = "ABCDEFGHIJKLMNOP"
27
+ ORDINAL = "0123456789ABCDEFGHIJKLMNOP"
28
+ CHOICE_SYSTEM = "You answer with exactly one character from the given list."
29
+
30
+
31
+ class Reader:
32
+ def __init__(self, model: str, device: str = "cuda", dtype=torch.bfloat16):
33
+ from huggingface_hub import hf_hub_download
34
+ from transformers import AutoModelForSequenceClassification, AutoTokenizer
35
+
36
+ path = Path(model)
37
+ cfg_file = (path / "prompt.json") if path.is_dir() else Path(
38
+ hf_hub_download(model, "prompt.json"))
39
+ self.cfg = json.loads(cfg_file.read_text(encoding="utf-8"))
40
+ self.tok = AutoTokenizer.from_pretrained(model)
41
+ self.tok.padding_side = "right"
42
+ self.model = AutoModelForSequenceClassification.from_pretrained(
43
+ model, dtype=dtype, device_map=device).eval()
44
+ self.model.requires_grad_(False)
45
+ self.trunk = self.model.model
46
+ self.names = [self.model.config.id2label[i]
47
+ for i in range(self.model.config.num_labels)]
48
+ self.col = {n: i for i, n in enumerate(self.names)}
49
+ # The head emits 29 raw logits: three answers and one per mark. A softmax over all of
50
+ # them is meaningless — each mode normalises its own subset.
51
+ self.cal = self.cfg.get("calibration", {})
52
+
53
+ # -- prompt ---------------------------------------------------------------------------------
54
+ def _fold(self, state: str) -> str:
55
+ """Fold a state longer than `max_state_tokens` in the middle."""
56
+ ids = self.tok.encode(state, add_special_tokens=False)
57
+ n = self.cfg["max_state_tokens"]
58
+ if len(ids) <= n:
59
+ return state
60
+ return (self.tok.decode(ids[: n // 2], skip_special_tokens=True) + " […] "
61
+ + self.tok.decode(ids[-(n // 2):], skip_special_tokens=True))
62
+
63
+ def _chat(self, system: str, body: str, tail: str) -> str:
64
+ return self.tok.apply_chat_template(
65
+ [{"role": "system", "content": system}, {"role": "user", "content": body}],
66
+ tokenize=False, add_generation_prompt=True, enable_thinking=False) + tail
67
+
68
+ def _text(self, state: str, question: str, true: str, false: str) -> str:
69
+ C = self.cfg
70
+ body = (C["format_two"] + C["body"].format(state=state, question=question)
71
+ + C["means_two"].format(true=true, false=false))
72
+ return self._chat(C["system"], body, C["tail"])
73
+
74
+ # -- marks ----------------------------------------------------------------------------------
75
+ def _has_mark(self, mark: str) -> bool:
76
+ """True if the head carries a row for this mark."""
77
+ return f"mark_{mark}" in self.col
78
+
79
+ def _marks_for(self, ids: list[str], ordinal: bool) -> list[str]:
80
+ """Marks for the options: a rubric\'s own digits when usable, else letters or `0..9A..P`."""
81
+ if ordinal:
82
+ own = [str(x) for x in ids]
83
+ if len(own) <= 10 and all(x.isdigit() and len(x) == 1 and self._has_mark(x)
84
+ for x in own):
85
+ return own
86
+ row = [c for c in ORDINAL if self._has_mark(c)]
87
+ else:
88
+ row = [c for c in LETTERS if self._has_mark(c)]
89
+ if len(row) < len(ids):
90
+ raise ValueError(f"{len(ids)} options but only {len(row)} single-token marks")
91
+ return row[: len(ids)]
92
+
93
+ # -- reading --------------------------------------------------------------------------------
94
+ @torch.no_grad()
95
+ def _last(self, text: str) -> torch.Tensor:
96
+ enc = self.tok(text, return_tensors="pt", add_special_tokens=False).to(self.model.device)
97
+ return self.trunk(**enc).last_hidden_state[0, -1].float()
98
+
99
+ def _temp(self, mode: str) -> float:
100
+ return float((self.cal.get(mode) or {}).get("temperature", 1.0))
101
+
102
+ def _verdict(self, z: list[float]) -> dict:
103
+ """Three raw logits to an answer. `p` uses the `verdict` temperature, `p_unsure` the
104
+ `three_answers` one; the two are fitted separately and are not interchangeable."""
105
+ yes, no = (v / self._temp("verdict") for v in z[:2])
106
+ m = max(yes, no)
107
+ ey, en = math.exp(yes - m), math.exp(no - m)
108
+ t3 = self._temp("three_answers")
109
+ w = [v / t3 for v in z]
110
+ m3 = max(w)
111
+ e3 = [math.exp(v - m3) for v in w]
112
+ s3 = sum(e3)
113
+ return dict(logits=dict(zip(self.names[:3], z)),
114
+ p={"yes": ey / (ey + en), "no": en / (ey + en)},
115
+ p_unsure=e3[2] / s3)
116
+
117
+ @torch.no_grad()
118
+ def ask(self, state: str, question: str, true: str, false: str) -> dict:
119
+ """One closed question. `p` is over yes/no; `p_unsure` is a separate signal to threshold."""
120
+ h = self._last(self._text(self._fold(state), question, true, false))
121
+ z = self.model.score(h.to(self.model.score.weight.dtype)).float()[:3].tolist()
122
+ return self._verdict(z)
123
+
124
+ @torch.no_grad()
125
+ def _pick(self, state: str, question: str, options: list[tuple[str, str]], ask: str,
126
+ ordinal: bool) -> dict:
127
+ marks = self._marks_for([o[0] for o in options], ordinal)
128
+ body = (f"<state>\n{self._fold(state)}\n</state>\n\n{question}\n\n"
129
+ + "\n".join(f"{m}. {desc}" for m, (_, desc) in zip(marks, options))
130
+ + f"\n\n{ask}")
131
+ h = self._last(self._chat(CHOICE_SYSTEM, body, "Answer: "))
132
+ full = self.model.score(h.to(self.model.score.weight.dtype)).float()
133
+ z = [float(full[self.col[f"mark_{m}"]]) for m in marks]
134
+ z = [v / self._temp("scale" if ordinal else "choice") for v in z]
135
+ e = [math.exp(v - max(z)) for v in z]
136
+ s = sum(e)
137
+ ids = [o[0] for o in options]
138
+ return dict(p=dict(zip(ids, (v / s for v in e))), logits=dict(zip(ids, z)),
139
+ marks=dict(zip(ids, marks)))
140
+
141
+ def choice(self, state: str, question: str, options: list[tuple[str, str]]) -> dict:
142
+ """Pick one of 2..16 options. `options` are `(name, description)`; names are yours and do
143
+ not reach the prompt."""
144
+ return self._pick(state, question, options, "Answer with one letter.", ordinal=False)
145
+
146
+ def scale(self, state: str, question: str, levels: list[tuple[str, str]]) -> dict:
147
+ """Pick a level of an ordinal rubric. `levels` are `(name, description)` in order."""
148
+ marks_are_digits = all(str(k).isdigit() for k, _ in levels)
149
+ ask = ("Answer with the number of the level that rates it." if marks_are_digits
150
+ else "Answer with the mark of the level that rates it.")
151
+ return self._pick(state, question, levels, ask, ordinal=True)
152
+
153
+ @staticmethod
154
+ def _snapshot(cache) -> dict:
155
+ """Copy of the cache state after the prefix.
156
+
157
+ A hybrid trunk keeps two kinds of state: attention layers hold keys and values per token
158
+ and could be cropped; linear-attention layers hold a recurrent summary that cannot. Both
159
+ are restored from a copy instead.
160
+ """
161
+ return {name: [None if t is None else t.clone() for t in getattr(cache, name)]
162
+ for name in ("conv_states", "recurrent_states", "key_cache", "value_cache")}
163
+
164
+ @staticmethod
165
+ def _restore(cache, snap: dict) -> None:
166
+ for name, tensors in snap.items():
167
+ setattr(cache, name, [None if t is None else t.clone() for t in tensors])
168
+
169
+ @torch.no_grad()
170
+ def ask_many(self, state: str, questions: list[tuple[str, str, str]]) -> list[dict]:
171
+ """Several `(question, true, false)` about one state. The shared prefix is computed once;
172
+ results equal calling `ask` separately. The split is taken in token space, not in text."""
173
+ state = self._fold(state)
174
+ fulls = [self.tok(self._text(state, q, t, f), add_special_tokens=False)["input_ids"]
175
+ for q, t, f in questions]
176
+ n_min = min(len(f) for f in fulls)
177
+ n = 0
178
+ while n < n_min and len({f[n] for f in fulls}) == 1:
179
+ n += 1
180
+ # Sharing the prefix across a hybrid trunk is not implemented: restoring the recurrent
181
+ # state of the linear-attention layers from a copy does not reproduce a clean pass.
182
+ return [self.ask(state, q, t, f) for q, t, f in questions]
183
+
184
+ cache = self._cache()
185
+ dev = self.model.device
186
+ pre = torch.tensor([fulls[0][:n]], device=dev)
187
+ self.trunk(input_ids=pre, attention_mask=torch.ones_like(pre),
188
+ past_key_values=cache, use_cache=True,
189
+ cache_position=torch.arange(n, device=dev))
190
+ snap = self._snapshot(cache)
191
+ out = []
192
+ for full in fulls:
193
+ self._restore(cache, snap)
194
+ ids = torch.tensor([full[n:]], device=dev)
195
+ h = self.trunk(input_ids=ids, past_key_values=cache, use_cache=True,
196
+ cache_position=torch.arange(n, len(full), device=dev),
197
+ attention_mask=torch.ones(1, len(full), device=dev, dtype=torch.long))
198
+ z = self.model.score(h.last_hidden_state[0, -1].to(self.model.score.weight.dtype))
199
+ out.append(self._verdict(z.float()[:3].tolist()))
200
+ return out
201
+
202
+ def _cache(self):
203
+ """The cache class this trunk needs. A plain `DynamicCache` does not fit a hybrid one."""
204
+ from transformers import DynamicCache
205
+
206
+ mod = type(self.model).__module__
207
+ for name in ("Qwen3_5DynamicCache", "Qwen3NextDynamicCache"):
208
+ cls = getattr(__import__(mod, fromlist=[name]), name, None)
209
+ if cls is not None:
210
+ return cls(self.model.config.get_text_config())
211
+ return DynamicCache()
typecastlm/remote.py ADDED
@@ -0,0 +1,227 @@
1
+ """HTTP client: a question about a state, an answer in numbers.
2
+
3
+ noul(state, instructions, true, false) -> Answer(prob, margin, unknown, logits, p)
4
+ tfu(state, instructions, true, false) -> Ternary(p, verdict, confidence, logits)
5
+ choice(state, instructions, options) -> Choice(p, verdict, confidence, logits)
6
+ scale(state, instructions, levels) -> Choice, for an ordinal rubric
7
+ ask(state, questions) -> the service body, any number of questions
8
+
9
+ `noul` matches the hosted service field for field. The rest are extensions: `unknown` is returned
10
+ whether or not the question asks for it, `logits` are raw, and option names in `choice` and
11
+ `scale` come back as the keys of `p`.
12
+
13
+ The service holds the weights; this package depends on `requests` only.
14
+ """
15
+ from __future__ import annotations
16
+
17
+ import math
18
+ import os
19
+ import time
20
+ from dataclasses import dataclass
21
+
22
+ DEFAULT_ENDPOINT = "" # no hosted service; set one or run `typecastlm-serve`
23
+ RETRY_CODES = (408, 429, 500, 502, 503, 504, 529)
24
+
25
+
26
+ @dataclass
27
+ class Ternary:
28
+ """What `tfu` returns: the three probabilities and which one leads."""
29
+
30
+ p: dict[str, float]
31
+ verdict: str
32
+ confidence: float
33
+ logits: dict[str, float]
34
+ ms: float
35
+ input_tokens: int
36
+ model: str
37
+
38
+
39
+ @dataclass
40
+ class Choice:
41
+ """What `choice` and `scale` return: a probability per option and which one leads. Keys are
42
+ the option names from the request."""
43
+
44
+ p: dict[str, float]
45
+ verdict: str
46
+ confidence: float
47
+ logits: dict[str, float]
48
+ ms: float
49
+ input_tokens: int
50
+ model: str
51
+
52
+
53
+ @dataclass
54
+ class Answer:
55
+ """What `noul` returns. `prob`, `margin`, `ms`, `input_tokens`, `model` are the service\'s
56
+ fields; `unknown`, `logits` and `p` are extensions."""
57
+
58
+ prob: float
59
+ margin: float
60
+ ms: float
61
+ input_tokens: int
62
+ model: str
63
+ unknown: float = 0.0
64
+ logits: dict[str, float] | None = None
65
+ p: dict[str, float] | None = None
66
+
67
+
68
+ class Client:
69
+ """One question, one state, three numbers — answered by the service.
70
+
71
+ The key is read from `TYPECASTLM_API_KEY` unless one is passed. A connection is kept open for
72
+ the life of the client: a fresh one per request means a name lookup per request, and a few
73
+ dozen of those at once is how a fast service starts looking slow.
74
+ """
75
+
76
+ def __init__(self, endpoint: str | None = None, api_key: str | None = None,
77
+ model: str = "typecastlm-qwen3.5-3.8b", timeout: float = 10.0, retries: int = 5,
78
+ pool: int = 32, temperature: float | None = None, calibrated: bool = True):
79
+ """`temperature` overrides the checkpoint\'s per-mode temperatures; `calibrated=False`
80
+ disables them. Temperatures are applied here because the service returns logits, so an
81
+ answer can be re-read at another temperature without asking again. They change no
82
+ ordering, only confidence."""
83
+ import requests
84
+ from requests.adapters import HTTPAdapter
85
+
86
+ self.endpoint = endpoint or os.environ.get("TYPECASTLM_ENDPOINT", DEFAULT_ENDPOINT)
87
+ if not self.endpoint:
88
+ raise ValueError(
89
+ "no endpoint. Pass Client(endpoint=...), set TYPECASTLM_ENDPOINT, or run the "
90
+ "model yourself: `pip install \"typecastlm[server]\"` and `typecastlm-serve`, or "
91
+ "`pip install \"typecastlm[local]\"` and use typecastlm.Reader in this process")
92
+ self.key = api_key or os.environ.get("TYPECASTLM_API_KEY", "")
93
+ self.model, self.timeout, self.retries = model, timeout, retries
94
+ self._s = requests.Session()
95
+ self._s.mount("https://", HTTPAdapter(pool_connections=pool, pool_maxsize=pool))
96
+ self._requests = requests
97
+ # The calibration travels with the service's answer: it belongs to the checkpoint, not to
98
+ # the client. A temperature passed here overrides it, `calibrated=False` turns it off.
99
+ if temperature is not None and temperature <= 0:
100
+ raise ValueError(f"temperature must be positive, got {temperature}")
101
+ self.calibrated = calibrated
102
+ self.temperature = temperature
103
+ self.calibration: dict = {}
104
+
105
+ def _temp(self, mode: str) -> float:
106
+ """The temperature for a mode: the caller's, then the checkpoint's, then none."""
107
+ if self.temperature is not None:
108
+ return float(self.temperature)
109
+ if not self.calibrated:
110
+ return 1.0
111
+ return float(((self.calibration or {}).get(mode) or {}).get("temperature", 1.0))
112
+
113
+ def _post(self, body: dict) -> tuple[dict, float]:
114
+ delay = 1.0
115
+ for attempt in range(self.retries + 1):
116
+ t0 = time.perf_counter()
117
+ try:
118
+ headers = {"Authorization": f"Bearer {self.key}"} if self.key else {}
119
+ r = self._s.post(self.endpoint, json=body, headers=headers, timeout=self.timeout)
120
+ ms = (time.perf_counter() - t0) * 1000.0
121
+ if r.status_code == 200:
122
+ return r.json(), ms
123
+ if r.status_code not in RETRY_CODES or attempt == self.retries:
124
+ raise RuntimeError(f"HTTP {r.status_code}: {r.text[:300]}")
125
+ except self._requests.RequestException:
126
+ if attempt == self.retries:
127
+ raise
128
+ ra = None
129
+ time.sleep(float(ra) if ra else delay)
130
+ delay = min(delay * 2, 30.0)
131
+ raise RuntimeError("unreachable")
132
+
133
+ def _soft(self, logits: dict, mode: str) -> dict[str, float]:
134
+ """Softmax over one mode\'s outputs, at that mode\'s temperature."""
135
+ v = {n: z / self._temp(mode) for n, z in logits.items()}
136
+ m = max(v.values())
137
+ e = {n: math.exp(x - m) for n, x in v.items()}
138
+ s = sum(e.values())
139
+ return {n: x / s for n, x in e.items()}
140
+
141
+ def choice(self, state: str, instructions: str, options: dict[str, str]) -> "Choice":
142
+ """Pick one of several options. Option names are yours and do not reach the prompt."""
143
+ return self._pick(state, instructions, options, "choice")
144
+
145
+ def scale(self, state: str, instructions: str, levels: dict[str, str]) -> "Choice":
146
+ """Pick a level of an ordinal rubric."""
147
+ return self._pick(state, instructions, levels, "score")
148
+
149
+ def _pick(self, state: str, instructions: str, options: dict[str, str], kind: str) -> "Choice":
150
+ if len(options) < 2:
151
+ raise ValueError(f"{kind} takes at least two options")
152
+ q = {"type": kind, "instructions": instructions, "criteria": dict(options)}
153
+ data, ms = self._post({"state": state, "model": self.model, "questions": {"q": q}})
154
+ self.calibration = data.get("calibration", self.calibration)
155
+ z = data["answers"]["q"]["logits"]
156
+ p = self._soft(z, "scale" if kind == "score" else "choice")
157
+ lead = max(p, key=p.get)
158
+ return Choice(p=p, verdict=lead, confidence=p[lead], logits=z, ms=ms,
159
+ input_tokens=int(data.get("usage", {}).get("input_tokens", 0)),
160
+ model=str(data.get("model", self.model)))
161
+
162
+ def _three(self, logits: dict, mode: str = "ternary") -> dict[str, float]:
163
+ """Three probabilities from three logits, calibrated for the mode asked for."""
164
+ key = {"ternary": "three_answers", "binary": "verdict"}.get(mode, mode)
165
+ v = {n: z / self._temp(key) for n, z in logits.items()}
166
+ m = max(v.values())
167
+ e = {n: math.exp(x - m) for n, x in v.items()}
168
+ s = sum(e.values())
169
+ return {n: x / s for n, x in e.items()}
170
+
171
+ def tfu(self, state: str, instructions: str, true: str | None = None,
172
+ false: str | None = None) -> Ternary:
173
+ """The three answers, unreduced — our type, beside the compatible one.
174
+
175
+ The criteria are still two: `unknown` is not something anyone states, it is what is left
176
+ when neither of the two fits. The temperature of this client applies here as well.
177
+ """
178
+ q: dict = {"type": "tfu", "instructions": instructions}
179
+ if true is not None or false is not None:
180
+ q["criteria"] = {"true": true or "", "false": false or ""}
181
+ data, ms = self._post({"state": state, "model": self.model, "questions": {"q": q}})
182
+ self.calibration = data.get("calibration", self.calibration)
183
+ a = data["answers"]["q"]
184
+ z = a["logits"]
185
+ p = self._three(z, "ternary")
186
+ lead = max(p, key=p.get)
187
+ return Ternary(p=p, verdict=lead, confidence=p[lead], logits=z, ms=ms,
188
+ input_tokens=int(data.get("usage", {}).get("input_tokens", 0)),
189
+ model=str(data.get("model", self.model)))
190
+
191
+ def ask(self, state: str, questions: dict) -> dict:
192
+ """Any number of questions about one state. The state is read once for the bundle;
193
+ results equal asking one by one."""
194
+ if not questions:
195
+ raise ValueError("no questions")
196
+ data, _ = self._post({"state": state, "model": self.model, "questions": questions})
197
+ self.calibration = data.get("calibration", self.calibration)
198
+ return {"answers": data["answers"],
199
+ "input_tokens": int(data.get("usage", {}).get("input_tokens", 0))}
200
+
201
+ def noul(self, state: str, instructions: str, true: str | None = None,
202
+ false: str | None = None) -> Answer:
203
+ """One closed question. `margin` is the log-odds, so a saturated probability still ranks."""
204
+ q: dict = {"type": "noul", "instructions": instructions}
205
+ if true is not None or false is not None:
206
+ q["criteria"] = {"true": true or "", "false": false or ""}
207
+ data, ms = self._post({"state": state, "model": self.model, "questions": {"q": q}})
208
+ self.calibration = data.get("calibration", self.calibration)
209
+ a = data["answers"]["q"]
210
+ z = a.get("logits")
211
+ if z:
212
+ # From logits everything follows exactly: the margin is a difference, not the logarithm of
213
+ # a rounded probability, so on confident documents it does not hit a clamp.
214
+ names = list(z)
215
+ p3 = self._three(z, "binary")
216
+ decided = p3[names[0]] + p3[names[1]]
217
+ prob = p3[names[0]] / decided if decided > 0 else 0.5
218
+ margin = (z[names[0]] - z[names[1]]) / self._temp("verdict")
219
+ unknown = p3[names[2]] if len(names) > 2 else 0.0
220
+ else: # the service sent a probability and nothing else
221
+ prob = float(a["noul"])
222
+ pc = min(max(prob, 1e-6), 1 - 1e-6)
223
+ margin, unknown, p3 = math.log(pc / (1 - pc)), float(a.get("unknown", 0.0)), None
224
+ return Answer(prob=prob, margin=margin, ms=ms,
225
+ input_tokens=int(data.get("usage", {}).get("input_tokens", 0)),
226
+ model=str(data.get("model", self.model)),
227
+ unknown=unknown, logits=z, p=p3)