diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-10 10:47:10 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-10 10:47:10 -0500 |
| commit | b5eab6cd911e074ae9268303389dfbe52e636e26 (patch) | |
| tree | a650bd25467f1484d49c9360f851361c0b038247 /experiments/babyai_population_predictor_screen.py | |
| parent | bc93d2a673f5f971415715dad11c5e8baf2a6202 (diff) | |
exp: freeze BabyAI population predictor screen
Diffstat (limited to 'experiments/babyai_population_predictor_screen.py')
| -rw-r--r-- | experiments/babyai_population_predictor_screen.py | 186 |
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() |
