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 +8 -0
- brier/__main__.py +73 -0
- brier/_math.py +57 -0
- brier/_version.py +8 -0
- brier/artifacts.py +519 -0
- brier/backends/__init__.py +1 -0
- brier/backends/base.py +50 -0
- brier/backends/fake.py +125 -0
- brier/backends/hf.py +264 -0
- brier/bench/__init__.py +1 -0
- brier/bench/__main__.py +104 -0
- brier/bench/features.py +97 -0
- brier/bench/run.py +335 -0
- brier/bench/tasks.py +166 -0
- brier/calibrate/__init__.py +1 -0
- brier/calibrate/temperature.py +160 -0
- brier/check.py +287 -0
- brier/debias.py +113 -0
- brier/decider.py +441 -0
- brier/decision.py +97 -0
- brier/errors.py +29 -0
- brier/heads/__init__.py +1 -0
- brier/heads/_features.py +59 -0
- brier/heads/fitted.py +130 -0
- brier/heads/lda.py +105 -0
- brier/heads/ridge.py +89 -0
- brier/heads/select.py +215 -0
- brier/metrics.py +244 -0
- brier/prompts.py +134 -0
- brier/py.typed +0 -0
- brier/questions.py +125 -0
- brier/readout.py +55 -0
- brier-0.1.0.dist-info/METADATA +211 -0
- brier-0.1.0.dist-info/RECORD +38 -0
- brier-0.1.0.dist-info/WHEEL +4 -0
- brier-0.1.0.dist-info/entry_points.txt +2 -0
- brier-0.1.0.dist-info/licenses/LICENSE +202 -0
- brier-0.1.0.dist-info/licenses/NOTICE +4 -0
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)."""
|