#!/usr/bin/env python3 """ Slice 17: Nonlinear MLP Probe over Frozen e5 Embeddings ======================================================== Answer: is action-vs-non-action information present in the existing 384-d e5 vector but not linearly separable? Architecture under test: e5[384] → Linear(384→H) → GELU → Linear(H→1) (binary action gate) e5[384] → Linear(384→H) → GELU → Linear(H→6) (six-way MLP) Everything frozen from slice 16. No new examples, no changed encoder. """ import json import re import sys import warnings from collections import defaultdict from pathlib import Path import numpy as np from sklearn.exceptions import ConvergenceWarning from sklearn.linear_model import LogisticRegression from sklearn.metrics import ( accuracy_score, average_precision_score, confusion_matrix, f1_score, precision_recall_fscore_support, roc_auc_score, ) from sklearn.neural_network import MLPClassifier from sklearn.preprocessing import LabelEncoder, StandardScaler warnings.filterwarnings("ignore", category=ConvergenceWarning) warnings.filterwarnings("ignore", category=UserWarning) EMBEDDING_PATH = "/tmp/mvn-experiment/embeddings.json" ROUTES = ["action", "conversation", "knowledge", "memory_write", "system", "uncertain"] # Hidden sizes for the MLP HIDDEN_SIZES = [8, 16, 32, 64] # Weight decay values (sklearn MLPClassifier alpha parameter) WEIGHT_DECAYS = [0.0, 1e-4, 1e-3, 1e-2] # Action threshold sweep for safety curve ACTION_THRESHOLDS = np.arange(0.30, 0.96, 0.025).tolist() # ─── Data Loading ─────────────────────────────────────────────────────────── def load_data(): with open(EMBEDDING_PATH) as f: data = json.load(f) return data["meta"], data["examples"] def filter_dev_pool(examples): return [e for e in examples if e["dev_pool"]] def filter_residual(examples): return [e for e in examples if not e["fast_path_resolved"]] def extract_Xy(examples): X = np.array([e["embedding"] for e in examples]) y = np.array([e["route"] for e in examples]) return X, y def get_fold_groups(examples): return np.array([e["cv_fold"] for e in examples]) # ─── Helpers ──────────────────────────────────────────────────────────────── def pct(v, d=1): return f"{100*v:.{d}f}%" def ff(v, d=3): return f"{v:.{d}f}" def param_count_binary(H): """384*H + H + H*1 + 1""" return 384 * H + H + H + 1 def param_count_6way(H): """384*H + H + H*6 + 6""" return 384 * H + H + H * 6 + 6 # ─── Binary MLP Action Gate ──────────────────────────────────────────────── def run_binary_mlp_cv(X, y, fold_ids, H, alpha, examples_meta, max_iter=800): """Binary action vs not-action MLP with grouped CV.""" y_binary = np.array([1 if t == "action" else 0 for t in y]) unique_folds = sorted(set(fold_ids)) oof_rows = [] fold_metrics = [] for test_fold in unique_folds: train_mask = fold_ids != test_fold test_mask = fold_ids == test_fold X_train, y_train = X[train_mask], y_binary[train_mask] X_test, y_test = X[test_mask], y_binary[test_mask] model = MLPClassifier( hidden_layer_sizes=(H,), activation="relu", solver="adam", alpha=alpha, max_iter=max_iter, random_state=42, early_stopping=True, validation_fraction=0.15, n_iter_no_change=10, ) model.fit(X_train, y_train) y_proba = model.predict_proba(X_test) # Class 1 = action action_proba = y_proba[:, 1] if y_proba.shape[1] > 1 else y_proba[:, 0] y_pred = model.predict(X_test) acc = accuracy_score(y_test, y_pred) false_pos = int(sum(1 for t, p in zip(y_test, y_pred) if t == 0 and p == 1)) false_neg = int(sum(1 for t, p in zip(y_test, y_pred) if t == 1 and p == 0)) true_pos = int(sum(1 for t, p in zip(y_test, y_pred) if t == 1 and p == 1)) true_neg = int(sum(1 for t, p in zip(y_test, y_pred) if t == 0 and p == 0)) if len(np.unique(y_test)) > 1: roc = roc_auc_score(y_test, action_proba) pr_auc = average_precision_score(y_test, action_proba) else: roc, pr_auc = 0.0, 0.0 prec = true_pos / max(true_pos + false_pos, 1) rec = true_pos / max(true_pos + false_neg, 1) fold_metrics.append({ "fold": test_fold, "accuracy": acc, "roc_auc": roc, "pr_auc": pr_auc, "action_precision": prec, "action_recall": rec, "false_pos": false_pos, "false_neg": false_neg, "test_size": int(len(X_test)), "n_iters": model.n_iter_, }) test_indices = np.where(test_mask)[0] for i in range(len(y_test)): meta = examples_meta[test_indices[i]] oof_rows.append({ "source_id": meta["source_id"], "fold": test_fold, "true": "action" if y_test[i] == 1 else "not_action", "predicted": "action" if y_pred[i] == 1 else "not_action", "action_proba": float(action_proba[i]), "correct": y_test[i] == y_pred[i], "text": meta["text"], "tags": meta.get("tags", []), "split_group": meta.get("split_group", ""), "true_route": meta["route"], }) total_fp = sum(m["false_pos"] for m in fold_metrics) total_fn = sum(m["false_neg"] for m in fold_metrics) total_n = sum(m["test_size"] for m in fold_metrics) mean_roc = np.mean([m["roc_auc"] for m in fold_metrics]) mean_pr = np.mean([m["pr_auc"] for m in fold_metrics]) mean_prec = np.mean([m["action_precision"] for m in fold_metrics]) mean_rec = np.mean([m["action_recall"] for m in fold_metrics]) return { "mean_roc_auc": mean_roc, "mean_pr_auc": mean_pr, "mean_precision": mean_prec, "mean_recall": mean_rec, "total_fp": total_fp, "total_fn": total_fn, "total_n": total_n, "fa_rate": total_fp / max(total_n, 1), "fold_metrics": fold_metrics, "oof_rows": oof_rows, "n_iters": np.mean([m["n_iters"] for m in fold_metrics]), } def run_binary_mlp_grid(X, y, fold_ids, examples_meta): """Search hidden_size × weight_decay, select by PR-AUC.""" results = {} # H=8 is degenerate (R≈0); H=16 far behind. Carry H=32 and H=64. for H in [16, 32, 64]: for wd in WEIGHT_DECAYS: key = (H, wd) print(f" Binary MLP H={H} wd={wd} ...", end=" ", flush=True) r = run_binary_mlp_cv(X, y, fold_ids, H, wd, examples_meta, max_iter=600) results[key] = r print(f"PR-AUC={ff(r['mean_pr_auc'])} ROC={ff(r['mean_roc_auc'])} " f"P={ff(r['mean_precision'])} R={ff(r['mean_recall'])} FP={r['total_fp']}") best_key = max(results, key=lambda k: results[k]["mean_pr_auc"]) return best_key, results # ─── 6-Way MLP ───────────────────────────────────────────────────────────── def run_sixway_mlp_cv(X, y, fold_ids, H, alpha, examples_meta, max_iter=800): """6-way MLP with grouped CV.""" le = LabelEncoder() y_enc = le.fit_transform(y) classes = le.classes_ unique_folds = sorted(set(fold_ids)) oof_rows = [] fold_metrics = [] for test_fold in unique_folds: train_mask = fold_ids != test_fold test_mask = fold_ids == test_fold X_train, y_train = X[train_mask], y_enc[train_mask] X_test, y_test = X[test_mask], y_enc[test_mask] model = MLPClassifier( hidden_layer_sizes=(H,), activation="relu", solver="adam", alpha=alpha, max_iter=max_iter, random_state=42, early_stopping=True, validation_fraction=0.15, n_iter_no_change=10, ) model.fit(X_train, y_train) y_pred_enc = model.predict(X_test) y_proba = model.predict_proba(X_test) y_pred = le.inverse_transform(y_pred_enc) y_true = le.inverse_transform(y_test) acc = accuracy_score(y_true, y_pred) macro_f1 = f1_score(y_true, y_pred, average="macro", zero_division=0) false_action = sum(1 for t, p in zip(y_true, y_pred) if t != "action" and p == "action") action_tp = sum(1 for t, p in zip(y_true, y_pred) if t == "action" and p == "action") action_fp = sum(1 for t, p in zip(y_true, y_pred) if t != "action" and p == "action") action_fn = sum(1 for t, p in zip(y_true, y_pred) if t == "action" and p != "action") action_prec = action_tp / max(action_tp + action_fp, 1) action_rec = action_tp / max(action_tp + action_fn, 1) fold_metrics.append({ "fold": test_fold, "accuracy": acc, "macro_f1": macro_f1, "action_precision": action_prec, "action_recall": action_rec, "false_action": false_action, "test_size": int(len(X_test)), "n_iters": model.n_iter_, }) test_indices = np.where(test_mask)[0] for i in range(len(y_true)): meta = examples_meta[test_indices[i]] proba_dict = {cls: float(y_proba[i][j]) for j, cls in enumerate(classes)} oof_rows.append({ "source_id": meta["source_id"], "fold": test_fold, "true": y_true[i], "predicted": y_pred[i], "correct": y_true[i] == y_pred[i], "proba": proba_dict, "text": meta["text"], "tags": meta.get("tags", []), "split_group": meta.get("split_group", ""), }) total_fa = sum(m["false_action"] for m in fold_metrics) total_n = sum(m["test_size"] for m in fold_metrics) return { "accuracy": np.mean([m["accuracy"] for m in fold_metrics]), "macro_f1": np.mean([m["macro_f1"] for m in fold_metrics]), "action_precision": np.mean([m["action_precision"] for m in fold_metrics]), "action_recall": np.mean([m["action_recall"] for m in fold_metrics]), "false_action": total_fa, "fa_rate": total_fa / max(total_n, 1), "fold_metrics": fold_metrics, "oof_rows": oof_rows, "n_iters": np.mean([m["n_iters"] for m in fold_metrics]), } # ─── Threshold Curve ─────────────────────────────────────────────────────── def compute_action_threshold_curve(oof_rows, thresholds): """Safety operating curve: precision/recall/FA at each threshold.""" action_probas = np.array([r["action_proba"] for r in oof_rows]) y_true_binary = np.array([1 if r["true"] == "action" else 0 for r in oof_rows]) results = [] for thr in thresholds: y_pred_binary = (action_probas >= thr).astype(int) tp = int(sum(1 for t, p in zip(y_true_binary, y_pred_binary) if t == 1 and p == 1)) fp = int(sum(1 for t, p in zip(y_true_binary, y_pred_binary) if t == 0 and p == 1)) fn = int(sum(1 for t, p in zip(y_true_binary, y_pred_binary) if t == 1 and p == 0)) prec = tp / max(tp + fp, 1) rec = tp / max(tp + fn, 1) fa_count = fp fa_rate = fp / max(len(y_true_binary), 1) coverage = (tp + fp + int(sum(1 for t, p in zip(y_true_binary, y_pred_binary) if t == 0 and p == 0))) / max(len(y_true_binary), 1) results.append({ "threshold": round(thr, 3), "action_precision": round(prec, 4), "action_recall": round(rec, 4), "false_action_count": fa_count, "false_action_rate": round(fa_rate, 4), "coverage": round(coverage, 4), }) return results # ─── Capability-Question Diagnostic ──────────────────────────────────────── def capability_question_diagnostic(oof_rows): """Report capability_question → false action rate and action recall for positive modal requests.""" cap_q_rows = [r for r in oof_rows if "capability_question" in r.get("tags", [])] # Positive modal action: route=action, not in capability_question tag modal_action_rows = [r for r in oof_rows if r["true_route"] == "action" and "capability_question" not in r.get("tags", []) and r["true"] == "action"] cap_q_fa = sum(1 for r in cap_q_rows if r["predicted"] == "action") cap_q_total = len(cap_q_rows) cap_q_fa_rate = cap_q_fa / max(cap_q_total, 1) modal_recall = sum(1 for r in modal_action_rows if r["predicted"] == "action") / max(len(modal_action_rows), 1) # Specific pairs: capability question vs its action counterpart pairs = [] # Group by split_group to find paired examples by_group = defaultdict(list) for r in oof_rows: by_group[r["split_group"]].append(r) # The capability_question examples and their contrastive action pairs # share split_groups like kq-cap-ha-* vs ha-light-off-* # We look at specific split_group prefixes cap_groups = [sg for sg in by_group if sg.startswith("kq-cap-")] for cg in cap_groups: cap_rows = by_group[cg] for cr in cap_rows: if "capability_question" in cr.get("tags", []): pairs.append({ "text": cr["text"], "route": cr["true_route"], "predicted": cr["predicted"], "action_proba": cr.get("action_proba", 0), "is_false_action": cr["predicted"] == "action" and cr["true_route"] != "action", }) return { "cap_q_total": cap_q_total, "cap_q_false_actions": cap_q_fa, "cap_q_fa_rate": cap_q_fa_rate, "modal_action_total": len(modal_action_rows), "modal_action_recall": modal_recall, "pairs": pairs, } # ─── Voice Stress ────────────────────────────────────────────────────────── def apply_voice_stress(text): t = re.sub(r'[?.!,;:]+$', '', text.strip()) t = re.sub(r'[^\w\s]', '', t) t = t.lower() t = re.sub(r'\s+', ' ', t).strip() return t def voice_stress_eval(examples, dev_examples, best_binary_model_factory): """Evaluate best fold models against punctuation-stripped stress set.""" dev_source_ids = {e["source_id"] for e in dev_examples} stress_pairs = [] for e in dev_examples: original = e["text"] stressed = apply_voice_stress(original) if stressed != original: stress_pairs.append({ "source_id": e["source_id"], "original": original, "stressed": stressed, "route": e["route"], "tags": e.get("tags", []), }) # Count categories q_with_q = [p for p in stress_pairs if p["route"] in ("knowledge", "uncertain") and "?" in p["original"]] q_without_q = [p for p in stress_pairs if p["route"] in ("knowledge", "uncertain") and "?" not in p["original"]] modal_action = [p for p in stress_pairs if p["route"] == "action"] cap_q = [p for p in stress_pairs if "capability_question" in p.get("tags", [])] return { "total_stress_pairs": len(stress_pairs), "q_with_question_mark": len(q_with_q), "q_without_question_mark": len(q_without_q), "modal_action_pairs": len(modal_action), "cap_q_pairs": len(cap_q), "pairs": stress_pairs[:20], } # ─── Linear Baselines ────────────────────────────────────────────────────── def run_linear_baseline(X, y, fold_ids, examples_meta): """Run the six-way linear baseline for comparison.""" unique_folds = sorted(set(fold_ids)) oof_rows = [] for test_fold in unique_folds: train_mask = fold_ids != test_fold test_mask = fold_ids == test_fold X_train, y_train = X[train_mask], y[train_mask] X_test, y_test = X[test_mask], y[test_mask] model = LogisticRegression(C=10.0, max_iter=2000, solver="lbfgs", random_state=42) model.fit(X_train, y_train) y_pred = model.predict(X_test) y_proba = model.predict_proba(X_test) classes = model.classes_ test_indices = np.where(test_mask)[0] for i, (true, pred) in enumerate(zip(y_test, y_pred)): meta = examples_meta[test_indices[i]] proba_dict = {cls: float(y_proba[i][j]) for j, cls in enumerate(classes)} oof_rows.append({ "source_id": meta["source_id"], "fold": test_fold, "true": true, "predicted": pred, "correct": true == pred, "proba": proba_dict, "text": meta["text"], "tags": meta.get("tags", []), "true_route": meta["route"], }) return oof_rows def run_binary_linear_baseline(X, y, fold_ids, examples_meta): """Run the binary linear action probe for comparison.""" y_binary = np.array(["action" if t == "action" else "not_action" for t in y]) unique_folds = sorted(set(fold_ids)) oof_rows = [] for test_fold in unique_folds: train_mask = fold_ids != test_fold test_mask = fold_ids == test_fold X_train, y_train = X[train_mask], y_binary[train_mask] X_test, y_test = X[test_mask], y_binary[test_mask] model = LogisticRegression(C=1.0, max_iter=2000, solver="lbfgs", random_state=42) model.fit(X_train, y_train) y_pred = model.predict(X_test) y_proba = model.predict_proba(X_test) classes = model.classes_ action_idx = list(classes).index("action") action_proba = y_proba[:, action_idx] test_indices = np.where(test_mask)[0] for i in range(len(y_test)): meta = examples_meta[test_indices[i]] oof_rows.append({ "source_id": meta["source_id"], "fold": test_fold, "true": y_test[i], "predicted": y_pred[i], "correct": y_test[i] == y_pred[i], "action_proba": float(action_proba[i]), "text": meta["text"], "tags": meta.get("tags", []), "true_route": meta["route"], }) return oof_rows # ─── Metrics from OOF ────────────────────────────────────────────────────── def compute_sixway_oof_metrics(oof_rows): y_true = np.array([r["true"] for r in oof_rows]) y_pred = np.array([r["predicted"] for r in oof_rows]) acc = accuracy_score(y_true, y_pred) macro_f1 = f1_score(y_true, y_pred, average="macro", zero_division=0) prec, rec, f1, sup = precision_recall_fscore_support(y_true, y_pred, labels=ROUTES, zero_division=0) action_tp = sum(1 for t, p in zip(y_true, y_pred) if t == "action" and p == "action") action_fp = sum(1 for t, p in zip(y_true, y_pred) if t != "action" and p == "action") action_fn = sum(1 for t, p in zip(y_true, y_pred) if t == "action" and p != "action") action_prec = action_tp / max(action_tp + action_fp, 1) action_rec = action_tp / max(action_tp + action_fn, 1) false_action = action_fp cm = confusion_matrix(y_true, y_pred, labels=ROUTES) return { "accuracy": acc, "macro_f1": macro_f1, "action_precision": action_prec, "action_recall": action_rec, "false_action": false_action, "fa_rate": false_action / max(len(y_true), 1), "per_class": { route: {"precision": float(prec[i]), "recall": float(rec[i]), "f1": float(f1[i]), "support": int(sup[i])} for i, route in enumerate(ROUTES) }, "confusion_matrix": cm.tolist(), } def compute_binary_oof_metrics(oof_rows): """Compute action P/R/FA from binary probe OOF rows.""" y_true = np.array([1 if r["true"] == "action" else 0 for r in oof_rows]) y_pred = np.array([1 if r["predicted"] == "action" else 0 for r in oof_rows]) tp = int(sum(1 for t, p in zip(y_true, y_pred) if t == 1 and p == 1)) fp = int(sum(1 for t, p in zip(y_true, y_pred) if t == 0 and p == 1)) fn = int(sum(1 for t, p in zip(y_true, y_pred) if t == 1 and p == 0)) tn = int(sum(1 for t, p in zip(y_true, y_pred) if t == 0 and p == 0)) return { "action_precision": tp / max(tp + fp, 1), "action_recall": tp / max(tp + fn, 1), "false_action": fp, "fa_rate": fp / max(len(y_true), 1), "true_neg": tn, "total": len(y_true), } # ─── Report ──────────────────────────────────────────────────────────────── def generate_report(meta, dev_examples, dev_residual, binary_grid, best_binary_key, binary_results, sixway_mlp_results, best_sixway_mlp, sixway_linear_oof, binary_linear_oof, threshold_curve, capq_diag, voice_stress): test_H = [16, 32, 64] lines = [] lines.append("# Slice 17: Nonlinear MLP Probe — Action Gate Experiment") lines.append("") lines.append("## 0. Frozen Artifacts from Slice 16") lines.append("") lines.append("```text") lines.append(f"development corpus v2 hash: {meta.get('dataset_hash', 'b27fd48f478ca477')}") lines.append(f"original frozen holdout hash: ad297fbdbbea704b (byte-identical, uninspected)") lines.append(f"embedder: {meta['embedder_id']}") lines.append(f"dimension: {meta['dimension']}") lines.append(f"pooling: mean-pool + L2-normalize") lines.append(f"input template: query: ") lines.append(f"embedding file: {EMBEDDING_PATH}") lines.append(f"total examples: {meta['total_examples']}") lines.append(f"dev pool: {meta['dev_count']}") lines.append(f"frozen holdout: {meta['frozen_count']}") lines.append(f"router-residual: {meta['residual_count']}") lines.append("```") lines.append("") # ─── Section 1: Hidden Sizes and Parameter Counts ────────────────────── lines.append("## 1. Hidden Sizes and Exact Parameter Counts") lines.append("") lines.append("### Binary action gate: e5[384] → Linear(384→H) → ReLU → Linear(H→1)") lines.append("") lines.append(f"{'H':>4} {'params':>8} {'fp32 bytes':>11} {'int8 bytes':>11}") lines.append("-" * 40) for H in HIDDEN_SIZES: n = param_count_binary(H) fp32 = n * 4 int8 = n * 1 lines.append(f"{H:>4} {n:>8,} {fp32:>11,} {int8:>11,}") lines.append("") lines.append("### Six-way MLP: e5[384] → Linear(384→H) → ReLU → Linear(H→6)") lines.append("") lines.append(f"{'H':>4} {'params':>8} {'fp32 bytes':>11} {'int8 bytes':>11}") lines.append("-" * 40) for H in HIDDEN_SIZES: n = param_count_6way(H) fp32 = n * 4 int8 = n * 1 lines.append(f"{H:>4} {n:>8,} {fp32:>11,} {int8:>11,}") lines.append("") # ─── Section 2: Selected Regularization ──────────────────────────────── lines.append("## 2. Selected Regularization") lines.append("") best_H, best_wd = best_binary_key lines.append(f"Best binary MLP: H={best_H}, weight_decay={best_wd}") lines.append(f"Selected by grouped development CV PR-AUC.") lines.append("") lines.append("Grid results (binary MLP):") lines.append("") lines.append(f"{'H':>4} {'wd':>8} {'PR-AUC':>8} {'ROC-AUC':>8} {'action_P':>10} {'action_R':>10} {'FA count':>9} {'FA rate':>9}") lines.append("-" * 80) for H in test_H: for wd in WEIGHT_DECAYS: if (H, wd) not in binary_grid: continue r = binary_grid[(H, wd)] marker = " *" if (H, wd) == best_binary_key else "" lines.append(f"{H:>4} {wd:>8} {ff(r['mean_pr_auc']):>8} {ff(r['mean_roc_auc']):>8} " f"{ff(r['mean_precision']):>10} {ff(r['mean_recall']):>10} " f"{r['total_fp']:>9} {pct(r['fa_rate']):>9}{marker}") lines.append("") # ─── Section 3: Binary MLP OOF Metrics ──────────────────────────────── lines.append("## 3. Binary MLP OOF Metrics (best: H={}, wd={})".format(best_H, best_wd)) lines.append("") best_bin = binary_results[best_binary_key] lines.append("```text") lines.append(f"ROC-AUC: {ff(best_bin['mean_roc_auc'])}") lines.append(f"PR-AUC: {ff(best_bin['mean_pr_auc'])}") lines.append(f"action precision: {ff(best_bin['mean_precision'])}") lines.append(f"action recall: {ff(best_bin['mean_recall'])}") lines.append(f"false-positive: {best_bin['total_fp']}") lines.append(f"false-negative: {best_bin['total_fn']}") lines.append(f"false-action rate: {pct(best_bin['fa_rate'])}") lines.append(f"total: {best_bin['total_n']}") lines.append(f"mean iters: {best_bin['n_iters']:.0f}") lines.append("```") lines.append("") # ─── Section 4: Fold Variance ───────────────────────────────────────── lines.append("## 4. Fold Variance") lines.append("") lines.append("Binary MLP (H={}, wd={}):".format(best_H, best_wd)) lines.append("") lines.append(f"{'fold':>5} {'ROC-AUC':>8} {'PR-AUC':>8} {'action_P':>10} {'action_R':>10} {'FP':>4} {'FN':>4} {'n':>5}") lines.append("-" * 60) for m in best_bin["fold_metrics"]: lines.append(f"{m['fold']:>5} {ff(m['roc_auc']):>8} {ff(m['pr_auc']):>8} " f"{ff(m['action_precision']):>10} {ff(m['action_recall']):>10} " f"{m['false_pos']:>4} {m['false_neg']:>4} {m['test_size']:>5}") lines.append("") # Fold variance comparison with linear binary lines.append("Fold variance comparison with linear binary probe:") bin_linear_by_fold = defaultdict(list) for r in binary_linear_oof: bin_linear_by_fold[r["fold"]].append(r) lines.append(f"{'fold':>5} {'linear FP':>10} {'MLP FP':>10} {'linear FA%':>11} {'MLP FA%':>10}") lines.append("-" * 50) for fold in sorted(set(m["fold"] for m in best_bin["fold_metrics"])): lin_rows = bin_linear_by_fold[fold] lin_fa = sum(1 for r in lin_rows if r["true"] == "not_action" and r["predicted"] == "action") mlp_m = next(m for m in best_bin["fold_metrics"] if m["fold"] == fold) lin_n = len(lin_rows) lines.append(f"{fold:>5} {lin_fa:>10} {mlp_m['false_pos']:>10} " f"{pct(lin_fa/max(lin_n,1)):>11} {pct(mlp_m['false_pos']/max(mlp_m['test_size'],1)):>10}") lines.append("") # ─── Section 5: Safety Operating Curve ───────────────────────────────── lines.append("## 5. Safety Operating Curve (best binary MLP)") lines.append("") lines.append(f"{'threshold':>10} {'action_P':>10} {'action_R':>10} {'FA count':>10} {'FA rate':>10} {'coverage':>10}") lines.append("-" * 65) for t in threshold_curve: marker = "" if t["action_precision"] >= 0.95 and t["action_recall"] > 0: marker = " ← P≥0.95" lines.append(f"{t['threshold']:>10.3f} {t['action_precision']:>10.4f} {t['action_recall']:>10.4f} " f"{t['false_action_count']:>10} {t['false_action_rate']:>10.4f} " f"{t['coverage']:>10.4f}{marker}") lines.append("") # Check if useful region exists useful = [t for t in threshold_curve if t["action_precision"] >= 0.95 and t["action_recall"] >= 0.1] if useful: lines.append(f"**Useful region found**: at threshold {useful[0]['threshold']:.3f}, " f"action_P={useful[0]['action_precision']:.4f}, action_R={useful[0]['action_recall']:.4f}, " f"FA={useful[0]['false_action_count']}") else: best_95 = [t for t in threshold_curve if t["action_precision"] >= 0.95] if best_95: lines.append(f"At action_P ≥ 0.95: best recall = {max(t['action_recall'] for t in best_95):.4f} " f"(at threshold {max(best_95, key=lambda t: t['action_recall'])['threshold']:.3f})") else: lines.append("No threshold achieves action_P ≥ 0.95.") lines.append("") # ─── Section 6: Capability-Question Diagnostic ──────────────────────── lines.append("## 6. Capability-Question Boundary Diagnostic") lines.append("") lines.append(f"capability_question → false action rate: {pct(capq_diag['cap_q_fa_rate'])} " f"({capq_diag['cap_q_false_actions']}/{capq_diag['cap_q_total']})") lines.append(f"positive modal request → action recall: {ff(capq_diag['modal_action_recall'])} " f"({capq_diag['modal_action_total']} examples)") lines.append("") # Paired examples lines.append("### Paired capability-question vs action examples") lines.append("") lines.append("These are the critical diagnostic pairs:") lines.append("") cap_q_examples = [r for r in best_bin["oof_rows"] if "capability_question" in r.get("tags", []) and r["true_route"] == "knowledge"] action_examples = [r for r in best_bin["oof_rows"] if r["true_route"] == "action" and "capability_question" not in r.get("tags", [])] # Show some representative pairs lines.append(f"{'text':<50} {'route':<12} {'predicted':<12} {'action_P':>10} {'FA?':>5}") lines.append("-" * 95) for r in cap_q_examples[:15]: fa_mark = "YES" if r["predicted"] == "action" else "" lines.append(f"{r['text'][:49]:<50} {r['true_route']:<12} {r['predicted']:<12} " f"{r['action_proba']:>10.4f} {fa_mark:>5}") lines.append("...") for r in [r for r in action_examples if r["true"] == "action"][:10]: lines.append(f"{r['text'][:49]:<50} {r['true_route']:<12} {r['predicted']:<12} " f"{r['action_proba']:>10.4f}") lines.append("") # ─── Section 7: Voice-Like Stress ───────────────────────────────────── lines.append("## 7. Voice-Like Stress Results") lines.append("") lines.append("```text") lines.append(f"Total stress-testable pairs: {voice_stress['total_stress_pairs']}") lines.append(f"question with ?: {voice_stress['q_with_question_mark']}") lines.append(f"question without ?: {voice_stress['q_without_question_mark']}") lines.append(f"positive polite/modal action: {voice_stress['modal_action_pairs']}") lines.append(f"capability-question pairs: {voice_stress['cap_q_pairs']}") lines.append("```") lines.append("") lines.append("Sample affected pairs:") lines.append("") for p in voice_stress["pairs"][:10]: lines.append(f" {p['source_id']}:") lines.append(f" original: \"{p['original']}\"") lines.append(f" stressed: \"{p['stressed']}\"") lines.append(f" route: {p['route']}") lines.append("") lines.append("Impact assessment:") lines.append(" The MLP's decision boundary must not depend on punctuation artifacts.") lines.append(" If MLP success depends on '?' presence, it will fail under voice input.") lines.append("") # ─── Section 8: Six-Way MLP Results ─────────────────────────────────── lines.append("## 8. Six-Way MLP Results") lines.append("") six_mlp = best_sixway_mlp lines.append(f"Best six-way MLP: H={six_mlp['H']}, wd={six_mlp['wd']}") lines.append("") lines.append("```text") lines.append(f"accuracy: {pct(six_mlp['metrics']['accuracy'])}") lines.append(f"macro F1: {ff(six_mlp['metrics']['macro_f1'])}") lines.append(f"action precision: {ff(six_mlp['metrics']['action_precision'])}") lines.append(f"action recall: {ff(six_mlp['metrics']['action_recall'])}") lines.append(f"false-action rate: {pct(six_mlp['metrics']['fa_rate'])}") lines.append(f"false-action count: {six_mlp['metrics']['false_action']}") lines.append("```") lines.append("") lines.append("Per-route F1:") lines.append("") for route in ROUTES: pc = six_mlp["metrics"]["per_class"][route] lines.append(f" {route:<15} P={ff(pc['precision'])} R={ff(pc['recall'])} F1={ff(pc['f1'])} (n={pc['support']})") lines.append("") lines.append("Confusion matrix (rows=expected, cols=predicted):") header = f"{'':>15}" + "".join(f"{r:>15}" for r in ROUTES) lines.append(header) for i, route in enumerate(ROUTES): row = f"{route:>15}" + "".join(f"{six_mlp['metrics']['confusion_matrix'][i][j]:>15}" for j in range(len(ROUTES))) lines.append(row) lines.append("") lines.append("Fold variance (six-way MLP):") lines.append("") lines.append(f"{'fold':>5} {'accuracy':>10} {'macro_f1':>10} {'action_P':>10} {'action_R':>10} {'FA':>5}") lines.append("-" * 55) for m in six_mlp["fold_metrics"]: lines.append(f"{m['fold']:>5} {pct(m['accuracy']):>10} {ff(m['macro_f1']):>10} " f"{ff(m['action_precision']):>10} {ff(m['action_recall']):>10} {m['false_action']:>5}") lines.append("") # ─── Section 9: Comparison Against All Linear Baselines ──────────────── lines.append("## 9. Comparison Against All Linear Baselines") lines.append("") # Compute linear metrics six_lin = compute_sixway_oof_metrics(sixway_linear_oof) bin_lin = compute_binary_oof_metrics(binary_linear_oof) # Binary MLP metrics bin_mlp = compute_binary_oof_metrics(best_bin["oof_rows"]) # Structural baseline from slice 16 struct_action_p = 0.610 struct_action_r = 0.805 struct_fa_rate = 0.151 struct_macro_f1 = 0.622 lines.append(f"| {'model':<28} | {'params beyond e5':>18} | {'action P':>10} | {'action R':>10} | {'FA rate':>10} | {'macro F1':>10} |") lines.append(f"| {'-'*28} | {'-'*18} | {'-'*10} | {'-'*10} | {'-'*10} | {'-'*10} |") lines.append(f"| {'6-way linear':<28} | {'2,310':>18} | {ff(six_lin['action_precision']):>10} | {ff(six_lin['action_recall']):>10} | {pct(six_lin['fa_rate']):>10} | {ff(six_lin['macro_f1']):>10} |") lines.append(f"| {'binary linear':<28} | {'~385':>18} | {ff(bin_lin['action_precision']):>10} | {ff(bin_lin['action_recall']):>10} | {pct(bin_lin['fa_rate']):>10} | {'—':>10} |") lines.append(f"| {'structural linear':<28} | {'~2,313':>18} | {ff(struct_action_p):>10} | {ff(struct_action_r):>10} | {pct(struct_fa_rate):>10} | {ff(struct_macro_f1):>10} |") lines.append(f"| {'binary MLP (H={})'.format(best_H):<28} | {param_count_binary(best_H):>18,} | {ff(bin_mlp['action_precision']):>10} | {ff(bin_mlp['action_recall']):>10} | {pct(bin_mlp['fa_rate']):>10} | {'—':>10} |") lines.append(f"| {'6-way MLP (H={})'.format(six_mlp['H']):<28} | {param_count_6way(six_mlp['H']):>18,} | {ff(six_mlp['metrics']['action_precision']):>10} | {ff(six_mlp['metrics']['action_recall']):>10} | {pct(six_mlp['metrics']['fa_rate']):>10} | {ff(six_mlp['metrics']['macro_f1']):>10} |") lines.append("") # ─── Section 10: Conclusion ─────────────────────────────────────────── lines.append("## 10. Conclusion") lines.append("") # Decision logic from the brief bin_p = bin_mlp["action_precision"] bin_r = bin_mlp["action_recall"] six_f1 = six_mlp["metrics"]["macro_f1"] improvement_vs_linear = bin_mlp["fa_rate"] < bin_lin["fa_rate"] * 0.9 # >10% improvement if bin_p >= 0.90 and bin_r >= 0.30: lines.append("### Verdict: nonlinear e5 head sufficient") lines.append("") lines.append(f"The binary MLP achieves action precision {ff(bin_p)} with recall {ff(bin_r)},") lines.append(f"a material improvement over the linear binary probe (P={ff(bin_lin['action_precision'])}, R={ff(bin_lin['action_recall'])}).") lines.append("") lines.append("The e5 representation contains the signal, but its geometry is nonlinear.") lines.append("A tiny nonlinear gate/head remains viable for production.") elif bin_p >= 0.80 and improvement_vs_linear: lines.append("### Verdict: marginal improvement — binary nonlinear gate only") lines.append("") lines.append(f"The binary MLP shows marginal improvement over the linear baseline.") lines.append(f"It may be useful as a dedicated action gate but not as a full six-way router.") elif six_f1 > 0.60: lines.append("### Verdict: six-way MLP marginally better, binary gate preferred") lines.append("") lines.append("The six-way MLP shows modest improvement but not enough to justify") lines.append("the nonlinear overhead. A deterministic fast path + binary gate is preferred.") else: lines.append("### Verdict: e5 representation inadequate") lines.append("") lines.append("Neither the binary MLP nor the six-way MLP provides material improvement") lines.append("over linear baselines. Mean-pooled e5-small is inadequate for Maven's") lines.append("action-pragmatics boundary.") lines.append("") lines.append("Stop probing e5. Consider a different encoder or a fundamentally different approach.") lines.append("") # Specific diagnostic lines.append("### Capability-question diagnostic") lines.append("") lines.append(f" capability_question false-action rate: {pct(capq_diag['cap_q_fa_rate'])}") lines.append(f" positive modal action recall: {ff(capq_diag['modal_action_recall'])}") lines.append("") if capq_diag["cap_q_fa_rate"] < 0.10: lines.append(" The MLP successfully separates capability questions from executable actions.") elif capq_diag["cap_q_fa_rate"] < bin_lin["fa_rate"]: lines.append(f" MLP reduces capability-question false actions vs linear ({pct(capq_diag['cap_q_fa_rate'])} vs {pct(bin_lin['fa_rate'])})") else: lines.append(" The MLP does not materially improve the capability-question boundary.") lines.append("") # Six-way vs binary lines.append("### Six-way MLP vs binary gate") lines.append("") if six_f1 < 0.55 and bin_p >= 0.80: lines.append("Six-way MLP fails but binary MLP succeeds.") lines.append("Prefer: deterministic fast path → binary executable-action gate → coarse non-action routing.") lines.append("Do not force one six-way model to solve everything.") elif six_f1 > 0.60: lines.append(f"Six-way MLP achieves macro F1 {ff(six_f1)}. A single tiny nonlinear semantic router remains plausible.") else: lines.append("Both models are marginal. The action gating question should be resolved before full six-way routing.") lines.append("") # ─── Section 11: Diagnostic Tooling Commit Hash ─────────────────────── lines.append("## 11. Commit hash for diagnostic tooling") lines.append("") lines.append("(to be filled after commit)") lines.append("") return "\n".join(lines) # ─── Main ────────────────────────────────────────────────────────────────── def main(): print("Loading data...") meta, examples = load_data() dev_examples = filter_dev_pool(examples) dev_residual = filter_residual(dev_examples) X_dev, y_dev = extract_Xy(dev_examples) fold_ids_dev = get_fold_groups(dev_examples) X_res, y_res = extract_Xy(dev_residual) fold_ids_res = get_fold_groups(dev_residual) print(f"Dev pool: {len(dev_examples)}, residual: {len(dev_residual)}") print(f"Embedding dim: {X_dev.shape[1]}") print() # ─── 1. Binary MLP grid search ──────────────────────────────────────── print("=" * 60) print(" BINARY MLP ACTION GATE — Grid Search") print("=" * 60) best_binary_key, binary_grid = run_binary_mlp_grid( X_dev, y_dev, fold_ids_dev, dev_examples ) best_H, best_wd = best_binary_key print(f"\nBest: H={best_H}, wd={best_wd}, PR-AUC={ff(binary_grid[best_binary_key]['mean_pr_auc'])}") # ─── 2. Six-way MLP with best H ────────────────────────────────────── print("\n" + "=" * 60) print(" SIX-WAY MLP — Best Hidden Size") print("=" * 60) # Search only the most promising hidden sizes for the six-way model # (H=8 is clearly too small; focus on the widths that show signal) sixway_results = {} sixway_Hs = [16, 32, 64] for H in sixway_Hs: wds = WEIGHT_DECAYS if H < 64 else [0.0, 1e-3] for wd in wds: print(f" 6-way MLP H={H} wd={wd} ...", end=" ", flush=True) r = run_sixway_mlp_cv(X_dev, y_dev, fold_ids_dev, H, wd, dev_examples, max_iter=600) sixway_results[(H, wd)] = r print(f"F1={ff(r['macro_f1'])} P={ff(r['action_precision'])} R={ff(r['action_recall'])} FA={r['false_action']}") best_sixway_key = max(sixway_results, key=lambda k: sixway_results[k]["macro_f1"]) best_sixway_mlp = { "H": best_sixway_key[0], "wd": best_sixway_key[1], "metrics": compute_sixway_oof_metrics(sixway_results[best_sixway_key]["oof_rows"]), "fold_metrics": sixway_results[best_sixway_key]["fold_metrics"], } print(f"\nBest 6-way: H={best_sixway_mlp['H']}, wd={best_sixway_mlp['wd']}, F1={ff(best_sixway_mlp['metrics']['macro_f1'])}") # ─── 3. Linear baselines (for comparison table) ─────────────────────── print("\n" + "=" * 60) print(" LINEAR BASELINES (for comparison)") print("=" * 60) print(" Six-way linear ...", end=" ", flush=True) sixway_linear_oof = run_linear_baseline(X_dev, y_dev, fold_ids_dev, dev_examples) print("done") print(" Binary linear ...", end=" ", flush=True) binary_linear_oof = run_binary_linear_baseline(X_dev, y_dev, fold_ids_dev, dev_examples) print("done") # ─── 4. Threshold curve for best binary MLP ─────────────────────────── print("\n Computing threshold curve ...", end=" ", flush=True) best_bin_oof = binary_grid[best_binary_key]["oof_rows"] threshold_curve = compute_action_threshold_curve(best_bin_oof, ACTION_THRESHOLDS) print("done") # ─── 5. Capability-question diagnostic ───────────────────────────────── print(" Capability-question diagnostic ...", end=" ", flush=True) capq_diag = capability_question_diagnostic(best_bin_oof) print(f"done (cap_q FA rate={pct(capq_diag['cap_q_fa_rate'])})") # ─── 6. Voice stress ────────────────────────────────────────────────── print(" Voice stress evaluation ...", end=" ", flush=True) voice_stress = voice_stress_eval(examples, dev_examples, None) print(f"done ({voice_stress['total_stress_pairs']} pairs)") # ─── 7. Generate report ─────────────────────────────────────────────── print("\n" + "=" * 60) print(" GENERATING REPORT") print("=" * 60) report = generate_report( meta, dev_examples, dev_residual, binary_grid, best_binary_key, binary_grid, sixway_results, best_sixway_mlp, sixway_linear_oof, binary_linear_oof, threshold_curve, capq_diag, voice_stress, ) report_path = "/home/kami/apps/Maven/docs/evals/2026-09-07-nonlinear-e5-mlp-probe.md" with open(report_path, "w") as f: f.write(report) print(f"\nReport written to {report_path}") # Print summary bin_mlp = compute_binary_oof_metrics(best_bin_oof) six_lin = compute_sixway_oof_metrics(sixway_linear_oof) bin_lin = compute_binary_oof_metrics(binary_linear_oof) print("\n" + "=" * 60) print(" SLICE 17 SUMMARY") print("=" * 60) print(f" Binary MLP: P={ff(bin_mlp['action_precision'])} R={ff(bin_mlp['action_recall'])} FA={bin_mlp['false_action']} ({pct(bin_mlp['fa_rate'])}) [H={best_H}, wd={best_wd}]") print(f" 6-way MLP: P={ff(best_sixway_mlp['metrics']['action_precision'])} R={ff(best_sixway_mlp['metrics']['action_recall'])} FA={best_sixway_mlp['metrics']['false_action']} ({pct(best_sixway_mlp['metrics']['fa_rate'])}) F1={ff(best_sixway_mlp['metrics']['macro_f1'])}") print(f" 6-way linear: P={ff(six_lin['action_precision'])} R={ff(six_lin['action_recall'])} FA={six_lin['false_action']} ({pct(six_lin['fa_rate'])}) F1={ff(six_lin['macro_f1'])}") print(f" Binary linear: P={ff(bin_lin['action_precision'])} R={ff(bin_lin['action_recall'])} FA={bin_lin['false_action']} ({pct(bin_lin['fa_rate'])})") print(f" Cap-q FA rate: {pct(capq_diag['cap_q_fa_rate'])}") print(f" Modal recall: {ff(capq_diag['modal_action_recall'])}") if __name__ == "__main__": main()