summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--BABYAI_SHARED_FEEDBACK.md116
-rw-r--r--experiments/babyai_shared_run.py260
-rw-r--r--experiments/babyai_shared_smoke.py150
-rw-r--r--experiments/prepare_babyai_shared.py136
-rw-r--r--sdil/babyai_shared.py299
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}