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 +44 -0
- typecastlm/calibrate.py +82 -0
- typecastlm/cli.py +59 -0
- typecastlm/local.py +211 -0
- typecastlm/remote.py +227 -0
- typecastlm/server.py +357 -0
- typecastlm-1.0.0.dist-info/METADATA +168 -0
- typecastlm-1.0.0.dist-info/RECORD +12 -0
- typecastlm-1.0.0.dist-info/WHEEL +4 -0
- typecastlm-1.0.0.dist-info/entry_points.txt +3 -0
- typecastlm-1.0.0.dist-info/licenses/LICENSE +202 -0
- typecastlm-1.0.0.dist-info/licenses/NOTICE +12 -0
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"
|
typecastlm/calibrate.py
ADDED
|
@@ -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)
|