summaryrefslogtreecommitdiff
path: root/experiments/babyai_shared_smoke.py
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 /experiments/babyai_shared_smoke.py
parent17fd6e1843adfa11a9958fd21e056c304805bde2 (diff)
feat: add basal history to BabyAI policy
Diffstat (limited to 'experiments/babyai_shared_smoke.py')
-rw-r--r--experiments/babyai_shared_smoke.py26
1 files changed, 20 insertions, 6 deletions
diff --git a/experiments/babyai_shared_smoke.py b/experiments/babyai_shared_smoke.py
index e5bbb55..16d3e90 100644
--- a/experiments/babyai_shared_smoke.py
+++ b/experiments/babyai_shared_smoke.py
@@ -10,10 +10,12 @@ 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,
+ BabyAISharedConfig, BabyAISharedNet, build_history_index, build_vocabulary,
+ encode_history_visual, history_input_dim, manual_step, missions_to_bow,
+ select_teaching_signal,
)
from prepare_babyai_shared import generate_split
+from babyai_shared_run import rollout_policy
def maximum_difference(left, right):
@@ -31,12 +33,18 @@ def main():
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"))
+ 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=visual_input_dim(*cardinalities),
+ 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)
@@ -129,6 +137,11 @@ def main():
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),
@@ -143,6 +156,7 @@ def main():
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"],
})