summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-10 10:13:57 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-10 10:13:57 -0500
commit6d99f411cf7e3f271c85976f27e3ad780f9ecbc8 (patch)
tree39e320b983592e713802d6e55bc31296b23673d5 /experiments
parentc18e939242d6fd88fc07aa35c4d11d17d663fce0 (diff)
exp: freeze BabyAI shared-feedback protocol
Diffstat (limited to 'experiments')
-rw-r--r--experiments/babyai_shared_run.py260
-rw-r--r--experiments/babyai_shared_smoke.py150
-rw-r--r--experiments/prepare_babyai_shared.py136
3 files changed, 546 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()
diff --git a/experiments/babyai_shared_smoke.py b/experiments/babyai_shared_smoke.py
new file mode 100644
index 0000000..e5bbb55
--- /dev/null
+++ b/experiments/babyai_shared_smoke.py
@@ -0,0 +1,150 @@
+#!/usr/bin/env python3
+"""Single mechanics smoke test for the BabyAI shared-feedback path."""
+
+from pathlib import Path
+import sys
+
+import torch
+import torch.nn.functional as F
+from minigrid.core.constants import COLOR_TO_IDX, OBJECT_TO_IDX, STATE_TO_IDX
+
+sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
+from sdil.babyai_shared import (
+ BabyAISharedConfig, BabyAISharedNet, build_vocabulary, encode_visual,
+ manual_step, missions_to_bow, select_teaching_signal, visual_input_dim,
+)
+from prepare_babyai_shared import generate_split
+
+
+def maximum_difference(left, right):
+ return max(float((a - b).abs().max()) for a, b in zip(left, right))
+
+
+def main():
+ torch.set_num_threads(1)
+ data = generate_split("BabyAI-GoToObjS6-v1", range(32))
+ vocabulary = build_vocabulary(data["mission_text"])
+ missions = torch.from_numpy(missions_to_bow(
+ data["mission_text"], vocabulary)).to(torch.float64)
+ cardinalities = (
+ max(OBJECT_TO_IDX.values()) + 1,
+ max(COLOR_TO_IDX.values()) + 1,
+ max(STATE_TO_IDX.values()) + 1,
+ )
+ features = encode_visual(
+ data["image"], data["direction"], *cardinalities,
+ device="cpu", dtype=torch.float64)
+ actions = torch.from_numpy(data["action"].astype("int64"))
+ config = BabyAISharedConfig(
+ input_dim=visual_input_dim(*cardinalities),
+ mission_dim=len(vocabulary), width=16, hidden_layers=2,
+ learning_rate=0.01, reciprocal_learning_rate=0.01,
+ momentum=0.0, weight_decay=0.0)
+ base = BabyAISharedNet(config, seed=4101, dtype=torch.float64)
+
+ clones = {name: base.clone() for name in (
+ "bp", "clean_kp", "raw_shared", "sdil")}
+ reference_logits = clones["bp"].forward_features(
+ features, missions)["logits"]
+ forward_identity_error = max(float((
+ reference_logits
+ - clones[name].forward_features(features, missions)["logits"]
+ ).abs().max()) for name in clones if name != "bp")
+ assert forward_identity_error == 0.0
+ lesioned_logits = base.forward_features(
+ features, missions, context_enabled=False)["logits"]
+ mission_lesion_logit_change = float(
+ (reference_logits - lesioned_logits).abs().max())
+ assert mission_lesion_logit_change > 0
+
+ predictor_reports = clones["sdil"].fit_neutral_predictor(
+ features, missions)
+ assert all(row["action_observations"] == 0 for row in predictor_reports)
+ assert all(row["teaching_observations"] == 0 for row in predictor_reports)
+ state = clones["sdil"].forward_features(features, missions)
+ instruction = torch.randn_like(state["context"][0])
+ used, raw, innovation = select_teaching_signal(
+ clones["sdil"], 0, instruction, state["context"][0],
+ state["h"][1], "sdil")
+ raw_identity_error = float((
+ raw - instruction - state["context"][0]).abs().max())
+ innovation_identity_error = float((used - innovation).abs().max())
+ assert raw_identity_error < 1e-14
+ assert innovation_identity_error == 0.0
+
+ # The manual BP direction must exactly match autograd on the same fixed
+ # mission-conditioned network.
+ bp = base.clone()
+ before_w = [value.clone() for value in bp.W]
+ before_b = [value.clone() for value in bp.b]
+ manual_step(bp, features, missions, actions, "bp")
+ auto_w = [value.clone().requires_grad_(True) for value in before_w]
+ auto_b = [value.clone().requires_grad_(True) for value in before_b]
+ hidden = features
+ for layer in range(config.hidden_layers):
+ context = missions @ base.C[layer].t()
+ hidden = torch.tanh(hidden @ auto_w[layer].t()
+ + auto_b[layer] + context)
+ logits = hidden @ auto_w[-1].t() + auto_b[-1]
+ F.cross_entropy(logits, actions).backward()
+ bp_direction_error = max(
+ max(float(((after - before) / config.learning_rate
+ + parameter.grad).abs().max())
+ for after, before, parameter in zip(bp.W, before_w, auto_w)),
+ max(float(((after - before) / config.learning_rate
+ + parameter.grad).abs().max())
+ for after, before, parameter in zip(bp.b, before_b, auto_b)),
+ )
+ assert bp_direction_error < 1e-12
+
+ # With a zero mission path and zero predictor, all shared KP rules coincide.
+ zero = base.clone()
+ for value in zero.C + zero.P + zero.P_bias:
+ value.zero_()
+ clean = zero.clone()
+ raw_net = zero.clone()
+ sdil = zero.clone()
+ manual_step(clean, features, missions, actions, "clean_kp")
+ manual_step(raw_net, features, missions, actions, "raw_shared")
+ manual_step(sdil, features, missions, actions, "sdil")
+ zero_context_error = max(
+ maximum_difference(clean.W, raw_net.W),
+ maximum_difference(clean.W, sdil.W),
+ maximum_difference(clean.Q[1:], raw_net.Q[1:]),
+ maximum_difference(clean.Q[1:], sdil.Q[1:]),
+ )
+ assert zero_context_error < 1e-14
+
+ # KP reciprocal parameters receive the same local correlation increment,
+ # without reading the updated forward tensor.
+ kp = base.clone()
+ kp_w = [value.clone() for value in kp.W]
+ kp_q = [None] + [value.clone() for value in kp.Q[1:]]
+ manual_step(kp, features, missions, actions, "clean_kp")
+ reciprocal_direction_error = max(float((
+ (kp.W[layer] - kp_w[layer]) / config.learning_rate
+ - (kp.Q[layer] - kp_q[layer]) / config.reciprocal_learning_rate
+ ).abs().max()) for layer in range(1, len(kp.W)))
+ assert reciprocal_direction_error < 1e-12
+ assert all(not value.requires_grad for values in (
+ kp.W, kp.b, kp.Q[1:], kp.C, kp.P, kp.P_bias) for value in values)
+
+ print({
+ "expert_episodes": 32,
+ "expert_steps": len(actions),
+ "forward_identity_error": forward_identity_error,
+ "mission_lesion_logit_change": mission_lesion_logit_change,
+ "raw_identity_error": raw_identity_error,
+ "innovation_identity_error": innovation_identity_error,
+ "bp_direction_error": bp_direction_error,
+ "zero_context_update_error": zero_context_error,
+ "reciprocal_direction_error": reciprocal_direction_error,
+ "predictor_action_observations": max(
+ row["action_observations"] for row in predictor_reports),
+ "predictor_teaching_observations": max(
+ row["teaching_observations"] for row in predictor_reports),
+ })
+
+
+if __name__ == "__main__":
+ main()
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()