inventio 0.4.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.
inventio/__init__.py ADDED
@@ -0,0 +1 @@
1
+ __version__ = "0.4.0"
@@ -0,0 +1,128 @@
1
+ """Kev's serving path, vendored so a System One checkpoint needs one install and no second repository.
2
+
3
+ Source: https://github.com/jaredpalmer/kev at f1535963cea021439370c23127bc970b6788e730 (2026-09-26),
4
+ Apache-2.0, (c) the Kev authors. Copied: `model.py`, `checkpoint.py`, `api.py`, `device.py` — the four
5
+ files `kev.serve` turns a checkpoint into answered questions with (Kev's own Hugging Face Space vendors
6
+ the first three next to each other, so the layout is Kev's, not ours). Everything a serving pass
7
+ touches is their code unchanged: the packed block-causal encode, the pointer head and its temperature,
8
+ the readouts, the request/answer schema.
9
+
10
+ Changed here, and nothing else:
11
+
12
+ - `api.py`: the option cap is 512, Kev ships 255 (a code pool's median line count is 407; see
13
+ `inventio/systemone.py` for what 255 cost on SWE-bench Lite).
14
+ - `model.py`: the CUDA-graph serving batch path (`kev.cuda_graphs`, `EAGER_STATES`) is gone;
15
+ `probs_batch` runs Kev's eager path, the one every measured number in this repository used.
16
+ - `checkpoint.py`: the trainer-only and other-backend parts are gone — `mlx_available`,
17
+ `fused_available`, `backend`, `_load_mlx`, `warm_start`, `_load_backbone_into`,
18
+ `_load_adapter_into`, `weights_sha256`, `release_date`, `hybrid_base`, `LoadOptions.from_env` and
19
+ the `backend`/`cuda_graphs`/`fused` fields. `_load_torch` is the torch path alone. `resolve_run` reads a
20
+ Hub checkpoint from the cache when it is there instead of asking the Hub for the current revision on every
21
+ load (inventio's release check is `inventio update`'s, once a day). `Checkpoint.load` reads a full-weight
22
+ checkpoint's own tokenizer when it carries one, so the published model needs neither the base repo nor the network.
23
+ - `device.py`: as published.
24
+ - This file: `load()` (the device and dtype defaults `kev.serve` serves with, the Windows attention
25
+ fallback, and the dependency check that names the missing extra instead of raising an ImportError
26
+ from inside a forward pass).
27
+
28
+ `load()` is the only entry point inventio uses; a checkpoint answers to any model name, so what
29
+ identifies weights is the run directory or Hub id it was loaded from.
30
+ """
31
+
32
+ import os
33
+ import re
34
+
35
+ __all__ = ["Checkpoint", "LoadOptions", "REQUIRES", "is_hub_id", "load", "missing", "patch_attention"]
36
+
37
+ # What a checkpoint needs, and what an environment without it has to be told. `transformers>=5.17` is
38
+ # the version the hybrid (Qwen3.5) cache classes came from; an older one imports and then fails inside
39
+ # the first forward pass, which is a worse error than this one.
40
+ REQUIRES = ("torch>=2.6", "transformers>=5.17,<6", "peft>=0.21", "pydantic>=2.9")
41
+
42
+
43
+ def missing() -> str | None:
44
+ """What this environment still has to install to load a checkpoint, or None when it can.
45
+
46
+ Reads installed distributions, never imports them: `default_ranker()` asks this while the command line
47
+ is being built, and importing torch there would make every `inventio` command pay for a model runtime
48
+ it may not use. A package that answers here and then fails to import is caught at load time, where the
49
+ same message comes back.
50
+ """
51
+ import importlib.util
52
+ from importlib.metadata import version
53
+
54
+ for name, need in (("torch", "2.6"), ("transformers", "5.17"), ("peft", "0.21"), ("pydantic", "2.9")):
55
+ if importlib.util.find_spec(name) is None:
56
+ return f"{name} is not installed — pip install 'inventio[dispositio]'"
57
+ try:
58
+ if tuple(int(x) for x in re.findall(r"\d+", version("transformers"))[:2]) < (5, 17):
59
+ return (f"transformers {version('transformers')} is older than 5.17 (the hybrid cache classes "
60
+ f"the encoder uses) — pip install -U 'inventio[dispositio]'")
61
+ except Exception: # a distribution without version metadata: let the import speak
62
+ pass
63
+ return None
64
+
65
+
66
+ def patch_attention() -> None:
67
+ """Windows: torch's SDPA GQA path falls back to the math kernel for the hybrid (Qwen3.5) bases, which
68
+ holds the whole L x L score matrix — out of memory past ~8k tokens, twice the time at 3k. Repeating
69
+ the kv heads instead lets SDPA take its memory-efficient kernel; measured on this machine as the
70
+ same numbers, a different kernel (benchmarks/kev_win.py, where the patch was found)."""
71
+ if os.name == "nt":
72
+ import transformers.integrations.sdpa_attention as sdpa
73
+
74
+ sdpa.use_gqa_in_sdpa = lambda *a, **k: False
75
+
76
+
77
+ def load(run, device=None, *, dtype=None, merge: bool = True, attn: str | None = None,
78
+ lora_scale: float = 1.0, temperature: float | None = None):
79
+ """A checkpoint as `kev.serve` loads it -> (tokenizer, model).
80
+
81
+ What `kev.serve` serves with: bf16 on an accelerator (half the memory; probabilities within ~0.01 of
82
+ the exact path, the same argmax), fp32 on CPU; a checkpoint trained on a bf16 backbone loads in bf16
83
+ wherever it runs, as it was trained. INVENTIO_DTYPE=fp32 asks for the exact path — the one
84
+ every number in this repository was reported from — which costs about twice the wall time.
85
+ INVENTIO_DEVICE=cpu|cuda|mps leaves the accelerator to something else.
86
+ """
87
+ from .device import default_device
88
+
89
+ dev = str(device or os.environ.get("INVENTIO_DEVICE") or default_device())
90
+ if dev == "cpu":
91
+ # The hybrid (Qwen3.5) DeltaNet layers take a compiled kernel when one can be imported: from the
92
+ # Hugging Face kernel hub, or from flash-linear-attention (`fla`) and `causal_conv1d` if those are
93
+ # installed. Each is a GPU kernel, and on CPU it raises inside a Triton launch ("cannot be accessed
94
+ # from Triton (cpu tensor?)", measured with fla in the environment). transformers picks the
95
+ # implementation when its model module is imported, so the pure-torch path has to be asked for
96
+ # before that: the hub flag off, and the two packages made unimportable in this process.
97
+ import sys
98
+
99
+ os.environ.setdefault("USE_HUB_KERNELS", "NO")
100
+ for pkg in ("fla", "causal_conv1d"):
101
+ sys.modules.setdefault(pkg, None) # `import fla` now raises ImportError: transformers falls back
102
+ if "transformers.models.qwen3_5.modeling_qwen3_5" in sys.modules and sys.modules.get("fla"):
103
+ raise RuntimeError("this process already loaded the GPU kernels for the model's DeltaNet layers; "
104
+ "run on CPU in a fresh process (INVENTIO_DEVICE=cpu)")
105
+ miss = missing()
106
+ if miss:
107
+ raise RuntimeError(miss)
108
+ from .checkpoint import Checkpoint, LoadOptions
109
+
110
+ import torch
111
+
112
+ patch_attention()
113
+ device = dev
114
+ if dtype is None and device != "cpu" and os.environ.get("INVENTIO_DTYPE", "").lower() != "fp32":
115
+ dtype = torch.bfloat16
116
+ opts = LoadOptions(dtype=dtype, merge=merge, attn=attn, lora_scale=lora_scale, temperature=temperature)
117
+ return Checkpoint(run).load(device, opts)
118
+
119
+
120
+ from .ids import is_hub_id # noqa: E402 (torch-free: `inventio model --use` validates a name without the extra)
121
+
122
+
123
+ def __getattr__(name): # PEP 562: Checkpoint/LoadOptions/resolve_run pull torch in, so they load on first use
124
+ if name in ("Checkpoint", "LoadOptions", "resolve_run"):
125
+ from . import checkpoint
126
+
127
+ return getattr(checkpoint, name)
128
+ raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
@@ -0,0 +1,165 @@
1
+ """TypeSafe-compatible request/response shapes (POST /v1/systemone) mapped onto the single pointer primitive.
2
+
3
+ Noul -> 2 options [false, true]; answer = p(true)
4
+ Choice -> options 'name' or 'name: desc'; answer = argmax, probabilities by name, confidence
5
+ Score -> options = ordered level descriptions; answer = expected level, legend, probabilities by index
6
+ """
7
+ import json
8
+ import re
9
+ from datetime import datetime
10
+ from typing import Any, Literal, Union
11
+ from pydantic import BaseModel, Field, model_validator
12
+
13
+ JSONContent = Union[str, dict, list, int, float, bool, None]
14
+ MAX_OPTIONS = 512 # Kev ships 255; a code pool's median line count is 407 (inventio/systemone.py)
15
+
16
+
17
+ class Noul(BaseModel):
18
+ type: Literal["noul"]
19
+ instructions: JSONContent = None
20
+ criteria: dict[str, JSONContent] | None = None
21
+
22
+
23
+ class Choice(BaseModel):
24
+ type: Literal["choice"]
25
+ instructions: JSONContent = None
26
+ criteria: dict[str, JSONContent]
27
+
28
+ @model_validator(mode="after")
29
+ def _check(self):
30
+ if not 1 <= len(self.criteria) <= MAX_OPTIONS: raise ValueError(f"criteria must have 1..{MAX_OPTIONS} options")
31
+ return self
32
+
33
+
34
+ class Score(BaseModel):
35
+ type: Literal["score"]
36
+ instructions: JSONContent = None
37
+ criteria: list[JSONContent] = Field(min_length=1, max_length=MAX_OPTIONS)
38
+
39
+
40
+ Question = Union[Noul, Choice, Score]
41
+
42
+
43
+ class SystemOneRequest(BaseModel):
44
+ state: JSONContent
45
+ model: str = "kev-latest"
46
+ questions: dict[str, Question] = Field(min_length=1)
47
+
48
+
49
+ def render(v: JSONContent, indent: int = 0) -> str:
50
+ """Flatten str | object | array into text the model sees. Field names are kept as labels."""
51
+ pad = " " * indent
52
+ if v is None: return ""
53
+ if isinstance(v, (str, int, float, bool)): return str(v)
54
+ if isinstance(v, list): return "\n".join(f"{pad}- {render(x, indent + 1).lstrip()}" for x in v)
55
+ return "\n".join(f"{pad}{k}:\n{render(x, indent + 1)}" if isinstance(x, (dict, list)) else f"{pad}{k}: {render(x)}" for k, x in v.items())
56
+
57
+
58
+ def option_text(name: str, desc: JSONContent) -> str:
59
+ return name if desc is None or desc == "" else f"{name}: {render(desc)}"
60
+
61
+
62
+ MONTHS = "January|February|March|April|May|June|July|August|September|October|November|December"
63
+ _DATE = re.compile(rf"\b(?:{MONTHS}) \d{{1,2}}, \d{{4}}\b|\b\d{{4}}-\d{{2}}-\d{{2}}\b")
64
+
65
+
66
+ def date_facts(text: str) -> str:
67
+ """Deterministic date arithmetic for the model: every pair of absolute dates found in `text`, as one sentence each
68
+ ("August 3, 2026 is 12 days after July 22, 2026."). The model cannot subtract dates reliably (issue #8); it can use a
69
+ stated day count. Returns "" when fewer than two dates are found. Dates are listed in order of first appearance."""
70
+ found = []
71
+ for m in _DATE.finditer(text):
72
+ raw = m.group(0)
73
+ try: d = datetime.strptime(raw, "%B %d, %Y") if "," in raw else datetime.strptime(raw, "%Y-%m-%d")
74
+ except ValueError: continue
75
+ if raw not in [r for r, _ in found]: found.append((raw, d))
76
+ facts = []
77
+ for i in range(len(found)):
78
+ for j in range(i + 1, len(found)):
79
+ n = (found[j][1] - found[i][1]).days
80
+ facts.append(f"{found[j][0]} is {abs(n)} day{'s' if abs(n) != 1 else ''} {'after' if n > 0 else 'before'} {found[i][0]}." if n else f"{found[j][0]} is the same day as {found[i][0]}.")
81
+ return " ".join(facts)
82
+
83
+
84
+ def with_date_facts(state):
85
+ """State with a `date_facts` field (object states) or an appended paragraph (string states) when two or more absolute
86
+ dates appear. Opt-in preprocessing (KEV_DATE_FACTS=1 in kev.serve, --date_facts in kev.benchmark)."""
87
+ facts = date_facts(render(state))
88
+ if not facts: return state
89
+ if isinstance(state, dict): return {**state, "date_facts": facts}
90
+ if isinstance(state, list): return state + [{"date_facts": facts}]
91
+ return f"{state}\n\ndate_facts: {facts}"
92
+
93
+
94
+ def question_keys(qtype: str, criteria) -> list[str]:
95
+ """The keys a question's probabilities are reported under, in option order: the criteria names (choice),
96
+ ["false", "true"] (noul), the level indices as strings (score). Labels, targets and anchors use the same keys."""
97
+ if qtype == "choice": return list(criteria)
98
+ if qtype == "noul": return ["false", "true"]
99
+ return [str(i) for i in range(len(criteria))]
100
+
101
+
102
+ def to_record(req: SystemOneRequest):
103
+ """-> internal record for encode(), plus per-question metadata ({"id", "type", "keys", "legend" for score}) to map
104
+ probabilities back."""
105
+ qs, meta = [], []
106
+ for qid, q in req.questions.items():
107
+ m = {"id": qid, "type": q.type, "keys": question_keys(q.type, q.criteria)}
108
+ if q.type == "noul":
109
+ c = q.criteria or {}
110
+ opts = [option_text("no", c.get("false")), option_text("yes", c.get("true"))]
111
+ elif q.type == "choice":
112
+ opts = [option_text(k, v) for k, v in q.criteria.items()]
113
+ else:
114
+ opts = [render(x) for x in q.criteria]
115
+ m["legend"] = dict(zip(m["keys"], opts))
116
+ qs.append({"instr": render(q.instructions), "options": opts, "label": 0}); meta.append(m)
117
+ return {"state": render(req.state), "questions": qs}, meta
118
+
119
+
120
+ # Both confidence formulas mirror TypeSafe's reference adapter, system-one-adapter 0.2.1
121
+ # (src/system_one_adapter/_utils/confidence_metrics.py): p is normalised to sum 1 (all zeros -> uniform), one option -> 1.
122
+ def _normalize(p: list[float]) -> list[float]:
123
+ t = sum(p)
124
+ return [1 / len(p)] * len(p) if t == 0 else [x / t for x in p]
125
+
126
+
127
+ def choice_confidence(p: list[float]) -> float:
128
+ """(p_max - 1/K) / (1 - 1/K): 0 at uniform, 1 at certainty."""
129
+ K = len(p)
130
+ return 1.0 if K == 1 else (max(_normalize(p)) - 1 / K) / (1 - 1 / K)
131
+
132
+
133
+ def score_confidence(p: list[float]) -> float:
134
+ """max(0, 1 - E|level - mode| / D), D = mean absolute deviation of a uniform distribution over the L levels around its
135
+ mean (L-1)/2; mode = first most likely level. 1 when all mass is on one level, 0 at uniform or anything as spread."""
136
+ L = len(p)
137
+ if L == 1: return 1.0
138
+ p = _normalize(p); mode = max(range(L), key=p.__getitem__)
139
+ D = sum(abs(i - (L - 1) / 2) for i in range(L)) / L
140
+ return max(0.0, 1.0 - sum(pi * abs(i - mode) for i, pi in enumerate(p)) / D)
141
+
142
+
143
+ def round_prob(x: float) -> float:
144
+ """Serialization precision for probabilities and derived scalars. 4 decimals keeps the sum of a rounded distribution
145
+ within TypeSafe's tolerance (|sum - 1| < 0.02) at the 512-option maximum: 512 * 0.00005 < 0.02."""
146
+ return round(float(x), 4)
147
+
148
+
149
+ def to_answers(probs: list[list[float]], meta: list[dict]) -> dict[str, Any]:
150
+ out = {}
151
+ for p, m in zip(probs, meta):
152
+ if m["type"] == "noul":
153
+ out[m["id"]] = {"type": "noul", "noul": round_prob(p[1])}
154
+ elif m["type"] == "choice":
155
+ dist = {k: round_prob(v) for k, v in zip(m["keys"], p)}
156
+ out[m["id"]] = {"type": "choice", "choice": m["keys"][max(range(len(p)), key=lambda i: p[i])], "confidence": round_prob(choice_confidence(p)), "probabilities": dist}
157
+ else:
158
+ score = sum(i * pi for i, pi in enumerate(p))
159
+ out[m["id"]] = {"type": "score", "score": round_prob(score), "legend": m["legend"], "probabilities": {str(i): round_prob(v) for i, v in enumerate(p)}, "confidence": round_prob(score_confidence(p))}
160
+ return out
161
+
162
+
163
+ def output_tokens(tok, answers: dict) -> int:
164
+ """Billing-style figure: tokens of the serialised answers. Not a measure of generation (there is none)."""
165
+ return len(tok(json.dumps(answers), add_special_tokens=False).input_ids)
@@ -0,0 +1,192 @@
1
+ """Trained checkpoints: a run directory or a Hub repo holding a LoRA adapter (or, for a full-weight run, the whole bf16
2
+ backbone), `head.pt` and the tokenizer.
3
+
4
+ Loader rule: `adapter_config.json` present -> a LoRA adapter on `meta.base` at `meta.base_revision`; no adapter and
5
+ `config.json` + `model*.safetensors` (save_pretrained of the backbone, `meta.weights == "full"`) -> the backbone is loaded from
6
+ the checkpoint directory itself, nothing is merged, in the dtype head.pt's `weights_dtype` names (it must match the `dtype`
7
+ save_pretrained wrote to config.json). The tokenizer always comes from the base (both layouts carry a copy).
8
+
9
+ This is the one place that knows the layout of `head.pt` and how a checkpoint becomes a `DecisionModel`:
10
+ `kev.serve`, `kev.benchmark`, `kev.train --init_from`, `kev.publish`, the scripts and the Hugging Face Space all go
11
+ through it. The Space vendors this file next to `model.py` and `api.py` (scripts/publish_space.sh), so it must not
12
+ import the data or suite modules at import time.
13
+
14
+ ck = Checkpoint("jaredpalmer/kev-4b") # or a local run directory; `@tag` pins a Hub revision
15
+ tok, model = ck.load("mps", LoadOptions.from_env())
16
+ ck.meta.temperature # the calibration the checkpoint carries
17
+ """
18
+ import json
19
+ import os
20
+ from dataclasses import dataclass, field
21
+ from pathlib import Path
22
+
23
+ import torch
24
+
25
+ from .model import DecisionModel, load_tokenizer
26
+
27
+ from .ids import is_hub_id # vendored here: no torch needed to answer a question about a name
28
+
29
+
30
+ def resolve_run(run):
31
+ """Local run directory as given, or a Hub repo id like jaredpalmer/kev-4b, optionally pinned to a revision or tag
32
+ with `@` (jaredpalmer/kev-4b@qwen3), downloaded to the HF cache. Returns a str path."""
33
+ if os.path.isdir(run):
34
+ return str(run)
35
+ from huggingface_hub import snapshot_download
36
+ repo, _, revision = str(run).partition("@")
37
+ pats = ["*.json", "*.safetensors", "*.pt", "*.txt", "*.jinja"]
38
+ try: # the cache first: a load must not ask the network which revision is current (inventio: `update` does)
39
+ return snapshot_download(repo, revision=revision or None, allow_patterns=pats, local_files_only=True)
40
+ except Exception: # not downloaded yet: the first load fetches it
41
+ return snapshot_download(repo, revision=revision or None, allow_patterns=pats)
42
+
43
+
44
+ @dataclass
45
+ class Meta:
46
+ """Contents of `head.pt`. Every reader gets the same defaults for fields older checkpoints did not write.
47
+ `extra` keeps the rest of the file (training args, suite hash, init provenance, temperature fit) so a
48
+ read-modify-write round trip loses nothing."""
49
+ base: str
50
+ head: dict | None = None
51
+ base_revision: str | None = None
52
+ lora: int = 0
53
+ head_dim: int = 256
54
+ option_isolation: bool = False
55
+ special_embeddings: bool = False
56
+ weights_dtype: str = "fp32"
57
+ temperature: float = 1.0
58
+ holdout: list = field(default_factory=list)
59
+ weights: str = "lora" # "lora": an adapter on the base; "full": the whole backbone is in the checkpoint (kev.train --full_ft)
60
+ extra: dict = field(default_factory=dict)
61
+
62
+ KNOWN = ("base", "head", "base_revision", "lora", "head_dim", "option_isolation", "special_embeddings", "weights_dtype", "temperature", "holdout", "weights")
63
+
64
+ @classmethod
65
+ def from_dict(cls, d):
66
+ return cls(**{k: d[k] for k in cls.KNOWN if k in d}, extra={k: v for k, v in d.items() if k not in cls.KNOWN})
67
+
68
+ def to_dict(self):
69
+ return {**self.extra, **{k: getattr(self, k) for k in self.KNOWN}} # known fields win over a stray key in extra
70
+
71
+
72
+ def read_meta(run):
73
+ return Meta.from_dict(torch.load(f"{run}/head.pt", map_location="cpu"))
74
+
75
+
76
+ def write_meta(run, meta):
77
+ torch.save(meta.to_dict(), f"{run}/head.pt")
78
+
79
+
80
+ @dataclass(frozen=True)
81
+ class LoadOptions:
82
+ """How a checkpoint is turned into a model. Defaults are the exact path every reported number uses.
83
+
84
+ dtype None = fp32, the exact path every reported number uses (bf16 when the checkpoint was trained with a bf16
85
+ backbone). Serving defaults to bf16 on CUDA and MPS instead: half the memory, 2-4.5x lower latency on an
86
+ L4 (Kev-4B: 209 -> 118 ms at 101 tokens, 850 -> 189 ms at 330 tokens), probabilities within ~0.01 and
87
+ the same argmax on the checks run so far. The exact path is INVENTIO_DTYPE=fp32.
88
+ merge fold the LoRA into the base weights: the delta is computed from the fp32 adapter and added in fp32 with
89
+ one rounding to the load dtype, so a bf16 model holds exactly round(W + delta), the same bits as merging
90
+ an fp32 copy and casting, without the fp32 copy (Kev-9B needed 36 GB of GPU memory to load for that).
91
+ Identical because the Qwen bases are stored in bf16; a base stored in fp32 would be rounded twice.
92
+ Exact in fp32; in bf16 it is faster (~15%) and closer to the fp32 numbers than the unmerged adapter
93
+ (kev-4b, 24 dev records: max |dp| 0.017 vs 0.029, 0 vs 1 argmax flips). Ignored for adapters that carry
94
+ trained token embeddings. A checkpoint trained on a bf16 backbone keeps its adapter unmerged, as it was
95
+ trained (the fused serving path that folded it is not vendored).
96
+ attn attention backend; None = the model default (SDPA on CUDA, eager elsewhere). "sdpa" on MPS measured
97
+ parity with eager and is a few percent faster.
98
+ lora_scale WiSE-FT-style interpolation between base (0) and fine-tuned weights (1), at inference.
99
+ temperature None = the temperature the checkpoint carries (fitted by scripts/calibrate_checkpoint.py); 1.0 = raw logits.
100
+ """
101
+ dtype: torch.dtype | None = None
102
+ merge: bool = True
103
+ attn: str | None = None
104
+ lora_scale: float = 1.0
105
+ temperature: float | None = None
106
+
107
+
108
+
109
+ class Checkpoint:
110
+ def __init__(self, run):
111
+ self.requested = str(run) # what the caller asked for (a Hub id stays a Hub id in labels)
112
+ self.path = resolve_run(run)
113
+ self.meta = read_meta(self.path)
114
+
115
+ def file(self, name):
116
+ return Path(self.path) / name
117
+
118
+ def adapter_config(self):
119
+ return json.loads(self.file("adapter_config.json").read_text(encoding="utf-8"))
120
+
121
+ @property
122
+ def full(self):
123
+ """The loader rule (module docstring): True for a full-weight checkpoint, False for a LoRA adapter; head.pt's
124
+ `weights` must agree with the files."""
125
+ found = "lora" if self.file("adapter_config.json").exists() else "full" if self.file("config.json").exists() and self.shards() else None
126
+ if found != self.meta.weights:
127
+ raise ValueError(f"{self.path}: head.pt says weights={self.meta.weights!r} but the directory holds "
128
+ f"{ {'lora': 'an adapter', 'full': 'backbone weights'}.get(found, 'neither an adapter nor backbone weights') }")
129
+ return found == "full"
130
+
131
+ def shards(self):
132
+ """The backbone's safetensors files of a full-weight checkpoint (model.safetensors or model-*-of-*.safetensors)."""
133
+ return sorted(Path(self.path).glob("model*.safetensors"))
134
+
135
+ def load(self, device, opts=LoadOptions()):
136
+ """-> (tokenizer, model) in eval mode with the LoRA applied (or the full backbone loaded) and the pointer head
137
+ loaded. The model is a DecisionModel, ready for encode()/probs()."""
138
+ meta = self.meta
139
+ # a checkpoint that carries its tokenizer is read from itself: a published full-weight model then loads with
140
+ # no second repository and no network (inventio's offline promise); otherwise the base's, as Kev does
141
+ own = self.file("tokenizer.json").exists() and self.full
142
+ tok = load_tokenizer(self.path) if own else load_tokenizer(meta.base, revision=meta.base_revision)
143
+ m = self._load_torch(tok, device, opts)
144
+ m.head.load_state_dict(meta.head); m.eval()
145
+ m.head.temperature = meta.temperature if opts.temperature is None else opts.temperature
146
+ return tok, m
147
+
148
+ def _load_torch(self, tok, device, opts):
149
+ return self._full_torch(tok, device, opts)[0] if self.full else self._adapted_torch(tok, device, opts)[0]
150
+
151
+ SAVED_DTYPES = {"bf16": "bfloat16", "fp32": "float32"} # head.pt weights_dtype -> the dtype save_pretrained writes to config.json
152
+
153
+ def _full_torch(self, tok, device, opts):
154
+ """-> (model, True). Full weights load in the dtype head.pt's `weights_dtype` names (bf16 for every kev.train
155
+ --full_ft run: the dtype they were trained in), which must be the dtype save_pretrained recorded in config.json;
156
+ otherwise a mislabelled export would be silently cast (fp32 weights rounded to bf16, or bf16 upcast to twice the
157
+ memory). An explicit dtype still casts on purpose (fp32: the same values computed in fp32). Nothing to merge."""
158
+ if opts.lora_scale != 1: raise ValueError("lora_scale interpolates an adapter; a full-weight checkpoint has none")
159
+ meta = self.meta
160
+ cfg = json.loads(self.file("config.json").read_text(encoding="utf-8"))
161
+ expected, saved = self.SAVED_DTYPES.get(meta.weights_dtype), cfg.get("dtype") or cfg.get("torch_dtype")
162
+ if expected is None or saved not in (None, expected):
163
+ raise ValueError(f"{self.path}: config.json records the weights as {saved} but head.pt says weights_dtype={meta.weights_dtype!r}")
164
+ return DecisionModel(meta.base, tok, device, head_dim=meta.head_dim, option_isolation=meta.option_isolation,
165
+ dtype=opts.dtype or getattr(torch, expected), attn=opts.attn, weights=self.path), True
166
+
167
+ def _adapted_torch(self, tok, device, opts):
168
+ """-> (model, whether the adapter was merged): the base with this checkpoint's LoRA."""
169
+ from peft import PeftModel
170
+ meta = self.meta
171
+ dtype, merge = opts.dtype or torch.float32, opts.merge
172
+ if meta.weights_dtype == "bf16":
173
+ # trained with a bf16 backbone: load it the same way. The exact path keeps the fp32 adapter unmerged.
174
+ dtype, merge = torch.bfloat16, merge
175
+ merge = merge and not self.adapter_config().get("trainable_token_indices") # token-trained adapters stay unmerged
176
+ m = DecisionModel(meta.base, tok, device, lora=None, revision=meta.base_revision, head_dim=meta.head_dim,
177
+ option_isolation=meta.option_isolation, dtype=dtype, attn=opts.attn)
178
+ m.lm = PeftModel.from_pretrained(m.lm, self.path, torch_device=str(device)).to(device) # trainable token embeddings, if any, live in the adapter
179
+ if opts.lora_scale != 1:
180
+ for module in m.lm.modules():
181
+ if isinstance(getattr(module, "scaling", None), dict):
182
+ for k in module.scaling: module.scaling[k] *= opts.lora_scale
183
+ m.lora_scale = opts.lora_scale
184
+ if merge: m.lm = m.lm.merge_and_unload() # W += delta: fp32 math, one rounding (see LoadOptions.merge)
185
+ if dtype != torch.float32: m.lm = m.lm.to(dtype)
186
+ return m, merge
187
+
188
+ COMPAT_FIELDS = ("base", "base_revision", "lora", "head_dim", "option_isolation", "special_embeddings", "weights")
189
+
190
+ def load(run, device, opts=LoadOptions()):
191
+ """Convenience: Checkpoint(run).load(device, opts)."""
192
+ return Checkpoint(run).load(device, opts)
@@ -0,0 +1,30 @@
1
+ """The accelerator this process uses: cuda, then mps, then cpu."""
2
+ import torch
3
+
4
+
5
+ def default_device():
6
+ return "cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu"
7
+
8
+
9
+ def sync(device):
10
+ """Wait for queued kernels, so wall-clock timings around a forward pass are real."""
11
+ if device == "mps": torch.mps.synchronize()
12
+ elif device == "cuda": torch.cuda.synchronize()
13
+
14
+
15
+ def empty_cache(device):
16
+ if device == "mps": torch.mps.empty_cache()
17
+ elif device == "cuda": torch.cuda.empty_cache()
18
+
19
+
20
+ def out_of_memory(e):
21
+ """Whether a torch allocator ran out: CUDA raises torch.OutOfMemoryError, MPS a plain RuntimeError with this message.
22
+ The MLX backend's Metal errors are neither."""
23
+ return isinstance(e, torch.OutOfMemoryError) or isinstance(e, RuntimeError) and str(e).startswith("MPS backend out of memory")
24
+
25
+
26
+ def allocated_bytes(device):
27
+ """Bytes currently allocated on the device (MPS) or the peak since the process started (CUDA); 0 on CPU."""
28
+ if device == "mps": return torch.mps.current_allocated_memory()
29
+ if device == "cuda": return torch.cuda.max_memory_allocated()
30
+ return 0
@@ -0,0 +1,15 @@
1
+ """What a checkpoint can be named by: a directory, or a Hub repo id (optionally pinned `@revision`).
2
+
3
+ Kev's `checkpoint.py` keeps the same rule; it lives here too because `inventio model --use` has to
4
+ validate a name on a machine that has not installed the extra, and importing `checkpoint` would import
5
+ torch and transformers to answer a question about a string.
6
+ """
7
+
8
+ import os
9
+ import re
10
+
11
+ HUB_ID = re.compile(r"[\w.-]+/[\w.-]+(@[\w.-]+)?")
12
+
13
+
14
+ def is_hub_id(run) -> bool:
15
+ return not os.path.isdir(run) and HUB_ID.fullmatch(str(run)) is not None