From a30c297eb8abf6f15f97dd4f35a2f530c9728100 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Mon, 10 Aug 2026 10:32:10 -0500 Subject: feat: add basal history to BabyAI policy --- experiments/babyai_shared_run.py | 120 ++++++++++++++++++++++++++++++--------- 1 file changed, 93 insertions(+), 27 deletions(-) (limited to 'experiments/babyai_shared_run.py') diff --git a/experiments/babyai_shared_run.py b/experiments/babyai_shared_run.py index 6b6ddd2..0b7c72a 100644 --- a/experiments/babyai_shared_run.py +++ b/experiments/babyai_shared_run.py @@ -16,8 +16,9 @@ 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, + BabyAISharedConfig, BabyAISharedNet, CONDITIONS, build_history_index, + encode_history_visual, encode_visual, evaluate_actions, history_input_dim, + manual_step, missions_to_bow, visual_input_dim, ) @@ -35,16 +36,16 @@ 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")} + "image", "direction", "mission_bow", "action", "episode_offset")} validation = {name: archive[f"validation_{name}"] for name in ( - "image", "direction", "mission_bow", "action")} + "image", "direction", "mission_bow", "action", "episode_offset")} 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): + context_enabled=True, chunk_size=64, history_steps=1): successes = 0 returns = [] lengths = [] @@ -52,6 +53,8 @@ def rollout_policy(net, env_id, seeds, vocabulary, cardinalities, chunk = seeds[chunk_start:chunk_start + chunk_size] envs = [gym.make(env_id) for _ in chunk] observations = [] + observation_histories = [] + previous_action_histories = [] active = [] chunk_returns = [0.0] * len(envs) chunk_lengths = [0] * len(envs) @@ -59,20 +62,51 @@ def rollout_policy(net, env_id, seeds, vocabulary, cardinalities, for index, (env, seed) in enumerate(zip(envs, chunk)): observation, _ = env.reset(seed=int(seed)) observations.append(observation) + observation_histories.append([observation]) + previous_action_histories.append([-1]) 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) + if history_steps == 1: + 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) + else: + image_shape = observations[active[0]]["image"].shape + images = np.zeros( + (len(active), history_steps, *image_shape), + dtype=np.uint8) + directions = np.zeros( + (len(active), history_steps), dtype=np.uint8) + previous_actions = np.full( + (len(active), history_steps), -1, dtype=np.int64) + history_mask = np.zeros( + (len(active), history_steps), dtype=bool) + for row, index in enumerate(active): + obs_history = observation_histories[index][-history_steps:] + action_history = previous_action_histories[index][ + -history_steps:] + offset = history_steps - len(obs_history) + images[row, offset:] = np.asarray([ + value["image"] for value in obs_history]) + directions[row, offset:] = np.asarray([ + value["direction"] for value in obs_history]) + previous_actions[row, offset:] = action_history + history_mask[row, offset:] = True 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) + if history_steps == 1: + features = encode_visual( + images, directions, *cardinalities, + device=net.device, dtype=net.dtype) + else: + features = encode_history_visual( + images, directions, previous_actions, history_mask, + *cardinalities, net.config.action_dim, + device=net.device, dtype=net.dtype) mission_tensor = torch.as_tensor( mission_bow, device=net.device, dtype=net.dtype) actions = net.forward_features( @@ -84,6 +118,9 @@ def rollout_policy(net, env_id, seeds, vocabulary, cardinalities, observation, reward, terminated, truncated, _ = ( envs[index].step(int(actions[position]))) observations[index] = observation + observation_histories[index].append(observation) + previous_action_histories[index].append( + int(actions[position])) chunk_returns[index] += float(reward) chunk_lengths[index] += 1 if terminated or truncated: @@ -104,6 +141,21 @@ def rollout_policy(net, env_id, seeds, vocabulary, cardinalities, } +def encode_compact_batch(split, indices, cardinalities, net, history): + if history is None: + return encode_visual( + split["image"][indices], split["direction"][indices], + *cardinalities, device=net.device, dtype=net.dtype) + history_indices, previous_actions, mask = history + selected_history = history_indices[indices] + return encode_history_visual( + split["image"][selected_history], + split["direction"][selected_history], + previous_actions[indices], mask[indices], + *cardinalities, net.config.action_dim, + device=net.device, dtype=net.dtype) + + def main(): parser = argparse.ArgumentParser() parser.add_argument("--data", type=Path, default=DEFAULT_DATA) @@ -115,8 +167,10 @@ def main(): 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("--history-steps", type=int, default=1) parser.add_argument("--model-seed", type=int, default=4101) parser.add_argument("--shuffle-seed", type=int, default=4101) + parser.add_argument("--stage") parser.add_argument("--device", default="cuda") parser.add_argument("--out", type=Path, required=True) args = parser.parse_args() @@ -138,13 +192,24 @@ def main(): cardinalities = ( metadata["object_cardinality"], metadata["color_cardinality"], metadata["state_cardinality"]) + if args.history_steps < 1: + raise ValueError("history steps must be positive") + input_dim = (visual_input_dim(*cardinalities) + if args.history_steps == 1 else history_input_dim( + *cardinalities, args.history_steps)) config = BabyAISharedConfig( - input_dim=visual_input_dim(*cardinalities), + input_dim=input_dim, 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) + train_history = (None if args.history_steps == 1 else build_history_index( + train["episode_offset"], train["action"], args.history_steps)) + validation_history = ( + None if args.history_steps == 1 else build_history_index( + validation["episode_offset"], validation["action"], + args.history_steps)) 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) @@ -155,10 +220,8 @@ def main(): 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_features = encode_compact_batch( + train, neutral_indices, cardinalities, net, train_history) neutral_missions = torch.as_tensor( train["mission_bow"][neutral_indices], device=device, dtype=net.dtype) @@ -169,9 +232,8 @@ def main(): 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) + features = encode_compact_batch( + train, indices, cardinalities, net, train_history) missions = torch.as_tensor( train["mission_bow"][indices], device=device, dtype=net.dtype) @@ -189,23 +251,26 @@ def main(): training_seconds = time.time() - started validation_metrics = evaluate_actions( net, validation["image"], validation["direction"], - validation["mission_bow"], validation["action"], cardinalities) + validation["mission_bow"], validation["action"], cardinalities, + history=validation_history) lesion_validation_metrics = evaluate_actions( net, validation["image"], validation["direction"], validation["mission_bow"], validation["action"], cardinalities, - context_enabled=False) + context_enabled=False, history=validation_history) rollout = rollout_policy( net, metadata["env_id"], rollout_seeds, metadata["vocabulary"], - cardinalities) + cardinalities, history_steps=args.history_steps) lesion_rollout = rollout_policy( net, metadata["env_id"], rollout_seeds, metadata["vocabulary"], - cardinalities, context_enabled=False) + cardinalities, context_enabled=False, + history_steps=args.history_steps) 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", + "stage": (args.stage or ( + "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, @@ -218,6 +283,7 @@ def main(): "shuffle_seed": args.shuffle_seed, "neutral_examples_per_epoch": ( neutral_count if args.condition == "sdil" else 0), + "history_steps": args.history_steps, "parameter_count_including_reciprocal_context_predictor": ( parameter_count), "training_wall_seconds": training_seconds, -- cgit v1.2.3