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 +20 -0
- plumbify/artifact.py +102 -0
- plumbify/branch.py +370 -0
- plumbify/calibration.py +109 -0
- plumbify/cli.py +61 -0
- plumbify/core/__init__.py +6 -0
- plumbify/core/answers.py +75 -0
- plumbify/core/decision_head.py +178 -0
- plumbify/core/render.py +107 -0
- plumbify/core/row.py +45 -0
- plumbify/core/suffix_lora.py +124 -0
- plumbify/core/taps.py +83 -0
- plumbify/core/targets.py +161 -0
- plumbify/core/tool.py +61 -0
- plumbify/metrics.py +73 -0
- plumbify/plumbed.py +390 -0
- plumbify/py.typed +0 -0
- plumbify/serving/__init__.py +1 -0
- plumbify/serving/vllm/__init__.py +89 -0
- plumbify/serving/vllm/client.py +419 -0
- plumbify/serving/vllm/decisions.py +103 -0
- plumbify/serving/vllm/hooks.py +78 -0
- plumbify/serving/vllm/model.py +364 -0
- plumbify/serving/vllm/openai.py +384 -0
- plumbify/system1.py +277 -0
- plumbify/training/__init__.py +1 -0
- plumbify/training/data.py +37 -0
- plumbify/training/eval_head.py +142 -0
- plumbify/training/train.py +84 -0
- plumbify/training/train_head.py +264 -0
- plumbify-0.2.0.dist-info/METADATA +212 -0
- plumbify-0.2.0.dist-info/RECORD +36 -0
- plumbify-0.2.0.dist-info/WHEEL +5 -0
- plumbify-0.2.0.dist-info/entry_points.txt +5 -0
- plumbify-0.2.0.dist-info/licenses/LICENSE +202 -0
- plumbify-0.2.0.dist-info/top_level.txt +1 -0
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
|
plumbify/calibration.py
ADDED
|
@@ -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())
|