sdm-learn 0.1.0.dev0__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.
- sdm/__init__.py +71 -0
- sdm/__main__.py +51 -0
- sdm/data.py +238 -0
- sdm/format.md +84 -0
- sdm/formats.py +1182 -0
- sdm/gateway.py +169 -0
- sdm/learn.py +699 -0
- sdm/signals.py +386 -0
- sdm_learn-0.1.0.dev0.dist-info/METADATA +35 -0
- sdm_learn-0.1.0.dev0.dist-info/RECORD +14 -0
- sdm_learn-0.1.0.dev0.dist-info/WHEEL +5 -0
- sdm_learn-0.1.0.dev0.dist-info/entry_points.txt +2 -0
- sdm_learn-0.1.0.dev0.dist-info/licenses/LICENSE +201 -0
- sdm_learn-0.1.0.dev0.dist-info/top_level.txt +1 -0
sdm/learn.py
ADDED
|
@@ -0,0 +1,699 @@
|
|
|
1
|
+
"""The learning algorithm: the gate, the control, and the run artifacts.
|
|
2
|
+
|
|
3
|
+
Everything here is format-agnostic. The gate calls into a ``SkillFormat`` for
|
|
4
|
+
anything that depends on how the hypothesis is written down, so adding a
|
|
5
|
+
format never touches this file.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import json
|
|
11
|
+
import os
|
|
12
|
+
import random
|
|
13
|
+
from pathlib import Path
|
|
14
|
+
|
|
15
|
+
from .data import (Dataset, Report, _inputs, _revealed, content, describe,
|
|
16
|
+
labeled)
|
|
17
|
+
from .formats import (DOCUMENT_RE, GATE_LINE_RE, NO_FORMAT, SkillFormat,
|
|
18
|
+
_clean_line, _copy, _rules_by_id, format_of, get_format,
|
|
19
|
+
parse_predictions, record, render, skill_filename)
|
|
20
|
+
from . import gateway
|
|
21
|
+
from .gateway import EFFORT, MODEL, SetupNeeded, _borrowed
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def _atomic(path: Path, text: str) -> None:
|
|
26
|
+
path.parent.mkdir(parents=True, exist_ok=True)
|
|
27
|
+
temporary = path.with_suffix(path.suffix + ".tmp")
|
|
28
|
+
temporary.write_text(text)
|
|
29
|
+
os.replace(temporary, path)
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def save_state(state: dict, run_dir) -> None:
|
|
33
|
+
"""Write ``state.json`` and the rendered skill file side by side.
|
|
34
|
+
|
|
35
|
+
The baseline control has no document, so it writes no skill file rather
|
|
36
|
+
than an empty one.
|
|
37
|
+
"""
|
|
38
|
+
run_dir = Path(run_dir)
|
|
39
|
+
_atomic(run_dir / "state.json", json.dumps(state, indent=1))
|
|
40
|
+
document = render(state)
|
|
41
|
+
if document:
|
|
42
|
+
_atomic(run_dir / skill_filename(state), document)
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def archive(run_dir, relative: str, raw: str) -> str | None:
|
|
46
|
+
"""Persist one raw model response so a finished run can be audited."""
|
|
47
|
+
if run_dir is None:
|
|
48
|
+
return None
|
|
49
|
+
path = Path(run_dir) / relative
|
|
50
|
+
path.parent.mkdir(parents=True, exist_ok=True)
|
|
51
|
+
path.write_text(raw)
|
|
52
|
+
return relative
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
def _resume(run_dir, config: dict, what: str) -> dict | None:
|
|
56
|
+
if run_dir is None:
|
|
57
|
+
return None
|
|
58
|
+
path = Path(run_dir) / "state.json"
|
|
59
|
+
if not path.exists():
|
|
60
|
+
return None
|
|
61
|
+
state = json.loads(path.read_text())
|
|
62
|
+
if state["config"] != config:
|
|
63
|
+
raise SetupNeeded(
|
|
64
|
+
f"{run_dir} already holds a {what} with different settings, so "
|
|
65
|
+
"resuming it would mix two experiments.\nPoint at a fresh "
|
|
66
|
+
"directory, or start this one over with resume=False."
|
|
67
|
+
)
|
|
68
|
+
return state
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
# --------------------------------------------------------------------------
|
|
72
|
+
|
|
73
|
+
class Skill:
|
|
74
|
+
"""A learned skill document, plus the state that produced it."""
|
|
75
|
+
|
|
76
|
+
def __init__(self, state: dict, run_dir=None, report: Report | None = None):
|
|
77
|
+
self.state = state
|
|
78
|
+
self.run_dir = Path(run_dir) if run_dir else None
|
|
79
|
+
self.report = report
|
|
80
|
+
|
|
81
|
+
@property
|
|
82
|
+
def text(self) -> str:
|
|
83
|
+
return render(self.state)
|
|
84
|
+
|
|
85
|
+
@property
|
|
86
|
+
def format(self) -> str | None:
|
|
87
|
+
return self.state.get("config", {}).get("format", "markdown")
|
|
88
|
+
|
|
89
|
+
@property
|
|
90
|
+
def rules(self) -> list[dict]:
|
|
91
|
+
return list(self.state.get("rules", []))
|
|
92
|
+
|
|
93
|
+
@property
|
|
94
|
+
def accuracy(self) -> float | None:
|
|
95
|
+
return self.report.accuracy if self.report else None
|
|
96
|
+
|
|
97
|
+
def save(self, path) -> Path:
|
|
98
|
+
"""Write the document to a file, or the whole run to a directory."""
|
|
99
|
+
path = Path(path)
|
|
100
|
+
if path.suffix:
|
|
101
|
+
path.parent.mkdir(parents=True, exist_ok=True)
|
|
102
|
+
path.write_text(self.text)
|
|
103
|
+
return path
|
|
104
|
+
save_state(self.state, path)
|
|
105
|
+
return path / skill_filename(self.state)
|
|
106
|
+
|
|
107
|
+
@classmethod
|
|
108
|
+
def load(cls, path) -> "Skill":
|
|
109
|
+
path = Path(path)
|
|
110
|
+
state_path = path / "state.json" if path.is_dir() else path
|
|
111
|
+
return cls(json.loads(state_path.read_text()), state_path.parent)
|
|
112
|
+
|
|
113
|
+
def __str__(self):
|
|
114
|
+
return self.text
|
|
115
|
+
|
|
116
|
+
def __repr__(self):
|
|
117
|
+
parts = [f"format={self.format!r}"]
|
|
118
|
+
if self.rules:
|
|
119
|
+
parts.append(f"rules={len(self.rules)}")
|
|
120
|
+
if self.accuracy is not None:
|
|
121
|
+
parts.append(f"accuracy={self.accuracy:.1%}")
|
|
122
|
+
return f"Skill({', '.join(parts)})"
|
|
123
|
+
|
|
124
|
+
def predict_batch(session, state: dict, rows: list[dict], description: str,
|
|
125
|
+
representation: str, image_size: int = 224,
|
|
126
|
+
attempts: int = 3, fmt: SkillFormat | None = None,
|
|
127
|
+
cache=None):
|
|
128
|
+
"""Classify one batch from the current skill. Labels stay hidden."""
|
|
129
|
+
fmt = fmt or format_of(state)
|
|
130
|
+
classes = state["config"]["classes"]
|
|
131
|
+
view_text, view_images = fmt.view(state, cache)
|
|
132
|
+
prompt = fmt.predict_prompt(description, classes, view_text, rows,
|
|
133
|
+
representation)
|
|
134
|
+
|
|
135
|
+
for attempt in range(1, attempts + 1):
|
|
136
|
+
raw = gateway.complete(session, [{"role": "user", "content": content(
|
|
137
|
+
prompt, rows, representation, "CURRENT_INPUT", image_size,
|
|
138
|
+
labeled("SKILL_VIEW", view_images))}], 8000)
|
|
139
|
+
parsed, errors = parse_predictions(raw, len(rows), classes,
|
|
140
|
+
fmt.cites_rules)
|
|
141
|
+
if not errors:
|
|
142
|
+
known = set(_rules_by_id(state))
|
|
143
|
+
for prediction in parsed.values():
|
|
144
|
+
prediction["rules_used"] = [
|
|
145
|
+
rule_id for rule_id in prediction["rules_used"]
|
|
146
|
+
if rule_id in known]
|
|
147
|
+
return [parsed[i] for i in range(1, len(rows) + 1)], raw, attempt
|
|
148
|
+
problem = "; ".join(errors)
|
|
149
|
+
print(f" [prediction format retry {attempt}/{attempts}: {problem}]",
|
|
150
|
+
flush=True)
|
|
151
|
+
prompt += (f"\n\nYour previous response was invalid: {problem}. "
|
|
152
|
+
"Return a complete corrected <predictions> block only.")
|
|
153
|
+
raise RuntimeError("prediction output never satisfied the batch schema")
|
|
154
|
+
|
|
155
|
+
|
|
156
|
+
# --------------------------------------------------------------------------
|
|
157
|
+
|
|
158
|
+
DEFAULT_ORDER_SEED = 20260818
|
|
159
|
+
ACCEPT_CRITERION = (
|
|
160
|
+
"strict held-out accuracy gain or 5pct prediction-preserving compression")
|
|
161
|
+
|
|
162
|
+
|
|
163
|
+
def stratified_holdout(rows: list[dict], per_class: int, seed: int,
|
|
164
|
+
order_seed: int, *, thin_classes_feedback_only=False
|
|
165
|
+
) -> tuple[list[int], list[int]]:
|
|
166
|
+
"""Split training positions into (protected, feedback) index lists.
|
|
167
|
+
|
|
168
|
+
Both halves are drawn per class and then shuffled, so every request the
|
|
169
|
+
model sees is class balanced rather than sorted by label.
|
|
170
|
+
"""
|
|
171
|
+
by_label: dict[str, list[int]] = {}
|
|
172
|
+
for position, row in enumerate(rows):
|
|
173
|
+
by_label.setdefault(str(row["label"]), []).append(position)
|
|
174
|
+
|
|
175
|
+
rng = random.Random(seed)
|
|
176
|
+
for positions in by_label.values():
|
|
177
|
+
rng.shuffle(positions)
|
|
178
|
+
|
|
179
|
+
thin = {label: len(positions) for label, positions in by_label.items()
|
|
180
|
+
if len(positions) <= per_class}
|
|
181
|
+
if thin and not thin_classes_feedback_only:
|
|
182
|
+
raise ValueError(
|
|
183
|
+
f"need more than {per_class} training example(s) per class to hold "
|
|
184
|
+
f"any out; these have too few: {thin}")
|
|
185
|
+
|
|
186
|
+
held, feedback = [], []
|
|
187
|
+
for _, positions in sorted(by_label.items()):
|
|
188
|
+
if len(positions) <= per_class:
|
|
189
|
+
# A singleton cannot sit on both sides of the gate without
|
|
190
|
+
# duplicating the same labeled input, so keep it feedback-only.
|
|
191
|
+
feedback += positions
|
|
192
|
+
continue
|
|
193
|
+
held += positions[:per_class]
|
|
194
|
+
feedback += positions[per_class:]
|
|
195
|
+
random.Random(order_seed).shuffle(held)
|
|
196
|
+
rng.shuffle(feedback)
|
|
197
|
+
return held, feedback
|
|
198
|
+
|
|
199
|
+
|
|
200
|
+
def cycle_points(feedback: int, batch_size: int, cycles: int) -> tuple[int, ...]:
|
|
201
|
+
"""Cumulative feedback counts at which a candidate revision is proposed.
|
|
202
|
+
|
|
203
|
+
A cycle is at least one whole batch, so asking for more cycles than there
|
|
204
|
+
are batches yields one per batch, and any remainder joins the final cycle
|
|
205
|
+
rather than becoming a short cycle of its own.
|
|
206
|
+
"""
|
|
207
|
+
batches = feedback // batch_size
|
|
208
|
+
if batches == 0:
|
|
209
|
+
raise ValueError(
|
|
210
|
+
f"{feedback} feedback example(s) cannot fill one batch of "
|
|
211
|
+
f"{batch_size}; lower batch_size or hold out fewer rows")
|
|
212
|
+
cycles = max(1, min(cycles, batches))
|
|
213
|
+
per_cycle, extra = divmod(batches, cycles)
|
|
214
|
+
sizes = [per_cycle] * cycles
|
|
215
|
+
sizes[-1] += extra
|
|
216
|
+
points, total = [], 0
|
|
217
|
+
for size in sizes:
|
|
218
|
+
total += size * batch_size
|
|
219
|
+
points.append(total)
|
|
220
|
+
return tuple(points)
|
|
221
|
+
|
|
222
|
+
|
|
223
|
+
def gate_prompt(fmt: SkillFormat, base: dict, candidate: dict,
|
|
224
|
+
rows: list[dict], order: list[str], description: str,
|
|
225
|
+
classes: list[str], representation: str, image_size: int,
|
|
226
|
+
cache=None):
|
|
227
|
+
"""One request that scores both documents, with the letters shuffled."""
|
|
228
|
+
states = {"base": base, "candidate": candidate}
|
|
229
|
+
documents = {letter: states[name] for letter, name in zip("AB", order)}
|
|
230
|
+
views = {}
|
|
231
|
+
for letter in ("A", "B"):
|
|
232
|
+
views[letter] = fmt.view(documents[letter], cache)
|
|
233
|
+
prompt = f"""You are a deterministic evaluation component for two frozen
|
|
234
|
+
{fmt.name} skill documents. {description}
|
|
235
|
+
|
|
236
|
+
Classify every input independently once with document A and once with document
|
|
237
|
+
B. Do not merge, compare, vote between, or transfer rules across documents. The
|
|
238
|
+
inputs are identical for both documents and their true labels are absent.
|
|
239
|
+
Document letters do not indicate recency or quality.
|
|
240
|
+
|
|
241
|
+
<document id="A">
|
|
242
|
+
{fmt.skill_section(views["A"][0])}
|
|
243
|
+
</document>
|
|
244
|
+
|
|
245
|
+
<document id="B">
|
|
246
|
+
{fmt.skill_section(views["B"][0])}
|
|
247
|
+
</document>
|
|
248
|
+
|
|
249
|
+
{_inputs(rows, representation, "current_input")}
|
|
250
|
+
|
|
251
|
+
Return only these two complete blocks, with exactly {len(rows)} lines each:
|
|
252
|
+
<predictions document="A">
|
|
253
|
+
P01 | class | evidence of at most 20 words
|
|
254
|
+
</predictions>
|
|
255
|
+
<predictions document="B">
|
|
256
|
+
P01 | class | evidence of at most 20 words
|
|
257
|
+
</predictions>
|
|
258
|
+
|
|
259
|
+
The class must be exactly one of: {", ".join(classes)}.
|
|
260
|
+
"""
|
|
261
|
+
return content(prompt, rows, representation, "CURRENT_INPUT", image_size,
|
|
262
|
+
labeled("DOCUMENT_A_VIEW", views["A"][1])
|
|
263
|
+
+ labeled("DOCUMENT_B_VIEW", views["B"][1]))
|
|
264
|
+
|
|
265
|
+
|
|
266
|
+
def candidate_prompt(fmt: SkillFormat, state: dict, rows: list[dict],
|
|
267
|
+
predictions: list[dict], description: str,
|
|
268
|
+
classes: list[str], representation: str, limits: dict,
|
|
269
|
+
image_size: int, cache=None):
|
|
270
|
+
"""Ask for one conservative revision, without the protected rows."""
|
|
271
|
+
view_text, view_images = fmt.view(state, cache)
|
|
272
|
+
prompt = f"""You maintain a {fmt.noun} for a supervised learner.
|
|
273
|
+
{description}
|
|
274
|
+
Classes: {", ".join(classes)}.
|
|
275
|
+
|
|
276
|
+
The feedback inputs below were predicted before their labels were revealed.
|
|
277
|
+
Propose one conservative candidate revision. It will be accepted or rejected
|
|
278
|
+
only by Python accuracy on separate held-out training inputs whose labels and
|
|
279
|
+
errors are never shown to you.
|
|
280
|
+
|
|
281
|
+
<current_source>
|
|
282
|
+
{fmt.source(state)}
|
|
283
|
+
</current_source>
|
|
284
|
+
|
|
285
|
+
<what_the_model_sees>
|
|
286
|
+
{view_text}
|
|
287
|
+
</what_the_model_sees>
|
|
288
|
+
|
|
289
|
+
{_revealed(rows, predictions, representation, fmt.cites_rules)}
|
|
290
|
+
|
|
291
|
+
{fmt.update_instructions(limits)}"""
|
|
292
|
+
return content(prompt, rows, representation, "FEEDBACK_INPUT", image_size,
|
|
293
|
+
labeled("SKILL_VIEW", view_images))
|
|
294
|
+
|
|
295
|
+
|
|
296
|
+
def parse_paired(raw: str, expected: int) -> tuple[dict | None, list[str]]:
|
|
297
|
+
"""Read the two per-document prediction blocks out of one response."""
|
|
298
|
+
parsed: dict[str, dict[int, str]] = {}
|
|
299
|
+
errors: list[str] = []
|
|
300
|
+
for document, body in DOCUMENT_RE.findall(raw):
|
|
301
|
+
document = document.upper()
|
|
302
|
+
if document in parsed:
|
|
303
|
+
errors.append(f"duplicate document {document}")
|
|
304
|
+
continue
|
|
305
|
+
values: dict[int, str] = {}
|
|
306
|
+
for raw_line in body.splitlines():
|
|
307
|
+
match = GATE_LINE_RE.match(_clean_line(raw_line))
|
|
308
|
+
if not match:
|
|
309
|
+
continue
|
|
310
|
+
position = int(match.group(1))
|
|
311
|
+
if not 1 <= position <= expected:
|
|
312
|
+
errors.append(f"out-of-range P{position:02d} in {document}")
|
|
313
|
+
elif position in values:
|
|
314
|
+
errors.append(f"duplicate P{position:02d} in {document}")
|
|
315
|
+
else:
|
|
316
|
+
values[position] = match.group(2).strip()
|
|
317
|
+
missing = [position for position in range(1, expected + 1)
|
|
318
|
+
if position not in values]
|
|
319
|
+
if missing:
|
|
320
|
+
errors.append(f"document {document} missing {missing}")
|
|
321
|
+
parsed[document] = values
|
|
322
|
+
for document in ("A", "B"):
|
|
323
|
+
if document not in parsed:
|
|
324
|
+
errors.append(f"missing document {document}")
|
|
325
|
+
if errors:
|
|
326
|
+
return None, errors
|
|
327
|
+
return {document: [parsed[document][position]
|
|
328
|
+
for position in range(1, expected + 1)]
|
|
329
|
+
for document in ("A", "B")}, []
|
|
330
|
+
|
|
331
|
+
|
|
332
|
+
def decide(base_labels: list, candidate_labels: list, true_labels: list,
|
|
333
|
+
base_words: int, candidate_words: int, valid: bool) -> dict:
|
|
334
|
+
"""Keep a candidate only on a strict gain, or a safe compression.
|
|
335
|
+
|
|
336
|
+
A candidate wins by raising protected accuracy, or by shortening the
|
|
337
|
+
document at least 5% while predicting every protected row identically.
|
|
338
|
+
"""
|
|
339
|
+
base_correct = sum(prediction == label for prediction, label
|
|
340
|
+
in zip(base_labels, true_labels))
|
|
341
|
+
candidate_correct = sum(prediction == label for prediction, label
|
|
342
|
+
in zip(candidate_labels, true_labels))
|
|
343
|
+
fixed = sum(before != label and after == label for before, after, label
|
|
344
|
+
in zip(base_labels, candidate_labels, true_labels))
|
|
345
|
+
broken = sum(before == label and after != label for before, after, label
|
|
346
|
+
in zip(base_labels, candidate_labels, true_labels))
|
|
347
|
+
accuracy_accept = candidate_correct > base_correct
|
|
348
|
+
compression_accept = (base_labels == candidate_labels
|
|
349
|
+
and candidate_words * 20 <= base_words * 19)
|
|
350
|
+
return {
|
|
351
|
+
"accepted": bool(valid and (accuracy_accept or compression_accept)),
|
|
352
|
+
"candidate_valid": bool(valid), "criterion": ACCEPT_CRITERION,
|
|
353
|
+
"base_correct": int(base_correct),
|
|
354
|
+
"candidate_correct": int(candidate_correct),
|
|
355
|
+
"total": len(true_labels), "fixed": int(fixed), "broken": int(broken),
|
|
356
|
+
"base_words": base_words, "candidate_words": candidate_words,
|
|
357
|
+
"reason": ("accuracy_gain" if accuracy_accept else
|
|
358
|
+
"safe_compression" if compression_accept else
|
|
359
|
+
"invalid_candidate" if not valid else "no_qualifying_gain"),
|
|
360
|
+
}
|
|
361
|
+
|
|
362
|
+
|
|
363
|
+
def _propose(session, fmt: SkillFormat, state: dict, rows: list[dict],
|
|
364
|
+
predictions: list[dict], description: str, classes: list[str],
|
|
365
|
+
representation: str, limits: dict, image_size: int, cache=None):
|
|
366
|
+
"""Ask for one bounded revision and apply it to a copy of the skill."""
|
|
367
|
+
raw = gateway.complete(session, [{"role": "user", "content": candidate_prompt(
|
|
368
|
+
fmt, state, rows, predictions, description, classes, representation,
|
|
369
|
+
limits, image_size, cache)}], 16000)
|
|
370
|
+
candidate = _copy(state)
|
|
371
|
+
changes, errors = fmt.propose(candidate, raw, limits)
|
|
372
|
+
if errors:
|
|
373
|
+
candidate = _copy(state)
|
|
374
|
+
return candidate, changes, raw, errors
|
|
375
|
+
|
|
376
|
+
|
|
377
|
+
def fit_gated(data: Dataset, *, format=None, format_file=None,
|
|
378
|
+
fmt: SkillFormat | None = None, batch_size: int = 10,
|
|
379
|
+
holdout_per_class=None, cycles: int = 4, max_words=None,
|
|
380
|
+
max_rules=None, max_chars=None, max_semantic_edits=None,
|
|
381
|
+
max_support_edits=None, run_dir=None, resume: bool = True,
|
|
382
|
+
seed: int = 0, order_seed: int = DEFAULT_ORDER_SEED,
|
|
383
|
+
thin_classes_feedback_only: bool = False, session=None) -> Skill:
|
|
384
|
+
"""Learn a skill document under the accept/reject gate."""
|
|
385
|
+
if not data.train:
|
|
386
|
+
raise ValueError("data.train is empty")
|
|
387
|
+
if not data.classes:
|
|
388
|
+
raise ValueError("the gate needs data.classes to hold out per class")
|
|
389
|
+
|
|
390
|
+
fmt = fmt or get_format(format, format_file)
|
|
391
|
+
limits = fmt.limits(max_words=max_words, max_rules=max_rules,
|
|
392
|
+
max_chars=max_chars, max_semantic=max_semantic_edits,
|
|
393
|
+
max_support=max_support_edits)
|
|
394
|
+
if holdout_per_class is None:
|
|
395
|
+
holdout_per_class = max(1, len(data.train) // (10 * len(data.classes)))
|
|
396
|
+
held, feedback = stratified_holdout(
|
|
397
|
+
data.train, holdout_per_class, seed, order_seed,
|
|
398
|
+
thin_classes_feedback_only=thin_classes_feedback_only)
|
|
399
|
+
points = cycle_points(len(feedback), batch_size, cycles)
|
|
400
|
+
description = describe(data)
|
|
401
|
+
classes = list(data.classes)
|
|
402
|
+
canonical = {label.lower(): label for label in classes}
|
|
403
|
+
config = {
|
|
404
|
+
"dataset": data.name, "title": data.title,
|
|
405
|
+
"algorithm": "accept/reject paired gate", "classes": classes,
|
|
406
|
+
"format": fmt.name, "format_md_sha": getattr(fmt, "digest", ""),
|
|
407
|
+
"skill_file": fmt.skill_file,
|
|
408
|
+
"representation": data.representation, "image_size": data.image_size,
|
|
409
|
+
"description": description, "train": len(data.train),
|
|
410
|
+
"feedback_examples": points[-1], "held_out_examples": len(held),
|
|
411
|
+
"held_out_positions": held,
|
|
412
|
+
"feedback_positions": feedback[:points[-1]],
|
|
413
|
+
"unused_feedback_positions": feedback[points[-1]:],
|
|
414
|
+
"thin_classes_feedback_only": thin_classes_feedback_only,
|
|
415
|
+
"held_out_labels_never_prompted": True, "batch_size": batch_size,
|
|
416
|
+
"candidate_update_points": list(points),
|
|
417
|
+
"max_semantic_edits_per_candidate": limits["max_semantic"],
|
|
418
|
+
"max_support_edits_per_candidate": limits["max_support"],
|
|
419
|
+
"max_words": limits["max_words"], "max_rules": limits["max_rules"],
|
|
420
|
+
"max_chars": limits["max_chars"],
|
|
421
|
+
"acceptance_criterion": ACCEPT_CRITERION, "seed": seed,
|
|
422
|
+
"gate_order_seed": order_seed, "model": MODEL, "effort": EFFORT,
|
|
423
|
+
}
|
|
424
|
+
run_dir = Path(run_dir) if run_dir else None
|
|
425
|
+
cache = (run_dir / "view_cache") if run_dir else None
|
|
426
|
+
state = _resume(run_dir if resume else None, config, "gate run")
|
|
427
|
+
if state is None:
|
|
428
|
+
state = fmt.new_state(config)
|
|
429
|
+
state["gate"] = {"cycles_completed": 0}
|
|
430
|
+
state["candidates"] = []
|
|
431
|
+
if run_dir:
|
|
432
|
+
save_state(state, run_dir)
|
|
433
|
+
else:
|
|
434
|
+
print(f"resuming after cycle {state['gate']['cycles_completed']}"
|
|
435
|
+
f"/{len(points)}", flush=True)
|
|
436
|
+
|
|
437
|
+
gate_rows = [data.train[position] for position in held]
|
|
438
|
+
truth = [str(row["label"]) for row in gate_rows]
|
|
439
|
+
with _borrowed(session) as session:
|
|
440
|
+
completed = int(state["gate"]["cycles_completed"])
|
|
441
|
+
for cycle, (start, stop) in enumerate(
|
|
442
|
+
zip((0,) + points[:-1], points), 1):
|
|
443
|
+
if cycle <= completed:
|
|
444
|
+
continue
|
|
445
|
+
|
|
446
|
+
# (a) Predict this cycle's feedback batches and reveal their labels.
|
|
447
|
+
cycle_rows, cycle_predictions = [], []
|
|
448
|
+
for offset in range(0, stop - start, batch_size):
|
|
449
|
+
positions = feedback[start + offset:
|
|
450
|
+
start + offset + batch_size]
|
|
451
|
+
rows = [data.train[position] for position in positions]
|
|
452
|
+
predictions, raw, tries = predict_batch(
|
|
453
|
+
session, state, rows, description, data.representation,
|
|
454
|
+
data.image_size, fmt=fmt, cache=cache)
|
|
455
|
+
archive(run_dir, f"raw/train/cycle_{cycle:02d}_batch_"
|
|
456
|
+
f"{offset // batch_size:02d}.txt", raw)
|
|
457
|
+
record(state, "feedback_prediction", cycle=cycle,
|
|
458
|
+
examples=len(rows), attempts=tries)
|
|
459
|
+
fmt.credit(state, predictions,
|
|
460
|
+
[str(row["label"]) for row in rows])
|
|
461
|
+
for position, row, prediction in zip(positions, rows,
|
|
462
|
+
predictions):
|
|
463
|
+
state["train"].append({
|
|
464
|
+
"step": len(state["train"]) + 1, "position": position,
|
|
465
|
+
"index": row.get("index", position),
|
|
466
|
+
"true": str(row["label"]), "pred": prediction["pred"],
|
|
467
|
+
"correct": prediction["pred"] == str(row["label"]),
|
|
468
|
+
"rules_used": prediction["rules_used"],
|
|
469
|
+
"evidence": prediction["evidence"]})
|
|
470
|
+
cycle_rows += rows
|
|
471
|
+
cycle_predictions += predictions
|
|
472
|
+
|
|
473
|
+
# (b) Ask for one conservative revision, applied only to a copy.
|
|
474
|
+
base_state = _copy(state)
|
|
475
|
+
candidate_state, changes, raw, candidate_errors = _propose(
|
|
476
|
+
session, fmt, base_state, cycle_rows, cycle_predictions,
|
|
477
|
+
description, classes, data.representation, limits,
|
|
478
|
+
data.image_size, cache)
|
|
479
|
+
archive(run_dir,
|
|
480
|
+
f"raw/candidates/cycle_{cycle:02d}_update.txt", raw)
|
|
481
|
+
record(state, "candidate_update", cycle=cycle,
|
|
482
|
+
examples=len(cycle_rows),
|
|
483
|
+
candidate_valid=not candidate_errors)
|
|
484
|
+
|
|
485
|
+
# (c) Score both documents head to head on the protected rows.
|
|
486
|
+
paired = {"base": [], "candidate": []}
|
|
487
|
+
orders, gate_errors = [], []
|
|
488
|
+
for gate_batch, gate_start in enumerate(
|
|
489
|
+
range(0, len(gate_rows), batch_size)):
|
|
490
|
+
slice_rows = gate_rows[gate_start:gate_start + batch_size]
|
|
491
|
+
order = ["base", "candidate"]
|
|
492
|
+
random.Random(order_seed + 10 * cycle + gate_batch).shuffle(
|
|
493
|
+
order)
|
|
494
|
+
raw = gateway.complete(session, [{"role": "user",
|
|
495
|
+
"content": gate_prompt(
|
|
496
|
+
fmt, base_state,
|
|
497
|
+
candidate_state, slice_rows,
|
|
498
|
+
order, description, classes,
|
|
499
|
+
data.representation,
|
|
500
|
+
data.image_size, cache)}], 8000)
|
|
501
|
+
archive(run_dir, f"raw/gate/cycle_{cycle:02d}_batch_"
|
|
502
|
+
f"{gate_batch:02d}.txt", raw)
|
|
503
|
+
record(state, "paired_gate", cycle=cycle,
|
|
504
|
+
gate_batch=gate_batch, examples=len(slice_rows),
|
|
505
|
+
documents=2, document_order=order)
|
|
506
|
+
part, part_errors = parse_paired(raw, len(slice_rows))
|
|
507
|
+
orders.append(order)
|
|
508
|
+
gate_errors += part_errors
|
|
509
|
+
if part is None:
|
|
510
|
+
paired = None
|
|
511
|
+
elif paired is not None:
|
|
512
|
+
for letter, name in zip("AB", order):
|
|
513
|
+
paired[name] += [
|
|
514
|
+
canonical.get(value.lower(), value)
|
|
515
|
+
for value in part[letter]]
|
|
516
|
+
|
|
517
|
+
# (d) Decide in Python. The protected labels never left this file.
|
|
518
|
+
if paired is None:
|
|
519
|
+
decision = {"accepted": False, "candidate_valid": False,
|
|
520
|
+
"criterion": "valid paired prediction required",
|
|
521
|
+
"base_correct": None, "candidate_correct": None,
|
|
522
|
+
"total": len(gate_rows), "fixed": None,
|
|
523
|
+
"broken": None, "reason": "gate_format_error"}
|
|
524
|
+
else:
|
|
525
|
+
decision = decide(
|
|
526
|
+
paired["base"], paired["candidate"], truth,
|
|
527
|
+
fmt.words(base_state), fmt.words(candidate_state),
|
|
528
|
+
not candidate_errors and not gate_errors)
|
|
529
|
+
|
|
530
|
+
# The audit is a cost ledger, so it survives a rejected candidate.
|
|
531
|
+
audit = _copy(state["request_audit"])
|
|
532
|
+
state.clear()
|
|
533
|
+
state.update(candidate_state if decision["accepted"]
|
|
534
|
+
else base_state)
|
|
535
|
+
state["request_audit"] = audit
|
|
536
|
+
state["candidates"].append({
|
|
537
|
+
"cycle": cycle, "feedback_through": stop,
|
|
538
|
+
"edits": [list(change) for change in changes],
|
|
539
|
+
"candidate_errors": candidate_errors,
|
|
540
|
+
"gate_errors": gate_errors, "document_orders": orders,
|
|
541
|
+
"gate_predictions_without_labels": paired,
|
|
542
|
+
"decision": decision})
|
|
543
|
+
state["gate"]["cycles_completed"] = cycle
|
|
544
|
+
if run_dir:
|
|
545
|
+
_atomic(run_dir / "candidate_skills" /
|
|
546
|
+
f"candidate_{cycle:02d}{Path(fmt.skill_file).suffix}",
|
|
547
|
+
render(candidate_state))
|
|
548
|
+
_atomic(run_dir / "candidates.json",
|
|
549
|
+
json.dumps(state["candidates"], indent=2))
|
|
550
|
+
save_state(state, run_dir)
|
|
551
|
+
fmt.render_media(state, run_dir)
|
|
552
|
+
print(f"cycle {cycle}: candidate {decision['candidate_correct']}/"
|
|
553
|
+
f"{len(gate_rows)} vs base {decision['base_correct']}/"
|
|
554
|
+
f"{len(gate_rows)} -> "
|
|
555
|
+
f"{'ACCEPT' if decision['accepted'] else 'REJECT'}",
|
|
556
|
+
flush=True)
|
|
557
|
+
return Skill(state, run_dir)
|
|
558
|
+
|
|
559
|
+
|
|
560
|
+
# --------------------------------------------------------------------------
|
|
561
|
+
|
|
562
|
+
def baseline(data: Dataset, *, batch_size: int = 10, run_dir=None,
|
|
563
|
+
**_ignored) -> Skill:
|
|
564
|
+
"""The zero-training control: no document, no requests spent."""
|
|
565
|
+
state = NO_FORMAT.new_state({
|
|
566
|
+
"dataset": data.name, "title": data.title,
|
|
567
|
+
"algorithm": "zero-training baseline", "classes": list(data.classes),
|
|
568
|
+
"format": None, "skill_file": "skill.txt",
|
|
569
|
+
"representation": data.representation, "image_size": data.image_size,
|
|
570
|
+
"description": describe(data), "train": 0, "batch_size": batch_size,
|
|
571
|
+
"model": MODEL, "effort": EFFORT,
|
|
572
|
+
})
|
|
573
|
+
if run_dir:
|
|
574
|
+
save_state(state, run_dir)
|
|
575
|
+
return Skill(state, run_dir)
|
|
576
|
+
|
|
577
|
+
|
|
578
|
+
def score(skill: Skill, data: Dataset, *, batch_size: int = 10, run_dir=None,
|
|
579
|
+
session=None) -> Report:
|
|
580
|
+
"""Score a frozen skill on ``data.test`` (or ``data.train`` if absent)."""
|
|
581
|
+
rows = data.test or data.train
|
|
582
|
+
if not rows:
|
|
583
|
+
raise ValueError("dataset has no rows to score")
|
|
584
|
+
state = skill.state
|
|
585
|
+
config = state.get("config", {})
|
|
586
|
+
fmt = format_of(state)
|
|
587
|
+
description = config.get("description") or describe(data)
|
|
588
|
+
representation = config.get("representation", data.representation)
|
|
589
|
+
image_size = int(config.get("image_size", data.image_size))
|
|
590
|
+
cache = (Path(run_dir) / "view_cache") if run_dir else None
|
|
591
|
+
results: list[dict] = []
|
|
592
|
+
requests = 0
|
|
593
|
+
with _borrowed(session) as session:
|
|
594
|
+
for start in range(0, len(rows), batch_size):
|
|
595
|
+
batch = rows[start:start + batch_size]
|
|
596
|
+
predictions, _, _ = predict_batch(
|
|
597
|
+
session, state, batch, description, representation,
|
|
598
|
+
image_size, fmt=fmt, cache=cache)
|
|
599
|
+
requests += 1
|
|
600
|
+
results += [
|
|
601
|
+
{"index": row["index"], "true": str(row["label"]),
|
|
602
|
+
"pred": prediction["pred"],
|
|
603
|
+
"correct": prediction["pred"] == str(row["label"]),
|
|
604
|
+
"rules_used": prediction["rules_used"],
|
|
605
|
+
"evidence": prediction["evidence"]}
|
|
606
|
+
for row, prediction in zip(batch, predictions)
|
|
607
|
+
]
|
|
608
|
+
correct = sum(int(row["correct"]) for row in results)
|
|
609
|
+
print(f" test {len(results):4d}/{len(rows)}: "
|
|
610
|
+
f"{correct}/{len(results)}", flush=True)
|
|
611
|
+
report = Report(correct=sum(int(row["correct"]) for row in results),
|
|
612
|
+
total=len(results), rows=results, requests=requests)
|
|
613
|
+
if run_dir:
|
|
614
|
+
state.setdefault("tests", {})[str(len(state["train"]))] = results
|
|
615
|
+
save_state(state, run_dir)
|
|
616
|
+
return report
|
|
617
|
+
|
|
618
|
+
|
|
619
|
+
# --------------------------------------------------------------------------
|
|
620
|
+
|
|
621
|
+
ALGOS = {"gate": fit_gated, "baseline": baseline}
|
|
622
|
+
|
|
623
|
+
# The prequential learner this file used to carry. Silently remapping it to
|
|
624
|
+
# the gate would run a different experiment under the old name.
|
|
625
|
+
RETIRED_ALGOS = ("markdown",)
|
|
626
|
+
|
|
627
|
+
|
|
628
|
+
def _coerce(data) -> Dataset:
|
|
629
|
+
if isinstance(data, Dataset):
|
|
630
|
+
return data
|
|
631
|
+
if isinstance(data, dict) and "train" in data:
|
|
632
|
+
return Dataset(**data)
|
|
633
|
+
if isinstance(data, (list, tuple)):
|
|
634
|
+
return Dataset(train=list(data))
|
|
635
|
+
raise TypeError(
|
|
636
|
+
"train() expects a Dataset, a list of labeled rows, or a mapping with "
|
|
637
|
+
f"a 'train' key; got {type(data).__name__}")
|
|
638
|
+
|
|
639
|
+
|
|
640
|
+
def _algo(algo, method) -> str:
|
|
641
|
+
chosen = method if method is not None else algo
|
|
642
|
+
if chosen in RETIRED_ALGOS:
|
|
643
|
+
raise ValueError(
|
|
644
|
+
f"the {chosen!r} learning algorithm was removed; it accepted "
|
|
645
|
+
"every validated revision without a held-out check. Use "
|
|
646
|
+
"algo='gate' instead.")
|
|
647
|
+
if chosen not in ALGOS:
|
|
648
|
+
raise ValueError(
|
|
649
|
+
f"unknown algo {chosen!r}; choose one of {', '.join(ALGOS)}")
|
|
650
|
+
return chosen
|
|
651
|
+
|
|
652
|
+
|
|
653
|
+
def train(data, *, algo: str = "gate", format=None, representation=None,
|
|
654
|
+
batch_size: int = 10, max_words=None, max_rules=None,
|
|
655
|
+
run_dir=None, resume: bool = True, evaluate_after: bool = False,
|
|
656
|
+
method=None, format_file=None, **options) -> Skill:
|
|
657
|
+
"""Learn a skill document from labeled data.
|
|
658
|
+
|
|
659
|
+
``algo`` is ``"gate"`` (keep a revision only when it beats the current
|
|
660
|
+
document on protected rows) or ``"baseline"`` (the zero-training control,
|
|
661
|
+
which is shown no document at all and spends no requests).
|
|
662
|
+
|
|
663
|
+
``format`` names a section of ``format.md``; omitting it uses the file's
|
|
664
|
+
own default. The input representation is a property of the data, not of
|
|
665
|
+
the algorithm: pass rows carrying ``text`` for serialized features, or
|
|
666
|
+
``image`` for pixels.
|
|
667
|
+
"""
|
|
668
|
+
chosen = _algo(algo, method)
|
|
669
|
+
prepared = _coerce(data)
|
|
670
|
+
if representation is not None and representation != prepared.representation:
|
|
671
|
+
prepared = Dataset(
|
|
672
|
+
train=prepared.train, test=prepared.test,
|
|
673
|
+
classes=prepared.classes, name=prepared.name,
|
|
674
|
+
title=prepared.title, description=prepared.description,
|
|
675
|
+
representation=representation, image_size=prepared.image_size)
|
|
676
|
+
skill = ALGOS[chosen](
|
|
677
|
+
prepared, format=format, format_file=format_file,
|
|
678
|
+
batch_size=batch_size, max_words=max_words, max_rules=max_rules,
|
|
679
|
+
run_dir=run_dir, resume=resume, **options)
|
|
680
|
+
if evaluate_after and prepared.test:
|
|
681
|
+
skill.report = score(skill, prepared, batch_size=batch_size,
|
|
682
|
+
run_dir=run_dir)
|
|
683
|
+
return skill
|
|
684
|
+
|
|
685
|
+
|
|
686
|
+
def evaluate(skill: Skill, data, *, batch_size: int = 10,
|
|
687
|
+
run_dir=None) -> Report:
|
|
688
|
+
"""Score a frozen skill on ``data.test`` (or ``data.train`` if absent)."""
|
|
689
|
+
report = score(skill, _coerce(data), batch_size=batch_size,
|
|
690
|
+
run_dir=run_dir)
|
|
691
|
+
skill.report = report
|
|
692
|
+
return report
|
|
693
|
+
|
|
694
|
+
def load(path) -> Skill:
|
|
695
|
+
"""Reopen a saved run directory or ``state.json``."""
|
|
696
|
+
return Skill.load(path)
|
|
697
|
+
|
|
698
|
+
|
|
699
|
+
# --------------------------------------------------------------------------
|