From 6d99f411cf7e3f271c85976f27e3ad780f9ecbc8 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Mon, 10 Aug 2026 10:13:57 -0500 Subject: exp: freeze BabyAI shared-feedback protocol --- experiments/babyai_shared_smoke.py | 150 +++++++++++++++++++++++++++++++++++++ 1 file changed, 150 insertions(+) create mode 100644 experiments/babyai_shared_smoke.py (limited to 'experiments/babyai_shared_smoke.py') 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() -- cgit v1.2.3