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/prepare_babyai_shared.py | |
| parent | c18e939242d6fd88fc07aa35c4d11d17d663fce0 (diff) | |
exp: freeze BabyAI shared-feedback protocol
Diffstat (limited to 'experiments/prepare_babyai_shared.py')
| -rw-r--r-- | experiments/prepare_babyai_shared.py | 136 |
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() |
