alphaXiv's picture
download
raw
3.41 kB
# /// script
# requires-python = ">=3.10"
# dependencies = ["torch", "numpy", "gymnasium", "scipy", "pandas"]
# ///
"""Minimal diagnostic: confirm train.py runs in the HF env with flushed output."""
import subprocess, sys, os, time
REPO = "/workspace/Rationality"
print("flush-test start", flush=True)
os.makedirs("/workspace", exist_ok=True)
if not os.path.isdir(os.path.join(REPO, ".git")):
subprocess.run(["git", "clone", "--quiet", "https://github.com/EVIEHub/Rationality", REPO], check=True)
# patch taxi.py + cliffwalking.py + train.py
def patch(path, old, new):
s = open(path).read()
if old in s:
open(path, "w").write(s.replace(old, new))
print(f"patched {os.path.basename(path)} ({old[:40]})", flush=True)
else:
print(f"NO-MATCH {os.path.basename(path)} ({old[:40]})", flush=True)
patch(os.path.join(REPO, "src/env/taxi.py"),
'def get_base_P(env_id="Taxi-v3"):\n env = gym.make(env_id)\n P = env.unwrapped.P',
'def get_base_P(env_id="Taxi-v3"):\n try:\n env = gym.make(env_id)\n except Exception:\n env = gym.make("Taxi-v4")\n P = env.unwrapped.P')
patch(os.path.join(REPO, "src/env/taxi.py"),
' env0 = gym.make("Taxi-v3")\n d0_train = make_d0_gym_valid_uniform(env0, nS)',
' try:\n env0 = gym.make("Taxi-v3")\n except Exception:\n env0 = gym.make("Taxi-v4")\n d0_train = make_d0_gym_valid_uniform(env0, nS)')
patch(os.path.join(REPO, "src/env/cliffwalking.py"),
'def get_base_P(env_id="CliffWalking-v0"):\n env = gym.make(env_id)\n P = env.unwrapped.P',
'def get_base_P(env_id="CliffWalking-v0"):\n try:\n env = gym.make(env_id)\n except Exception:\n env = gym.make("CliffWalking-v1")\n P = env.unwrapped.P')
patch(os.path.join(REPO, "src/env/cliffwalking.py"),
' env0 = gym.make(env_id)\n d0_train = make_d0_cliff_start(env0, nS)',
' try:\n env0 = gym.make(env_id)\n except Exception:\n env0 = gym.make("CliffWalking-v1")\n d0_train = make_d0_cliff_start(env0, nS)')
patch(os.path.join(REPO, "train.py"),
'LOG_DIR = Path("~/rational_exp/logs").expanduser()',
'LOG_DIR = Path(os.environ.get("RATIONALITY_LOG_DIR", "~/rational_exp/logs")).expanduser()')
# add import os to train.py if missing
tp = os.path.join(REPO, "train.py")
ts = open(tp).read()
if "import os" not in ts:
open(tp, "w").write(ts.replace("import argparse", "import argparse\nimport os", 1))
print("added import os to train.py", flush=True)
env = dict(os.environ)
env.update({"OMP_NUM_THREADS":"1","MKL_NUM_THREADS":"1","TORCH_NUM_THREADS":"1","RATIONALITY_LOG_DIR":"/workspace/logs"})
print("starting train.py (taxi, 200 eps)", flush=True)
t0 = time.time()
cmd = [sys.executable, "train.py", "--algo","dqn","--env","taxi","--device","cpu",
"--num_episodes","200","--eval_every","20","--eps_train","0.25","--horizon","500",
"--seed","1","--experiment","diag","--variable","baseline"]
r = subprocess.run(cmd, cwd=REPO, env=env, timeout=900)
print(f"train.py exit={r.returncode} in {time.time()-t0:.0f}s", flush=True)
csv = "/workspace/logs/taxi/diag/baseline/result_1.csv"
print("csv exists:", os.path.exists(csv), "size:", os.path.getsize(csv) if os.path.exists(csv) else 0, flush=True)
if os.path.exists(csv):
print("last line:", open(csv).read().strip().splitlines()[-1], flush=True)

Xet Storage Details

Size:
3.41 kB
·
Xet hash:
5346fd591ff74dc26955315ac43afbf4bef3f78ac6a62120e8f69eb717ee08c1

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