File size: 2,226 Bytes
ea3a71e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
"""Topology arm only: DecDPO on DistilGPT-2 across four 5-node graphs."""
import json, time, numpy as np, torch
from transformers import AutoModelForCausalLM
from dpo_real import DEV, MODEL
import fed_real as FR

RES = json.load(open("fed_real_results.json"))

def main():
    t0=time.time()
    clients,names,tok = FR.build_clients()
    base = AutoModelForCausalLM.from_pretrained(MODEL)
    ref = AutoModelForCausalLM.from_pretrained(MODEL).to(DEV).eval()
    for p in ref.parameters(): p.requires_grad_(False)
    pad=tok.pad_token_id; n=FR.N_CLIENTS
    tops={}
    ring=np.zeros((n,n),int)
    for i in range(n): ring[i,(i+1)%n]=ring[(i+1)%n,i]=1
    tops["ring"]=ring
    star=np.zeros((n,n),int); star[0,1:]=star[1:,0]=1; tops["star"]=star
    path=np.zeros((n,n),int)
    for i in range(n-1): path[i,i+1]=path[i+1,i]=1
    tops["path"]=path
    tops["complete"]=np.ones((n,n),int)-np.eye(n,dtype=int)
    rows=[]
    for nm,adj in tops.items():
        W,rho=FR.metropolis(adj)
        m,cons=FR.dec_run(base,ref,clients,tok,W)
        l,a=FR.evaluate(m,ref,clients,pad)
        rows.append({"topology":nm,"rho":round(rho,4),
                     "one_over_1_minus_rho2":(round(1/(1-rho**2),3) if rho<0.999999 else None),
                     "consensus_error":cons,"loss":l,"acc":a})
        print("  %-9s rho=%.4f  1/(1-rho^2)=%s  cons=%.4e  loss=%.5f  (%.0fs)"%(
            nm,rho,rows[-1]["one_over_1_minus_rho2"],cons,l,time.time()-t0),flush=True)
        json.dump({**RES,"claim5_topology":{"rows":rows}},open("fed_real_results.json","w"),indent=1)
    ok=[r for r in rows if r["consensus_error"]>1e-7 and r["one_over_1_minus_rho2"]]
    out={"rows":rows,"n_in_fit":len(ok)}
    if len(ok)>=3:
        x=np.log([r["one_over_1_minus_rho2"] for r in ok]); y=np.log([r["consensus_error"] for r in ok])
        sl,ic=np.polyfit(x,y,1)
        out.update({"loglog_slope":round(float(sl),4),
                    "r2":round(float(1-np.var(y-(sl*x+ic))/np.var(y)),4),
                    "spearman_positive":bool(np.corrcoef(x,y)[0,1]>0)})
    RES["claim5_topology"]=out
    json.dump(RES,open("fed_real_results.json","w"),indent=1)
    print("DONE %.0fs"%(time.time()-t0),flush=True)

if __name__=="__main__": main()