summaryrefslogtreecommitdiff
path: root/experiments/babyai_shared_smoke.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/babyai_shared_smoke.py')
-rw-r--r--experiments/babyai_shared_smoke.py15
1 files changed, 14 insertions, 1 deletions
diff --git a/experiments/babyai_shared_smoke.py b/experiments/babyai_shared_smoke.py
index 16d3e90..1b50a67 100644
--- a/experiments/babyai_shared_smoke.py
+++ b/experiments/babyai_shared_smoke.py
@@ -11,7 +11,8 @@ 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_history_index, build_vocabulary,
- encode_history_visual, history_input_dim, manual_step, missions_to_bow,
+ encode_history_visual, fit_population_predictor, history_input_dim,
+ manual_step, missions_to_bow, population_predictor_metrics,
select_teaching_signal,
)
from prepare_babyai_shared import generate_split
@@ -79,6 +80,16 @@ def main():
innovation_identity_error = float((used - innovation).abs().max())
assert raw_identity_error < 1e-14
assert innovation_identity_error == 0.0
+ population_state = base.forward_features(features, missions)
+ split = features.shape[0] // 2
+ coefficient, intercept = fit_population_predictor(
+ population_state["h"][1][:split],
+ population_state["context"][0][:split])
+ population_metrics = population_predictor_metrics(
+ population_state["h"][1][split:],
+ population_state["context"][0][split:], coefficient, intercept)
+ assert torch.isfinite(torch.tensor(
+ population_metrics["mean_per_cell_r2"]))
# The manual BP direction must exactly match autograd on the same fixed
# mission-conditioned network.
@@ -157,6 +168,8 @@ def main():
"predictor_teaching_observations": max(
row["teaching_observations"] for row in predictor_reports),
"history_rollout_episodes": rollout["episodes"],
+ "population_predictor_holdout_r2": population_metrics[
+ "mean_per_cell_r2"],
})