plumbify 0.2.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.
plumbify/__init__.py ADDED
@@ -0,0 +1,20 @@
1
+ """Plumbify: a fast decision path (System 1) for open language models, served natively by vLLM.
2
+
3
+ plumbify train --base Qwen/Qwen3.5-9B --rows train.jsonl --dev dev.jsonl --out plumbed-qwen3.5-9b
4
+ vllm serve plumbed-qwen3.5-9b
5
+
6
+ In Python, ``System1`` runs a plumb with transformers (training, evaluation, the reference implementation).
7
+ """
8
+
9
+ __version__ = "0.2.0"
10
+
11
+
12
+ def __getattr__(name): # System1 pulls in torch and transformers: load it on first use
13
+ if name == "System1":
14
+ from .system1 import System1
15
+
16
+ return System1
17
+ raise AttributeError(name)
18
+
19
+
20
+ __all__ = ["System1", "__version__"]
plumbify/artifact.py ADDED
@@ -0,0 +1,102 @@
1
+ """The plumb: the trained decision branch of a model, stored next to (never inside) the base weights.
2
+
3
+ plumb.json spec: base model, taps, head shape, calibration
4
+ head.safetensors the decision head
5
+ suffix_adapter.json/.safetensors the suffix-only LoRA (absent when trained with --no_lora)
6
+
7
+ ``plumbify train-head`` writes this layout; ``plumbify package`` (plumbify/plumbed.py) combines it with the base model
8
+ into a plumbed model directory. Paths can be local directories or Hugging Face Hub repos.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import dataclasses
14
+ import json
15
+ import os
16
+ from dataclasses import dataclass, field
17
+
18
+ FORMAT = "plumb/2"
19
+ RENDER_VERSION = 2 # plumbify.core.render: bump when the decision text layout changes
20
+
21
+
22
+ @dataclass
23
+ class Conformal:
24
+ """Split-conformal threshold fitted on held-out rows: Set = {k : p_k >= 1 - qhat}, P(gold in Set) >= 1 - alpha."""
25
+
26
+ alpha: float
27
+ qhat: float
28
+ n: int
29
+
30
+
31
+ @dataclass
32
+ class Calibration:
33
+ temperature: float = 1.0
34
+ conformal: Conformal | None = None
35
+
36
+
37
+ @dataclass
38
+ class PlumbSpec:
39
+ base_model: str
40
+ d: int
41
+ d_proj: int = 512
42
+ has_adapter: bool = True
43
+ render_version: int = RENDER_VERSION
44
+ head_type: str = "decision"
45
+ head_config: dict = field(default_factory=dict)
46
+ taps: list = field(default_factory=list) # layers the head reads; -1 is the final normed output
47
+ calibration: Calibration = field(default_factory=Calibration)
48
+ name: str = ""
49
+ format: str = FORMAT
50
+ extra: dict = field(default_factory=dict)
51
+
52
+ def to_json(self) -> str:
53
+ return json.dumps(dataclasses.asdict(self), indent=1)
54
+
55
+ @staticmethod
56
+ def from_dict(d: dict) -> PlumbSpec:
57
+ cal = d.get("calibration") or {}
58
+ conf = cal.get("conformal")
59
+ d = {
60
+ **d,
61
+ "calibration": Calibration(
62
+ cal.get("temperature", 1.0), Conformal(**conf) if conf else None
63
+ ),
64
+ }
65
+ known = {f.name for f in dataclasses.fields(PlumbSpec)}
66
+ return PlumbSpec(**{k: v for k, v in d.items() if k in known})
67
+
68
+
69
+ def fetch(path: str, name: str) -> str:
70
+ """Local path of a file in a plumb directory or a Hub repo."""
71
+ if os.path.isdir(path):
72
+ return os.path.join(path, name)
73
+ from huggingface_hub import hf_hub_download
74
+
75
+ return hf_hub_download(path, name)
76
+
77
+
78
+ def read_spec(path: str) -> PlumbSpec:
79
+ try:
80
+ spec_file = fetch(path, "plumb.json")
81
+ except Exception as e: # Hub errors come in several types; the message says which
82
+ raise FileNotFoundError(f"{path}: cannot read plumb.json ({e})") from e
83
+ if not os.path.exists(spec_file):
84
+ raise FileNotFoundError(
85
+ f"{path}: no plumb.json; is this a plumb or plumbed model directory?"
86
+ )
87
+ spec = PlumbSpec.from_dict(json.load(open(spec_file)))
88
+ if spec.head_type != "decision":
89
+ raise ValueError(f"{path}: unsupported plumb head type {spec.head_type!r}")
90
+ return spec
91
+
92
+
93
+ def save(out: str, spec: PlumbSpec, head) -> None:
94
+ """Write the spec and the head (with its calibrated temperature)."""
95
+ from safetensors.torch import save_file
96
+
97
+ os.makedirs(out, exist_ok=True)
98
+ sd = {k: v.detach().contiguous().cpu() for k, v in head.state_dict().items()}
99
+ sd["temperature"] = sd["temperature"].new_tensor(spec.calibration.temperature)
100
+ save_file(sd, os.path.join(out, "head.safetensors"))
101
+ with open(os.path.join(out, "plumb.json"), "w") as f:
102
+ f.write(dataclasses.replace(spec, format=FORMAT).to_json())
plumbify/branch.py ADDED
@@ -0,0 +1,370 @@
1
+ """System 1 inside System 2's generation: a decision reads the live KV cache, then the cache is rolled back.
2
+
3
+ Qwen generating ... <tool_call>{"name": "plumb_decide", "arguments": {question, options}}</tool_call>
4
+ -> checkpoint the cache (clone the small recurrent / conv states; remember the attention length)
5
+ -> append the decision suffix (close the turn, "Decision: ... Options: ...", assistant-open), adapter ON
6
+ -> the head reads the suffix at the tapped layers -> probabilities
7
+ -> restore the cache exactly, append <tool_response>{answer}</tool_response>, keep generating
8
+
9
+ The tool call is only the trigger; nothing external is called. The decision costs its ~200 suffix tokens, never a
10
+ re-read of the conversation. That is valid because the plumb's adapter is suffix-only: the context in the live cache
11
+ was computed by the base weights, which is exactly what the adapter was trained on top of. Without a cache (a fresh
12
+ request) the same decision is ``System1.decide`` over the rendered text.
13
+ """
14
+
15
+ from __future__ import annotations
16
+
17
+ import json
18
+ import re
19
+ from collections.abc import Callable
20
+ from dataclasses import dataclass, field
21
+
22
+ import torch
23
+
24
+ from .calibration import prediction_set
25
+ from .core.answers import answer
26
+ from .core.decision_head import SuffixBatch
27
+ from .core.render import LEADS
28
+ from .core.row import Row
29
+ from .core.tool import PLUMB_TOOL, row_from_tool_args, tool_schema
30
+
31
+ _A, _D, _T = "\x00ASSISTANT\x00", "\x00DECISION\x00", "\x00TOOL\x00"
32
+ _XML_FN = re.compile(r"<function=([^>\s]+)>(.*?)</function>", re.S)
33
+ _XML_PARAM = re.compile(r"<parameter=([^>\s]+)>\n?(.*?)\n?</parameter>", re.S)
34
+
35
+
36
+ _GEMMA_CALL = re.compile(r"^call:([^\s{]+)\s*(\{.*\})$", re.S)
37
+ _GEMMA_KEY = re.compile(r"([{,]\s*)([A-Za-z_][\w\-]*)\s*:")
38
+ _GEMMA_BARE = re.compile(r"(:\s*)([A-Za-z_][\w\-]*)(\s*[,}\]])")
39
+ _GEMMA_STR = '<|"|>'
40
+
41
+
42
+ @dataclass(frozen=True)
43
+ class Markup:
44
+ """The model's own tool-call and thinking markup, read from its chat template."""
45
+
46
+ call_open: str
47
+ call_close: str
48
+ think_open: str | None
49
+ think_close: str | None
50
+
51
+ # Gemma 4; Hermes / Qwen
52
+ CALLS = (
53
+ ("<|tool_call>", "<tool_call|>"),
54
+ ("<tool_call>", "</tool_call>"),
55
+ )
56
+ # Gemma 4; Qwen, DeepSeek, GLM
57
+ THINKS = (
58
+ ("<|channel>thought", "<channel|>"),
59
+ ("<think>", "</think>"),
60
+ )
61
+
62
+ @staticmethod
63
+ def of(tok) -> Markup:
64
+ tmpl = getattr(tok, "chat_template", None) or ""
65
+ if isinstance(tmpl, dict):
66
+ tmpl = " ".join(tmpl.values())
67
+ call = next((c for c in Markup.CALLS if c[0] in tmpl and c[1] in tmpl), Markup.CALLS[1])
68
+ think = next((t for t in Markup.THINKS if t[0] in tmpl), (None, None))
69
+ return Markup(*call, *think)
70
+
71
+ def calls(self, text: str) -> list[str]:
72
+ """Bodies of the complete tool calls in ``text``."""
73
+ pat = re.escape(self.call_open) + "(.*?)" + re.escape(self.call_close)
74
+ return re.findall(pat, text, re.S)
75
+
76
+ def opens_thinking(self, prompt: str) -> bool:
77
+ """The prompt ends inside an open thinking block (generation starts as reasoning)."""
78
+ return bool(self.think_open) and prompt.rstrip().endswith(self.think_open.rstrip())
79
+
80
+ def final(self, text: str) -> str:
81
+ """The reply after the last tool call and the last thinking block, without later call markup."""
82
+ text = text.split(self.call_close)[-1]
83
+ if self.think_close:
84
+ text = text.split(self.think_close)[-1]
85
+ return text.split(self.call_open)[0]
86
+
87
+
88
+ def _gemma_args(s: str) -> dict | None:
89
+ """Gemma 4's argument syntax -> dict: bare keys, strings between ``<|"|>`` marks, JSON-like nesting."""
90
+ parts = s.split(_GEMMA_STR)
91
+ if len(parts) % 2 == 0:
92
+ return None
93
+ code = []
94
+ for i, p in enumerate(parts):
95
+ if i % 2:
96
+ code.append(json.dumps(p))
97
+ else:
98
+ p = _GEMMA_KEY.sub(r'\1"\2":', p)
99
+ p = _GEMMA_BARE.sub(
100
+ lambda m: m[0] if m[2] in ("true", "false", "null") else f'{m[1]}"{m[2]}"{m[3]}', p
101
+ )
102
+ code.append(p)
103
+ try:
104
+ v = json.loads("".join(code))
105
+ except ValueError:
106
+ return None
107
+ return v if isinstance(v, dict) else None
108
+
109
+
110
+ def parse_tool_call(body: str) -> tuple[str, dict] | None:
111
+ """One tool-call body (between the model's call markers) -> (name, arguments). Formats open models emit:
112
+ Hermes JSON ``{"name": ..., "arguments": {...}}`` (Qwen3, most templates), XML
113
+ ``<function=NAME><parameter=KEY>VALUE</parameter>...</function>`` (Qwen3.5, Qwen3-Coder) and Gemma 4's
114
+ ``call:NAME{key:<|"|>value<|"|>,...}``. XML values that parse as JSON (lists, numbers) are decoded; others stay
115
+ strings."""
116
+ body = body.strip()
117
+ g = _GEMMA_CALL.match(body)
118
+ if g:
119
+ args = _gemma_args(g[2])
120
+ return (g[1], args) if args is not None else None
121
+ if body.startswith("{"):
122
+ try:
123
+ c = json.loads(body)
124
+ except ValueError:
125
+ return None
126
+ args = c.get("arguments") or {}
127
+ if isinstance(args, str):
128
+ try:
129
+ args = json.loads(args)
130
+ except ValueError:
131
+ return None
132
+ return c.get("name"), args
133
+ m = _XML_FN.search(body)
134
+ if not m:
135
+ return None
136
+ args = {}
137
+ for k, v in _XML_PARAM.findall(m.group(2)):
138
+ try:
139
+ args[k] = json.loads(v)
140
+ except ValueError:
141
+ args[k] = v.strip()
142
+ return m.group(1), args
143
+
144
+
145
+ class CacheCheckpoint:
146
+ """Everything needed to put a transformers cache back as it was: the attention length (keys/values grow by
147
+ concatenation, so a slice undoes them) and clones of linear-attention / Mamba states (updated in place)."""
148
+
149
+ def __init__(self, cache):
150
+ self.length = cache.get_seq_length()
151
+ self.states = []
152
+ for layer in cache.layers:
153
+ for attr in ("conv_states", "recurrent_states"):
154
+ d = getattr(layer, attr, None)
155
+ if isinstance(d, dict):
156
+ self.states += [(d, i, t.clone()) for i, t in d.items() if t is not None]
157
+
158
+ def restore(self, cache) -> None:
159
+ for d, i, t in self.states:
160
+ if d[i].shape == t.shape:
161
+ d[i].copy_(t) # keep the tensor's address (static for cudagraphs)
162
+ else:
163
+ d[i] = t
164
+ for layer in cache.layers:
165
+ keys = getattr(layer, "keys", None)
166
+ if isinstance(keys, torch.Tensor) and keys.dim() == 4 and keys.shape[-2] > self.length:
167
+ layer.keys = keys[..., : self.length, :]
168
+ layer.values = layer.values[..., : self.length, :]
169
+
170
+
171
+ @dataclass
172
+ class Frames:
173
+ """Chat-template text that follows an assistant turn in progress, derived from the tokenizer's own template."""
174
+
175
+ to_decision: str # close the assistant turn, open a user turn
176
+ after_decision: str # close the user turn, open the assistant (non-thinking, as in training)
177
+ to_tool: str # close the assistant turn, open the tool response
178
+ after_tool: str # close the tool response, open the assistant
179
+
180
+ @staticmethod
181
+ def of(tok, template_kwargs: dict | None = None, markup: Markup | None = None) -> Frames:
182
+ kw = {"enable_thinking": False, **(template_kwargs or {})}
183
+ markup = markup or Markup.of(tok)
184
+ u = {"role": "user", "content": "hi"}
185
+ a = {"role": "assistant", "content": _A}
186
+ t1 = tok.apply_chat_template(
187
+ [u, a, {"role": "user", "content": _D}],
188
+ tokenize=False,
189
+ add_generation_prompt=True,
190
+ **kw,
191
+ )
192
+ # the tool response follows a real tool call: some templates (Gemma 4) render tool messages only after one,
193
+ # inside the same model turn
194
+ call = {
195
+ "role": "assistant",
196
+ "content": "",
197
+ "tool_calls": [
198
+ {
199
+ "id": "call_0",
200
+ "type": "function",
201
+ "function": {"name": PLUMB_TOOL, "arguments": {"question": "q"}},
202
+ }
203
+ ],
204
+ }
205
+ tool = {"role": "tool", "tool_call_id": "call_0", "name": PLUMB_TOOL, "content": _T}
206
+ t2 = tok.apply_chat_template(
207
+ [u, call, tool], tokenize=False, add_generation_prompt=True, **kw
208
+ )
209
+ close = t2.rfind(markup.call_close, 0, t2.index(_T)) if _T in t2 else -1
210
+ if close >= 0:
211
+ start = close + len(markup.call_close)
212
+ else: # a template that doesn't render tool calls: the tool turn follows plain assistant text
213
+ t2 = tok.apply_chat_template(
214
+ [u, a, {"role": "tool", "content": _T}],
215
+ tokenize=False,
216
+ add_generation_prompt=True,
217
+ **kw,
218
+ )
219
+ start = t2.index(_A) + len(_A)
220
+ return Frames(
221
+ t1[t1.index(_A) + len(_A) : t1.index(_D)],
222
+ t1[t1.index(_D) + len(_D) :],
223
+ t2[start : t2.index(_T)],
224
+ t2[t2.index(_T) + len(_T) :],
225
+ )
226
+
227
+
228
+ def decision_suffix(tok, frames: Frames, row: Row) -> tuple[list[int], list[tuple[int, int]], int]:
229
+ """Token ids of the decision suffix after a live assistant turn, with option spans and DECIDE relative to it.
230
+ Same pieces, same order and same instruction as ``render_decision`` (training)."""
231
+ lead, instruction = LEADS[row.qtype]
232
+ ids = tok(
233
+ frames.to_decision + f"Decision: {row.question}\n{lead}\n", add_special_tokens=False
234
+ ).input_ids
235
+ spans = []
236
+ for o in row.options:
237
+ s = len(ids)
238
+ ids += tok(
239
+ f"- {o.name}: {o.desc}\n" if o.desc else f"- {o.name}\n", add_special_tokens=False
240
+ ).input_ids
241
+ spans.append((s, len(ids)))
242
+ ids += tok(instruction + frames.after_decision, add_special_tokens=False).input_ids
243
+ return ids, spans, len(ids) - 1
244
+
245
+
246
+ @torch.no_grad()
247
+ def decide_on_cache(s1, cache, row: Row, frames: Frames) -> list[float]:
248
+ """Probabilities over ``row.options`` from the live cache; the cache is left exactly as it was."""
249
+ ids, spans, decide = decision_suffix(s1.tok, frames, row)
250
+ ck = CacheCheckpoint(cache)
251
+ dev = next(s1.trunk.parameters()).device
252
+ x = torch.tensor([ids], device=dev)
253
+ try:
254
+ s1._set_mask(torch.ones(1, len(ids), device=dev)) # every token here is suffix: adapter on
255
+ feats = s1.reader.run(
256
+ x,
257
+ None,
258
+ [0],
259
+ [len(ids)],
260
+ past_key_values=cache,
261
+ use_cache=True,
262
+ cache_position=torch.arange(ck.length, ck.length + len(ids), device=dev),
263
+ )
264
+ finally:
265
+ s1._set_mask(None)
266
+ ck.restore(cache)
267
+ item = {"feats": feats[0], "opt_spans": spans, "decide": decide}
268
+ logits = s1.head(SuffixBatch.collate([item], device=next(s1.head.parameters()).device))
269
+ return torch.softmax(logits[0, : len(row.options)].float(), -1).tolist()
270
+
271
+
272
+ @dataclass
273
+ class Turn:
274
+ text: str
275
+ decisions: list[dict] = field(default_factory=list)
276
+ tokens: int = 0
277
+
278
+
279
+ class Assistant:
280
+ """Generation with the untouched base model; ``plumb_decide`` tool calls are answered by System 1 on the live
281
+ cache. ``sample(logits) -> token id`` defaults to greedy (a scripted sampler drives the tests)."""
282
+
283
+ def __init__(
284
+ self, s1, conformal_qhat: float | None = None, template_kwargs: dict | None = None
285
+ ):
286
+ self.s1, self.tok, self.lm = s1, s1.tok, s1.lm
287
+ self.qhat = conformal_qhat
288
+ self.template_kwargs = template_kwargs or {}
289
+ self.markup = Markup.of(self.tok)
290
+ self.frames = Frames.of(self.tok, s1.template_kwargs, self.markup)
291
+
292
+ def _feed(self, ids: list[int], cache):
293
+ dev = next(self.lm.parameters()).device
294
+ start = cache.get_seq_length() if cache is not None else 0
295
+ out = self.lm(
296
+ input_ids=torch.tensor([ids], device=dev),
297
+ past_key_values=cache,
298
+ use_cache=True,
299
+ cache_position=torch.arange(start, start + len(ids), device=dev),
300
+ logits_to_keep=1,
301
+ )
302
+ return out.logits[0, -1], out.past_key_values
303
+
304
+ @torch.no_grad()
305
+ def chat(
306
+ self,
307
+ messages: list[dict],
308
+ max_new_tokens: int = 512,
309
+ tools: list | None = None,
310
+ sample: Callable[[torch.Tensor], int] | None = None,
311
+ max_decisions: int = 8,
312
+ ) -> Turn:
313
+ sample = sample or (lambda z: int(z.argmax()))
314
+ gen_eos = getattr(getattr(self.lm, "generation_config", None), "eos_token_id", None)
315
+ eos = (
316
+ {
317
+ self.tok.convert_tokens_to_ids(t)
318
+ for t in ("<|im_end|>", "<|endoftext|>")
319
+ if hasattr(self.tok, "convert_tokens_to_ids")
320
+ }
321
+ | {getattr(self.tok, "eos_token_id", None)}
322
+ | set(gen_eos if isinstance(gen_eos, list) else [gen_eos])
323
+ )
324
+ prompt = self.tok.apply_chat_template(
325
+ messages,
326
+ tools=[tool_schema(), *(tools or [])],
327
+ tokenize=False,
328
+ add_generation_prompt=True,
329
+ **self.template_kwargs,
330
+ )
331
+ logits, cache = self._feed(self.tok(prompt, add_special_tokens=False).input_ids, None)
332
+ turn, gen, handled, sampled = Turn(""), [], 0, 0
333
+ while sampled < max_new_tokens: # injected tool responses don't count against the budget
334
+ t = sample(logits)
335
+ if t in eos:
336
+ break
337
+ gen.append(t)
338
+ sampled += 1
339
+ logits, cache = self._feed([t], cache)
340
+ text = self.tok.decode(gen, skip_special_tokens=False)
341
+ calls = self.markup.calls(text)
342
+ if len(calls) > handled and len(turn.decisions) < max_decisions:
343
+ handled = len(calls)
344
+ parsed = parse_tool_call(calls[-1])
345
+ if parsed is None:
346
+ continue
347
+ name, args = parsed
348
+ if name != PLUMB_TOOL:
349
+ break # a client tool: hand the turn back
350
+ try:
351
+ row = row_from_tool_args(args)
352
+ except (KeyError, TypeError, ValueError):
353
+ continue # malformed arguments: let generation carry on
354
+ probs = decide_on_cache(self.s1, cache, row, self.frames)
355
+ ans = answer(row, probs)
356
+ if self.qhat is not None:
357
+ ans["set"] = [row.options[k].name for k in prediction_set(probs, self.qhat)]
358
+ turn.decisions.append(
359
+ {
360
+ "arguments": args,
361
+ "result": ans,
362
+ "context_tokens": cache.get_seq_length(),
363
+ }
364
+ )
365
+ resp = self.frames.to_tool + json.dumps(ans) + self.frames.after_tool
366
+ resp_ids = self.tok(resp, add_special_tokens=False).input_ids
367
+ gen += resp_ids
368
+ logits, cache = self._feed(resp_ids, cache)
369
+ turn.text, turn.tokens = self.tok.decode(gen, skip_special_tokens=True), len(gen)
370
+ return turn
@@ -0,0 +1,109 @@
1
+ """Calibration: one temperature, a split-conformal threshold, and the System 1 -> System 2 escalation it implies.
2
+
3
+ Everything works on prediction rows ``{"probs": [...], "gold": int}`` as written by ``plumbify eval-head`` (probs at T = 1).
4
+
5
+ - ``fit_temperature``: NLL-optimal T; moves confidence, never the winning option.
6
+ - ``fit_conformal``: split conformal with the LAC score s = 1 - p_gold. ``prediction_set`` then contains the gold
7
+ option with probability >= 1 - alpha on exchangeable data, distribution-free.
8
+ - ``escalation_curve``: for each confidence threshold, the share of traffic System 1 keeps and its accuracy there.
9
+ A request whose conformal set is not a singleton is exactly one System 1 should hand to System 2.
10
+ """
11
+
12
+ from __future__ import annotations
13
+
14
+ import math
15
+ from collections.abc import Sequence
16
+
17
+
18
+ def rescale(probs: Sequence[float], temperature: float) -> list[float]:
19
+ z = [math.log(max(p, 1e-12)) / temperature for p in probs]
20
+ m = max(z)
21
+ e = [math.exp(x - m) for x in z]
22
+ s = sum(e)
23
+ return [x / s for x in e]
24
+
25
+
26
+ def fit_temperature(rows: list[dict], lo: float = 0.25, hi: float = 8.0, iters: int = 60) -> float:
27
+ """Golden-section search of the mean NLL over log T (the NLL is unimodal in log T)."""
28
+ rows = [r for r in rows if r.get("gold") is not None]
29
+
30
+ def nll(log_t: float) -> float:
31
+ t = math.exp(log_t)
32
+ return -sum(math.log(max(rescale(r["probs"], t)[r["gold"]], 1e-12)) for r in rows) / max(
33
+ len(rows), 1
34
+ )
35
+
36
+ a, b = math.log(lo), math.log(hi)
37
+ g = (math.sqrt(5) - 1) / 2
38
+ c, d = b - g * (b - a), a + g * (b - a)
39
+ fc, fd = nll(c), nll(d)
40
+ for _ in range(iters):
41
+ if fc < fd:
42
+ b, d, fd = d, c, fc
43
+ c = b - g * (b - a)
44
+ fc = nll(c)
45
+ else:
46
+ a, c, fc = c, d, fd
47
+ d = a + g * (b - a)
48
+ fd = nll(d)
49
+ return math.exp((a + b) / 2)
50
+
51
+
52
+ def fit_conformal(rows: list[dict], alpha: float = 0.1, temperature: float = 1.0) -> dict:
53
+ """qhat = the ceil((n+1)(1-alpha))/n empirical quantile of s = 1 - p_gold on held-out rows."""
54
+ scores = sorted(
55
+ 1.0 - rescale(r["probs"], temperature)[r["gold"]] for r in rows if r.get("gold") is not None
56
+ )
57
+ n = len(scores)
58
+ if n == 0:
59
+ raise ValueError("conformal calibration needs labelled rows")
60
+ k = math.ceil((n + 1) * (1 - alpha))
61
+ qhat = 1.0 if k > n else scores[k - 1]
62
+ return {"alpha": alpha, "qhat": qhat, "n": n}
63
+
64
+
65
+ def prediction_set(probs: Sequence[float], qhat: float) -> list[int]:
66
+ """Options whose probability clears 1 - qhat; never empty (falls back to the argmax)."""
67
+ s = [k for k, p in enumerate(probs) if p >= 1.0 - qhat]
68
+ return s or [max(range(len(probs)), key=probs.__getitem__)]
69
+
70
+
71
+ def coverage(rows: list[dict], qhat: float, temperature: float = 1.0) -> dict:
72
+ """Empirical coverage and mean set size of the conformal sets on another labelled set."""
73
+ hit = size = singles = 0
74
+ rows = [r for r in rows if r.get("gold") is not None]
75
+ for r in rows:
76
+ s = prediction_set(rescale(r["probs"], temperature), qhat)
77
+ hit += r["gold"] in s
78
+ size += len(s)
79
+ singles += len(s) == 1
80
+ n = max(len(rows), 1)
81
+ return {
82
+ "coverage": hit / n,
83
+ "mean_set_size": size / n,
84
+ "singleton_rate": singles / n,
85
+ "n": len(rows),
86
+ }
87
+
88
+
89
+ def escalation_curve(
90
+ rows: list[dict],
91
+ temperature: float = 1.0,
92
+ thresholds: Sequence[float] = (0.0, 0.5, 0.6, 0.7, 0.8, 0.9, 0.95, 0.99),
93
+ ) -> list[dict]:
94
+ """For each p_max threshold: share kept by System 1, System 1 accuracy on what it keeps."""
95
+ rows = [r for r in rows if r.get("gold") is not None]
96
+ out = []
97
+ for t in thresholds:
98
+ kept = [r for r in rows if max(rescale(r["probs"], temperature)) >= t]
99
+ acc = sum(
100
+ max(range(len(r["probs"])), key=r["probs"].__getitem__) == r["gold"] for r in kept
101
+ )
102
+ out.append(
103
+ {
104
+ "threshold": t,
105
+ "system1_share": len(kept) / max(len(rows), 1),
106
+ "system1_accuracy": acc / max(len(kept), 1),
107
+ }
108
+ )
109
+ return out
plumbify/cli.py ADDED
@@ -0,0 +1,61 @@
1
+ """``plumbify <command>``: each command is a module with ``main()``."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import argparse
6
+ import importlib
7
+ import sys
8
+
9
+ COMMANDS = {
10
+ "train": (
11
+ "plumbify.training.train",
12
+ "train a plumb on an open model and write a plumbed model directory for vLLM",
13
+ ),
14
+ "package": (None, "turn a trained plumb into a plumbed model directory"),
15
+ "train-head": (
16
+ "plumbify.training.train_head",
17
+ "train and calibrate a plumb only (the first step of train)",
18
+ ),
19
+ "eval-head": ("plumbify.training.eval_head", "evaluate a plumb, or the base model zero-shot"),
20
+ }
21
+
22
+
23
+ def _usage() -> str:
24
+ w = max(map(len, COMMANDS))
25
+ return "usage: plumbify <command> [args]\n\n" + "\n".join(
26
+ f" {k:<{w}} {v[1]}" for k, v in COMMANDS.items()
27
+ )
28
+
29
+
30
+ def _package(rest: list[str]) -> int:
31
+ from .plumbed import package
32
+
33
+ ap = argparse.ArgumentParser("plumbify package")
34
+ ap.add_argument("plumb", help="trained plumb directory (plumb.json, head, suffix adapter)")
35
+ ap.add_argument("out", help="plumbed model directory to write")
36
+ ap.add_argument("--base", default=None, help="override the base model recorded in plumb.json")
37
+ ap.add_argument(
38
+ "--copy", action="store_true", help="copy the base weights instead of linking them"
39
+ )
40
+ a = ap.parse_args(rest)
41
+ print(package(a.plumb, a.out, a.base, a.copy))
42
+ return 0
43
+
44
+
45
+ def main(argv=None) -> int:
46
+ argv = list(sys.argv[1:] if argv is None else argv)
47
+ if not argv or argv[0] in ("-h", "--help"):
48
+ print(_usage())
49
+ return 0
50
+ cmd, rest = argv[0], argv[1:]
51
+ if cmd not in COMMANDS:
52
+ print(f"unknown command {cmd!r}\n\n{_usage()}")
53
+ return 2
54
+ if cmd == "package":
55
+ return _package(rest)
56
+ sys.argv = [f"plumbify {cmd}", *rest] # command modules parse sys.argv
57
+ return importlib.import_module(COMMANDS[cmd][0]).main() or 0
58
+
59
+
60
+ if __name__ == "__main__":
61
+ raise SystemExit(main())
@@ -0,0 +1,6 @@
1
+ """The pieces shared by training and serving: the decision row, its rendering, the suffix-only LoRA, the decision
2
+ head and the layer taps."""
3
+
4
+ from .row import QTYPES, Option, Row
5
+
6
+ __all__ = ["QTYPES", "Option", "Row"]