diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-10 10:47:10 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-10 10:47:10 -0500 |
| commit | b5eab6cd911e074ae9268303389dfbe52e636e26 (patch) | |
| tree | a650bd25467f1484d49c9360f851361c0b038247 /sdil/babyai_shared.py | |
| parent | bc93d2a673f5f971415715dad11c5e8baf2a6202 (diff) | |
exp: freeze BabyAI population predictor screen
Diffstat (limited to 'sdil/babyai_shared.py')
| -rw-r--r-- | sdil/babyai_shared.py | 37 |
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 |
