MI-Guided Latent Protection / mi_latent_protection.py
Failed on benchmark
1"""MVP for MI-guided latent protection.
2
3Run with the configured Python interpreter. The toy checks are deliberately
4first and print predicted versus observed quantities before the learning test.
5"""
6import json, math, random
7from pathlib import Path
8import numpy as np
9
10SEED = 2392
11
12def seed_all(seed=SEED):
13 random.seed(seed); np.random.seed(seed)
14 try:
15 import torch
16 torch.manual_seed(seed)
17 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
18 except Exception:
19 pass
20
21def allocate_variance(scores, mean_var=0.25, delta=1e-8):
22 scores = np.asarray(scores, dtype=np.float64)
23 q = (scores + delta) / (scores.mean() + delta)
24 raw = 1.0 / q
25 var = mean_var * len(scores) * raw / raw.sum()
26 return q, var
27
28def math_checks():
29 # Prediction 1: normalization preserves exactly K * average variance.
30 rng = np.random.default_rng(SEED)
31 budget_errors = []
32 for _ in range(200):
33 s = np.exp(rng.normal(0, 2, 16))
34 _, v = allocate_variance(s, .37)
35 budget_errors.append(abs(v.mean() - .37))
36 p1 = {"predicted_mean_variance": .37,
37 "observed_mean_variance": float(np.mean([np.mean(allocate_variance(np.exp(rng.normal(0,2,16)), .37)[1]) for _ in range(200)])),
38 "max_abs_sweep_error": float(max(budget_errors))}
39
40 # Prediction 2: a high-score coordinate gets inverse score-ratio variance.
41 ratios = np.array([1, 2, 4, 8, 16, 32], dtype=float)
42 observed = []
43 for r in ratios:
44 _, v = allocate_variance([1, r], .5)
45 observed.append(v[1] / v[0])
46 predicted = 1.0 / ratios
47 p2 = {"score_ratios": ratios.tolist(), "predicted_high_to_low_variance": predicted.tolist(),
48 "observed_high_to_low_variance": observed,
49 "max_abs_error": float(np.max(np.abs(np.asarray(observed)-predicted)))}
50
51 # Prediction 3: if local task-noise sensitivity equals s, the quadratic
52 # penalty ratio is K^2/(sum(s)*sum(1/s)); it decreases as heterogeneity grows.
53 penalty_ratios=[]
54 predicted_penalty=[]
55 for r in ratios:
56 s=np.array([1., r])
57 _, v=allocate_variance(s, .5)
58 adaptive=float(np.sum(s*v)); uniform=float(np.sum(s*.5))
59 penalty_ratios.append(adaptive/uniform)
60 predicted_penalty.append(4/(np.sum(s)*np.sum(1/s)))
61 p3={"score_ratios":ratios.tolist(), "predicted_adaptive_over_uniform_penalty":predicted_penalty,
62 "observed_adaptive_over_uniform_penalty":penalty_ratios,
63 "max_abs_error":float(np.max(np.abs(np.asarray(penalty_ratios)-predicted_penalty)))}
64 return {"budget_conservation":p1,"inverse_scaling":p2,"quadratic_penalty":p3}
65
66def learning_experiment(steps=700):
67 try:
68 import torch
69 from torch import nn
70 device=torch.device("cuda" if torch.cuda.is_available() else "cpu")
71 torch.set_num_threads(4)
72 seed_all(SEED)
73 n_train,n_test=6000,2500; k=8; batch=128
74 # First two coordinates contain the target; remaining dimensions are nuisance.
75 g=torch.Generator().manual_seed(SEED)
76 xtr=torch.randn(n_train,k,generator=g); xte=torch.randn(n_test,k,generator=g)
77 ytr=((xtr[:,0]+0.8*xtr[:,1]+0.35*torch.randn(n_train,generator=g))>0).float()
78 yte=((xte[:,0]+0.8*xte[:,1]+0.35*torch.randn(n_test,generator=g))>0).float()
79 def run(mode):
80 seed_all(SEED+({"uniform":0,"mi":1,"random":2,"magnitude":3}[mode]))
81 model=nn.Linear(k,1).to(device); opt=torch.optim.Adam(model.parameters(),lr=.01)
82 score_ema=torch.ones(k,device=device); random_scores=torch.rand(k,device=device); mean_var=.45; beta=.90
83 perm=torch.arange(n_train)
84 for step in range(steps):
85 if step% (n_train//batch)==0: perm=torch.randperm(n_train)
86 ix=perm[(step*batch)%n_train:((step+1)*batch)%n_train]
87 z=xtr[ix].to(device); y=ytr[ix].to(device)
88 if mode=="uniform": var=torch.full((k,),mean_var,device=device)
89 else:
90 # Sensitivity proxy: gradient of a detached supervised MI-style
91 # critic loss with respect to z; scores do not train the model.
92 zz=z.detach().requires_grad_(True)
93 critic=nn.functional.binary_cross_entropy_with_logits(model(zz).squeeze(1),y)
94 grad=torch.autograd.grad(critic,zz)[0].abs().mean(0).detach()
95 if mode=="mi": s=beta*score_ema+(1-beta)*grad; score_ema=s
96 elif mode=="random": s=random_scores
97 else: s=z.abs().mean(0)
98 # Stabilize noisy minibatch sensitivities: retain a 10% floor.
99 s=torch.nan_to_num(s, nan=1.0, posinf=1.0, neginf=0.0)
100 s=s + 0.10*s.mean()
101 s=torch.nan_to_num(s, nan=1.0, posinf=1.0, neginf=1.0).clamp_min(1e-6)
102 q=(s+1e-6)/(s.mean()+1e-6); raw=1/q
103 var=mean_var*k*raw/raw.sum()
104 noisy=z+torch.randn_like(z)*var.sqrt()
105 loss=nn.functional.binary_cross_entropy_with_logits(model(noisy).squeeze(1),y)
106 opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0); opt.step()
107 z=xte.to(device)
108 if mode=="uniform":
109 var=torch.full((k,),mean_var,device=device)
110 else:
111 # Sensitivity needs autograd; only this small block is enabled.
112 zz=z[:batch].detach().requires_grad_(True)
113 yy=yte[:batch].to(device)
114 cr=nn.functional.binary_cross_entropy_with_logits(model(zz).squeeze(1),yy)
115 gr=torch.autograd.grad(cr,zz)[0].abs().mean(0).detach()
116 if mode=="mi": s=score_ema
117 elif mode=="random": s=random_scores
118 else: s=z.abs().mean(0)
119 s=torch.nan_to_num(s, nan=1.0, posinf=1.0, neginf=0.0)
120 s=s + 0.10*s.mean()
121 s=torch.nan_to_num(s, nan=1.0, posinf=1.0, neginf=1.0).clamp_min(1e-6)
122 q=(s+1e-6)/(s.mean()+1e-6); var=mean_var*k*(1/q)/(1/q).sum()
123 with torch.no_grad():
124 logits=model(z+torch.randn_like(z)*var.sqrt()).squeeze(1)
125 acc=((logits>0)==(yte.to(device)>0.5)).float().mean().item()
126 clean=((model(z).squeeze(1)>0)==(yte.to(device)>0.5)).float().mean().item()
127 return acc,clean,float(var.mean().item())
128 result={m:run(m) for m in ["uniform","mi","random","magnitude"]}
129 return {"device":str(device),"steps":steps,"results":result}
130 except Exception as e:
131 # Required safe fallback if CUDA/runtime setup fails.
132 return {"error":repr(e),"fallback_note":"analytic checks completed; learning run unavailable"}
133
134def main():
135 seed_all()
136 out={"math_checks":math_checks(),"learning":learning_experiment()}
137 Path("results.json").write_text(json.dumps(out,indent=2))
138 print(json.dumps(out,indent=2))
139
140if __name__ == '__main__': main()