#!/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()