thx01 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.
thx01/__init__.py ADDED
@@ -0,0 +1,60 @@
1
+ """THX-01: fast, non-autoregressive decision engine with calibrated probabilities."""
2
+
3
+ from .agent import Agent, RLAgent, load
4
+ from .decide import decide as _decide
5
+ from .common import (
6
+ QTYPES,
7
+ QTYPE_NAMES,
8
+ confidence_from_probs,
9
+ ece_score,
10
+ proper_reward,
11
+ render_options,
12
+ td_lambda_targets,
13
+ )
14
+ from .email import clean_email_body, email_state
15
+ from .lang import analyse as detect_language
16
+ from .lang import detect_script, is_english
17
+ from .presets import (
18
+ email_questions,
19
+ guard_questions,
20
+ moderation_questions,
21
+ router_questions,
22
+ triage_questions,
23
+ )
24
+ from .shortlist import embed_fn_from_agent, predict_shortlist, shortlist_choice
25
+
26
+ __version__ = "1.0.0"
27
+ __all__ = [
28
+ "Agent",
29
+ "RLAgent",
30
+ "load",
31
+ "shortlist_choice",
32
+ "predict_shortlist",
33
+ "embed_fn_from_agent",
34
+ "detect_language",
35
+ "detect_script",
36
+ "is_english",
37
+ "clean_email_body",
38
+ "email_questions",
39
+ "email_state",
40
+ "guard_questions",
41
+ "moderation_questions",
42
+ "router_questions",
43
+ "triage_questions",
44
+ "proper_reward",
45
+ "td_lambda_targets",
46
+ "ece_score",
47
+ "confidence_from_probs",
48
+ "render_options",
49
+ "QTYPES",
50
+ "QTYPE_NAMES",
51
+ "__version__",
52
+ ]
53
+
54
+
55
+ def _agent_decide(self, state, questions):
56
+ """Answer any THX-01 question set (choice, noul, score, number, excerpt, cite) exactly like the HTTP API."""
57
+ return _decide(self, state, questions)
58
+
59
+
60
+ Agent.decide = _agent_decide
thx01/agent.py ADDED
@@ -0,0 +1,519 @@
1
+ """High-level inference runtime for thx01 System 1 decision models."""
2
+ import json
3
+ import os
4
+ import warnings
5
+ from typing import Any, Dict, Optional, Union
6
+
7
+ import numpy as np
8
+ import torch
9
+
10
+ from .common import (
11
+ QTYPES,
12
+ TEMP_MAX,
13
+ TEMP_MIN,
14
+ amp_dtype,
15
+ build_model,
16
+ build_sequence,
17
+ clamp_temperature,
18
+ collate_items,
19
+ confidence_from_probs,
20
+ render_options,
21
+ temp_bucket,
22
+ )
23
+
24
+
25
+ def _fix_tokenizer_config(path: str):
26
+ """Ensure tokenizer_config.json can be loaded across all transformers versions."""
27
+ cfg_file = os.path.join(path, "tokenizer", "tokenizer_config.json")
28
+ if not os.path.exists(cfg_file):
29
+ return
30
+ try:
31
+ with open(cfg_file) as f:
32
+ tcfg = json.load(f)
33
+ changed = False
34
+ if tcfg.get("tokenizer_class") in (None, "TokenizersBackend"):
35
+ tcfg["tokenizer_class"] = "PreTrainedTokenizerFast"
36
+ tcfg.pop("backend", None)
37
+ tcfg.pop("is_local", None)
38
+ changed = True
39
+ # Checkpoints built on the mmBERT/Gemma tokenizer store extra_special_tokens as a list;
40
+ # transformers expects a mapping and raises "'list' object has no attribute 'keys'",
41
+ # which makes AutoTokenizer -- and so the whole model -- fail to load.
42
+ extra = tcfg.get("extra_special_tokens")
43
+ if isinstance(extra, list):
44
+ tcfg["extra_special_tokens"] = {"extra_%d" % i: t for i, t in enumerate(extra)}
45
+ changed = True
46
+ if changed:
47
+ with open(cfg_file, "w") as f:
48
+ json.dump(tcfg, f, indent=2)
49
+ except Exception:
50
+ pass
51
+
52
+
53
+ def _verify_compatibility(model: torch.nn.Module, cfg: Dict, weights: Dict[str, torch.Tensor], model_id: str):
54
+ """Verify that the loaded checkpoint weights and config strictly match the expected architecture."""
55
+ # 1. Verify required configuration attributes
56
+ required_cfg = ["encoder", "head_layers"]
57
+ missing_cfg = [k for k in required_cfg if k not in cfg]
58
+ if missing_cfg:
59
+ raise ValueError(
60
+ f"Incompatible model config for {model_id!r}: missing configuration keys {missing_cfg}. "
61
+ f"Ensure this is a valid RL Agent decision model."
62
+ )
63
+
64
+ # 2. Check for required component prefixes
65
+ required_prefixes = ("encoder.", "type_emb.", "scorer.", "act_head.")
66
+ for prefix in required_prefixes:
67
+ if not any(k.startswith(prefix) for k in weights.keys()):
68
+ raise ValueError(
69
+ f"Incompatible model weights for {model_id!r}: checkpoint is missing '{prefix}' parameters. "
70
+ f"Expected an RL Agent decision model with encoder and decision heads."
71
+ )
72
+
73
+ # 3. Check for parameter shape mismatches
74
+ model_sd = model.state_dict()
75
+ shape_mismatches = []
76
+ missing_keys = []
77
+
78
+ for name, param in model.named_parameters():
79
+ if name not in weights:
80
+ missing_keys.append(name)
81
+ elif tuple(weights[name].shape) != tuple(param.shape):
82
+ shape_mismatches.append(f" - {name}: expected {tuple(param.shape)}, found {tuple(weights[name].shape)}")
83
+
84
+ if shape_mismatches:
85
+ err_details = "\n".join(shape_mismatches[:5])
86
+ if len(shape_mismatches) > 5:
87
+ err_details += f"\n ... and {len(shape_mismatches) - 5} more mismatched layers."
88
+ raise ValueError(
89
+ f"Model architecture mismatch for {model_id!r}:\n{err_details}\n"
90
+ f"The checkpoint weights do not match the configured model architecture."
91
+ )
92
+
93
+ if missing_keys:
94
+ raise ValueError(
95
+ f"Model weights incomplete for {model_id!r}: missing {len(missing_keys)} parameter tensors "
96
+ f"(e.g. {missing_keys[:3]})."
97
+ )
98
+
99
+
100
+ class Agent:
101
+ """System 1 decision model runtime: fast, non-autoregressive, calibrated decisions."""
102
+
103
+ def __init__(
104
+ self,
105
+ model_id_or_path: str = "doofz/THX-01",
106
+ device: Optional[str] = None,
107
+ token: Optional[str] = None,
108
+ subfolder: Optional[str] = None,
109
+ half: bool = True,
110
+ ):
111
+ """Load a THX-01 checkpoint.
112
+
113
+ `subfolder` selects one checkpoint from a repo that bundles several, e.g.
114
+ `Agent("doofz/THX-01", subfolder="multilingual")`. Only that subfolder is
115
+ downloaded, so bundling does not cost every user the whole family.
116
+ """
117
+ from safetensors.torch import load_file
118
+ from transformers import AutoTokenizer
119
+
120
+ model_dir = model_id_or_path
121
+ if not os.path.exists(model_dir):
122
+ if model_id_or_path.startswith(("/", "./", "../")) or os.path.isabs(model_id_or_path):
123
+ raise FileNotFoundError(
124
+ f"Local model path not found: {model_id_or_path!r}. "
125
+ f"Check that the directory exists and that training saved the model successfully."
126
+ )
127
+ from huggingface_hub import snapshot_download
128
+
129
+ # Restrict root checkpoints too: the default repo also contains sibling
130
+ # checkpoints, which an unfiltered snapshot would unnecessarily download.
131
+ prefix = f"{subfolder}/" if subfolder else ""
132
+ kw = {
133
+ "token": token or os.environ.get("HF_TOKEN"),
134
+ "allow_patterns": [prefix + name for name in (
135
+ "config.json", "rl_agent_config.json", "model.safetensors", "tokenizer/*", "encoder/*",
136
+ )],
137
+ }
138
+ model_dir = snapshot_download(model_id_or_path, **kw)
139
+
140
+ if subfolder:
141
+ model_dir = os.path.join(model_dir, subfolder)
142
+ if not os.path.isdir(model_dir):
143
+ raise FileNotFoundError(
144
+ f"Subfolder {subfolder!r} not found in {model_id_or_path!r}."
145
+ )
146
+
147
+ _fix_tokenizer_config(model_dir)
148
+
149
+ cfg_path = os.path.join(model_dir, "rl_agent_config.json")
150
+ if not os.path.exists(cfg_path):
151
+ raise FileNotFoundError(
152
+ f"Incompatible model: {model_id_or_path!r} does not contain 'rl_agent_config.json'. "
153
+ f"That file ships with the weights of a THX-01 checkpoint, so load one of those "
154
+ f"(e.g. 'doofz/THX-01') or a directory your own training run wrote."
155
+ )
156
+
157
+ with open(cfg_path) as f:
158
+ self.cfg = json.load(f)
159
+
160
+ weights_path = os.path.join(model_dir, "model.safetensors")
161
+ if not os.path.exists(weights_path):
162
+ raise FileNotFoundError(
163
+ f"Incompatible model: 'model.safetensors' not found in {model_id_or_path!r}."
164
+ )
165
+
166
+ # 1. Device resolution with automatic fallback
167
+ if device is not None:
168
+ target_device = torch.device(device)
169
+ if target_device.type == "cuda" and not torch.cuda.is_available():
170
+ print("Warning: CUDA requested but not available. Falling back to CPU.")
171
+ self.device = torch.device("cpu")
172
+ elif target_device.type == "mps" and not (hasattr(torch.backends, "mps") and torch.backends.mps.is_available()):
173
+ print("Warning: MPS requested but not available. Falling back to CPU.")
174
+ self.device = torch.device("cpu")
175
+ else:
176
+ self.device = target_device
177
+ else:
178
+ if torch.cuda.is_available():
179
+ self.device = torch.device("cuda")
180
+ elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
181
+ self.device = torch.device("mps")
182
+ else:
183
+ self.device = torch.device("cpu")
184
+
185
+ tok_dir = os.path.join(model_dir, "tokenizer")
186
+ self.tok = AutoTokenizer.from_pretrained(tok_dir if os.path.exists(tok_dir) else self.cfg.get("encoder"))
187
+
188
+ enc_dir = os.path.join(model_dir, "encoder")
189
+ self.model = build_model(self.cfg, encoder_dir=enc_dir if os.path.exists(enc_dir) else None)
190
+
191
+ # Load weights and verify architectural compatibility
192
+ weights = load_file(weights_path)
193
+ _verify_compatibility(self.model, self.cfg, weights, model_id_or_path)
194
+
195
+ self.model.load_state_dict(weights, strict=True)
196
+
197
+ # ModernBERT's reference_compile defaults to "auto" and will torch.compile the encoder.
198
+ # That is a loss for the batch sizes THX-01 runs (a handful of questions per call) and can
199
+ # hang on some platforms, so keep the eager path.
200
+ try:
201
+ self.model.encoder.config.reference_compile = False
202
+ except Exception:
203
+ pass
204
+
205
+ # Keep what the checkpoint shipped for inspection, but only ever apply clamped values:
206
+ # some buckets are fitted to sharpen rather than soften (see clamp_temperature).
207
+ self.temperature_raw = self.cfg.get("temperature", [1.0, 1.0, 1.0])
208
+ self.temperature_by_options_raw = self.cfg.get("temperature_by_options", {})
209
+ self.temperature = [clamp_temperature(t) for t in self.temperature_raw]
210
+ self.temperature_by_options = {k: clamp_temperature(v)
211
+ for k, v in self.temperature_by_options_raw.items()}
212
+ rejected = ["%s=%.4g" % (k, float(v)) for k, v in self.temperature_by_options_raw.items()
213
+ if clamp_temperature(v) != float(v)]
214
+ rejected += ["temperature[%d]=%.4g" % (i, float(t)) for i, t in enumerate(self.temperature_raw)
215
+ if clamp_temperature(t) != float(t)]
216
+ if rejected:
217
+ warnings.warn(
218
+ "thx01: this checkpoint ships temperatures outside [%g, %g] which would distort "
219
+ "confidence; clamping %s. Treat confidence from the affected buckets as uncalibrated."
220
+ % (TEMP_MIN, TEMP_MAX, ", ".join(rejected)),
221
+ RuntimeWarning, stacklevel=2)
222
+ self.dtype = amp_dtype(self.cfg.get("amp_dtype", "fp16"))
223
+
224
+ if self.device.type == "cuda" and torch.cuda.get_device_capability(self.device)[0] < 8:
225
+ self.dtype = torch.float16
226
+ elif self.device.type in ("cpu", "mps"):
227
+ self.dtype = torch.float32
228
+
229
+ # 2. Place on device with graceful fallback to CPU on memory error
230
+ fell_back_from = fell_back_why = None
231
+ try:
232
+ self.model.to(self.device).eval()
233
+ if self.device.type == "cuda" and half:
234
+ # Store the weights in the autocast dtype instead of fp32: halves VRAM
235
+ # (421M -> ~0.85 GB, 322M -> ~0.65 GB) with no change to what autocast computes.
236
+ self.model.to(self.dtype)
237
+ except (RuntimeError, torch.cuda.OutOfMemoryError) as e:
238
+ if self.device.type != "cpu":
239
+ # Record what actually went wrong: the reason matters more than the symptom,
240
+ # and it is the only place the underlying exception is ever surfaced.
241
+ fell_back_from, fell_back_why = self.device, e
242
+ self.device = torch.device("cpu")
243
+ self.dtype = torch.float32
244
+ self.model.to(self.device).eval()
245
+ else:
246
+ raise e
247
+
248
+ if fell_back_from is not None:
249
+ print(
250
+ "\n[thx01] Warning: could not place the model on %s, so it is running on CPU.\n"
251
+ " Reason: %s\n"
252
+ " Inference will be roughly 10-15x slower (~200-500 ms rather than ~35 ms).\n"
253
+ " If this is a newer NVIDIA GPU (Blackwell / RTX 50-series), your PyTorch build\n"
254
+ " may not support its CUDA architecture:\n"
255
+ " pip install --pre torch --index-url https://download.pytorch.org/whl/nightly/cu128\n"
256
+ " See https://pytorch.org/get-started/locally/\n"
257
+ % (fell_back_from, fell_back_why), flush=True)
258
+
259
+ @staticmethod
260
+ def _to_internal(qdef: Dict) -> Dict:
261
+ t = qdef["type"]
262
+ crit = qdef.get("criteria")
263
+ if t == "choice" and isinstance(crit, list):
264
+ crit = {c: None for c in crit}
265
+ ins = qdef["instructions"]
266
+ if not isinstance(ins, str):
267
+ ins = json.dumps(ins)
268
+ return {"t": t, "ins": ins, "crit": crit}
269
+
270
+ def _head_budget(self, q: Dict) -> int:
271
+ """Head (instruction + options) token budget for one question.
272
+
273
+ The checkpoint default is kept for ordinary questions. When the options need more room
274
+ (e.g. 77 intents), the budget grows to fit them -- up to half the context -- instead of
275
+ truncating every label to 3-4 tokens, which is what made large label sets collapse.
276
+ """
277
+ base = self.cfg.get("head_max_len", 192)
278
+ if not self.cfg.get("auto_head_budget", True):
279
+ return base
280
+ opts = render_options(q)
281
+ need = 24 + sum(min(49, len(self.tok(" " + o, add_special_tokens=False)["input_ids"]) + 1) for o in opts)
282
+ return int(min(max(base, need), self.cfg.get("max_len", 512) // 2))
283
+
284
+ def _decode(self, questions: Dict[str, Dict[str, Any]], items, logits: np.ndarray, act: np.ndarray,
285
+ offset: int = 0) -> Dict[str, Any]:
286
+ """Turn raw marker logits for one state's questions into calibrated typed answers."""
287
+ ids = list(questions.keys())
288
+ answers = {}
289
+ for r, qid in enumerate(ids):
290
+ q = self._to_internal(questions[qid])
291
+ k = len(items[r]["markers"])
292
+ qt = QTYPES[q["t"]]
293
+ t_scale = self.temperature_by_options.get(temp_bucket(qt, k), self.temperature[qt])
294
+ z = logits[offset + r, :k] / t_scale
295
+ p = np.exp(z - z.max())
296
+ p = p / p.sum()
297
+
298
+ conf_score = round(confidence_from_probs(p, k), 4)
299
+ ext = {"act_probability": round(float(act[offset + r, 0]), 4)}
300
+
301
+ if q["t"] == "choice":
302
+ keys = list(q["crit"].keys())
303
+ answers[qid] = {
304
+ "type": "choice",
305
+ "choice": keys[int(p.argmax())],
306
+ "probabilities": {kk: round(float(v), 4) for kk, v in zip(keys, p)},
307
+ "confidence": conf_score,
308
+ "action": ext,
309
+ }
310
+ elif q["t"] == "score":
311
+ exp_score = float((np.arange(k) * p).sum())
312
+ answers[qid] = {
313
+ "type": "score",
314
+ "score": round(exp_score, 4),
315
+ "legend": {str(i): c for i, c in enumerate(q["crit"])},
316
+ "probabilities": {str(i): round(float(v), 4) for i, v in enumerate(p)},
317
+ "confidence": conf_score,
318
+ "action": ext,
319
+ }
320
+ else:
321
+ answers[qid] = {
322
+ "type": "noul",
323
+ "noul": round(float(p[1]), 4),
324
+ "confidence": round(max(float(p[1]), 1.0 - float(p[1])), 4),
325
+ "action": ext,
326
+ }
327
+
328
+ return answers
329
+
330
+ def _tournament(self, states, qdefs, chunk: int = 20, keep: int = 3):
331
+ """Coarse-to-fine answer for choice questions with many options (two batched passes).
332
+
333
+ Round 1 scores each state against every chunk of ~`chunk` options and keeps the top `keep`
334
+ of each; round 2 chooses among those finalists. Returns one answer dict per state, with
335
+ probabilities over the finalists (every other option gets 0).
336
+ """
337
+ rng = __import__("random").Random(0)
338
+ s1, q1, spans = [], [], []
339
+ for st, q in zip(states, qdefs):
340
+ keys = list(q["criteria"])
341
+ rng.shuffle(keys)
342
+ chunks = [keys[i:i + chunk] for i in range(0, len(keys), chunk)]
343
+ spans.append(len(chunks))
344
+ for c in chunks:
345
+ s1.append(st)
346
+ q1.append({"q": {**q, "criteria": {k: q["criteria"][k] for k in c}}})
347
+ o1 = self.predict_batch(s1, q1, _tournament=False) if s1 else []
348
+ s2, q2, off = [], [], 0
349
+ for st, q, n in zip(states, qdefs, spans):
350
+ fin = []
351
+ for o in o1[off:off + n]:
352
+ pr = o["answers"]["q"]["probabilities"]
353
+ fin += sorted(pr, key=pr.get, reverse=True)[:keep]
354
+ off += n
355
+ s2.append(st)
356
+ q2.append({"q": {**q, "criteria": {k: q["criteria"][k] for k in fin}}})
357
+ o2 = self.predict_batch(s2, q2, _tournament=False)
358
+ out = []
359
+ for q, o in zip(qdefs, o2):
360
+ a = o["answers"]["q"]
361
+ a["probabilities"] = {k: a["probabilities"].get(k, 0.0) for k in q["criteria"]}
362
+ a["mode"] = "tournament"
363
+ out.append(a)
364
+ return out
365
+
366
+ @torch.no_grad()
367
+ def predict_batch(self, states, questions, batch_size: int = 64, _tournament: bool = True):
368
+ """Answer many states at once (throughput mode for evals and bulk jobs).
369
+
370
+ `questions` is either one schema shared by every state or a list with one schema per
371
+ state. Returns one `system_one`-style result per state, in order.
372
+ """
373
+ from .common import collate_items as _collate
374
+ qlist = questions if isinstance(questions, list) else [questions] * len(states)
375
+ limit = self.cfg.get("tournament_above", 24)
376
+ if _tournament and limit:
377
+ big = [(i, qid) for i, qs in enumerate(qlist) for qid, q in qs.items()
378
+ if q.get("type") == "choice" and isinstance(q.get("criteria"), (dict, list)) and len(q["criteria"]) > limit]
379
+ if big:
380
+ norm = lambda q: {**q, "criteria": q["criteria"] if isinstance(q["criteria"], dict) else {c: None for c in q["criteria"]}}
381
+ answers = self._tournament([states[i] for i, _ in big], [norm(qlist[i][qid]) for i, qid in big])
382
+ rest = [{k: v for k, v in qs.items() if (i, k) not in set(big)} for i, qs in enumerate(qlist)]
383
+ keep = [i for i, r in enumerate(rest) if r]
384
+ sub = self.predict_batch([states[i] for i in keep], [rest[i] for i in keep], batch_size, _tournament=False) if keep else []
385
+ res = [{"model": "THX-01", "answers": {}, "usage": {"input_tokens": 0, "output_tokens": 0}} for _ in states]
386
+ for i, r in zip(keep, sub):
387
+ res[i] = r
388
+ for (i, qid), a in zip(big, answers):
389
+ res[i]["answers"][qid] = a
390
+ for i, qs in enumerate(qlist): # restore the caller's question order
391
+ res[i]["answers"] = {k: res[i]["answers"][k] for k in qs}
392
+ return res
393
+ max_len = self.cfg.get("max_len", 512)
394
+ head_max_len = self.cfg.get("head_max_len", 192)
395
+ per_state = []
396
+ flat = []
397
+ for st, qs in zip(states, qlist):
398
+ its = []
399
+ for qid in qs:
400
+ q = self._to_internal(qs[qid])
401
+ head_max_len = self._head_budget(q)
402
+ seq, markers = build_sequence(self.tok, st, q, max_len, head_max_len)
403
+ if len(markers) != len(render_options(q)):
404
+ raise ValueError("question %r options exceed head_max_len=%d" % (qid, head_max_len))
405
+ its.append({"ids": seq, "markers": markers, "qtype": QTYPES[q["t"]]})
406
+ per_state.append(its)
407
+ flat.extend(its)
408
+ # sort by length so each chunk pads as little as possible
409
+ order = sorted(range(len(flat)), key=lambda i: len(flat[i]["ids"]))
410
+ kmax = max(len(it["markers"]) for it in flat)
411
+ all_logits = np.full((len(flat), kmax), -1e4, dtype=np.float32)
412
+ all_act = np.zeros((len(flat), 2), dtype=np.float32)
413
+ use_amp = self.device.type == "cuda"
414
+ for c in range(0, len(order), batch_size):
415
+ idx = order[c:c + batch_size]
416
+ bt = _collate([[flat[i] for i in idx]], self.tok.pad_token_id)
417
+ with torch.autocast(device_type=self.device.type, dtype=self.dtype, enabled=use_amp):
418
+ lg, ac = self.model(bt["input_ids"].to(self.device), bt["attention_mask"].to(self.device),
419
+ bt["marker_pos"].to(self.device), bt["marker_mask"].to(self.device),
420
+ bt["qtype"].to(self.device))
421
+ lg = lg.float().cpu().numpy()
422
+ ac = torch.softmax(ac.float(), -1).cpu().numpy()
423
+ for j, i in enumerate(idx):
424
+ all_logits[i, :lg.shape[1]] = lg[j]
425
+ all_act[i, :ac.shape[1]] = ac[j]
426
+ out, off = [], 0
427
+ for qs, its in zip(qlist, per_state):
428
+ out.append({"model": "THX-01",
429
+ "answers": self._decode(qs, its, all_logits, all_act, offset=off),
430
+ "usage": {"input_tokens": sum(len(it["ids"]) for it in its), "output_tokens": 0}})
431
+ off += len(its)
432
+ return out
433
+
434
+ @torch.no_grad()
435
+ def system_one(self, state: Union[str, dict, list], questions: Dict[str, Dict[str, Any]]) -> Dict[str, Any]:
436
+ """Evaluate typed questions across state in a single, parallel forward pass.
437
+
438
+ Args:
439
+ state: Text string, JSON dict, or conversation turn list.
440
+ questions: Dictionary mapping question_id -> question definition.
441
+ - choice: {"type": "choice", "instructions": "...", "criteria": {"optA": "...", ...}}
442
+ - score: {"type": "score", "instructions": "...", "criteria": ["lvl0", "lvl1", ...]}
443
+ - noul: {"type": "noul", "instructions": "..."}
444
+
445
+ Returns:
446
+ Dictionary with answers, probabilities, calibrated confidence, and token usage.
447
+ """
448
+ limit = self.cfg.get("tournament_above", 24)
449
+ if limit and any(q.get("type") == "choice" and isinstance(q.get("criteria"), (dict, list))
450
+ and len(q["criteria"]) > limit for q in questions.values()):
451
+ return self.predict_batch([state], questions)[0]
452
+ ids = list(questions.keys())
453
+ items = []
454
+ max_len = self.cfg.get("max_len", 512)
455
+ head_max_len = self.cfg.get("head_max_len", 192)
456
+
457
+ for qid in ids:
458
+ q = self._to_internal(questions[qid])
459
+ head_max_len = self._head_budget(q)
460
+ seq, markers = build_sequence(self.tok, state, q, max_len, head_max_len)
461
+ if len(markers) != len(render_options(q)):
462
+ raise ValueError("question %r options exceed head_max_len=%d" % (qid, head_max_len))
463
+ items.append({"ids": seq, "markers": markers, "qtype": QTYPES[q["t"]]})
464
+
465
+ b = collate_items([items], self.tok.pad_token_id)
466
+ use_amp = self.device.type == "cuda"
467
+
468
+ try:
469
+ with torch.autocast(device_type=self.device.type, dtype=self.dtype, enabled=use_amp):
470
+ logits, act = self.model(
471
+ b["input_ids"].to(self.device),
472
+ b["attention_mask"].to(self.device),
473
+ b["marker_pos"].to(self.device),
474
+ b["marker_mask"].to(self.device),
475
+ b["qtype"].to(self.device),
476
+ )
477
+ except (RuntimeError, torch.cuda.OutOfMemoryError) as e:
478
+ if self.device.type != "cpu" and ("memory" in str(e).lower() or "cuda" in str(e).lower()):
479
+ print("Warning: GPU memory exceeded during inference. Falling back to CPU...")
480
+ self.device = torch.device("cpu")
481
+ self.dtype = torch.float32
482
+ self.model.to(self.device)
483
+ logits, act = self.model(
484
+ b["input_ids"].to(self.device),
485
+ b["attention_mask"].to(self.device),
486
+ b["marker_pos"].to(self.device),
487
+ b["marker_mask"].to(self.device),
488
+ b["qtype"].to(self.device),
489
+ )
490
+ else:
491
+ raise e
492
+
493
+ logits = logits.float().cpu().numpy()
494
+ act = torch.softmax(act.float(), -1).cpu().numpy()
495
+ n_tokens = int(b["attention_mask"].sum())
496
+ answers = self._decode(questions, items, logits, act)
497
+
498
+ return {
499
+ "model": "THX-01",
500
+ "answers": answers,
501
+ "usage": {"input_tokens": n_tokens, "output_tokens": 0},
502
+ }
503
+
504
+ predict = system_one
505
+
506
+
507
+ RLAgent = Agent
508
+
509
+
510
+ def load(model_id_or_path: str = "doofz/THX-01", device: Optional[str] = None,
511
+ token: Optional[str] = None, subfolder: Optional[str] = None, half: bool = True) -> Agent:
512
+ """Load a THX-01 agent.
513
+
514
+ `subfolder` picks one checkpoint out of a repo that bundles several:
515
+
516
+ thx01.load("doofz/THX-01") # English (repo root)
517
+ thx01.load("doofz/THX-01", subfolder="multilingual")
518
+ """
519
+ return Agent(model_id_or_path, device=device, token=token, subfolder=subfolder, half=half)