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()
|