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