specialists: switch training to PyTorch, add capacity sweep infrastructure

- train_specialist_pt.py: pure PyTorch trainer replacing Fabricate binary.
  No SQLite/WAL — saves checkpoint.pt + history.json only. Fabricate was
  causing D-state IO blocking on the HDD due to continuous WAL writes.
- eval_specialist_pt.py: PyTorch eval with confusion matrix, threshold
  analysis, guardrail assessment, and deployment verdict.
- capacity_sweep.py: updated to call PyTorch scripts; TIERS now include
  heads field (4/6/8) since PyTorch supports multi-head unlike Fabricate.
- prereq_op_binary_split.py: evaluate_combined rewritten in PyTorch.
- Performance fix in train_specialist_pt.py: load training data onto GPU
  once and sample via torch.randint instead of DataLoader (eliminates
  Python/CPU overhead on tiny datasets). Fix applied, not yet timed.

Next session: time the fix, then run sweep_all_gates.sh --pilot-only.
See HANDOFF-2026-03-31.md for full context.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
Bill
2026-03-31 22:55:51 -06:00
parent 410bb2d78c
commit 3fc9deac5b
61 changed files with 9168 additions and 0 deletions

View File

@@ -0,0 +1,332 @@
#!/usr/bin/env python3
"""
gen_from_projects.py
Runs all project benchmark specs through the whetstone_mcp pipeline
(architect_intake → generate_taskitems) and extracts labeled training
examples for each specialist from the real pipeline output.
Why this matters:
- Training data from sprint plans uses step description text.
- Specialists run at inference time on taskitem TITLES + INTENT.
- These are similar but not identical formats. Running real specs through
the pipeline produces training examples in the exact inference-time format.
- The pipeline output is the "correct answer" — workerType from inferWorkerType,
prerequisiteOps from inferPrerequisiteOps, verificationType from the contract.
Input files:
datasets/project_benchmarks/common_projects_100.jsonl (100 projects)
datasets/project_benchmarks/challenging_*.jsonl (25 projects)
Output TSVs (appended to existing data or written fresh):
specialists/data/generated/prereq_op_projects_train.tsv
specialists/data/generated/prereq_op_projects_eval.tsv
specialists/data/generated/worker_type_projects_train.tsv
specialists/data/generated/worker_type_projects_eval.tsv
specialists/data/generated/verification_type_projects_train.tsv
specialists/data/generated/verification_type_projects_eval.tsv
Text format for all specialists:
"<title> | <intent[:120]>"
Using pipe separator so title and intent are both visible but the
model sees them as one continuous text sequence.
Usage:
python3 specialists/scripts/gen_from_projects.py \\
--whetstone-bin editor/build-native/whetstone_mcp \\
--workspace .
"""
import argparse
import json
import random
import subprocess
import sys
from collections import Counter
from pathlib import Path
# ---------------------------------------------------------------------------
# Label maps
# ---------------------------------------------------------------------------
# Text-based prereq classifier — same tokens as gen_prereq_op_data.py so that
# project-spec task text is labeled consistently with the sprint-plan data.
# Using text rather than the C++ heuristic's prerequisiteOps because the heuristic
# never emits "manual-approval" (circular labels = all "standard").
_PREREQ_REVIEW_TOKENS = {
"introduce", "phase", "framework", "architecture", "design", "refactor",
"planning", "taxonomy", "grammar", "interface", "api surface", "subsystem",
"extension point", "abstraction", "decompos", "scaffold", "polyglot",
"cross-language", "cross-project", "multi-language", "orchestrat",
"migration plan", "phase plan", "protocol design", "new layer",
"foundation", "bootstrap", "initial", "phase 1", "phase 2", "phase 3",
"sprint plan", "roadmap",
}
_PREREQ_APPROVAL_TOKENS = {
"release", "deploy", "publish", "migrate", "security", "compliance",
"upgrade", "breaking change", "production", "rollout", "finalize",
"sign off", "certification", "approval", "cutover", "go-live",
"data migration", "schema migration", "breaking", "deprecat",
"audit", "regulatory", "gdpr", "pii", "secret", "credential", "vault",
"access control", "rbac", "acl", "permission", "fraud", "encryption",
}
def prereq_ops_to_class(ops: list[str]) -> str:
"""Legacy: classify by actual ops list (unused — text_to_prereq_class is preferred)."""
s = set(ops)
has_review = "architect-review" in s
has_approval = "manual-approval" in s
if has_review and has_approval:
return "full_gates"
if has_review:
return "needs_review"
if has_approval:
return "needs_approval"
return "standard"
def text_to_prereq_class(text: str) -> str:
"""Label prereq class from task title+intent text (same signal the specialist sees)."""
t = text.lower()
has_review = any(tok in t for tok in _PREREQ_REVIEW_TOKENS)
has_approval = any(tok in t for tok in _PREREQ_APPROVAL_TOKENS)
if has_review and has_approval:
return "full_gates"
if has_review:
return "needs_review"
if has_approval:
return "needs_approval"
return "standard"
PREREQ_LABEL_INDEX = {"standard": 0, "needs_review": 1, "needs_approval": 2, "full_gates": 3}
WORKER_LABEL_INDEX = {"implementer": 0, "reviewer": 1, "architect": 2, "qa": 3}
VERIF_LABEL_INDEX = {"unit": 0, "integration": 1, "schema": 2, "smoke": 3, "docs": 4}
# ---------------------------------------------------------------------------
# MCP pipeline helpers
# ---------------------------------------------------------------------------
def mcp_call(binary: str, workspace: str, tool: str, args: dict) -> dict:
payload = json.dumps({
"jsonrpc": "2.0", "id": 1,
"method": "tools/call",
"params": {"name": tool, "arguments": args}
})
try:
r = subprocess.run(
[binary, "--workspace", workspace, "--language", "cpp"],
input=payload, capture_output=True, text=True, timeout=20
)
content = (json.loads(r.stdout)
.get("result", {})
.get("content", [{}])[0]
.get("text", ""))
return json.loads(content) if content else {}
except Exception as e:
print(f" [WARN] mcp_call {tool}: {e}", file=sys.stderr)
return {}
def run_pipeline(binary: str, workspace: str, spec: str) -> list[dict]:
"""Run architect_intake → generate_taskitems for one spec. Return task list."""
intake = mcp_call(binary, workspace, "whetstone_architect_intake", {"markdown": spec})
reqs = intake.get("normalizedRequirements", [])
if not reqs:
return []
result = mcp_call(binary, workspace, "whetstone_generate_taskitems", {
"normalizedRequirements": reqs,
"conflicts": intake.get("conflicts", [])
})
return result.get("tasks", [])
def task_text(task: dict) -> str:
"""Combine title + intent for specialist input."""
title = task.get("title", "").strip()
intent = task.get("intent", "").strip()[:120]
if intent and intent != title:
return f"{title} | {intent}"
return title
# ---------------------------------------------------------------------------
# Load project specs
# ---------------------------------------------------------------------------
def load_projects(dataset_dir: Path) -> list[dict]:
projects = []
# Common 100
common = dataset_dir / "common_projects_100.jsonl"
if common.exists():
with open(common) as f:
for line in f:
projects.append(json.loads(line))
# All challenging sets
for path in sorted(dataset_dir.glob("challenging_*.jsonl")):
with open(path) as f:
for line in f:
p = json.loads(line)
if p.get("spec"):
projects.append(p)
print(f"Loaded {len(projects)} project specs")
return projects
def spec_markdown(project: dict) -> str:
"""Wrap a project spec into the sections architect_intake expects."""
return (
f"## Goals\n{project['spec']}\n\n"
f"## Constraints\n- Production quality\n- Well-tested\n\n"
f"## Acceptance Criteria\n- All core features implemented\n"
f"- Tests passing\n"
)
# ---------------------------------------------------------------------------
# Balance + write TSV
# ---------------------------------------------------------------------------
def balance_and_split(examples: list[tuple[str, str]],
label_index: dict[str, int],
label_order: list[str],
seed: int = 42,
eval_frac: float = 0.15,
min_per_class: int = 20) -> tuple[list, list]:
"""Return (train_rows, eval_rows) as list of (label_str, text) pairs."""
by_class: dict[str, list[str]] = {}
for lbl, text in examples:
by_class.setdefault(lbl, []).append(text)
counts = {lbl: len(items) for lbl, items in by_class.items()}
print(f" Raw class distribution: {counts}")
if not counts:
return [], []
sorted_vals = sorted(counts.values())
median = sorted_vals[len(sorted_vals) // 2]
target = max(median, min_per_class)
print(f" Balancing to {target} per class")
rng = random.Random(seed)
balanced: list[tuple[str, str]] = []
for lbl in label_order:
items = by_class.get(lbl, [])
if not items:
print(f" WARNING: no examples for '{lbl}'", file=sys.stderr)
continue
rng.shuffle(items)
chosen = items[:target]
while len(chosen) < target:
chosen += items
balanced.extend((lbl, t) for t in chosen[:target])
rng.shuffle(balanced)
n_eval = max(1, int(len(balanced) * eval_frac))
return balanced[n_eval:], balanced[:n_eval]
def write_tsv(path: Path, rows: list[tuple[str, str]], label_index: dict[str, int]) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with open(path, "w") as f:
for lbl, text in rows:
idx = label_index[lbl]
f.write(f"{idx}\t0\t{text}\n")
print(f" Wrote {len(rows)} rows → {path}")
# ---------------------------------------------------------------------------
# Main
# ---------------------------------------------------------------------------
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--whetstone-bin", default="editor/build-native/whetstone_mcp")
ap.add_argument("--workspace", default=".")
ap.add_argument("--dataset-dir", default="datasets/project_benchmarks")
ap.add_argument("--out-dir", default="specialists/data/generated")
ap.add_argument("--eval-frac", type=float, default=0.15)
ap.add_argument("--seed", type=int, default=42)
ap.add_argument("--dry-run", action="store_true",
help="Print extracted examples without writing files")
args = ap.parse_args()
binary = str(Path(args.workspace) / args.whetstone_bin)
workspace = str(Path(args.workspace).resolve())
out_dir = Path(args.out_dir)
projects = load_projects(Path(args.workspace) / args.dataset_dir)
prereq_examples: list[tuple[str, str]] = []
worker_examples: list[tuple[str, str]] = []
verif_examples: list[tuple[str, str]] = []
for i, project in enumerate(projects):
print(f"[{i+1:3}/{len(projects)}] {project.get('project','?')} ({project.get('category','?')})")
spec_md = spec_markdown(project)
tasks = run_pipeline(binary, workspace, spec_md)
if not tasks:
print(" -> no tasks generated, skipping")
continue
for task in tasks:
text = task_text(task)
if not text or len(text) < 8:
continue
# prereq_op_selector — use text-based classifier (same signal as sprint-plan
# data) to avoid circular heuristic labels (C++ heuristic never emits
# manual-approval, so ops-based labeling produces 100% "standard").
prereq_examples.append((text_to_prereq_class(text), text))
# worker_type
wt = task.get("workerType", "")
if wt in WORKER_LABEL_INDEX:
worker_examples.append((wt, text))
# verification_type
vt = task.get("executionContract", {}).get("verificationType", "")
if vt in VERIF_LABEL_INDEX:
verif_examples.append((vt, text))
print(f" -> {len(tasks)} tasks, prereq={len(prereq_examples)} "
f"worker={len(worker_examples)} verif={len(verif_examples)}")
if args.dry_run:
print("\n=== SAMPLE EXAMPLES ===")
for name, examples in [("prereq", prereq_examples[:5]),
("worker", worker_examples[:5]),
("verif", verif_examples[:5])]:
print(f"\n{name}:")
for lbl, text in examples:
print(f" [{lbl}] {text[:80]}")
return
PREREQ_ORDER = ["standard", "needs_review", "needs_approval", "full_gates"]
WORKER_ORDER = ["implementer", "reviewer", "architect", "qa"]
VERIF_ORDER = ["unit", "integration", "schema", "smoke", "docs"]
print("\n=== prereq_op from projects ===")
train, eval_ = balance_and_split(prereq_examples, PREREQ_LABEL_INDEX, PREREQ_ORDER,
seed=args.seed, eval_frac=args.eval_frac)
write_tsv(out_dir / "prereq_op_projects_train.tsv", train, PREREQ_LABEL_INDEX)
write_tsv(out_dir / "prereq_op_projects_eval.tsv", eval_, PREREQ_LABEL_INDEX)
print("\n=== worker_type from projects ===")
train, eval_ = balance_and_split(worker_examples, WORKER_LABEL_INDEX, WORKER_ORDER,
seed=args.seed, eval_frac=args.eval_frac)
write_tsv(out_dir / "worker_type_projects_train.tsv", train, WORKER_LABEL_INDEX)
write_tsv(out_dir / "worker_type_projects_eval.tsv", eval_, WORKER_LABEL_INDEX)
print("\n=== verification_type from projects ===")
train, eval_ = balance_and_split(verif_examples, VERIF_LABEL_INDEX, VERIF_ORDER,
seed=args.seed, eval_frac=args.eval_frac)
write_tsv(out_dir / "verification_type_projects_train.tsv", train, VERIF_LABEL_INDEX)
write_tsv(out_dir / "verification_type_projects_eval.tsv", eval_, VERIF_LABEL_INDEX)
print("\nDone. Combine with sprint-plan TSVs using:")
print(" cat specialists/data/generated/prereq_op_train.tsv \\")
print(" specialists/data/generated/prereq_op_projects_train.tsv \\")
print(" > specialists/data/combined/prereq_op_train.tsv")
if __name__ == "__main__":
main()