summaryrefslogtreecommitdiff
path: root/experiments/babyai_shared_run.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-10 10:13:57 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-10 10:13:57 -0500
commit6d99f411cf7e3f271c85976f27e3ad780f9ecbc8 (patch)
tree39e320b983592e713802d6e55bc31296b23673d5 /experiments/babyai_shared_run.py
parentc18e939242d6fd88fc07aa35c4d11d17d663fce0 (diff)
exp: freeze BabyAI shared-feedback protocol
Diffstat (limited to 'experiments/babyai_shared_run.py')
-rw-r--r--experiments/babyai_shared_run.py260
1 files changed, 260 insertions, 0 deletions
diff --git a/experiments/babyai_shared_run.py b/experiments/babyai_shared_run.py
new file mode 100644
index 0000000..1b460b7
--- /dev/null
+++ b/experiments/babyai_shared_run.py
@@ -0,0 +1,260 @@
+#!/usr/bin/env python3
+"""Train one BabyAI shared-feedback condition from the frozen protocol."""
+
+import argparse
+import json
+import os
+from pathlib import Path
+import subprocess
+import sys
+import time
+
+import gymnasium as gym
+import minigrid
+import numpy as np
+import torch
+
+sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
+from sdil.babyai_shared import (
+ BabyAISharedConfig, BabyAISharedNet, CONDITIONS, encode_visual,
+ evaluate_actions, manual_step, missions_to_bow, visual_input_dim,
+)
+
+
+ROOT = Path(__file__).resolve().parents[1]
+DEFAULT_DATA = ROOT / "data" / "babyai_shared" / "goto_obj_s6_b0.npz"
+
+
+def git_output(*args):
+ return subprocess.run(
+ ["git", *args], cwd=ROOT, check=True, capture_output=True,
+ text=True).stdout.strip()
+
+
+def load_data(path):
+ archive = np.load(path, allow_pickle=False)
+ metadata = json.loads(str(archive["metadata_json"]))
+ train = {name: archive[f"train_{name}"] for name in (
+ "image", "direction", "mission_bow", "action")}
+ validation = {name: archive[f"validation_{name}"] for name in (
+ "image", "direction", "mission_bow", "action")}
+ rollout_seeds = archive["rollout_seed"].copy()
+ return metadata, train, validation, rollout_seeds
+
+
+@torch.no_grad()
+def rollout_policy(net, env_id, seeds, vocabulary, cardinalities,
+ context_enabled=True, chunk_size=64):
+ successes = 0
+ returns = []
+ lengths = []
+ for chunk_start in range(0, len(seeds), chunk_size):
+ chunk = seeds[chunk_start:chunk_start + chunk_size]
+ envs = [gym.make(env_id) for _ in chunk]
+ observations = []
+ active = []
+ chunk_returns = [0.0] * len(envs)
+ chunk_lengths = [0] * len(envs)
+ try:
+ for index, (env, seed) in enumerate(zip(envs, chunk)):
+ observation, _ = env.reset(seed=int(seed))
+ observations.append(observation)
+ active.append(index)
+ while active:
+ images = np.asarray([
+ observations[index]["image"] for index in active],
+ dtype=np.uint8)
+ directions = np.asarray([
+ observations[index]["direction"] for index in active],
+ dtype=np.uint8)
+ mission_text = [
+ observations[index]["mission"] for index in active]
+ mission_bow = missions_to_bow(mission_text, vocabulary)
+ features = encode_visual(
+ images, directions, *cardinalities,
+ device=net.device, dtype=net.dtype)
+ mission_tensor = torch.as_tensor(
+ mission_bow, device=net.device, dtype=net.dtype)
+ actions = net.forward_features(
+ features, mission_tensor,
+ context_enabled=context_enabled)["logits"].argmax(1)
+ actions = actions.cpu().numpy()
+ next_active = []
+ for position, index in enumerate(active):
+ observation, reward, terminated, truncated, _ = (
+ envs[index].step(int(actions[position])))
+ observations[index] = observation
+ chunk_returns[index] += float(reward)
+ chunk_lengths[index] += 1
+ if terminated or truncated:
+ successes += int(chunk_returns[index] > 0)
+ returns.append(chunk_returns[index])
+ lengths.append(chunk_lengths[index])
+ else:
+ next_active.append(index)
+ active = next_active
+ finally:
+ for env in envs:
+ env.close()
+ return {
+ "success": successes / len(seeds),
+ "mean_return": float(np.mean(returns)),
+ "mean_length": float(np.mean(lengths)),
+ "episodes": int(len(seeds)),
+ }
+
+
+def main():
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--data", type=Path, default=DEFAULT_DATA)
+ parser.add_argument("--condition", choices=CONDITIONS, required=True)
+ parser.add_argument("--hidden-layers", type=int, default=2)
+ parser.add_argument("--width", type=int, default=256)
+ parser.add_argument("--learning-rate", type=float, default=0.03)
+ parser.add_argument("--context-gain", type=float, default=1.0)
+ parser.add_argument("--epochs", type=int, default=15)
+ parser.add_argument("--batch-size", type=int, default=256)
+ parser.add_argument("--neutral-examples", type=int, default=4096)
+ parser.add_argument("--model-seed", type=int, default=4101)
+ parser.add_argument("--shuffle-seed", type=int, default=4101)
+ 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("BabyAI endpoint requires clean tracked source")
+ if not args.data.exists():
+ raise FileNotFoundError(args.data)
+
+ device = torch.device(args.device)
+ if device.type == "cpu":
+ torch.set_num_threads(1)
+ if device.type == "cuda":
+ torch.cuda.set_device(device)
+ torch.cuda.reset_peak_memory_stats(device)
+ metadata, train, validation, rollout_seeds = load_data(args.data)
+ cardinalities = (
+ metadata["object_cardinality"], metadata["color_cardinality"],
+ metadata["state_cardinality"])
+ config = BabyAISharedConfig(
+ input_dim=visual_input_dim(*cardinalities),
+ mission_dim=len(metadata["vocabulary"]), width=args.width,
+ hidden_layers=args.hidden_layers, learning_rate=args.learning_rate,
+ reciprocal_learning_rate=args.learning_rate,
+ context_gain=args.context_gain)
+ net = BabyAISharedNet(config, seed=args.model_seed, device=device)
+ shuffle = torch.Generator(device="cpu").manual_seed(args.shuffle_seed)
+ neutral_count = min(args.neutral_examples, len(train["action"]))
+ neutral_indices = np.arange(neutral_count)
+ predictor_reports = []
+ epoch_history = []
+ first_nonfinite_epoch = None
+ started = time.time()
+
+ for epoch in range(args.epochs):
+ if args.condition == "sdil":
+ neutral_features = encode_visual(
+ train["image"][neutral_indices],
+ train["direction"][neutral_indices], *cardinalities,
+ device=device, dtype=net.dtype)
+ neutral_missions = torch.as_tensor(
+ train["mission_bow"][neutral_indices],
+ device=device, dtype=net.dtype)
+ predictor_reports = net.fit_neutral_predictor(
+ neutral_features, neutral_missions)
+ permutation = torch.randperm(
+ len(train["action"]), generator=shuffle).numpy()
+ losses = []
+ for start in range(0, len(permutation), args.batch_size):
+ indices = permutation[start:start + args.batch_size]
+ features = encode_visual(
+ train["image"][indices], train["direction"][indices],
+ *cardinalities, device=device, dtype=net.dtype)
+ 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, args.condition)
+ losses.append(loss)
+ mean_loss = float(np.mean(losses))
+ epoch_history.append({"epoch": epoch + 1, "train_loss": mean_loss})
+ if not np.isfinite(mean_loss):
+ first_nonfinite_epoch = epoch + 1
+ break
+
+ training_seconds = time.time() - started
+ validation_metrics = evaluate_actions(
+ net, validation["image"], validation["direction"],
+ validation["mission_bow"], validation["action"], cardinalities)
+ lesion_validation_metrics = evaluate_actions(
+ net, validation["image"], validation["direction"],
+ validation["mission_bow"], validation["action"], cardinalities,
+ context_enabled=False)
+ rollout = rollout_policy(
+ net, metadata["env_id"], rollout_seeds, metadata["vocabulary"],
+ cardinalities)
+ lesion_rollout = rollout_policy(
+ net, metadata["env_id"], rollout_seeds, metadata["vocabulary"],
+ cardinalities, context_enabled=False)
+ total_seconds = time.time() - started
+ parameter_count = sum(value.numel() for values in (
+ net.W, net.b, net.Q[1:], net.C, net.P, net.P_bias)
+ for value in values)
+ result = {
+ "stage": "babyai_shared_b0" if args.epochs == 15 else "babyai_shared_b1",
+ "condition": args.condition,
+ "finite": first_nonfinite_epoch is None,
+ "first_nonfinite_epoch": first_nonfinite_epoch,
+ "epochs_completed": len(epoch_history),
+ "epoch_history": epoch_history,
+ "config": config.to_dict(),
+ "training": {
+ "batch_size": args.batch_size,
+ "model_seed": args.model_seed,
+ "shuffle_seed": args.shuffle_seed,
+ "neutral_examples_per_epoch": (
+ neutral_count if args.condition == "sdil" else 0),
+ "parameter_count_including_reciprocal_context_predictor": (
+ parameter_count),
+ "training_wall_seconds": training_seconds,
+ "total_wall_seconds_including_rollouts": total_seconds,
+ },
+ "validation": validation_metrics,
+ "mission_lesion_validation": lesion_validation_metrics,
+ "rollout": rollout,
+ "mission_lesion_rollout": lesion_rollout,
+ "predictor": predictor_reports,
+ "data": metadata,
+ "provenance": {
+ "git_commit": git_output("rev-parse", "HEAD"),
+ "git_dirty_tracked": False,
+ "torch_version": torch.__version__,
+ "minigrid_version": minigrid.__version__,
+ "device": str(device),
+ "cuda_visible_devices": os.environ.get("CUDA_VISIBLE_DEVICES"),
+ "cuda_device_name": (
+ torch.cuda.get_device_name(device)
+ if device.type == "cuda" else None),
+ "cuda_peak_allocated_bytes": (
+ int(torch.cuda.max_memory_allocated(device))
+ if device.type == "cuda" else None),
+ },
+ }
+ args.out.parent.mkdir(parents=True, exist_ok=True)
+ with open(args.out, "w", encoding="utf-8") as handle:
+ json.dump(result, handle, indent=2, sort_keys=True)
+ handle.write("\n")
+ print(json.dumps({
+ "condition": args.condition,
+ "validation": validation_metrics,
+ "mission_lesion_validation": lesion_validation_metrics,
+ "rollout": rollout,
+ "mission_lesion_rollout": lesion_rollout,
+ "finite": result["finite"],
+ "out": str(args.out),
+ }, indent=2, sort_keys=True))
+
+
+if __name__ == "__main__":
+ main()