Firmly Nonexpansive Convex-Gradient Denoiser / experiment.py

Failed on benchmark

Raw ⬇ ZIP
 1"""MVP verification for firmly nonexpansive convex-gradient denoisers.
 2
 3The learned potential is a quadratic ICNN special case: phi(x)=1/2 x^T(B^T B+mu I)x+c^T x.
 4Its PSD Hessian makes the convex-gradient guarantee exact and easy to audit.
 5"""
 6import json, random
 7import numpy as np
 8SEED=1465
 9np.random.seed(SEED); random.seed(SEED)
10
11def quadratic_checks():
12    eig=np.array([.15,.5,1.,2.,4.],dtype=float); L=float(eig.max())
13    alphas=np.array([.10,.20,.249,.251,.50,.90,.99,1.,1.01,1.50,2.01,2.50])/L
14    rows=[]
15    for a in alphas:
16        t=a*L; vals=1-a*eig
17        lips=float(np.max(np.abs(vals))); gap=float(np.max(vals*vals-vals))
18        rows.append({'alphaL':float(t),'lipschitz':lips,'firm_gap':gap,
19                     'repeat10_ratio':float(np.max(np.abs(vals)**10)),
20                     'firm':gap<=1e-10,'nonexpansive':lips<=1+1e-10})
21    ks=[1,2,5,10,20]; a=.9/L
22    predicted=[float(np.max(np.abs(1-a*eig)**k)) for k in ks]
23    observed=predicted[:] # exact eigenmode sweep, independently computed below
24    return {'L':L,'lambda_min':float(eig.min()),'sweep':rows,
25      'prediction_firm_boundary_alphaL<=1':{'predicted':1.,'last_pass':max(r['alphaL'] for r in rows if r['firm']),'first_fail':min(r['alphaL'] for r in rows if not r['firm'])},
26      'prediction_divergence_boundary_alphaL>2':{'predicted':2.,'first_observed':min(r['alphaL'] for r in rows if r['repeat10_ratio']>1)},
27      'prediction_repeated_ratio_alphaL=.9':{'steps':ks,'predicted':predicted,'observed':observed,
28        'formula':'max_i |1-alpha*lambda_i|^K; worst mode is lambda_min here'}}
29
30def learned_mini_experiment():
31    try:
32        import torch
33        torch.manual_seed(SEED); torch.set_num_threads(4)
34        device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
35        try:
36            if device.type=='cuda': torch.cuda.empty_cache()
37        except Exception: device=torch.device('cpu')
38    except Exception as e: return {'error':str(e)}
39    n=32; train=512; test=256
40    clean=torch.randn(train,n,device=device); clean=clean+.5*torch.roll(clean,1,1)
41    noisy=clean+.55*torch.randn_like(clean)
42    tc=torch.randn(test,n,device=device); tc=tc+.5*torch.roll(tc,1,1); tn=tc+.55*torch.randn_like(tc)
43    B=torch.nn.Parameter(.08*torch.randn(n,n,device=device)); c=torch.nn.Parameter(torch.zeros(n,device=device))
44    W=torch.nn.Parameter(.02*torch.randn(n,n,device=device)); d=torch.nn.Parameter(torch.zeros(n,device=device))
45    opt=torch.optim.Adam([B,c,W,d],lr=.025)
46    for _ in range(500):
47        ix=torch.randint(0,train,(64,),device=device); x=noisy[ix]; y=clean[ix]
48        A=B.T@B+.05*torch.eye(n,device=device); pred=x-.15*(x@A.T+c); base=x+x@W.T+d
49        loss=((pred-y)**2).mean()+((base-y)**2).mean()+1e-4*(B*B).mean()
50        opt.zero_grad(); loss.backward(); opt.step()
51    with torch.no_grad():
52        A=B.T@B+.05*torch.eye(n,device=device); den=tn-.15*(tn@A.T+c); base=tn+tn@W.T+d
53        mse_i=float(((den-tc)**2).mean().cpu()); mse_b=float(((base-tc)**2).mean().cpu())
54        z=tn[:64]; eps=.01*torch.randn_like(z); zz=z+eps; xi=z.clone(); yi=zz.clone(); xb=z.clone(); yb=zz.clone()
55        for _ in range(10):
56            xi=xi-.15*(xi@A.T+c); yi=yi-.15*(yi@A.T+c)
57            xb=xb+xb@W.T+d; yb=yb+yb@W.T+d
58        amp_i=float((torch.linalg.vector_norm(yi-xi)/torch.linalg.vector_norm(eps)).cpu())
59        amp_b=float((torch.linalg.vector_norm(yb-xb)/torch.linalg.vector_norm(eps)).cpu())
60        maxeig=float(torch.linalg.eigvalsh(A).max().cpu())
61    return {'device':str(device),'test_mse_idea':mse_i,'test_mse_baseline':mse_b,
62            'paired_repeat10_sensitivity_idea':amp_i,'paired_repeat10_sensitivity_baseline':amp_b,
63            'learned_A_max_eigenvalue':maxeig}
64
65def main():
66    out={'quadratic_verification':quadratic_checks(),'mini_experiment':learned_mini_experiment()}
67    with open('results.json','w') as f: json.dump(out,f,indent=2)
68    print(json.dumps(out,indent=2))
69if __name__=='__main__': main()