Buckets:
| """Analyze + plot the FRESH independent runs (uploaded to HF dataset).""" | |
| import os, glob | |
| import numpy as np | |
| import pandas as pd | |
| import plotly.graph_objects as go | |
| from scipy import stats | |
| # double-nested path from upload | |
| BASE = {"taxi": "/root/fresh_logs/fresh_logs/taxi/taxi", | |
| "cliffwalking": "/root/fresh_logs/fresh_logs/cliffwalking/cliffwalking"} | |
| OUT = "/root/Rationality/repro_bundle/fresh_figures" | |
| os.makedirs(OUT, exist_ok=True) | |
| VAR_LABEL = {"baseline":"baseline DQN","ln_train":"LayerNorm","l2_train":"L2 reg","wn_train":"WeightNorm", | |
| "envrnd_train_25":"domain rand.","default":"shift 0%","eps_train_01":"shift 10%", | |
| "eps_train_03":"shift 30%","eps_train_05":"shift 50%","eps_train_07":"shift 70%"} | |
| def agg(env, exp, max_ep): | |
| d=os.path.join(BASE[env],exp); out={} | |
| for var in sorted(os.listdir(d)): | |
| vd=os.path.join(d,var) | |
| if not os.path.isdir(vd):continue | |
| cs=sorted(glob.glob(os.path.join(vd,"result_*.csv"))) | |
| sm=[] | |
| for f in cs: | |
| df=pd.read_csv(f); df=df[df["episode"]<=max_ep] | |
| if not df.empty: sm.append(df["rational_risk_gap"].mean()) | |
| if sm: out[var]=(np.mean(sm),np.std(sm,ddof=1),len(sm)) | |
| return out | |
| def show(name, a): | |
| print(f"\n{name}") | |
| for v,(m,s,n) in sorted(a.items()): | |
| print(f" {v:18s} {m:9.2f} ± {s:8.2f} (n={n})") | |
| print("="*60); print("FRESH RUNS — Claim 4: regularisation (ep<=1400)"); print("="*60) | |
| for env in ["taxi","cliffwalking"]: | |
| a=agg(env,"exp_reg",1400); show(f"{env}/exp_reg:",a) | |
| b=a.get("baseline") | |
| if b: | |
| for v in ["ln_train","l2_train","wn_train"]: | |
| if v in a: print(f" {v}: {(b[0]-a[v][0])/b[0]*100:+.1f}% vs baseline -> {'SUPPORTS' if a[v][0]<b[0] else 'refutes'}") | |
| print("\n"+"="*60); print("FRESH RUNS — Claim 5a: domain randomisation"); print("="*60) | |
| win={"taxi":800,"cliffwalking":1500} | |
| for env in ["taxi","cliffwalking"]: | |
| a=agg(env,"exp_domain_rand",win[env]); show(f"{env}/exp_domain_rand:",a) | |
| b=a.get("baseline") | |
| if b and "envrnd_train_25" in a: | |
| print(f" envrnd: {(b[0]-a['envrnd_train_25'][0])/b[0]*100:+.1f}% vs baseline -> {'SUPPORTS' if a['envrnd_train_25'][0]<b[0] else 'refutes'}") | |
| print("\n"+"="*60); print("FRESH RUNS — Claim 5b: env shift magnitude (ep<=900)"); print("="*60) | |
| lvl={"default":0.0,"eps_train_01":0.1,"eps_train_03":0.3,"eps_train_05":0.5,"eps_train_07":0.7} | |
| for env in ["taxi","cliffwalking"]: | |
| a=agg(env,"exp_environment_level",900); show(f"{env}/exp_environment_level:",a) | |
| xs=[lvl[v] for v in a if v in lvl]; ys=[a[v][0] for v in a if v in lvl] | |
| if len(xs)>=2: | |
| r,p=stats.pearsonr(xs,ys); print(f" Pearson r={r:.3f} (p={p:.4g}) -> {'SUPPORTS' if r>0 else 'refutes'}") | |
| # figures | |
| def make_fig(env, exp, title, max_ep, ymax, fn): | |
| d=os.path.join(BASE[env],exp); fig=go.Figure() | |
| for var in sorted(os.listdir(d)): | |
| vd=os.path.join(d,var) | |
| if not os.path.isdir(vd):continue | |
| cs=sorted(glob.glob(os.path.join(vd,"result_*.csv"))) | |
| if not cs:continue | |
| df=pd.concat([pd.read_csv(f)[["episode","rational_risk_gap"]] for f in cs],ignore_index=True) | |
| df=df[df["episode"]<=max_ep] | |
| st=df.groupby("episode")["rational_risk_gap"].agg(["mean","std"]).reset_index().sort_values("episode") | |
| st["mean"]=st["mean"].rolling(3,min_periods=1,center=True).mean() | |
| lab=VAR_LABEL.get(var,var) | |
| fig.add_trace(go.Scatter(x=st["episode"],y=st["mean"],name=lab,mode="lines",line=dict(width=2))) | |
| fig.update_layout(title=title,xaxis_title="episode",yaxis_title="rational risk gap",template="plotly_white", | |
| height=380,legend=dict(font=dict(size=11)),margin=dict(l=50,r=20,t=50,b=40)) | |
| if ymax: fig.update_yaxes(range=[0,ymax]) | |
| fig.write_html(os.path.join(OUT,fn),include_plotlyjs="cdn",full_html=False) | |
| make_fig("taxi","exp_reg","Fresh — Taxi: regularisation vs baseline (rational risk gap)",1400,400,"fresh_claim4_taxi_reg.html") | |
| make_fig("cliffwalking","exp_reg","Fresh — CliffWalking: regularisation vs baseline",1400,1400,"fresh_claim4_cliff_reg.html") | |
| make_fig("taxi","exp_domain_rand","Fresh — Taxi: domain randomisation vs baseline",800,500,"fresh_claim5a_taxi_dr.html") | |
| make_fig("cliffwalking","exp_domain_rand","Fresh — CliffWalking: domain randomisation vs baseline",1500,2000,"fresh_claim5a_cliff_dr.html") | |
| make_fig("taxi","exp_environment_level","Fresh — Taxi: env shift magnitude",900,700,"fresh_claim5b_taxi_shift.html") | |
| make_fig("cliffwalking","exp_environment_level","Fresh — CliffWalking: env shift magnitude",900,5500,"fresh_claim5b_cliff_shift.html") | |
| print("\nFresh figures written to",OUT) | |
Xet Storage Details
- Size:
- 4.69 kB
- Xet hash:
- 7977ac464550d37f70926d59a8f76a217fa8f9f2498b3e8dbf8cd46d7d7f54ed
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.