#!/usr/bin/env python3 """ prereq_op_binary_split.py Tests whether prereq_op is a factorizable gate: - Instead of a 4-way flat classifier, train two binary classifiers: 1. requires_architect_review (yes/no) 2. requires_manual_approval (yes/no) - Combine deterministically: no/no → standard (0) yes/no → needs_review (1) no/yes → needs_approval (2) yes/yes → full_gates (3) Then compare: - flat 4-way model accuracy - combined binary model accuracy Also measures: false-gate rate (how often either binary fires when neither should). Usage: python3 specialists/scripts/prereq_op_binary_split.py \\ --data-dir specialists/data/combined \\ --runs-dir /mnt/storage/fabricate_runs \\ --results-out specialists/eval/results/prereq_op_factorize.json """ import argparse import json import subprocess import sys from pathlib import Path ROOT = Path(__file__).parent.parent.parent FABRICATE = Path("/home/bill/Documents/WhetstoneAI_Fabricate") VENV_PYTHON = FABRICATE / ".venv" / "bin" / "python3" TRAIN_SCRIPT = ROOT / "specialists" / "scripts" / "train_specialist_pt.py" EVAL_SCRIPT = ROOT / "specialists" / "eval" / "eval_specialist_pt.py" # prereq_op label mapping LABEL_NAMES = ["standard", "needs_review", "needs_approval", "full_gates"] # Axis decomposition: # requires_review: needs_review(1), full_gates(3) → binary label 1 # requires_approval: needs_approval(2), full_gates(3) → binary label 1 REVIEW_LABELS = {1, 3} # 4-way indices that mean "needs review" APPROVAL_LABELS = {2, 3} # 4-way indices that mean "needs approval" def split_to_binary(src_tsv: Path, review_out: Path, approval_out: Path): """Write two binary TSVs from 4-way prereq_op data.""" review_rows = [] approval_rows = [] with open(src_tsv) as f: for line in f: parts = line.strip().split("\t", 2) if len(parts) < 3: continue label_idx = int(parts[0]) text = parts[2] review_label = 1 if label_idx in REVIEW_LABELS else 0 approval_label = 1 if label_idx in APPROVAL_LABELS else 0 review_rows.append(f"{review_label}\t0\t{text}") approval_rows.append(f"{approval_label}\t0\t{text}") review_out.write_text("\n".join(review_rows) + "\n") approval_out.write_text("\n".join(approval_rows) + "\n") print(f" review TSV: {len(review_rows)} rows → {review_out}") print(f" approval TSV: {len(approval_rows)} rows → {approval_out}") # Report class balance r1 = sum(1 for r in review_rows if r.startswith("1")) a1 = sum(1 for r in approval_rows if r.startswith("1")) print(f" review pos={r1} neg={len(review_rows)-r1}") print(f" approval pos={a1} neg={len(approval_rows)-a1}") def train_binary(run_dir: Path, train_tsv: Path, eval_tsv: Path, labels: str, history_out: Path) -> bool: run_dir.mkdir(parents=True, exist_ok=True) for f in run_dir.glob("*"): f.unlink() history_out.unlink(missing_ok=True) cmd = [ str(VENV_PYTHON), str(TRAIN_SCRIPT), "--dataset", str(train_tsv), "--eval-dataset", str(eval_tsv), "--out-dir", str(run_dir), "--labels", labels, "--lr", "0.001", "--weight-decay", "0.01", "--max-steps", "10000", "--hidden-dim", "256", "--layers", "2", "--heads", "4", "--batch-size", "256", "--grok-loss-threshold", "0.05", "--grok-acc-jump", "10.0", "--stop-after-grokking-blocks", "4", "--reset", "--history-out", str(history_out), ] r = subprocess.run(cmd) return r.returncode == 0 def eval_binary(run_dir: Path, eval_tsv: Path, labels: str) -> dict | None: ck = run_dir / "checkpoint.pt" if not ck.exists(): return None r = subprocess.run( [str(VENV_PYTHON), str(EVAL_SCRIPT), "--checkpoint", str(ck), "--dataset", str(eval_tsv), "--labels", labels, "--quiet"], capture_output=True, text=True, ) if r.returncode != 0: return None try: return json.loads(r.stdout.strip().splitlines()[-1] if r.stdout.strip() else "{}") except json.JSONDecodeError: return None def evaluate_combined(eval_tsv: Path, review_run: Path, approval_run: Path) -> dict: """Load both PyTorch binary models, combine predictions, compare to 4-way ground truth.""" import numpy as np import torch import torch.nn as nn # Inline model class (mirrors train_specialist_pt.py) def _tok(text): return text.lower().split() class _Model(nn.Module): def __init__(self, cfg): super().__init__() D, L, H, V, seq = cfg["hidden_dim"], cfg["layers"], cfg["heads"], cfg["vocab_size"], cfg["seq_len"] self.tok_emb = nn.Embedding(V, D, padding_idx=0) self.pos_emb = nn.Embedding(seq, D) enc = nn.TransformerEncoderLayer(d_model=D, nhead=H, dim_feedforward=D*4, dropout=0.0, batch_first=True) self.encoder = nn.TransformerEncoder(enc, num_layers=L) self.head = nn.Linear(D, cfg["n_labels"]) self.seq_len = seq def forward(self, x): pad_mask = (x == 0) pos = torch.arange(x.size(1), device=x.device).unsqueeze(0) h = self.tok_emb(x) + self.pos_emb(pos) h = self.encoder(h, src_key_padding_mask=pad_mask) lengths = (~pad_mask).float().sum(1, keepdim=True).clamp(min=1) h = (h * (~pad_mask).unsqueeze(-1).float()).sum(1) / lengths return self.head(h) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") def load_pt_model(run_dir): ckpt = torch.load(run_dir / "checkpoint.pt", map_location=device) m = _Model(ckpt["config"]).to(device) m.load_state_dict(ckpt["model"]) m.eval() vocab = ckpt.get("vocab_obj") if vocab is None: # reconstruct minimal vocab object class _V: UNK = 1 def __init__(self, d): self.str_to_id = d; self.seq_len = ckpt["config"]["seq_len"] def encode(self, text, seq_len): ids = [self.str_to_id.get(t, self.UNK) for t in _tok(text)] ids = ids[:seq_len]; ids += [0] * (seq_len - len(ids)); return ids vocab = _V(ckpt["vocab"]) return m, vocab, ckpt["config"]["seq_len"] def predict(model, vocab, seq_len, texts): preds = [] with torch.no_grad(): for start in range(0, len(texts), 256): batch = texts[start:start+256] ids = [vocab.encode(t, seq_len) for t in batch] x = torch.tensor(ids, dtype=torch.long, device=device) logits = model(x).cpu().numpy() preds.append(np.argmax(logits, axis=1)) return np.concatenate(preds) # Load data labels_gt, texts = [], [] with open(eval_tsv) as f: for line in f: parts = line.strip().split("\t", 2) if len(parts) < 3: continue labels_gt.append(int(parts[0])) texts.append(parts[2]) labels_gt = np.array(labels_gt) n = len(labels_gt) review_model, review_vocab, review_seq = load_pt_model(review_run) approval_model, approval_vocab, approval_seq = load_pt_model(approval_run) # Binary models: label 0 = no, label 1 = yes review_pred = predict(review_model, review_vocab, review_seq, texts) approval_pred = predict(approval_model, approval_vocab, approval_seq, texts) # Combine: review×approval → 4-way index combined_pred = np.where( (review_pred == 0) & (approval_pred == 0), 0, # standard np.where( (review_pred == 1) & (approval_pred == 0), 1, # needs_review np.where( (review_pred == 0) & (approval_pred == 1), 2, # needs_approval 3 # full_gates ) ) ) correct = (combined_pred == labels_gt) overall = float(correct.mean() * 100) per_class = {} for i, name in enumerate(LABEL_NAMES): mask = labels_gt == i if mask.sum() == 0: per_class[name] = {"acc": None, "n": 0} else: per_class[name] = {"acc": float(correct[mask].mean() * 100), "n": int(mask.sum())} from collections import Counter errors = Counter() for gt, pred in zip(labels_gt.tolist(), combined_pred.tolist()): if gt != pred: errors[(LABEL_NAMES[gt], LABEL_NAMES[pred])] += 1 return { "method": "binary_factorized", "overall_accuracy": round(overall, 1), "n_eval": n, "per_class": per_class, "top_errors": [{"gt": k[0], "pred": k[1], "count": v} for k, v in errors.most_common(6)], } def main(): ap = argparse.ArgumentParser() ap.add_argument("--data-dir", default="specialists/data/combined") ap.add_argument("--runs-dir", default="/mnt/storage/fabricate_runs") ap.add_argument("--results-out", default="specialists/eval/results/prereq_op_factorize.json") ap.add_argument("--skip-train", action="store_true", help="Skip training, just eval existing checkpoints") args = ap.parse_args() data_dir = Path(args.data_dir) runs_dir = Path(args.runs_dir) results = Path(args.results_out) results.parent.mkdir(parents=True, exist_ok=True) # Intermediate data paths tmp = Path("specialists/data/generated") tmp.mkdir(parents=True, exist_ok=True) review_train = tmp / "prereq_review_train.tsv" review_eval = tmp / "prereq_review_eval.tsv" approval_train = tmp / "prereq_approval_train.tsv" approval_eval = tmp / "prereq_approval_eval.tsv" # --- Split data --- print("=== Splitting 4-way data into two binary axes ===") split_to_binary(data_dir / "prereq_op_train.tsv", review_train, approval_train) split_to_binary(data_dir / "prereq_op_eval.tsv", review_eval, approval_eval) review_run = runs_dir / "whetstone_prereq_review_binary" approval_run = runs_dir / "whetstone_prereq_approval_binary" review_hist = ROOT / "specialists" / "runs" / "prereq_review_history.json" approval_hist = ROOT / "specialists" / "runs" / "prereq_approval_history.json" if not args.skip_train: # --- Train binary models --- print("\n=== Training: requires_architect_review (binary) ===") train_binary(review_run, review_train, review_eval, "no,yes", review_hist) print("\n=== Training: requires_manual_approval (binary) ===") train_binary(approval_run, approval_train, approval_eval, "no,yes", approval_hist) # --- Eval individual binary models --- print("\n=== Evaluating binary models independently ===") review_result = eval_binary(review_run, review_eval, "no,yes") approval_result = eval_binary(approval_run, approval_eval, "no,yes") if review_result: print(f" review binary: {review_result.get('overall_accuracy')}%") if approval_result: print(f" approval binary: {approval_result.get('overall_accuracy')}%") # --- Evaluate combined predictions on 4-way task --- print("\n=== Evaluating combined binary → 4-way accuracy ===") combined = evaluate_combined( data_dir / "prereq_op_eval.tsv", review_run, approval_run ) print(f" Combined 4-way: {combined['overall_accuracy']}%") print(" Per class:") for name, stat in combined["per_class"].items(): if stat["acc"] is not None: print(f" {name:20s}: {stat['acc']:.1f}% (n={stat['n']})") output = { "gate": "prereq_op", "factorized": { "requires_architect_review": review_result, "requires_manual_approval": approval_result, "combined_4way": combined, }, "note": "Compare combined_4way.overall_accuracy against flat 4-way model accuracy", } results.write_text(json.dumps(output, indent=2)) print(f"\nResults → {results}") if __name__ == "__main__": main()