Residual-Gated DRS Solver Layer / experiment.py

Mechanism failed

Raw ⬇ ZIP
 1import json
 2import numpy as np
 3
 4SEED = 2293
 5
 6def projK(x, identity=False): return x.copy() if identity else np.maximum(x, 0.0)
 7
 8def projector(A, b):
 9    A = np.asarray(A, float); b = np.asarray(b, float)
10    G = np.linalg.pinv(A @ A.T)
11    return lambda w: w - A.T @ (G @ (A @ w - b))
12
13def run(z0, A, b, c, gamma, beta, lam, T, identity=False):
14    P = projector(A, b); z = z0.copy(); out=[]
15    for _ in range(T):
16        u=projK(z, identity); v=P(2*u-z-gamma*beta*c)
17        out.append({'z':z.copy(),'u':u.copy(),'v':v.copy(),
18                    'r':float(np.linalg.norm(v-u)),
19                    'p':float(np.linalg.norm(A@v-b))})
20        z=z+lam*(v-u)
21    return out
22
23def controller(z0,A,b,c,gamma,T):
24    P=projector(A,b); z=z0.copy(); lam=1.; beta=1.; oldr=np.inf; oldp=np.inf; out=[]
25    for t in range(T):
26        u=projK(z); v=P(2*u-z-gamma*beta*c)
27        r=float(np.linalg.norm(v-u)); p=float(np.linalg.norm(A@v-b))
28        if np.isfinite(oldr):
29            decay=.96**t
30            lam=float(np.clip(lam*np.exp(.18*decay*np.clip(np.log((oldr+1e-12)/(r+1e-12)),-1,1)),.15,1.85))
31            if r < oldr and p <= oldp*(1+1e-6): beta=min(2.,beta*(1+.08*decay))
32            elif r > 1.1*oldr or p > 1.1*oldp:
33                beta=max(.1,beta*(1-.12*decay)); lam=max(.15,.8*lam)
34        out.append({'u':u,'v':v,'r':r,'p':p,'lam':lam,'beta':beta})
35        z=z+lam*(v-u); oldr=r; oldp=p
36    return out
37
38def affine_check():
39    rng=np.random.default_rng(SEED); A=rng.normal(size=(3,5)); b=rng.normal(size=3); P=projector(A,b)
40    w=rng.normal(size=5); x=P(w); _,_,vh=np.linalg.svd(A); null=vh[3:].T
41    return {'equality_error':float(np.linalg.norm(A@x-b)),
42            'idempotence_error':float(np.linalg.norm(P(x)-x)),
43            'nullspace_orthogonality':float(np.linalg.norm(null.T@(w-x)))}
44
45def exact_toy():
46    # K=R^n and A=I,b=0: u=z,v=0, hence z_next=(1-lambda)z.
47    n=3; A=np.eye(n); b=np.zeros(n); c=np.zeros(n); z=np.array([1.,-2.,.5]); T=12; ans=[]
48    for lam in [.25,.5,1.,1.5,1.75,2.,2.1,2.5]:
49        q=run(z,A,b,c,1.,1.,lam,T,identity=True)
50        rate=(q[-1]['r']/q[0]['r'])**(1/(T-1))
51        ans.append({'lambda':lam,'predicted_rate_abs_1_minus_lambda':abs(1-lam),
52                    'observed_rate':float(rate),'final_over_initial':float(q[-1]['r']/q[0]['r']),
53                    'contractive':bool(abs(1-lam)<1)})
54    return ans
55
56def finite_bound():
57    z=np.array([1.,-2.,.5]); D=float(z@z); N=40; ans=[]
58    for lam in [.25,.75,1.25,1.75]:
59        q=run(z,np.eye(3),np.zeros(3),np.zeros(3),1,1,lam,N,identity=True)
60        vals=np.array([x['r'] for x in q]); alpha=lam/2; bound=alpha/(1-alpha)*D/N
61        ans.append({'lambda':lam,'N_min_r2':float(N*np.min(vals**2)),
62                    'bound':float(bound),'ratio':float(N*np.min(vals**2)/bound)})
63    return ans
64
65def mini():
66    rng=np.random.default_rng(SEED); B=[]; C=[]; lams=[]; betas=[]
67    for _ in range(100):
68        n=8; A=np.ones((1,n)); b=np.ones(1); c=rng.uniform(.1,2.,n); z=rng.normal(size=n); opt=float(c.min())
69        f=run(z,A,b,c,.35,1.,1.,24)[-1]; cc=controller(z,A,b,c,.35,24); g=cc[-1]
70        def m(q):
71            u=q['u']; eq=abs(float((A@u-b)[0])); cone=float(np.linalg.norm(np.minimum(u,0)))
72            return [float(c@u-opt+10*eq),eq,cone,q['r']]
73        B.append(m(f)); C.append(m(g)); lams.append(g['lam']); betas.append(g['beta'])
74    B=np.array(B); C=np.array(C)
75    return {'trials':100,'T':24,'fixed_median_merit':float(np.median(B[:,0])),
76      'controller_median_merit':float(np.median(C[:,0])),'fixed_median_eq_violation':float(np.median(B[:,1])),
77      'controller_median_eq_violation':float(np.median(C[:,1])),'fixed_median_residual':float(np.median(B[:,3])),
78      'controller_median_residual':float(np.median(C[:,3])),'controller_final_lambda_median':float(np.median(lams)),
79      'controller_final_beta_median':float(np.median(betas)),'residual_win_fraction':float(np.mean(C[:,3]<B[:,3]))}
80
81def main():
82    print(json.dumps({'affine_projection':affine_check(),'exact_toy':exact_toy(),
83                      'finite_step_bound':finite_bound(),'mini_experiment':mini()},indent=2))
84if __name__=='__main__': main()