summaryrefslogtreecommitdiff
path: root/experiments/prepare_babyai_shared.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/prepare_babyai_shared.py')
-rw-r--r--experiments/prepare_babyai_shared.py136
1 files changed, 136 insertions, 0 deletions
diff --git a/experiments/prepare_babyai_shared.py b/experiments/prepare_babyai_shared.py
new file mode 100644
index 0000000..04a7cc2
--- /dev/null
+++ b/experiments/prepare_babyai_shared.py
@@ -0,0 +1,136 @@
+#!/usr/bin/env python3
+"""Generate deterministic BabyAI expert demonstrations for shared feedback."""
+
+import argparse
+import json
+from pathlib import Path
+import sys
+import time
+
+import gymnasium as gym
+import minigrid
+import numpy as np
+from minigrid.core.constants import COLOR_TO_IDX, OBJECT_TO_IDX, STATE_TO_IDX
+from minigrid.utils.baby_ai_bot import BabyAIBot
+
+sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
+from sdil.babyai_shared import build_vocabulary, missions_to_bow
+
+
+ROOT = Path(__file__).resolve().parents[1]
+DEFAULT_OUT = ROOT / "data" / "babyai_shared" / "goto_obj_s6_b0.npz"
+
+
+def generate_split(env_id, seeds):
+ images = []
+ directions = []
+ missions = []
+ actions = []
+ episode_offsets = [0]
+ episode_returns = []
+ env = gym.make(env_id)
+ try:
+ for seed in seeds:
+ observation, _ = env.reset(seed=int(seed))
+ bot = BabyAIBot(env)
+ previous_action = None
+ episode_return = 0.0
+ finished = False
+ for _ in range(env.unwrapped.max_steps):
+ action = bot.replan(previous_action)
+ images.append(observation["image"].copy())
+ directions.append(int(observation["direction"]))
+ missions.append(str(observation["mission"]))
+ actions.append(int(action))
+ observation, reward, terminated, truncated, _ = env.step(int(action))
+ episode_return += float(reward)
+ previous_action = action
+ if terminated or truncated:
+ finished = True
+ break
+ if not finished or episode_return <= 0:
+ raise RuntimeError(
+ f"BabyAIBot failed env={env_id} seed={seed} "
+ f"return={episode_return}")
+ episode_offsets.append(len(actions))
+ episode_returns.append(episode_return)
+ finally:
+ env.close()
+ return {
+ "image": np.asarray(images, dtype=np.uint8),
+ "direction": np.asarray(directions, dtype=np.uint8),
+ "mission_text": np.asarray(missions),
+ "action": np.asarray(actions, dtype=np.uint8),
+ "episode_offset": np.asarray(episode_offsets, dtype=np.int64),
+ "episode_return": np.asarray(episode_returns, dtype=np.float32),
+ "episode_seed": np.asarray(list(seeds), dtype=np.int64),
+ }
+
+
+def prefix(values, name):
+ return {f"{name}_{key}": value for key, value in values.items()}
+
+
+def main():
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--env-id", default="BabyAI-GoToObjS6-v1")
+ parser.add_argument("--train-episodes", type=int, default=20_000)
+ parser.add_argument("--validation-episodes", type=int, default=2_000)
+ parser.add_argument("--train-seed-start", type=int, default=0)
+ parser.add_argument("--validation-seed-start", type=int, default=100_000)
+ parser.add_argument("--rollout-seed-start", type=int, default=200_000)
+ parser.add_argument("--rollout-episodes", type=int, default=500)
+ parser.add_argument("--out", type=Path, default=DEFAULT_OUT)
+ args = parser.parse_args()
+ if args.train_seed_start + args.train_episodes > args.validation_seed_start:
+ raise ValueError("training and validation episode seeds overlap")
+ if (args.validation_seed_start + args.validation_episodes
+ > args.rollout_seed_start):
+ raise ValueError("validation demonstrations and rollouts overlap")
+
+ started = time.time()
+ train_seeds = range(
+ args.train_seed_start, args.train_seed_start + args.train_episodes)
+ validation_seeds = range(
+ args.validation_seed_start,
+ args.validation_seed_start + args.validation_episodes)
+ train = generate_split(args.env_id, train_seeds)
+ validation = generate_split(args.env_id, validation_seeds)
+ vocabulary = build_vocabulary(train["mission_text"])
+ train["mission_bow"] = missions_to_bow(
+ train.pop("mission_text"), vocabulary)
+ validation["mission_bow"] = missions_to_bow(
+ validation.pop("mission_text"), vocabulary)
+
+ metadata = {
+ "protocol": "babyai_shared_feedback_b0",
+ "env_id": args.env_id,
+ "minigrid_version": minigrid.__version__,
+ "train_episodes": args.train_episodes,
+ "validation_episodes": args.validation_episodes,
+ "train_steps": int(len(train["action"])),
+ "validation_steps": int(len(validation["action"])),
+ "rollout_seed_start": args.rollout_seed_start,
+ "rollout_episodes": args.rollout_episodes,
+ "object_cardinality": max(OBJECT_TO_IDX.values()) + 1,
+ "color_cardinality": max(COLOR_TO_IDX.values()) + 1,
+ "state_cardinality": max(STATE_TO_IDX.values()) + 1,
+ "vocabulary": list(vocabulary),
+ "elapsed_seconds": time.time() - started,
+ }
+ payload = {
+ **prefix(train, "train"),
+ **prefix(validation, "validation"),
+ "rollout_seed": np.arange(
+ args.rollout_seed_start,
+ args.rollout_seed_start + args.rollout_episodes,
+ dtype=np.int64),
+ "metadata_json": np.asarray(json.dumps(metadata, sort_keys=True)),
+ }
+ args.out.parent.mkdir(parents=True, exist_ok=True)
+ np.savez_compressed(args.out, **payload)
+ print(json.dumps({"out": str(args.out), **metadata}, indent=2))
+
+
+if __name__ == "__main__":
+ main()