summaryrefslogtreecommitdiff
path: root/experiments/babyai_shared_smoke.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-10 10:13:57 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-10 10:13:57 -0500
commit6d99f411cf7e3f271c85976f27e3ad780f9ecbc8 (patch)
tree39e320b983592e713802d6e55bc31296b23673d5 /experiments/babyai_shared_smoke.py
parentc18e939242d6fd88fc07aa35c4d11d17d663fce0 (diff)
exp: freeze BabyAI shared-feedback protocol
Diffstat (limited to 'experiments/babyai_shared_smoke.py')
-rw-r--r--experiments/babyai_shared_smoke.py150
1 files changed, 150 insertions, 0 deletions
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()