Buckets:
| # /// 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.