diff options
Diffstat (limited to 'sdil')
| -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 |
