diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-10 10:13:57 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-10 10:13:57 -0500 |
| commit | 6d99f411cf7e3f271c85976f27e3ad780f9ecbc8 (patch) | |
| tree | 39e320b983592e713802d6e55bc31296b23673d5 /experiments/babyai_shared_run.py | |
| parent | c18e939242d6fd88fc07aa35c4d11d17d663fce0 (diff) | |
exp: freeze BabyAI shared-feedback protocol
Diffstat (limited to 'experiments/babyai_shared_run.py')
| -rw-r--r-- | experiments/babyai_shared_run.py | 260 |
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() |
