Curvature-Guided Discrepancy Gradient Accumulation / experiment.py

Mechanism failed

Raw ⬇ ZIP
 1import json, random
 2import numpy as np
 3SEED=609
 4np.random.seed(SEED); random.seed(SEED)
 5
 6def choose_signs(cands,b,g,rho=0.05):
 7    n=cands.shape[0]; signs=np.ones(n); target=rho*float(g@g)
 8    def obj(s): return float(np.max(np.abs(b+s@cands)))
 9    def ok(s): return float(g@(s@cands))>=target-1e-12
10    if n<=8:
11        best=None
12        for mask in range(1<<n):
13            s=np.array([1. if (mask>>j)&1 else -1. for j in range(n)])
14            if ok(s):
15                v=obj(s)
16                if best is None or v<best[0]: best=(v,s)
17        if best is not None: return best[1]
18    if not ok(signs): return signs
19    cur=obj(signs)
20    for _ in range(3):
21        changed=False
22        for j in range(n):
23            q=signs.copy(); q[j]*=-1; v=obj(q)
24            if ok(q) and v<cur-1e-12: signs,cur,changed=q,v,True
25        if not changed: break
26    return signs
27
28def toy_scaling():
29    ms=[4,8,16,32,64]; T=150; trials=8; rows=[]
30    for m in ms:
31        vals={k:[] for k in ('guided','all_plus','random')}
32        for tr in range(trials):
33            bs={k:np.zeros(m) for k in vals}
34            for t in range(T):
35                a=np.random.RandomState(SEED+tr*1000+t+m).randn(m-1,m)
36                a/=np.maximum(np.linalg.norm(a,axis=1,keepdims=True),1e-12); g=a.mean(0)
37                bs['guided']+=choose_signs(a,bs['guided'],g)@a
38                bs['all_plus']+=a.sum(0); bs['random']+=np.random.choice([-1.,1.],m-1)@a
39            for k in vals: vals[k].append(np.max(np.abs(bs[k])))
40        rows.append({'m':m,**{k:float(np.median(v)) for k,v in vals.items()}})
41    m=32; checkpoints=np.array([15,30,60,120,150]); paths={k:[] for k in ('guided','all_plus','random')}
42    for tr in range(10):
43        bs={k:np.zeros(m) for k in paths}
44        for t in range(1,T+1):
45            a=np.random.RandomState(SEED+5000+tr*1000+t).randn(m-1,m); a/=np.linalg.norm(a,axis=1,keepdims=True); g=a.mean(0)
46            bs['guided']+=choose_signs(a,bs['guided'],g)@a; bs['all_plus']+=a.sum(0); bs['random']+=np.random.choice([-1.,1.],m-1)@a
47            if t in checkpoints:
48                for k in paths: paths[k].append((t,np.max(np.abs(bs[k]))))
49    exponents={k:float(np.polyfit(np.log([z[0] for z in paths[k]]),np.log([z[1] for z in paths[k]]),1)[0]) for k in paths}
50    return {'rows':rows,'growth_exponents_m32':exponents}
51
52def neural_run(guided,steps=80,K=8,n=4):
53    rng=np.random.RandomState(SEED+7); N=1024; X=rng.randn(N,2).astype('float32'); y=(X[:,0]*X[:,1]>0).astype('int64')
54    p=[rng.randn(2,16).astype('float32')*.4,np.zeros(16,'float32'),rng.randn(16,2).astype('float32')*.2,np.zeros(2,'float32')]
55    def grad(x,yy):
56        w,bb,v,cc=p; h=np.maximum(0,x@w+bb); z=h@v+cc; z-=z.max(1,keepdims=True); pr=np.exp(z); pr/=pr.sum(1,keepdims=True); dz=(pr-np.eye(2,dtype='float32')[yy])/len(x)
57        gv=h.T@dz; gc=dz.sum(0); dh=(dz@v.T)*(h>0); return np.r_[(x.T@dh).ravel(),dh.sum(0),gv.ravel(),gc]
58    b=np.zeros(sum(q.size for q in p)); peaks=[]; clips=0
59    for t in range(steps):
60        ids=rng.choice(N,n*32,False); cs=np.stack([grad(X[ids[j*32:(j+1)*32]],y[ids[j*32:(j+1)*32]]) for j in range(n)]); g=cs.mean(0)
61        s=choose_signs(cs,b,g) if guided else np.ones(n); b+=s@cs; peaks.append(np.max(np.abs(b)))
62        if (t+1)%K==0:
63            u=b/K
64            if np.linalg.norm(u)>2: clips+=1; u*=2/np.linalg.norm(u)
65            off=0
66            for q in p: q-=.25*u[off:off+q.size].reshape(q.shape); off+=q.size
67            b*=0
68    z=np.maximum(0,X@p[0]+p[1])@p[2]+p[3]; loss=float(np.mean(np.logaddexp(0,z.max(1))-z[np.arange(N),y])); acc=float(np.mean(z.argmax(1)==y))
69    return {'peak_residual_median':float(np.median(peaks)),'peak_residual_max':float(max(peaks)),'clip_events':clips,'final_accuracy':acc,'final_loss':loss}
70
71def main():
72    out={'seed':SEED,'toy':toy_scaling(),'neural':{'ordinary':neural_run(False),'guided':neural_run(True)}}
73    open('results.json','w').write(json.dumps(out,indent=2)); print(json.dumps(out,indent=2))
74if __name__=='__main__':main()