From b5eab6cd911e074ae9268303389dfbe52e636e26 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Mon, 10 Aug 2026 10:47:10 -0500 Subject: exp: freeze BabyAI population predictor screen --- BABYAI_SHARED_FEEDBACK.md | 16 ++ experiments/babyai_population_predictor_screen.py | 186 ++++++++++++++++++++++ experiments/babyai_shared_smoke.py | 15 +- sdil/babyai_shared.py | 37 +++++ 4 files changed, 253 insertions(+), 1 deletion(-) create mode 100644 experiments/babyai_population_predictor_screen.py diff --git a/BABYAI_SHARED_FEEDBACK.md b/BABYAI_SHARED_FEEDBACK.md index aa24b84..5749949 100644 --- a/BABYAI_SHARED_FEEDBACK.md +++ b/BABYAI_SHARED_FEEDBACK.md @@ -220,3 +220,19 @@ neutral prediction screen that reads no action, teaching, raw, SDIL, or test outcome. If it cannot explain at least 80% of held-out mission-field variance and leave at most 25% context RMS in every layer, no further PickupLoc endpoint is run. + +### P3 population-predictor capacity screen + +P3 does not train raw or SDIL and does not evaluate rollout or test outcomes. +For clean-KP networks with the frozen P2 architecture and seeds `4101--4103`, +it audits epochs `0, 1, 5, 10, 20, 40`. At each checkpoint, a full layer-local +linear map is ridge-fitted from 4,096 neutral `(h_l, C_l m)` pairs and evaluated +on 4,096 disjoint neutral pairs. The ridge coefficient is `1e-3` times the +mean somatic variance. The fit receives no action, teaching field, loss, +downstream weight, raw result, or SDIL result. + +The capacity screen passes only if every seed, checkpoint, and layer has +held-out mean per-cell `R^2 >= 0.8` and residual context RMS at most `0.25` of +the original. This ridge fit is an expressivity upper bound, not the proposed +hardware learning rule. Passing opens a separate local-delta predictor check; +failure closes the population-predictor rescue. diff --git a/experiments/babyai_population_predictor_screen.py b/experiments/babyai_population_predictor_screen.py new file mode 100644 index 0000000..34c2a9f --- /dev/null +++ b/experiments/babyai_population_predictor_screen.py @@ -0,0 +1,186 @@ +#!/usr/bin/env python3 +"""Neutral-only capacity screen for the original population P_l h_l map.""" + +import argparse +import json +import os +from pathlib import Path +import subprocess +import sys + +import numpy as np +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from sdil.babyai_shared import ( + BabyAISharedConfig, BabyAISharedNet, build_history_index, + fit_population_predictor, history_input_dim, manual_step, + population_predictor_metrics, +) +from babyai_shared_run import encode_compact_batch, load_data + + +ROOT = Path(__file__).resolve().parents[1] +DEFAULT_PICKUP_DATA = ROOT / "data" / "babyai_shared" / "pickup_loc_p0.npz" +AUDIT_EPOCHS = (0, 1, 5, 10, 20, 40) + + +def git_output(*args): + return subprocess.run( + ["git", *args], cwd=ROOT, check=True, capture_output=True, + text=True).stdout.strip() + + +@torch.no_grad() +def audit_predictor(net, calibration_features, calibration_missions, + holdout_features, holdout_missions, ridge): + calibration = net.forward_features( + calibration_features, calibration_missions) + holdout = net.forward_features(holdout_features, holdout_missions) + layers = [] + for layer in range(net.config.hidden_layers): + coefficient, intercept = fit_population_predictor( + calibration["h"][layer + 1], calibration["context"][layer], ridge) + train_metrics = population_predictor_metrics( + calibration["h"][layer + 1], calibration["context"][layer], + coefficient, intercept) + holdout_metrics = population_predictor_metrics( + holdout["h"][layer + 1], holdout["context"][layer], + coefficient, intercept) + layers.append({ + "layer": layer, + "calibration": train_metrics, + "holdout": holdout_metrics, + "action_observations": 0, + "teaching_observations": 0, + }) + return layers + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--data", type=Path, default=DEFAULT_PICKUP_DATA) + parser.add_argument("--model-seed", type=int, required=True) + parser.add_argument("--shuffle-seed", type=int) + parser.add_argument("--ridge", type=float, default=1e-3) + parser.add_argument("--device", default="cuda") + parser.add_argument("--out", type=Path, required=True) + args = parser.parse_args() + if git_output("status", "--porcelain", "--untracked-files=no"): + raise RuntimeError("population predictor screen requires clean tracked source") + shuffle_seed = (args.model_seed if args.shuffle_seed is None + else args.shuffle_seed) + device = torch.device(args.device) + if device.type == "cuda": + device_index = (torch.cuda.current_device() + if device.index is None else device.index) + torch.cuda.set_device(device_index) + device = torch.device("cuda", device_index) + else: + torch.set_num_threads(1) + metadata, train, _, _ = load_data(args.data) + cardinalities = ( + metadata["object_cardinality"], metadata["color_cardinality"], + metadata["state_cardinality"]) + history_steps = 4 + config = BabyAISharedConfig( + input_dim=history_input_dim(*cardinalities, history_steps), + mission_dim=len(metadata["vocabulary"]), width=256, hidden_layers=4, + learning_rate=0.03, reciprocal_learning_rate=0.03, + context_gain=1.0) + net = BabyAISharedNet(config, seed=args.model_seed, device=device) + history = build_history_index( + train["episode_offset"], train["action"], history_steps) + calibration_indices = np.arange(0, 4096) + holdout_indices = np.arange(4096, 8192) + calibration_features = encode_compact_batch( + train, calibration_indices, cardinalities, net, history) + holdout_features = encode_compact_batch( + train, holdout_indices, cardinalities, net, history) + calibration_missions = torch.as_tensor( + train["mission_bow"][calibration_indices], + device=device, dtype=net.dtype) + holdout_missions = torch.as_tensor( + train["mission_bow"][holdout_indices], + device=device, dtype=net.dtype) + shuffle = torch.Generator(device="cpu").manual_seed(shuffle_seed) + audits = [{ + "epoch": 0, + "layers": audit_predictor( + net, calibration_features, calibration_missions, + holdout_features, holdout_missions, args.ridge), + }] + epoch_losses = [] + for epoch in range(1, 41): + permutation = torch.randperm( + len(train["action"]), generator=shuffle).numpy() + losses = [] + for start in range(0, len(permutation), 256): + indices = permutation[start:start + 256] + features = encode_compact_batch( + train, indices, cardinalities, net, history) + missions = torch.as_tensor( + train["mission_bow"][indices], + device=device, dtype=net.dtype) + actions = torch.as_tensor( + train["action"][indices], device=device, dtype=torch.long) + loss, _ = manual_step( + net, features, missions, actions, "clean_kp") + losses.append(loss) + epoch_losses.append(float(np.mean(losses))) + if epoch in AUDIT_EPOCHS: + audits.append({ + "epoch": epoch, + "layers": audit_predictor( + net, calibration_features, calibration_missions, + holdout_features, holdout_missions, args.ridge), + }) + all_holdout = [layer["holdout"] + for audit in audits for layer in audit["layers"]] + checks = { + "holdout_r2_at_least_0p8_every_epoch_and_layer": all( + row["mean_per_cell_r2"] >= 0.8 for row in all_holdout), + "holdout_residual_ratio_at_most_0p25_every_epoch_and_layer": all( + row["residual_context_rms_ratio"] <= 0.25 + for row in all_holdout), + "zero_action_or_teaching_observations": all( + layer["action_observations"] == 0 + and layer["teaching_observations"] == 0 + for audit in audits for layer in audit["layers"]), + } + report = { + "stage": "babyai_population_predictor_capacity_screen", + "gate": "pass" if all(checks.values()) else "fail", + "checks": checks, + "model_seed": args.model_seed, + "shuffle_seed": shuffle_seed, + "config": config.to_dict(), + "ridge_relative_to_mean_soma_variance": args.ridge, + "calibration_examples": 4096, + "holdout_examples": 4096, + "audits": audits, + "epoch_clean_kp_train_loss": epoch_losses, + "raw_sdil_rollout_or_test_outcomes_read": False, + "provenance": { + "git_commit": git_output("rev-parse", "HEAD"), + "torch_version": torch.__version__, + "device": str(device), + "cuda_visible_devices": os.environ.get("CUDA_VISIBLE_DEVICES"), + }, + } + args.out.parent.mkdir(parents=True, exist_ok=True) + with open(args.out, "w", encoding="utf-8") as handle: + json.dump(report, handle, indent=2, sort_keys=True) + handle.write("\n") + print(json.dumps({ + "gate": report["gate"], "checks": checks, + "minimum_holdout_r2": min( + row["mean_per_cell_r2"] for row in all_holdout), + "maximum_holdout_residual_ratio": max( + row["residual_context_rms_ratio"] for row in all_holdout), + "out": str(args.out), + }, indent=2, sort_keys=True)) + + +if __name__ == "__main__": + main() diff --git a/experiments/babyai_shared_smoke.py b/experiments/babyai_shared_smoke.py index 16d3e90..1b50a67 100644 --- a/experiments/babyai_shared_smoke.py +++ b/experiments/babyai_shared_smoke.py @@ -11,7 +11,8 @@ from minigrid.core.constants import COLOR_TO_IDX, OBJECT_TO_IDX, STATE_TO_IDX sys.path.insert(0, str(Path(__file__).resolve().parents[1])) from sdil.babyai_shared import ( BabyAISharedConfig, BabyAISharedNet, build_history_index, build_vocabulary, - encode_history_visual, history_input_dim, manual_step, missions_to_bow, + encode_history_visual, fit_population_predictor, history_input_dim, + manual_step, missions_to_bow, population_predictor_metrics, select_teaching_signal, ) from prepare_babyai_shared import generate_split @@ -79,6 +80,16 @@ def main(): innovation_identity_error = float((used - innovation).abs().max()) assert raw_identity_error < 1e-14 assert innovation_identity_error == 0.0 + population_state = base.forward_features(features, missions) + split = features.shape[0] // 2 + coefficient, intercept = fit_population_predictor( + population_state["h"][1][:split], + population_state["context"][0][:split]) + population_metrics = population_predictor_metrics( + population_state["h"][1][split:], + population_state["context"][0][split:], coefficient, intercept) + assert torch.isfinite(torch.tensor( + population_metrics["mean_per_cell_r2"])) # The manual BP direction must exactly match autograd on the same fixed # mission-conditioned network. @@ -157,6 +168,8 @@ def main(): "predictor_teaching_observations": max( row["teaching_observations"] for row in predictor_reports), "history_rollout_episodes": rollout["episodes"], + "population_predictor_holdout_r2": population_metrics[ + "mean_per_cell_r2"], }) diff --git a/sdil/babyai_shared.py b/sdil/babyai_shared.py index 209ea36..238a210 100644 --- a/sdil/babyai_shared.py +++ b/sdil/babyai_shared.py @@ -169,6 +169,43 @@ class BabyAISharedNet: return reports +@torch.no_grad() +def fit_population_predictor(soma, target, ridge=1e-3): + """Fit target ~= soma @ coefficient + intercept for a capacity audit.""" + soma_mean = soma.mean(0) + target_mean = target.mean(0) + centered_soma = soma - soma_mean + centered_target = target - target_mean + gram = centered_soma.t() @ centered_soma / soma.shape[0] + cross = centered_soma.t() @ centered_target / soma.shape[0] + scale = gram.diagonal().mean().clamp_min(1e-8) + regularized = gram + ridge * scale * torch.eye( + gram.shape[0], device=gram.device, dtype=gram.dtype) + coefficient = torch.linalg.solve(regularized, cross) + intercept = target_mean - soma_mean @ coefficient + return coefficient, intercept + + +@torch.no_grad() +def population_predictor_metrics(soma, target, coefficient, intercept): + prediction = soma @ coefficient + intercept + residual = target - prediction + centered_target = target - target.mean(0) + target_ss = centered_target.square().sum(0) + residual_ss = residual.square().sum(0) + valid = target_ss > 1e-12 + r2 = 1.0 - residual_ss[valid] / target_ss[valid] + context_rms = target.square().mean().sqrt() + return { + "mean_per_cell_r2": float(r2.mean()) if r2.numel() else 0.0, + "context_rms": float(context_rms), + "residual_context_rms_ratio": float( + residual.square().mean().sqrt() + / context_rms.clamp_min(1e-12)), + "observations": int(soma.shape[0]), + } + + def visual_input_dim(object_cardinality, color_cardinality, state_cardinality, view_size=7): return (view_size * view_size -- cgit v1.2.3