diff options
Diffstat (limited to 'experiments/babyai_shared_smoke.py')
| -rw-r--r-- | experiments/babyai_shared_smoke.py | 26 |
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"], }) |
