summaryrefslogtreecommitdiff
path: root/sdil
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-10 10:32:10 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-10 10:32:10 -0500
commita30c297eb8abf6f15f97dd4f35a2f530c9728100 (patch)
tree7415a8cbaf58a72412a657e5eb69f9b6ccb3bce2 /sdil
parent17fd6e1843adfa11a9958fd21e056c304805bde2 (diff)
feat: add basal history to BabyAI policy
Diffstat (limited to 'sdil')
-rw-r--r--sdil/babyai_shared.py84
1 files changed, 80 insertions, 4 deletions
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(