iteratehack/ascender-rl-artifacts
0563
1#!/bin/bash2# Training entry point inside an HF Jobs container.3# Expects (env): ARTIFACTS_REPO, EXP_NAME, TRAIN_FLAGS; HF_TOKEN secret set.4# Mounted: the artifacts dataset repo read-only at /artifacts.5set -euo pipefail6 7mkdir -p /work/logs8LOG=/work/logs/job.log9exec > >(tee "$LOG") 2>&1 || true10 11echo "[job] host: $(hostname)"12nvidia-smi || echo "[job] WARNING: no GPU visible"13 14# --- deps: jax + CUDA 12, mujoco playground, wandb --------------------------15pip install --no-cache-dir \16 "jax[cuda12]==0.11.1" jaxlib==0.11.1 \17 mujoco==3.12.0 mujoco-mjx==3.12.0 playground==0.2.0 \18 brax==0.14.2 mediapy tensorboardX wandb19 20# W&B: authenticate from the HF secret if provided21if [[ -n "${WANDB_API_KEY:-}" ]]; then22 wandb login --relogin "$WANDB_API_KEY" > /dev/null 2>&1 || echo "[job] wandb login failed (continuing without W&B)"23fi24 25export MUJOCO_GL=egl26export XLA_PYTHON_CLIENT_PREALLOCATE=false27export XLA_PYTHON_CLIENT_MEM_FRACTION=0.8528export XLA_FLAGS="--xla_gpu_triton_gemm_any=True"29 30# --- unpack the code snapshot ------------------------------------------------31mkdir -p /work/repo && cd /work/repo32tar xzf /artifacts/code/code.tar.gz33 34# --- artifacts uploader: pushes new checkpoints/logs every 5 min -------------35(36 while true; do37 sleep 30038 python - <<'PY' || echo "[artifacts] sync failed this round; retrying in 5 min"39import os40from huggingface_hub import HfApi41api = HfApi()42api.upload_folder(43 repo_id=os.environ["ARTIFACTS_REPO"], repo_type="dataset",44 folder_path="/work/logs", path_in_repo=f"runs/{os.environ['EXP_NAME']}",45)46print("[job] artifacts synced to", os.environ["ARTIFACTS_REPO"], flush=True)47PY48 done49) &50ARTIFACT_UPLOADER_PID=$!51trap 'kill $ARTIFACT_UPLOADER_PID 2>/dev/null || true' EXIT52 53echo "[job] training starting: $TRAIN_FLAGS"54python rl/scripts/train_jax_ppo.py $TRAIN_FLAGS --logdir /work/logs || \55 echo "[job] training exited nonzero — uploading logs anyway"56 57# --- final full upload --------------------------------------------------------58python - <<'PY'59import os, traceback60from huggingface_hub import HfApi61try:62 api = HfApi()63 repo = os.environ["ARTIFACTS_REPO"]64 exp = os.environ["EXP_NAME"]65 api.upload_folder(66 repo_id=repo, repo_type="dataset",67 folder_path="/work/logs",68 path_in_repo=f"runs/{exp}",69 )70 print("[job] artifacts uploaded to", repo, "runs/", exp)71except Exception:72 import traceback; traceback.print_exc()73PY74echo "[job] done."75 