diff options
Diffstat (limited to 'experiments')
| -rw-r--r-- | experiments/babyai_shared_run.py | 260 | ||||
| -rw-r--r-- | experiments/babyai_shared_smoke.py | 150 | ||||
| -rw-r--r-- | experiments/prepare_babyai_shared.py | 136 |
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() |
