alphaXiv's picture
download
raw
4.69 kB
"""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.