SabaPivot's picture
Upgrade canonical logbook from stronger peer evidence with attribution
ea3a71e verified
Raw
History Blame Contribute Delete
2.23 kB
"""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()