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