brier 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.
brier/__init__.py ADDED
@@ -0,0 +1,8 @@
1
+ """brier: calibrated typed decisions from open LLMs."""
2
+
3
+ from brier._version import __version__
4
+ from brier.decider import Decider
5
+ from brier.decision import Decision
6
+ from brier.questions import Choice, Noul, Score
7
+
8
+ __all__ = ["Choice", "Decider", "Decision", "Noul", "Score", "__version__"]
brier/__main__.py ADDED
@@ -0,0 +1,73 @@
1
+ """Command line: ``brier check <model-id>`` (also ``python -m brier check``)."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import argparse
6
+ import json
7
+ import sys
8
+ from collections.abc import Sequence
9
+ from typing import Any
10
+
11
+ from brier.backends.base import Backend
12
+ from brier.check import check_backend, format_report
13
+
14
+
15
+ def build_parser() -> argparse.ArgumentParser:
16
+ """Argument parser for the ``brier`` command."""
17
+ parser = argparse.ArgumentParser(prog="brier")
18
+ sub = parser.add_subparsers(dest="command", required=True)
19
+ check = sub.add_parser("check", help="check whether a model works with brier")
20
+ check.add_argument("model", help="Hugging Face model id")
21
+ check.add_argument("--revision", default=None, help="model commit sha (recommended)")
22
+ check.add_argument("--dtype", default=None, choices=("float32", "bfloat16", "float16"))
23
+ check.add_argument("--device", default=None)
24
+ check.add_argument("--batch-size", type=int, default=8)
25
+ check.add_argument("--attn-implementation", default=None, help="e.g. eager (Gemma 3)")
26
+ check.add_argument("--json", action="store_true", help="print the report as JSON")
27
+ return parser
28
+
29
+
30
+ def _make_backend(args: Any) -> Backend:
31
+ from brier.backends.hf import HFBackend
32
+
33
+ return HFBackend(
34
+ args.model,
35
+ revision=args.revision,
36
+ dtype=args.dtype,
37
+ device=args.device,
38
+ batch_size=args.batch_size,
39
+ attn_implementation=args.attn_implementation,
40
+ )
41
+
42
+
43
+ def _safe_output() -> None:
44
+ """Escape characters the console cannot encode instead of crashing (e.g. cp1252 on Windows).
45
+
46
+ Error details quote tokenizer messages, which may contain characters like ``Ġ``.
47
+ """
48
+ for stream in (sys.stdout, sys.stderr):
49
+ reconfigure = getattr(stream, "reconfigure", None)
50
+ if reconfigure is not None:
51
+ reconfigure(errors="backslashreplace")
52
+
53
+
54
+ def main(argv: Sequence[str] | None = None) -> int:
55
+ """Run the CLI; returns the process exit code (1 if a required check fails)."""
56
+ args = build_parser().parse_args(argv)
57
+ _safe_output()
58
+ try:
59
+ backend = _make_backend(args)
60
+ except Exception as exc: # any load failure is the answer to "does it work"
61
+ message = f"{type(exc).__name__}: {exc}"[:500]
62
+ if args.json:
63
+ print(json.dumps({"model": {"id": args.model}, "ok": False, "load_error": message}))
64
+ else:
65
+ print(f"brier check: {args.model}\n\n FAIL load {message}\n\nResult: FAILED")
66
+ return 1
67
+ report = check_backend(backend)
68
+ print(json.dumps(report.to_dict(), indent=2) if args.json else format_report(report))
69
+ return 0 if report.ok else 1
70
+
71
+
72
+ if __name__ == "__main__":
73
+ sys.exit(main())
brier/_math.py ADDED
@@ -0,0 +1,57 @@
1
+ """Numerically stable log-space helpers (float64)."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import numpy as np
6
+ import numpy.typing as npt
7
+
8
+ FloatArray = npt.NDArray[np.float64]
9
+
10
+
11
+ def logsumexp(x: npt.ArrayLike, axis: int = -1) -> FloatArray:
12
+ """Stable ``log(sum(exp(x)))`` along ``axis``.
13
+
14
+ Parameters
15
+ ----------
16
+ x : array_like
17
+ Input values; ``-inf`` entries count as zero mass.
18
+ axis : int
19
+ Axis to reduce.
20
+
21
+ Returns
22
+ -------
23
+ numpy.ndarray
24
+ float64 array with ``axis`` removed (``-inf`` where every entry is ``-inf``).
25
+ """
26
+ a = np.asarray(x, dtype=np.float64)
27
+ m = np.max(a, axis=axis, keepdims=True)
28
+ m = np.where(np.isfinite(m), m, 0.0) # all -inf rows: avoid -inf - -inf = nan
29
+ with np.errstate(divide="ignore"):
30
+ out = np.log(np.sum(np.exp(a - m), axis=axis, keepdims=True)) + m
31
+ squeezed: FloatArray = np.squeeze(out, axis=axis)
32
+ return squeezed
33
+
34
+
35
+ def norm(x: npt.ArrayLike, axis: int = -1) -> FloatArray:
36
+ """Log-space renormalisation ``x - logsumexp(x)`` (METHODS.md ``norm``).
37
+
38
+ Parameters
39
+ ----------
40
+ x : array_like
41
+ Unnormalised log-probabilities.
42
+ axis : int
43
+ Axis along which the result sums to 1 in probability space.
44
+
45
+ Returns
46
+ -------
47
+ numpy.ndarray
48
+ float64 log-probabilities with the same shape as ``x``. A slice with no
49
+ mass (all ``-inf``) has no distribution and comes back as ``nan``; callers
50
+ reject it (``Decision`` refuses non-finite probabilities).
51
+ """
52
+ a = np.asarray(x, dtype=np.float64)
53
+ with np.errstate(invalid="ignore"):
54
+ return a - np.expand_dims(logsumexp(a, axis=axis), axis)
55
+
56
+
57
+ log_softmax = norm
brier/_version.py ADDED
@@ -0,0 +1,8 @@
1
+ """Installed package version (its own module so internal imports avoid import cycles)."""
2
+
3
+ from importlib.metadata import PackageNotFoundError, version
4
+
5
+ try:
6
+ __version__ = version("brier")
7
+ except PackageNotFoundError: # pragma: no cover - running from a source tree without install
8
+ __version__ = "0.0.0"
brier/artifacts.py ADDED
@@ -0,0 +1,519 @@
1
+ """Calibration artifacts: a directory with ``artifact.json`` + ``arrays.npz`` (ADR-0003).
2
+
3
+ Loading treats the files as untrusted (THREAT_MODEL T2): fixed file names (no paths from
4
+ the JSON); each file opened once, must be a regular non-symlink file, read with a size cap;
5
+ strict JSON schema (unknown keys rejected); SHA-256 of the npz verified; npz members
6
+ checked (count, names, compression, encryption, size) and each ``.npy`` header parsed by
7
+ hand so its shape, dtype and order are validated *before* any allocation (no pickle, no
8
+ ``np.load``); model id, revision and prompt-template hash must match. Every failure,
9
+ including unexpected ones, raises :class:`ArtifactError`.
10
+ """
11
+
12
+ from __future__ import annotations
13
+
14
+ import hashlib
15
+ import io
16
+ import json
17
+ import math
18
+ import os
19
+ import re
20
+ import stat
21
+ import zipfile
22
+ from collections.abc import Callable
23
+ from dataclasses import dataclass
24
+ from pathlib import Path
25
+ from typing import Any
26
+
27
+ import numpy as np
28
+
29
+ from brier import prompts
30
+ from brier._math import FloatArray
31
+ from brier._version import __version__
32
+ from brier.errors import ArtifactError, BrierError, QuestionError
33
+ from brier.heads.fitted import MAX_HIDDEN, L2Head
34
+ from brier.questions import Choice, Noul, Question, Score
35
+
36
+ SCHEMA_VERSION = 3 # ADR-0007: adds model.dtype; ADR-0006: L2 heads; 1 and 2 still load
37
+ _READABLE_VERSIONS = (1, 2, 3)
38
+ _DTYPE = re.compile(r"[a-z0-9_.+-]{1,32}")
39
+ MAX_ARTIFACT_QUESTIONS = 1024 # bounds parsing work (each head is four npz members)
40
+ JSON_FILE = "artifact.json"
41
+ ARRAYS_FILE = "arrays.npz"
42
+ DEFAULT_MAX_BYTES = 100_000_000
43
+ TEMPERATURE_RANGE = (1e-6, 1e6)
44
+ _SUM_TOL = 1e-6
45
+ _MAX_MEMBER_BYTES = 4096 # a prior has at most 26 float64s plus a .npy header
46
+ _MAX_HEAD_MEMBER_BYTES = MAX_HIDDEN * 26 * 8 + 4096 # (d, C) float64 weights plus header
47
+ _HEAD_KEYS = {
48
+ "layer",
49
+ "solver",
50
+ "alpha",
51
+ "temperature",
52
+ "temperature_at_bound",
53
+ "oof_nll",
54
+ "oof_accuracy",
55
+ }
56
+ _HEAD_ARRAYS = ("mean", "scale", "weights", "bias")
57
+ _F8 = np.dtype("<f8")
58
+ _NPY: Any = np.lib.format # untyped in older numpy stubs (Python 3.10 resolves numpy 2.2)
59
+ _TOP_KEYS = {
60
+ "schema_version",
61
+ "brier_version",
62
+ "model",
63
+ "template_hash",
64
+ "prior_strength",
65
+ "arrays_sha256",
66
+ "questions",
67
+ }
68
+ _QUESTION_KEYS = {
69
+ "choice": {"type", "name", "text", "options", "prior", "temperature"},
70
+ "noul": {"type", "name", "text", "prior", "temperature"},
71
+ "score": {"type", "name", "text", "levels", "labels", "prior", "temperature"},
72
+ }
73
+
74
+
75
+ def _n_answers(q: Question) -> int:
76
+ if isinstance(q, Choice):
77
+ return len(q.options)
78
+ return 2 if isinstance(q, Noul) else q.levels
79
+
80
+
81
+ def _is_number(x: object) -> bool:
82
+ return isinstance(x, (int, float)) and not isinstance(x, bool)
83
+
84
+
85
+ @dataclass(frozen=True, eq=False)
86
+ class Calibration:
87
+ """Fitted calibration of one question.
88
+
89
+ Parameters
90
+ ----------
91
+ question : Choice, Noul or Score
92
+ The question (calibration is keyed by the whole question).
93
+ prior : array_like or None
94
+ L0 prior: a strictly positive distribution over the question's answers.
95
+ temperature : float or None
96
+ L1 temperature in ``[1e-6, 1e6]``.
97
+ head : L2Head or None
98
+ Fitted L2 head (ADR-0006).
99
+ """
100
+
101
+ question: Question
102
+ prior: FloatArray | None = None
103
+ temperature: float | None = None
104
+ head: L2Head | None = None
105
+
106
+ def __post_init__(self) -> None:
107
+ if not isinstance(self.question, (Choice, Noul, Score)):
108
+ raise ArtifactError("question must be a Choice, Noul or Score")
109
+ name = self.question.name
110
+ if self.prior is not None:
111
+ p = np.array(self.prior, dtype=np.float64)
112
+ if p.shape != (_n_answers(self.question),):
113
+ raise ArtifactError(f"prior for {name!r} has the wrong shape")
114
+ if not np.all(np.isfinite(p)) or np.any(p <= 0) or abs(p.sum() - 1.0) > _SUM_TOL:
115
+ raise ArtifactError(f"prior for {name!r} is not a positive distribution")
116
+ p.flags.writeable = False
117
+ object.__setattr__(self, "prior", p)
118
+ if self.temperature is not None:
119
+ try:
120
+ t = float(self.temperature) if _is_number(self.temperature) else math.nan
121
+ except OverflowError:
122
+ t = math.nan
123
+ lo, hi = TEMPERATURE_RANGE
124
+ if not lo <= t <= hi: # NaN fails too
125
+ raise ArtifactError(f"temperature for {name!r} must be a number in [{lo}, {hi}]")
126
+ object.__setattr__(self, "temperature", t)
127
+ if self.head is not None:
128
+ if not isinstance(self.head, L2Head):
129
+ raise ArtifactError(f"head for {name!r} must be an L2Head")
130
+ if self.head.n_classes != _n_answers(self.question):
131
+ raise ArtifactError(f"head for {name!r} has the wrong number of classes")
132
+
133
+
134
+ @dataclass(frozen=True, eq=False)
135
+ class Artifact:
136
+ """Everything needed to reproduce calibrated decisions with the same model.
137
+
138
+ Parameters
139
+ ----------
140
+ model_id, revision : str, str or None
141
+ The model the calibration was fitted on.
142
+ prior_strength : float
143
+ λ used when applying L0 priors.
144
+ calibrations : tuple of Calibration
145
+ One entry per distinct question.
146
+ dtype : str or None
147
+ The backend's precision (ADR-0007); ``None`` if the backend exposes none, or for
148
+ version-1/2 artifacts, which did not record it.
149
+ """
150
+
151
+ model_id: str
152
+ revision: str | None
153
+ prior_strength: float
154
+ calibrations: tuple[Calibration, ...]
155
+ dtype: str | None = None
156
+
157
+ def __post_init__(self) -> None:
158
+ if not isinstance(self.model_id, str) or not self.model_id:
159
+ raise ArtifactError("model_id must be a non-empty string")
160
+ if self.revision is not None and not isinstance(self.revision, str):
161
+ raise ArtifactError("revision must be a string or None")
162
+ if self.dtype is not None and not _valid_dtype(self.dtype):
163
+ raise ArtifactError("dtype must be None or 1-32 characters from [a-z0-9_.+-]")
164
+ lam = self.prior_strength
165
+ if not _is_number(lam) or not 0 <= lam <= 1:
166
+ raise ArtifactError("prior_strength must be a number in [0, 1]")
167
+ object.__setattr__(self, "calibrations", tuple(self.calibrations))
168
+ if not all(isinstance(c, Calibration) for c in self.calibrations):
169
+ raise ArtifactError("calibrations must be Calibration instances")
170
+ questions = [c.question for c in self.calibrations]
171
+ if len(set(questions)) != len(questions):
172
+ raise ArtifactError("each question may appear only once")
173
+
174
+
175
+ def save_artifact(path: str | Path, artifact: Artifact) -> None:
176
+ """Write ``artifact`` to the directory ``path`` (created; must be empty if it exists).
177
+
178
+ Files are created exclusively, so nothing that appears concurrently is overwritten.
179
+
180
+ Raises
181
+ ------
182
+ ArtifactError
183
+ If ``path`` is a file or a non-empty directory.
184
+ """
185
+ d = Path(path)
186
+ if d.exists() and (not d.is_dir() or any(d.iterdir())):
187
+ raise ArtifactError(f"{d} exists and is not an empty directory")
188
+ arrays: dict[str, Any] = {
189
+ f"prior_{i}": c.prior for i, c in enumerate(artifact.calibrations) if c.prior is not None
190
+ }
191
+ for i, c in enumerate(artifact.calibrations):
192
+ if c.head is not None:
193
+ arrays.update({f"head_{i}_{a}": getattr(c.head, a) for a in _HEAD_ARRAYS})
194
+ buf = io.BytesIO()
195
+ np.savez(buf, **arrays)
196
+ data = buf.getvalue()
197
+ doc = {
198
+ "schema_version": SCHEMA_VERSION,
199
+ "brier_version": __version__,
200
+ "model": {"id": artifact.model_id, "revision": artifact.revision, "dtype": artifact.dtype},
201
+ "template_hash": prompts.template_hash(),
202
+ "prior_strength": float(artifact.prior_strength),
203
+ "arrays_sha256": hashlib.sha256(data).hexdigest(),
204
+ "questions": [_question_to_json(c) for c in artifact.calibrations],
205
+ }
206
+ d.mkdir(parents=True, exist_ok=True)
207
+ try:
208
+ with (d / ARRAYS_FILE).open("xb") as f:
209
+ f.write(data)
210
+ with (d / JSON_FILE).open("x", encoding="utf-8") as f:
211
+ f.write(json.dumps(doc, indent=2))
212
+ except FileExistsError as e:
213
+ raise ArtifactError(f"{d} is no longer empty") from e
214
+
215
+
216
+ def load_artifact(
217
+ path: str | Path,
218
+ *,
219
+ model_id: str,
220
+ revision: str | None,
221
+ dtype: str | None = None,
222
+ max_bytes: int = DEFAULT_MAX_BYTES,
223
+ ) -> Artifact:
224
+ """Load and fully validate an artifact for the given model.
225
+
226
+ Parameters
227
+ ----------
228
+ path : str or Path
229
+ Artifact directory.
230
+ model_id, revision : str, str or None
231
+ The backend's model; the artifact must have been fitted on exactly this.
232
+ dtype : str or None
233
+ The backend's precision; must equal the recorded one (version 3 artifacts only).
234
+ max_bytes : int
235
+ Cap on each file and on the uncompressed npz contents (default 100 MB).
236
+
237
+ Raises
238
+ ------
239
+ ArtifactError
240
+ On any problem with the files, including unexpected parser errors.
241
+ """
242
+ if not isinstance(max_bytes, int) or isinstance(max_bytes, bool) or max_bytes <= 0:
243
+ raise ArtifactError("max_bytes must be a positive int")
244
+ try:
245
+ return _load(Path(path), model_id, revision, dtype, max_bytes)
246
+ except ArtifactError:
247
+ raise
248
+ except Exception as e: # fail closed: hostile input must never escape as another type
249
+ raise ArtifactError(f"artifact could not be loaded safely ({type(e).__name__})") from e
250
+
251
+
252
+ def _load(
253
+ d: Path, model_id: str, revision: str | None, dtype: str | None, max_bytes: int
254
+ ) -> Artifact:
255
+ if not d.is_dir():
256
+ raise ArtifactError(f"{d} is not a directory")
257
+ raw_json = _read_capped(d / JSON_FILE, max_bytes)
258
+ raw_arrays = _read_capped(d / ARRAYS_FILE, max_bytes)
259
+ try:
260
+ doc = json.loads(raw_json.decode("utf-8"))
261
+ except (ValueError, RecursionError) as e: # covers UnicodeDecodeError, JSONDecodeError
262
+ raise ArtifactError(f"{JSON_FILE} is not valid JSON") from e
263
+ version = _check_header(doc, model_id, revision, dtype)
264
+ if hashlib.sha256(raw_arrays).hexdigest() != doc["arrays_sha256"]:
265
+ raise ArtifactError(f"{ARRAYS_FILE} does not match its recorded SHA-256")
266
+ entries = doc["questions"]
267
+ if not isinstance(entries, list):
268
+ raise ArtifactError("questions must be a list")
269
+ if len(entries) > MAX_ARTIFACT_QUESTIONS:
270
+ raise ArtifactError(f"at most {MAX_ARTIFACT_QUESTIONS} questions per artifact")
271
+ for entry in entries:
272
+ _check_question_json(entry, version)
273
+ try:
274
+ questions = [_question_from_json(entry) for entry in entries]
275
+ except QuestionError as err:
276
+ raise ArtifactError(f"invalid question in artifact: {err}") from err
277
+ specs: dict[str, _Spec] = {}
278
+ for i, (q, entry) in enumerate(zip(questions, entries, strict=True)):
279
+ k = _n_answers(q)
280
+ if entry["prior"]:
281
+ specs[f"prior_{i}"] = (_exact((k,)), _MAX_MEMBER_BYTES)
282
+ if entry.get("head") is not None:
283
+ specs[f"head_{i}_mean"] = (_hidden_vector, _MAX_HEAD_MEMBER_BYTES)
284
+ specs[f"head_{i}_scale"] = (_hidden_vector, _MAX_HEAD_MEMBER_BYTES)
285
+ specs[f"head_{i}_weights"] = (_hidden_matrix(k), _MAX_HEAD_MEMBER_BYTES)
286
+ specs[f"head_{i}_bias"] = (_exact((k,)), _MAX_MEMBER_BYTES)
287
+ arrays = _read_arrays(raw_arrays, specs, max_bytes)
288
+ calibrations = tuple(
289
+ Calibration(
290
+ q, arrays.get(f"prior_{i}"), entry["temperature"], _head(i, entry.get("head"), arrays)
291
+ )
292
+ for i, (q, entry) in enumerate(zip(questions, entries, strict=True))
293
+ )
294
+ model = doc["model"]
295
+ return Artifact(
296
+ model["id"], model["revision"], doc["prior_strength"], calibrations, model.get("dtype")
297
+ )
298
+
299
+
300
+ def _read_capped(p: Path, max_bytes: int) -> bytes:
301
+ """Open once, require a regular non-symlink file under the cap, read only its size."""
302
+ if p.is_symlink():
303
+ raise ArtifactError(f"{p.name} is a symlink")
304
+ flags = os.O_RDONLY | getattr(os, "O_BINARY", 0) | getattr(os, "O_NOFOLLOW", 0)
305
+ try:
306
+ fd = os.open(p, flags)
307
+ except OSError as e:
308
+ raise ArtifactError(f"{p.name} is missing or unreadable") from e
309
+ with os.fdopen(fd, "rb") as f:
310
+ st = os.fstat(f.fileno())
311
+ if not stat.S_ISREG(st.st_mode):
312
+ raise ArtifactError(f"{p.name} is not a regular file")
313
+ if st.st_size > max_bytes:
314
+ raise ArtifactError(f"{p.name} exceeds the size cap of {max_bytes} bytes")
315
+ # Read only what fstat reported (+1): read(n) pre-allocates n bytes, so never ask
316
+ # for max_bytes up front; an extra byte means the file grew after fstat.
317
+ data = f.read(st.st_size + 1)
318
+ if len(data) > st.st_size:
319
+ raise ArtifactError(f"{p.name} changed size while being read")
320
+ return data
321
+
322
+
323
+ def _valid_dtype(x: object) -> bool:
324
+ return isinstance(x, str) and _DTYPE.fullmatch(x) is not None
325
+
326
+
327
+ def _check_header(doc: Any, model_id: str, revision: str | None, dtype: str | None) -> int:
328
+ if not isinstance(doc, dict) or set(doc) != _TOP_KEYS:
329
+ raise ArtifactError(f"{JSON_FILE} must have exactly the keys {sorted(_TOP_KEYS)}")
330
+ version = doc["schema_version"]
331
+ if type(version) is not int or version not in _READABLE_VERSIONS:
332
+ raise ArtifactError(
333
+ f"unsupported schema_version {version!r} (readable: {_READABLE_VERSIONS})"
334
+ )
335
+ model = doc["model"]
336
+ model_keys = {"id", "revision", "dtype"} if version >= 3 else {"id", "revision"}
337
+ if (
338
+ not isinstance(model, dict)
339
+ or set(model) != model_keys
340
+ or not isinstance(model["id"], str)
341
+ or not (model["revision"] is None or isinstance(model["revision"], str))
342
+ or not (model.get("dtype") is None or _valid_dtype(model["dtype"]))
343
+ ):
344
+ raise ArtifactError(f"model must have exactly the keys {sorted(model_keys)}")
345
+ for key in ("brier_version", "template_hash", "arrays_sha256"):
346
+ if not isinstance(doc[key], str):
347
+ raise ArtifactError(f"{key} must be a string")
348
+ lam = doc["prior_strength"]
349
+ if not _is_number(lam) or not 0 <= lam <= 1:
350
+ raise ArtifactError("prior_strength must be a number in [0, 1]")
351
+ if model["id"] != model_id:
352
+ raise ArtifactError(f"artifact was fitted on model {model['id']!r}, not {model_id!r}")
353
+ if model["revision"] != revision:
354
+ raise ArtifactError(
355
+ f"artifact was fitted on revision {model['revision']!r}, not {revision!r}"
356
+ )
357
+ if version >= 3 and model["dtype"] != dtype:
358
+ raise ArtifactError(
359
+ f"artifact was fitted with dtype {model['dtype']!r}, not {dtype!r} "
360
+ "(re-fit the calibration for this precision)"
361
+ )
362
+ if doc["template_hash"] != prompts.template_hash():
363
+ raise ArtifactError("prompt template mismatch: the artifact was fitted with other prompts")
364
+ return int(version)
365
+
366
+
367
+ def _check_question_json(e: Any, version: int) -> None:
368
+ if not isinstance(e, dict):
369
+ raise ArtifactError("each question entry must be an object")
370
+ kind = e.get("type")
371
+ if not isinstance(kind, str) or kind not in _QUESTION_KEYS:
372
+ raise ArtifactError("each question needs a type of choice, noul or score")
373
+ keys = _QUESTION_KEYS[kind] | ({"head"} if version >= 2 else set())
374
+ if set(e) != keys:
375
+ raise ArtifactError(f"{kind} entry must have exactly {sorted(keys)}")
376
+ if version >= 2 and e["head"] is not None:
377
+ _check_head_json(e["head"])
378
+ if not isinstance(e["prior"], bool):
379
+ raise ArtifactError("prior must be true or false")
380
+ if e["temperature"] is not None and not _is_number(e["temperature"]):
381
+ raise ArtifactError("temperature must be a number or null")
382
+ if kind == "choice" and not isinstance(e["options"], list):
383
+ raise ArtifactError("options must be a list")
384
+ if kind == "score" and e["labels"] is not None and not isinstance(e["labels"], list):
385
+ raise ArtifactError("labels must be a list or null")
386
+
387
+
388
+ def _question_from_json(e: dict[str, Any]) -> Question:
389
+ if e["type"] == "choice":
390
+ return Choice(e["text"], e["options"], name=e["name"])
391
+ if e["type"] == "noul":
392
+ return Noul(e["text"], name=e["name"])
393
+ return Score(e["text"], e["levels"], name=e["name"], labels=e["labels"])
394
+
395
+
396
+ def _question_to_json(c: Calibration) -> dict[str, Any]:
397
+ q = c.question
398
+ entry: dict[str, Any] = {"name": q.name, "text": q.text}
399
+ if isinstance(q, Choice):
400
+ entry.update(type="choice", options=list(q.options))
401
+ elif isinstance(q, Noul):
402
+ entry.update(type="noul")
403
+ else:
404
+ entry.update(type="score", levels=q.levels, labels=list(q.labels) if q.labels else None)
405
+ entry.update(prior=c.prior is not None, temperature=c.temperature, head=None)
406
+ if c.head is not None:
407
+ h = c.head
408
+ entry["head"] = {
409
+ "layer": h.layer,
410
+ "solver": h.solver,
411
+ "alpha": h.alpha,
412
+ "temperature": h.temperature,
413
+ "temperature_at_bound": h.temperature_at_bound,
414
+ "oof_nll": h.oof_nll,
415
+ "oof_accuracy": h.oof_accuracy,
416
+ }
417
+ return entry
418
+
419
+
420
+ def _check_head_json(h: Any) -> None:
421
+ if not isinstance(h, dict) or set(h) != _HEAD_KEYS:
422
+ raise ArtifactError(f"head must be an object with exactly {sorted(_HEAD_KEYS)}")
423
+ if type(h["layer"]) is not int or not isinstance(h["solver"], str):
424
+ raise ArtifactError("head layer must be an int and solver a string")
425
+ if h["alpha"] is not None and not _is_number(h["alpha"]):
426
+ raise ArtifactError("head alpha must be a number or null")
427
+ if not all(_is_number(h[k]) for k in ("temperature", "oof_nll", "oof_accuracy")):
428
+ raise ArtifactError("head temperature and OOF metrics must be numbers")
429
+ if not isinstance(h["temperature_at_bound"], bool):
430
+ raise ArtifactError("head temperature_at_bound must be true or false")
431
+
432
+
433
+ def _head(i: int, h: dict[str, Any] | None, arrays: dict[str, FloatArray]) -> L2Head | None:
434
+ if h is None:
435
+ return None
436
+ try:
437
+ return L2Head(
438
+ layer=h["layer"],
439
+ solver=h["solver"],
440
+ alpha=h["alpha"],
441
+ temperature=h["temperature"],
442
+ temperature_at_bound=h["temperature_at_bound"],
443
+ oof_nll=h["oof_nll"],
444
+ oof_accuracy=h["oof_accuracy"],
445
+ mean=arrays[f"head_{i}_mean"],
446
+ scale=arrays[f"head_{i}_scale"],
447
+ weights=arrays[f"head_{i}_weights"],
448
+ bias=arrays[f"head_{i}_bias"],
449
+ )
450
+ except BrierError as e:
451
+ raise ArtifactError(f"invalid head for question {i}: {e}") from e
452
+
453
+
454
+ _ShapeCheck = Callable[[tuple[int, ...]], bool]
455
+ _Spec = tuple[_ShapeCheck, int] # (accepted shapes, per-member byte cap)
456
+
457
+
458
+ def _exact(shape: tuple[int, ...]) -> _ShapeCheck:
459
+ return lambda s: s == shape
460
+
461
+
462
+ def _hidden_vector(s: tuple[int, ...]) -> bool:
463
+ return len(s) == 1 and 1 <= s[0] <= MAX_HIDDEN
464
+
465
+
466
+ def _hidden_matrix(n_classes: int) -> _ShapeCheck:
467
+ return lambda s: len(s) == 2 and 1 <= s[0] <= MAX_HIDDEN and s[1] == n_classes
468
+
469
+
470
+ def _read_arrays(data: bytes, specs: dict[str, _Spec], max_bytes: int) -> dict[str, FloatArray]:
471
+ """Read exactly the arrays in ``specs`` from the npz without trusting its headers."""
472
+ try:
473
+ z = zipfile.ZipFile(io.BytesIO(data))
474
+ except zipfile.BadZipFile as e:
475
+ raise ArtifactError(f"{ARRAYS_FILE} is not a valid npz file") from e
476
+ with z:
477
+ infos = z.infolist()
478
+ expected = {f"{key}.npy": key for key in specs}
479
+ names = [info.filename for info in infos]
480
+ if len(names) != len(expected) or set(names) != set(expected):
481
+ raise ArtifactError(f"{ARRAYS_FILE} must contain exactly the referenced arrays")
482
+ if sum(info.file_size for info in infos) > max_bytes:
483
+ raise ArtifactError(f"{ARRAYS_FILE} contents exceed the size cap of {max_bytes} bytes")
484
+ out: dict[str, FloatArray] = {}
485
+ for info in infos:
486
+ if info.flag_bits & 0x1:
487
+ raise ArtifactError(f"{info.filename} is encrypted")
488
+ if info.compress_type not in (zipfile.ZIP_STORED, zipfile.ZIP_DEFLATED):
489
+ raise ArtifactError(f"{info.filename} uses an unsupported compression method")
490
+ key = expected[info.filename]
491
+ check, cap = specs[key]
492
+ if info.file_size > cap:
493
+ raise ArtifactError(f"{info.filename} is too large")
494
+ with z.open(info) as member: # bounded: never trust the declared file_size
495
+ raw = member.read(cap + 1)
496
+ if len(raw) > cap: # pragma: no cover - zipfile stops at file_size
497
+ raise ArtifactError(f"{info.filename} is too large")
498
+ out[key] = _parse_npy(raw, check, info.filename)
499
+ return out
500
+
501
+
502
+ def _parse_npy(raw: bytes, check: _ShapeCheck, name: str) -> FloatArray:
503
+ """Validate a ``.npy`` header (``<f8``, C order, accepted shape), then read the data."""
504
+ f = io.BytesIO(raw)
505
+ version = _NPY.read_magic(f)
506
+ if version == (1, 0):
507
+ shape, fortran_order, dtype = _NPY.read_array_header_1_0(f)
508
+ elif version == (2, 0):
509
+ shape, fortran_order, dtype = _NPY.read_array_header_2_0(f)
510
+ else:
511
+ raise ArtifactError(f"{name} has unsupported .npy version {version}")
512
+ exact_ints = isinstance(shape, tuple) and all(type(v) is int for v in shape)
513
+ if dtype != _F8 or fortran_order or not exact_ints or not check(shape):
514
+ raise ArtifactError(f"{name} must be little-endian float64, C order, of the expected shape")
515
+ count = math.prod(shape)
516
+ body = f.read()
517
+ if len(body) != count * _F8.itemsize:
518
+ raise ArtifactError(f"{name} has {len(body)} data bytes, expected {count * _F8.itemsize}")
519
+ return np.frombuffer(body, dtype=_F8).astype(np.float64).reshape(shape)
@@ -0,0 +1 @@
1
+ """Model backends. Only ``hf`` imports torch/transformers (lazily)."""