diff options
| -rw-r--r-- | BABYAI_SHARED_FEEDBACK.md | 116 | ||||
| -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 | ||||
| -rw-r--r-- | sdil/babyai_shared.py | 299 |
5 files changed, 961 insertions, 0 deletions
diff --git a/BABYAI_SHARED_FEEDBACK.md b/BABYAI_SHARED_FEEDBACK.md new file mode 100644 index 0000000..addd58a --- /dev/null +++ b/BABYAI_SHARED_FEEDBACK.md @@ -0,0 +1,116 @@ +# BabyAI shared-feedback protocol + +## Question + +The clean scaling experiments show that reciprocal Kolen--Pollack (KP) +feedback can train large networks, but do not show that SDIL adds anything to +KP. This experiment asks whether SDIL is useful when an ordinary top-down +signal is required during inference and the same physical pathway also carries +the teaching signal. + +BabyAI supplies this structure without an injected nuisance. The visual +observation is the basal input, while the language mission is required to +choose the correct object and enters each hidden population through a fixed +apical projection. During imitation learning, the measured apical signal is + +```text +a_l = C_l m + t_l, +``` + +where `C_l m` is the mission field used during inference and `t_l` is the +reciprocal KP teaching field. Raw shared-path learning uses `a_l`. SDIL fits +a per-cell affine prediction of `C_l m` from the local hidden activity and uses +`a_l - P_l(h_l)`. + +This is a fixed-distribution, randomly interleaved grounded-language task. It +is not a continual-learning experiment. The endpoint tests signal separation +and task learning, not catastrophic forgetting. + +## Data and task + +The first task is the official `BabyAI-GoToObjS6-v1` environment from +MiniGrid 3.1.0. Expert demonstrations are generated by the official +`BabyAIBot`. Every recorded trajectory must end with positive reward; failed +expert episodes abort data generation rather than being silently discarded. + +The basal observation is the standard partial `7 x 7 x 3` symbolic image plus +agent direction. Each of the object, color, and state channels is one-hot +encoded. The mission is a normalized bag of training-vocabulary tokens and is +not included in the basal input. Validation missions use the training +vocabulary, with an explicit unknown token retained for auditability. + +Development data contain 20,000 expert episodes for training and 2,000 +disjoint expert episodes for validation. A separate set of 500 validation +seeds is used for closed-loop policy rollouts. No test seeds are generated or +evaluated until the method configuration and mechanism comparison are frozen. + +## Model and local updates + +The policy is a tanh MLP with a fixed mission projection into every hidden +population and a seven-action linear readout: + +```text +u_l = W_l h_(l-1) + b_l + C_l m +h_l = tanh(u_l). +``` + +`C_l` is fixed after initialization. Forward weights and independently stored +reciprocal weights use the same modified-KP momentum, decay, and local +correlation rule used in the existing shared-feedback implementation. The +reciprocal tensor is updated from the same locally available activity product; +it is never copied from the forward tensor. SDIL predictor fitting receives +only the current hidden activity and ordinary mission field, never the expert +action, loss, output error, KP teaching field, or downstream weight. + +## B0 clean selector + +B0 exists only to choose a model that the task and clean learning rules can +train. It runs BP and clean KP, never raw shared KP or SDIL. The fixed grid is + +```text +hidden layers: 2, 4 +hidden width: 256 +learning rate: 0.01, 0.03 +context gain: 1.0 +epochs: 15 +batch size: 256 +model/data seed: 4101 +``` + +All candidates share demonstrations, minibatch order, initialization within a +model shape, and rollout seeds. A candidate is eligible only if BP and clean +KP each reach at least 80% validation rollout success and removing the mission +field reduces BP rollout success by at least 20 points. Among eligible +candidates, choose the highest clean-KP rollout success; ties are resolved by +action accuracy, then fewer layers, then the smaller learning rate. If no +candidate is eligible, B0 fails and the shared-path endpoint is not run. + +## B1 mechanism endpoint + +The selected depth and learning rate are frozen. B1 trains for 40 epochs on +three model seeds with identical data and minibatch orders across conditions: + +1. `bp`: exact backpropagation reference with the same fixed mission pathway. +2. `clean_kp`: instruction-only reciprocal KP; this receives a separate clean + teaching wire and is the local-learning upper bound. +3. `raw_shared`: the local teaching rule directly uses `C_l m + t_l`. +4. `sdil`: a per-cell neutral affine predictor is fitted on 4,096 training + observations before each epoch and subtracts the predicted mission field. + +The predictor's extra forward observations, wall time, and arithmetic are +reported. The main metrics are validation closed-loop mission success, +validation expert-action accuracy, mission-lesion success, cross-entropy, +predictor explained variance, and the residual mission-field RMS. Accuracy +and mission success, not residual size, decide the result. + +B1 supports the mechanism if clean KP is successful, mission removal damages +the task, raw shared KP is worse than clean KP, and SDIL recovers a substantial +part of that downstream success gap consistently across seeds. No fixed +five-point threshold is imposed before observing the natural effect size. +Norm-matched raw feedback is added only after a positive raw--SDIL difference, +because it diagnoses that difference but cannot create it. + +Passing B1 opens `PickupLoc` and `PutNextLocalS6N4` with the same method and +selection rule, followed by depth scaling. Failure is retained as evidence +that task-required shared feedback alone is insufficient to make +residualization useful in this setting. 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() diff --git a/sdil/babyai_shared.py b/sdil/babyai_shared.py new file mode 100644 index 0000000..a4c871a --- /dev/null +++ b/sdil/babyai_shared.py @@ -0,0 +1,299 @@ +"""Shared inference-and-teaching pathway for BabyAI imitation learning. + +The visual observation is basal. The language mission is projected into each +hidden population through a fixed apical path. Manual updates implement exact +BP, clean reciprocal KP, raw shared-path KP, or SDIL residualization. +""" + +from dataclasses import asdict, dataclass +import math +import re + +import numpy as np +import torch +import torch.nn.functional as F + + +CONDITIONS = ("bp", "clean_kp", "raw_shared", "sdil") +TOKEN_PATTERN = re.compile(r"[a-z0-9]+") + + +def tokenize_mission(mission): + return TOKEN_PATTERN.findall(str(mission).lower()) + + +def build_vocabulary(missions): + tokens = sorted({token for mission in missions + for token in tokenize_mission(mission)}) + return ("<unk>", *tokens) + + +def missions_to_bow(missions, vocabulary): + lookup = {token: index for index, token in enumerate(vocabulary)} + result = np.zeros((len(missions), len(vocabulary)), dtype=np.float32) + for row, mission in enumerate(missions): + tokens = tokenize_mission(mission) + for token in tokens: + result[row, lookup.get(token, 0)] += 1.0 + norm = np.linalg.norm(result[row]) + if norm > 0: + result[row] /= norm + return result + + +@dataclass(frozen=True) +class BabyAISharedConfig: + input_dim: int + mission_dim: int + action_dim: int = 7 + width: int = 256 + hidden_layers: int = 2 + learning_rate: float = 0.03 + reciprocal_learning_rate: float = 0.03 + momentum: float = 0.9 + weight_decay: float = 1e-4 + context_gain: float = 1.0 + + def to_dict(self): + return asdict(self) + + +class BabyAISharedNet: + """Tanh policy with fixed mission projections and local manual updates.""" + + def __init__(self, config, seed, device="cpu", dtype=torch.float32): + if config.hidden_layers < 1: + raise ValueError("BabyAI shared policy needs a hidden population") + self.config = config + self.device = torch.device(device) + self.dtype = dtype + sizes = ([config.input_dim] + + [config.width] * config.hidden_layers + + [config.action_dim]) + forward_generator = torch.Generator(device="cpu").manual_seed(seed) + reciprocal_generator = torch.Generator(device="cpu").manual_seed(seed + 1) + context_generator = torch.Generator(device="cpu").manual_seed(seed + 2) + + def normal(shape, scale, generator): + value = torch.randn(*shape, generator=generator) * scale + return value.to(device=self.device, dtype=dtype) + + self.W = [normal((sizes[i + 1], sizes[i]), + 1.0 / math.sqrt(sizes[i]), forward_generator) + for i in range(len(sizes) - 1)] + self.b = [torch.zeros(size, device=self.device, dtype=dtype) + for size in sizes[1:]] + # Q[i] has the same storage orientation as W[i]. Q[0] is not needed + # because no teaching signal is transported into the visual input. + self.Q = [None] + [normal(tuple(self.W[i].shape), + 1.0 / math.sqrt(sizes[i]), + reciprocal_generator) + for i in range(1, len(self.W))] + self.C = [normal((config.width, config.mission_dim), + config.context_gain + / math.sqrt(config.mission_dim), + context_generator) + for _ in range(config.hidden_layers)] + + self.P = [torch.zeros(config.width, device=self.device, dtype=dtype) + for _ in range(config.hidden_layers)] + self.P_bias = [torch.zeros_like(value) for value in self.P] + self.mW = [torch.zeros_like(value) for value in self.W] + self.mb = [torch.zeros_like(value) for value in self.b] + self.mQ = [None] + [torch.zeros_like(value) for value in self.Q[1:]] + + def clone(self): + copied = BabyAISharedNet( + self.config, seed=0, device=self.device, dtype=self.dtype) + for name in ("W", "b", "Q", "C", "P", "P_bias", "mW", "mb", "mQ"): + values = getattr(self, name) + setattr(copied, name, [ + None if value is None else value.clone() for value in values]) + return copied + + def context_fields(self, missions, enabled=True): + if not enabled: + return [torch.zeros( + (missions.shape[0], self.config.width), + device=self.device, dtype=self.dtype) for _ in self.C] + return [missions @ projection.t() for projection in self.C] + + def forward_features(self, features, missions, context_enabled=True): + contexts = self.context_fields(missions, enabled=context_enabled) + h = [features] + u = [] + for layer in range(self.config.hidden_layers): + value = (h[-1] @ self.W[layer].t() + self.b[layer] + + contexts[layer]) + u.append(value) + h.append(torch.tanh(value)) + logits = h[-1] @ self.W[-1].t() + self.b[-1] + h.append(logits) + return {"h": h, "u": u, "context": contexts, "logits": logits} + + def predictor(self, layer, soma): + return self.P[layer] * soma + self.P_bias[layer] + + @torch.no_grad() + def fit_neutral_predictor(self, features, missions): + """Fit the per-cell neutral relation without actions or teaching.""" + state = self.forward_features(features, missions) + reports = [] + for layer, target in enumerate(state["context"]): + soma = state["h"][layer + 1] + soma_centered = soma - soma.mean(0) + target_centered = target - target.mean(0) + variance = soma_centered.square().mean(0) + covariance = (soma_centered * target_centered).mean(0) + slope = covariance / variance.clamp_min(1e-8) + intercept = target.mean(0) - slope * soma.mean(0) + self.P[layer].copy_(slope) + self.P_bias[layer].copy_(intercept) + prediction = slope * soma + intercept + residual = target - prediction + target_ss = target_centered.square().sum(0) + residual_ss = residual.square().sum(0) + valid = target_ss > 1e-12 + r2 = 1.0 - residual_ss[valid] / target_ss[valid] + context_rms = target.square().mean().sqrt() + reports.append({ + "mean_per_cell_r2": float(r2.mean()) if r2.numel() else 0.0, + "context_rms": float(context_rms), + "residual_context_rms_ratio": float( + residual.square().mean().sqrt() + / context_rms.clamp_min(1e-12)), + "neutral_observations": int(features.shape[0]), + "action_observations": 0, + "teaching_observations": 0, + }) + return reports + + +def visual_input_dim(object_cardinality, color_cardinality, + state_cardinality, view_size=7): + return (view_size * view_size + * (object_cardinality + color_cardinality + state_cardinality) + 4) + + +def encode_visual(images, directions, object_cardinality, + color_cardinality, state_cardinality, device, + dtype=torch.float32): + """One-hot encode the compact MiniGrid symbolic observation.""" + images = torch.as_tensor(images, device=device, dtype=torch.long) + directions = torch.as_tensor(directions, device=device, dtype=torch.long) + object_features = F.one_hot( + images[..., 0], num_classes=object_cardinality) + color_features = F.one_hot( + images[..., 1], num_classes=color_cardinality) + state_features = F.one_hot( + images[..., 2], num_classes=state_cardinality) + grid = torch.cat( + (object_features, color_features, state_features), dim=-1) + grid = grid.reshape(grid.shape[0], -1).to(dtype) + direction = F.one_hot(directions, num_classes=4).to(dtype) + return torch.cat((grid, direction), dim=1) + + +def select_teaching_signal(net, layer, instruction, context, soma, condition): + raw = instruction + context + innovation = raw - net.predictor(layer, soma) + if condition == "clean_kp": + used = instruction + elif condition == "raw_shared": + used = raw + elif condition == "sdil": + used = innovation + else: + raise ValueError(f"no shared teaching rule for condition {condition}") + return used, raw, innovation + + +@torch.no_grad() +def manual_step(net, features, missions, actions, condition): + """Apply one exact-BP or reciprocal local-learning update.""" + if condition not in CONDITIONS: + raise ValueError(f"unknown BabyAI condition: {condition}") + state = net.forward_features(features, missions) + h, u, contexts = state["h"], state["u"], state["context"] + probabilities = torch.softmax(state["logits"], dim=1) + output_instruction = ( + F.one_hot(actions, num_classes=net.config.action_dim).to(net.dtype) + - probabilities) + batch = features.shape[0] + + hidden_deltas = [None] * net.config.hidden_layers + used_fields = [None] * net.config.hidden_layers + raw_fields = [None] * net.config.hidden_layers + innovation_fields = [None] * net.config.hidden_layers + child_delta = output_instruction + for layer in reversed(range(net.config.hidden_layers)): + if condition == "bp": + used = child_delta @ net.W[layer + 1] + raw = used + contexts[layer] + innovation = raw - net.predictor(layer, h[layer + 1]) + else: + instruction = child_delta @ net.Q[layer + 1] + used, raw, innovation = select_teaching_signal( + net, layer, instruction, contexts[layer], h[layer + 1], + condition) + delta = used * (1.0 - torch.tanh(u[layer]).square()) + hidden_deltas[layer] = delta + used_fields[layer] = used + raw_fields[layer] = raw + innovation_fields[layer] = innovation + child_delta = delta + + directions = [hidden_deltas[0].t() @ h[0] / batch] + for layer in range(1, net.config.hidden_layers): + directions.append(hidden_deltas[layer].t() @ h[layer] / batch) + directions.append(output_instruction.t() @ h[-2] / batch) + bias_directions = [delta.mean(0) for delta in hidden_deltas] + bias_directions.append(output_instruction.mean(0)) + + for layer, (direction, bias_direction) in enumerate( + zip(directions, bias_directions)): + net.mW[layer].mul_(net.config.momentum).add_( + direction - net.config.weight_decay * net.W[layer]) + net.mb[layer].mul_(net.config.momentum).add_(bias_direction) + net.W[layer].add_(net.mW[layer], alpha=net.config.learning_rate) + net.b[layer].add_(net.mb[layer], alpha=net.config.learning_rate) + if condition != "bp" and layer > 0: + net.mQ[layer].mul_(net.config.momentum).add_( + direction - net.config.weight_decay * net.Q[layer]) + net.Q[layer].add_( + net.mQ[layer], alpha=net.config.reciprocal_learning_rate) + + loss = F.cross_entropy(state["logits"], actions) + return float(loss), { + "directions": directions, + "bias_directions": bias_directions, + "used": used_fields, + "raw": raw_fields, + "innovation": innovation_fields, + "context": contexts, + } + + +@torch.no_grad() +def evaluate_actions(net, images, directions, missions, actions, + cardinalities, batch_size=1024, context_enabled=True): + device = net.device + total_loss = 0.0 + correct = 0 + total = len(actions) + for start in range(0, total, batch_size): + stop = min(start + batch_size, total) + features = encode_visual( + images[start:stop], directions[start:stop], *cardinalities, + device=device, dtype=net.dtype) + mission_batch = torch.as_tensor( + missions[start:stop], device=device, dtype=net.dtype) + action_batch = torch.as_tensor( + actions[start:stop], device=device, dtype=torch.long) + logits = net.forward_features( + features, mission_batch, + context_enabled=context_enabled)["logits"] + total_loss += float(F.cross_entropy( + logits, action_batch, reduction="sum")) + correct += int((logits.argmax(1) == action_batch).sum()) + return {"accuracy": correct / total, "loss": total_loss / total} |
