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 --- sdil/babyai_shared.py | 84 ++++++++++++++++++++++++++++++++++++++++++++++++--- 1 file changed, 80 insertions(+), 4 deletions(-) (limited to 'sdil') diff --git a/sdil/babyai_shared.py b/sdil/babyai_shared.py index a4c871a..209ea36 100644 --- a/sdil/babyai_shared.py +++ b/sdil/babyai_shared.py @@ -175,6 +175,16 @@ def visual_input_dim(object_cardinality, color_cardinality, * (object_cardinality + color_cardinality + state_cardinality) + 4) +def history_input_dim(object_cardinality, color_cardinality, + state_cardinality, history_steps, action_dim=7, + view_size=7): + if history_steps < 1: + raise ValueError("history_steps must be positive") + frame_dim = visual_input_dim( + object_cardinality, color_cardinality, state_cardinality, view_size) + return history_steps * (frame_dim + action_dim) + + def encode_visual(images, directions, object_cardinality, color_cardinality, state_cardinality, device, dtype=torch.float32): @@ -194,6 +204,62 @@ def encode_visual(images, directions, object_cardinality, return torch.cat((grid, direction), dim=1) +def build_history_index(episode_offsets, actions, history_steps): + """Build right-aligned history indices without crossing episodes.""" + if history_steps < 1: + raise ValueError("history_steps must be positive") + episode_offsets = np.asarray(episode_offsets, dtype=np.int64) + actions = np.asarray(actions, dtype=np.int64) + total = len(actions) + episode_start = np.empty(total, dtype=np.int64) + for start, stop in zip(episode_offsets[:-1], episode_offsets[1:]): + episode_start[start:stop] = start + positions = np.arange(total, dtype=np.int64) + indices = np.zeros((total, history_steps), dtype=np.int64) + previous_actions = np.full((total, history_steps), -1, dtype=np.int64) + mask = np.zeros((total, history_steps), dtype=bool) + for lag in range(history_steps): + column = history_steps - 1 - lag + source = positions - lag + valid = source >= episode_start + indices[valid, column] = source[valid] + mask[valid, column] = True + has_previous = valid & (source > episode_start) + previous_actions[has_previous, column] = actions[ + source[has_previous] - 1] + return indices, previous_actions, mask + + +def encode_history_visual(images, directions, previous_actions, mask, + object_cardinality, color_cardinality, + state_cardinality, action_dim, device, + dtype=torch.float32): + """Encode a padded sequence of symbolic frames and preceding actions.""" + images = np.asarray(images) + directions = np.asarray(directions) + previous_actions = np.asarray(previous_actions) + mask = np.asarray(mask) + if images.ndim != 5: + raise ValueError("history images must have shape [batch, time, H, W, 3]") + batch, history_steps = images.shape[:2] + frames = encode_visual( + images.reshape(batch * history_steps, *images.shape[2:]), + directions.reshape(batch * history_steps), object_cardinality, + color_cardinality, state_cardinality, device=device, dtype=dtype) + frame_dim = frames.shape[1] + mask_tensor = torch.as_tensor(mask, device=device, dtype=dtype) + frames = frames.reshape(batch, history_steps, frame_dim) + frames = frames * mask_tensor[:, :, None] + action_values = torch.as_tensor( + np.maximum(previous_actions, 0), device=device, dtype=torch.long) + action_valid = torch.as_tensor( + (previous_actions >= 0) & mask, device=device, dtype=dtype) + action_features = F.one_hot( + action_values, num_classes=action_dim).to(dtype) + action_features = action_features * action_valid[:, :, None] + return torch.cat((frames, action_features), dim=2).reshape(batch, -1) + + def select_teaching_signal(net, layer, instruction, context, soma, condition): raw = instruction + context innovation = raw - net.predictor(layer, soma) @@ -276,16 +342,26 @@ def manual_step(net, features, missions, actions, condition): @torch.no_grad() def evaluate_actions(net, images, directions, missions, actions, - cardinalities, batch_size=1024, context_enabled=True): + cardinalities, batch_size=1024, context_enabled=True, + history=None): 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) + if history is None: + features = encode_visual( + images[start:stop], directions[start:stop], *cardinalities, + device=device, dtype=net.dtype) + else: + indices, previous_actions, mask = history + batch_indices = indices[start:stop] + features = encode_history_visual( + images[batch_indices], directions[batch_indices], + previous_actions[start:stop], mask[start:stop], + *cardinalities, net.config.action_dim, + device=device, dtype=net.dtype) mission_batch = torch.as_tensor( missions[start:stop], device=device, dtype=net.dtype) action_batch = torch.as_tensor( -- cgit v1.2.3