#!/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_history_index, build_vocabulary, encode_history_visual, fit_population_predictor, history_input_dim, manual_step, missions_to_bow, population_predictor_metrics, select_teaching_signal, ) from prepare_babyai_shared import generate_split from babyai_shared_run import rollout_policy 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, ) actions = torch.from_numpy(data["action"].astype("int64")) history_steps = 4 history_indices, previous_actions, history_mask = build_history_index( data["episode_offset"], data["action"], history_steps) features = encode_history_visual( data["image"][history_indices], data["direction"][history_indices], previous_actions, history_mask, *cardinalities, 7, device="cpu", dtype=torch.float64) assert history_mask[:, -1].all() assert (previous_actions[data["episode_offset"][:-1], -1] == -1).all() config = BabyAISharedConfig( input_dim=history_input_dim(*cardinalities, history_steps), 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 population_state = base.forward_features(features, missions) split = features.shape[0] // 2 coefficient, intercept = fit_population_predictor( population_state["h"][1][:split], population_state["context"][0][:split]) population_metrics = population_predictor_metrics( population_state["h"][1][split:], population_state["context"][0][split:], coefficient, intercept) assert torch.isfinite(torch.tensor( population_metrics["mean_per_cell_r2"])) # 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) rollout = rollout_policy( base, "BabyAI-GoToObjS6-v1", [50_000, 50_001], vocabulary, cardinalities, history_steps=history_steps) assert rollout["episodes"] == 2 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), "history_rollout_episodes": rollout["episodes"], "population_predictor_holdout_r2": population_metrics[ "mean_per_cell_r2"], }) if __name__ == "__main__": main()