Zero-Augmented Double-Scoring / run_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import json, math, random
 2from pathlib import Path
 3import numpy as np
 4
 5SEED=1363
 6
 7def seed(s=SEED):
 8    random.seed(s); np.random.seed(s)
 9
10
11def topk_counts(M, rho, trials=20000, seed0=SEED):
12    rng=np.random.default_rng(seed0)
13    K=int(math.floor(rho*2*M))
14    counts=np.empty(trials,dtype=int)
15    for t in range(trials):
16        scores=rng.standard_normal(2*M)
17        counts[t]=np.sum(np.argpartition(scores,-K)[-K:] < M)
18    mean=K/2.0
19    var=K*(M/(2*M))*(1-M/(2*M))*((2*M-K)/(2*M-1)) if 2*M>1 else 0
20    return K, counts, mean, var
21
22
23def representation_check(M=12, K=9):
24    ok=True; examples=[]
25    for bits in range(1<<M):
26        real=np.array([(bits>>j)&1 for j in range(M)], dtype=int)
27        r=int(real.sum())
28        if r<=K:
29            aug=np.r_[real, np.zeros(M,dtype=int)]
30            aug[M:M+(K-r)]=1
31            ok &= (aug[:M].tolist()==real.tolist() and int(aug.sum())==K)
32            if r in (0,K) and len(examples)<2: examples.append((r,int(aug[M:].sum())))
33    return ok, examples
34
35
36def score_training(d=20, n=600, rho_aug=.30, steps=180, seed0=0, double=True):
37    import torch
38    torch.manual_seed(seed0); np.random.seed(seed0)
39    X=torch.randn(n,d)
40    teacher=torch.randn(d,1)
41    y=(X@teacher + .25*torch.randn(n,1)>0).float().squeeze()
42    W=torch.randn(d,1)/math.sqrt(d)
43    M=d
44    K=max(1,int(math.floor(rho_aug*(2*M if double else M))))
45    sr=torch.randn(M,requires_grad=True)
46    sd=torch.randn(M,requires_grad=True) if double else None
47    opt=torch.optim.SGD([sr]+([] if sd is None else [sd]),lr=.35)
48    for _ in range(steps):
49        opt.zero_grad()
50        sa=torch.cat([sr,sd]) if double else sr
51        hard=torch.zeros_like(sa); hard[torch.topk(sa,K).indices]=1
52        h=(hard-sa).detach()+sa
53        he=h[:M]
54        pred=(X*(W.squeeze()*he)).sum(1)
55        loss=torch.nn.functional.binary_cross_entropy_with_logits(pred,y)
56        loss.backward(); opt.step()
57    with torch.no_grad():
58        sa=torch.cat([sr,sd]) if double else sr
59        ind=torch.topk(sa,K).indices
60        real=(ind<M)
61        r=int(real.sum()); dummy=int((~real).sum())
62        mask=torch.zeros(M)
63        mask[ind[ind<M]]=1
64        pred=(X*(W.squeeze()*mask)).sum(1)
65        acc=float(((pred>0)==(y>.5)).float().mean())
66    return {'acc':acc,'r':r,'dummy':dummy,'eff_density':r/M,'K':K}
67
68
69def main():
70    seed(); out={'predictions':{},'representation':{},'mini_experiment':{}}
71    rows=[]
72    for M in (20,50,100):
73        for rho in (.2,.5,.8):
74            K,c,mu,var=topk_counts(M,rho)
75            rows.append({'M':M,'rho_aug':rho,'K':K,'pred_mean_r':mu,'obs_mean_r':float(c.mean()),
76                         'pred_std_r':math.sqrt(var),'obs_std_r':float(c.std()),
77                         'pred_eff_density':mu/M,'obs_eff_density':float(c.mean()/M),
78                         'pred_dummy_fraction':.5,'obs_dummy_fraction':float(1-c.mean()/K)})
79    out['predictions']['iid_topk_sweep']=rows
80    ok,ex=representation_check()
81    out['representation']={'M':12,'K':9,'all_masks_at_most_K_represented':ok,'examples_r_dummy':ex}
82    reps=[]
83    for s in range(5):
84        reps.append({'seed':s,'baseline':score_training(rho_aug=.30,seed0=s,double=False),
85                     'double':score_training(rho_aug=.30,seed0=s,double=True)})
86    out['mini_experiment']['runs']=reps
87    for kind in ('baseline','double'):
88        vals=[x[kind]['acc'] for x in reps]; dens=[x[kind]['eff_density'] for x in reps]
89        out['mini_experiment'][kind]={'mean_acc':float(np.mean(vals)),'std_acc':float(np.std(vals)),
90             'mean_eff_density':float(np.mean(dens)),'std_eff_density':float(np.std(dens)), 'K':reps[0][kind]['K']}
91    Path('results.json').write_text(json.dumps(out,indent=2))
92    print(json.dumps(out,indent=2))
93
94if __name__=='__main__': main()