summaryrefslogtreecommitdiff
path: root/experiments/babyai_population_predictor_screen.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/babyai_population_predictor_screen.py')
-rw-r--r--experiments/babyai_population_predictor_screen.py186
1 files changed, 186 insertions, 0 deletions
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()