Projector-Gap Trust Region for Shared Updates / experiment.py
Mechanism failed
1import json
2from pathlib import Path
3import numpy as np
4
5SEED=3157
6
7def projector(A,horizon=6):
8 if not np.all(np.isfinite(A)): return None
9 O=np.vstack([np.linalg.matrix_power(A,t) for t in range(horizon)])
10 if not np.all(np.isfinite(O)): return None
11 Q,_=np.linalg.qr(O,mode='reduced')
12 return Q@Q.T
13
14def opnorm(M):
15 if M is None or not np.all(np.isfinite(M)): return float('inf')
16 try: return float(np.linalg.svd(M,compute_uv=False)[0])
17 except np.linalg.LinAlgError: return float('inf')
18
19def gap(A,B,horizon=6):
20 P,Q=projector(A,horizon),projector(B,horizon)
21 return opnorm(None if P is None or Q is None else P-Q)
22
23def make_problem(seed=SEED,d=3,n=4):
24 rng=np.random.default_rng(seed); targets=[]
25 for _ in range(n):
26 X=rng.normal(size=(d,d)); X=X/(1.8*max(1.,opnorm(X)))
27 targets.append(X+.12*np.diag(rng.normal(size=d)))
28 return np.array([T+.20*rng.normal(size=T.shape) for T in targets]),np.array(targets)
29
30def fd_check():
31 rng=np.random.default_rng(11); A=rng.normal(size=(3,3)); A=A/(2*opnorm(A)); D=rng.normal(size=(3,3)); D=D/opnorm(D)
32 P=projector(A); href=1e-7; ref=(projector(A+href*D)-P)/href; rows=[]
33 for h in [1e-1,3e-2,1e-2,3e-3,1e-3,3e-4]: rows.append((h,opnorm((projector(A+h*D)-P)/h-ref)))
34 B=A+.03*rng.normal(size=A.shape); Delta=.04*D
35 pred=opnorm((P+(projector(A+Delta)-P))-projector(B)); actual=gap(A+Delta,B)
36 return rows,pred,actual
37
38def run(trust,seed,lr,epsilon,steps=100,horizon=6):
39 A,T=make_problem(seed); initial=A.copy(); losses=[]; gaps=[]; alphas=[]; spikes=0; diverged=False
40 for step in range(steps):
41 if not np.all(np.isfinite(A)) or np.max(np.abs(A))>1e100: diverged=True; break
42 grads=2*(A-T); proposed=np.zeros_like(A)
43 for inds in ([0,1],[2,3]):
44 D=-lr*np.mean(grads[list(inds)],axis=0)
45 for i in inds: proposed[i]=D
46 aa=[]
47 for inds in ([0,1],[2,3]):
48 inds=list(inds); leader=inds[0]
49 p=[projector(A[i]+proposed[i],horizon) for i in inds]
50 mg=opnorm(None if p[0] is None or p[1] is None else p[1]-p[0])
51 alpha=min(1.,epsilon/(mg+1e-12)) if trust else 1.
52 if trust:
53 while alpha>1e-6 and gap(A[inds[1]]+alpha*proposed[inds[1]],A[leader]+alpha*proposed[leader],horizon)>epsilon: alpha*=.5
54 aa.append(alpha)
55 for i in inds: A[i]+=alpha*proposed[i]
56 g=max(gap(A[1],A[0],horizon),gap(A[3],A[2],horizon))
57 loss=float(np.mean((A-T)**2)) if np.all(np.isfinite(A)) else float('inf')
58 if not np.isfinite(loss): diverged=True; break
59 if step and loss>1.5*losses[-1]: spikes+=1
60 losses.append(loss); gaps.append(g); alphas.append(min(aa))
61 return dict(final_loss=losses[-1] if losses else float('inf'),min_loss=min(losses) if losses else float('inf'),max_gap=max(gaps) if gaps else float('inf'),mean_gap=np.mean(gaps) if gaps else float('inf'),clipped=sum(a<.999999 for a in alphas),mean_alpha=np.mean(alphas) if alphas else 0.,spikes=spikes,diverged=diverged,initial_gap=max(gap(initial[1],initial[0],horizon),gap(initial[3],initial[2],horizon)))
62
63def summarize(vals):
64 keys=['final_loss','min_loss','max_gap','mean_gap','clipped','mean_alpha','spikes','diverged','initial_gap']
65 return {k:[float(v[k]) for v in vals] for k in keys}
66
67def main():
68 fd,pred,actual=fd_check(); print('FINITE_DIFFERENCE')
69 for h,e in fd: print(f'h={h:.1e} error={e:.6e}')
70 print(f'PREDICTED_ACTUAL_GAP pred={pred:.8f} actual={actual:.8f} abs_error={abs(pred-actual):.3e}')
71 out={'fd':fd,'predicted_gap':pred,'actual_gap':actual,'runs':{}}
72 for eps in [.90,.95,1.00]:
73 for lr in [.8,1.0,1.2]:
74 for trust in [False,True]:
75 vals=[run(trust,s,lr,eps) for s in [3157,3158,3159,3160,3161]]
76 key=f"{'trust' if trust else 'baseline'}_eps{eps}_lr{lr}"; out['runs'][key]=summarize(vals)
77 print('\n'+key)
78 for k in ['final_loss','min_loss','max_gap','mean_alpha','clipped','spikes','diverged']:
79 x=np.array([v[k] for v in vals]); print(f'{k}: mean={np.mean(x):.6g} std={np.std(x):.6g}')
80 Path('results.json').write_text(json.dumps(out,indent=2))
81if __name__=='__main__': main()