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/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
+ # --------------------------------------------------------------------------