summaryrefslogtreecommitdiff
path: root/sdil/babyai_shared.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-10 10:47:10 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-10 10:47:10 -0500
commitb5eab6cd911e074ae9268303389dfbe52e636e26 (patch)
treea650bd25467f1484d49c9360f851361c0b038247 /sdil/babyai_shared.py
parentbc93d2a673f5f971415715dad11c5e8baf2a6202 (diff)
exp: freeze BabyAI population predictor screen
Diffstat (limited to 'sdil/babyai_shared.py')
-rw-r--r--sdil/babyai_shared.py37
1 files changed, 37 insertions, 0 deletions
diff --git a/sdil/babyai_shared.py b/sdil/babyai_shared.py
index 209ea36..238a210 100644
--- a/sdil/babyai_shared.py
+++ b/sdil/babyai_shared.py
@@ -169,6 +169,43 @@ class BabyAISharedNet:
return reports
+@torch.no_grad()
+def fit_population_predictor(soma, target, ridge=1e-3):
+ """Fit target ~= soma @ coefficient + intercept for a capacity audit."""
+ soma_mean = soma.mean(0)
+ target_mean = target.mean(0)
+ centered_soma = soma - soma_mean
+ centered_target = target - target_mean
+ gram = centered_soma.t() @ centered_soma / soma.shape[0]
+ cross = centered_soma.t() @ centered_target / soma.shape[0]
+ scale = gram.diagonal().mean().clamp_min(1e-8)
+ regularized = gram + ridge * scale * torch.eye(
+ gram.shape[0], device=gram.device, dtype=gram.dtype)
+ coefficient = torch.linalg.solve(regularized, cross)
+ intercept = target_mean - soma_mean @ coefficient
+ return coefficient, intercept
+
+
+@torch.no_grad()
+def population_predictor_metrics(soma, target, coefficient, intercept):
+ prediction = soma @ coefficient + intercept
+ residual = target - prediction
+ centered_target = target - target.mean(0)
+ target_ss = centered_target.square().sum(0)
+ residual_ss = residual.square().sum(0)
+ valid = target_ss > 1e-12
+ r2 = 1.0 - residual_ss[valid] / target_ss[valid]
+ context_rms = target.square().mean().sqrt()
+ return {
+ "mean_per_cell_r2": float(r2.mean()) if r2.numel() else 0.0,
+ "context_rms": float(context_rms),
+ "residual_context_rms_ratio": float(
+ residual.square().mean().sqrt()
+ / context_rms.clamp_min(1e-12)),
+ "observations": int(soma.shape[0]),
+ }
+
+
def visual_input_dim(object_cardinality, color_cardinality,
state_cardinality, view_size=7):
return (view_size * view_size