alphaXiv's picture
download
raw
3.23 kB
"""Quantitative claim verification from pre-collected logs (exp_1 = evaluated policy)."""
import os, glob
import numpy as np
import pandas as pd
from scipy import stats
BASE = "rational_exp_1/logs"
def load_agg(env, experiment, y_col="rational_risk_gap", max_episode=None, last_n=None):
exp_dir = os.path.join(BASE, env, experiment)
out = {}
for var in sorted(os.listdir(exp_dir)):
vdir = os.path.join(exp_dir, var)
if not os.path.isdir(vdir): continue
csvs = sorted(glob.glob(os.path.join(vdir, "result_*.csv")))
if not csvs: continue
seed_means = []
for f in csvs:
d = pd.read_csv(f)
if max_episode is not None:
d = d[d["episode"] <= max_episode]
if last_n is not None and not d.empty:
eps = sorted(d["episode"].unique())
d = d[d["episode"].isin(eps[-last_n:])]
if not d.empty:
seed_means.append(d[y_col].mean())
if seed_means:
out[var] = (np.mean(seed_means), np.std(seed_means, ddof=1), len(seed_means))
return out
def show(name, agg):
print(f"\n=== {name} (mean ± std over last-15 episodes, n seeds) ===")
for v,(m,s,n) in agg.items():
print(f" {v:20s} {m:10.3f} ± {s:8.3f} (n={n})")
return agg
print("="*70)
print("CLAIM 4: Regularisation reduces rational risk gap vs baseline")
print("(mean over the paper's plotted training window: ep<=1400)")
print("="*70)
for env in ["taxi","cliffwalking"]:
agg = show(f"{env}/exp_reg", load_agg(env,"exp_reg", max_episode=1400))
base = agg.get("baseline")
if base:
for v in ["ln_train","l2_train","wn_train"]:
if v in agg:
red = base[0]-agg[v][0]
print(f" {v}: reduction vs baseline = {red:.3f} ({(red/base[0]*100):.1f}%) -> {'SUPPORTS' if red>0 else 'REFUTES'}")
print("\n" + "="*70)
print("CLAIM 5a: Domain randomisation decreases rational risk gap vs baseline")
print("(mean over the paper's plotted training window)")
print("="*70)
windows = {"taxi":800, "cliffwalking":1500}
for env in ["taxi","cliffwalking"]:
agg = show(f"{env}/exp_domain_rand", load_agg(env,"exp_domain_rand", max_episode=windows[env]))
base = agg.get("baseline")
if base and "envrnd_train_25" in agg:
red = base[0]-agg["envrnd_train_25"][0]
print(f" envrnd reduction vs baseline = {red:.3f} ({(red/base[0]*100):.1f}%) -> {'SUPPORTS' if red>0 else 'REFUTES'}")
print("\n" + "="*70)
print("CLAIM 5b: Env shift magnitude (eps_train) correlates positively with gap")
print("(mean over the paper's plotted window: ep<=900)")
print("="*70)
for env in ["taxi","cliffwalking"]:
agg = show(f"{env}/exp_environment_level", load_agg(env,"exp_environment_level", max_episode=900))
level_map = {"default":0.0,"eps_train_01":0.1,"eps_train_03":0.3,"eps_train_05":0.5,"eps_train_07":0.7}
xs, ys = [], []
for v,(m,s,n) in agg.items():
if v in level_map:
xs.append(level_map[v]); ys.append(m)
if len(xs)>=2:
r,p = stats.pearsonr(xs,ys)
print(f" Pearson r = {r:.4f} (p={p:.4g}) -> {'SUPPORTS' if r>0 else 'REFUTES'} positive correlation")

Xet Storage Details

Size:
3.23 kB
·
Xet hash:
9bd5c4960824fbf317692e09991604d3a589cfd8188be918f890164e9607230c

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.