#!/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, build_history_index, encode_history_visual, encode_visual, evaluate_actions, history_input_dim, 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", "episode_offset")} validation = {name: archive[f"validation_{name}"] for name in ( "image", "direction", "mission_bow", "action", "episode_offset")} 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, history_steps=1): 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 = [] observation_histories = [] previous_action_histories = [] 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) observation_histories.append([observation]) previous_action_histories.append([-1]) active.append(index) while active: if history_steps == 1: 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) else: image_shape = observations[active[0]]["image"].shape images = np.zeros( (len(active), history_steps, *image_shape), dtype=np.uint8) directions = np.zeros( (len(active), history_steps), dtype=np.uint8) previous_actions = np.full( (len(active), history_steps), -1, dtype=np.int64) history_mask = np.zeros( (len(active), history_steps), dtype=bool) for row, index in enumerate(active): obs_history = observation_histories[index][-history_steps:] action_history = previous_action_histories[index][ -history_steps:] offset = history_steps - len(obs_history) images[row, offset:] = np.asarray([ value["image"] for value in obs_history]) directions[row, offset:] = np.asarray([ value["direction"] for value in obs_history]) previous_actions[row, offset:] = action_history history_mask[row, offset:] = True mission_text = [ observations[index]["mission"] for index in active] mission_bow = missions_to_bow(mission_text, vocabulary) if history_steps == 1: features = encode_visual( images, directions, *cardinalities, device=net.device, dtype=net.dtype) else: features = encode_history_visual( images, directions, previous_actions, history_mask, *cardinalities, net.config.action_dim, 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 observation_histories[index].append(observation) previous_action_histories[index].append( int(actions[position])) 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 encode_compact_batch(split, indices, cardinalities, net, history): if history is None: return encode_visual( split["image"][indices], split["direction"][indices], *cardinalities, device=net.device, dtype=net.dtype) history_indices, previous_actions, mask = history selected_history = history_indices[indices] return encode_history_visual( split["image"][selected_history], split["direction"][selected_history], previous_actions[indices], mask[indices], *cardinalities, net.config.action_dim, device=net.device, dtype=net.dtype) 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("--history-steps", type=int, default=1) parser.add_argument("--model-seed", type=int, default=4101) parser.add_argument("--shuffle-seed", type=int, default=4101) parser.add_argument("--stage") 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"]) if args.history_steps < 1: raise ValueError("history steps must be positive") input_dim = (visual_input_dim(*cardinalities) if args.history_steps == 1 else history_input_dim( *cardinalities, args.history_steps)) config = BabyAISharedConfig( input_dim=input_dim, 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) train_history = (None if args.history_steps == 1 else build_history_index( train["episode_offset"], train["action"], args.history_steps)) validation_history = ( None if args.history_steps == 1 else build_history_index( validation["episode_offset"], validation["action"], args.history_steps)) 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_compact_batch( train, neutral_indices, cardinalities, net, train_history) 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_compact_batch( train, indices, cardinalities, net, train_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, 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, history=validation_history) lesion_validation_metrics = evaluate_actions( net, validation["image"], validation["direction"], validation["mission_bow"], validation["action"], cardinalities, context_enabled=False, history=validation_history) rollout = rollout_policy( net, metadata["env_id"], rollout_seeds, metadata["vocabulary"], cardinalities, history_steps=args.history_steps) lesion_rollout = rollout_policy( net, metadata["env_id"], rollout_seeds, metadata["vocabulary"], cardinalities, context_enabled=False, history_steps=args.history_steps) 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": (args.stage or ( "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), "history_steps": args.history_steps, "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()