"""Unitree G1 rough-terrain locomotion training, on the Space's GPU. Training runs in a background thread so the Gradio server stays responsive; the UI is a log tail plus start/stop. Brax PPO cannot be interrupted from outside, so the stop button sets a flag that the progress callback checks and raises on. """ from __future__ import annotations import os import pickle import threading import traceback from collections import deque from datetime import datetime from pathlib import Path import gradio as gr ENV_NAME = os.environ.get("ENV_NAME", "G1JoystickRoughTerrain") NUM_TIMESTEPS = int(os.environ.get("NUM_TIMESTEPS", 200_000_000)) SEED = int(os.environ.get("SEED", 0)) HF_REPO = os.environ.get("HF_REPO", "").strip() HF_TOKEN = os.environ.get("HF_TOKEN", "").strip() or None AUTO_START = os.environ.get("AUTO_START", "1") == "1" # Terrain: "himalaya" swaps Playground's 5 cm noise heightfield for crops of real # SRTM elevation of the Khumbu valley (see himalaya_terrain.py); "playground" # keeps the stock terrain. Every env trains on its own crop. TERRAIN = os.environ.get("TERRAIN", "himalaya") NUM_TERRAINS = int(os.environ.get("NUM_TERRAINS", 64)) HIMALAYA_RELIEF = float(os.environ.get("HIMALAYA_RELIEF", 0.3)) # metres over the 20 m arena HIMALAYA_PATCH = float(os.environ.get("HIMALAYA_PATCH", 600.0)) # metres of real ground per arena SNOW = os.environ.get("SNOW", "1") == "1" # snow on the terrain SNOW_FRICTION = tuple(float(v) for v in os.environ.get("SNOW_FRICTION", "0.3,0.7").split(",")) SNOW_DEPTH = tuple(float(v) for v in os.environ.get("SNOW_DEPTH", "0.0,0.08").split(",")) # "walk" = Playground's joystick task (the terrain/snow work above). # "getup" = fall recovery (g1_getup.py): most episodes start fallen, nothing # terminates on being down, and the model gains torso/pelvis collision geoms. TASK = os.environ.get("TASK", "walk").strip().lower() # "snow" retunes the walking gait for deep snow: higher swing, wider stance, # slower cadence, and a lateral foot-separation cost (see snow_gait.py). # Defaults to the snow gait: this Space trains a G1 for snow, and the street # gait is what the 200M-step runs plateaued on (the policy crossed its own feet # after ~45 steps and took the -100 termination every episode). Set GAIT=street # for Playground's stock pavement tuning. GAIT = os.environ.get("GAIT", "snow").strip().lower() FOOT_HEIGHT = float(os.environ.get("FOOT_HEIGHT", 0.22)) # swing reference, m FOOT_SEPARATION = float(os.environ.get("FOOT_SEPARATION", 0.20)) # min lateral gap, m GAIT_FREQ = tuple(float(v) for v in os.environ.get("GAIT_FREQ", "1.0,1.3").split(",")) SURFACE = (TERRAIN + ("-snow" if SNOW else "") + ("-getup" if TASK == "getup" else "") + ("-snowgait" if TASK == "walk" and GAIT == "snow" else "")) NUM_EVALS = int(os.environ.get("NUM_EVALS", 40)) # evals (and checkpoints) per run # /data exists only when persistent storage is attached; fall back to /tmp. OUT = Path("/data/ckpt") if Path("/data").is_dir() else Path("/tmp/ckpt") LOG: deque[str] = deque(maxlen=800) STATE = {"thread": None, "status": "idle", "stop": False, "step": 0, "target": NUM_TIMESTEPS} def log(msg: str) -> None: LOG.append(f"[{datetime.now():%H:%M:%S}] {msg}") print(msg, flush=True) class Stopped(Exception): pass UPLOAD = {"enabled": False, "fails": 0} def init_upload() -> None: """Validate credentials once, up front, instead of failing on every eval.""" if not HF_REPO: log(f"HF_REPO unset -- checkpoints stay in {OUT} and are lost on restart") return if not HF_TOKEN: log("HF_TOKEN secret unset -- cannot push checkpoints") return try: from huggingface_hub import HfApi api = HfApi(token=HF_TOKEN) who = api.whoami().get("name") api.create_repo(HF_REPO, repo_type="model", exist_ok=True) UPLOAD["enabled"] = True sub = f"/tree/main/{SURFACE}" if SURFACE != "playground" else "" log(f"checkpoints -> https://huggingface.co/{HF_REPO}{sub} (as {who})") except Exception as e: log(f"cannot write to HF_REPO={HF_REPO}: {type(e).__name__}: " f"{str(e).splitlines()[0][:200]}") log("fix: HF_TOKEN needs write scope on that namespace. A token scoped to " "your own user cannot write to an org repo -- point HF_REPO at a " "namespace the token owns, e.g. /g1-rough-terrain.") def _upload(path: Path) -> None: if not UPLOAD["enabled"]: return try: from huggingface_hub import HfApi prefix = f"{SURFACE}/" if SURFACE != "playground" else "" HfApi(token=HF_TOKEN).upload_file( path_or_fileobj=str(path), path_in_repo=prefix + path.name, repo_id=HF_REPO, repo_type="model") log(f"uploaded {path.name}") UPLOAD["fails"] = 0 except Exception as e: # never let upload kill training UPLOAD["fails"] += 1 log(f"upload failed ({type(e).__name__}): {str(e).splitlines()[0][:160]}") if UPLOAD["fails"] >= 3: UPLOAD["enabled"] = False log("disabling uploads after 3 failures; training continues, " f"checkpoints remain in {OUT}") def install_jax_pmap_shims() -> None: """Re-add jax.device_put_replicated / device_put_sharded if this JAX removed them. JAX deprecated both in 0.8.1 and removed them in 0.10.0 (April 2026), but the current brax *release* (0.14.2, which playground 0.2.0 requires) still calls device_put_replicated; only brax main has the sharding-based replacement. Shimming the two functions is a smaller intervention than pinning JAX down, which would risk the MJX/Warp stack that already compiles cleanly here. """ import jax import jax.numpy as jnp import numpy as np def _sharding(devices): mesh = jax.sharding.Mesh(np.array(list(devices)), axis_names=("i",)) return jax.sharding.NamedSharding(mesh, jax.sharding.PartitionSpec("i")) def device_put_replicated(x, devices): sharding, n = _sharding(devices), len(devices) def rep(leaf): stack = jnp.stack if isinstance(leaf, jax.Array) else np.stack return jax.device_put(stack([leaf] * n), sharding) return jax.tree_util.tree_map(rep, x) def device_put_sharded(shards, devices): sharding = _sharding(devices) def put(*leaves): stack = jnp.stack if isinstance(leaves[0], jax.Array) else np.stack return jax.device_put(stack(list(leaves)), sharding) return jax.tree_util.tree_map(put, *shards) for name, fn in (("device_put_replicated", device_put_replicated), ("device_put_sharded", device_put_sharded)): try: getattr(jax, name) except AttributeError: setattr(jax, name, fn) log(f"shimmed jax.{name} (removed in this JAX version)") ENVS = {} def build_envs(): """(env, eval_env, randomization_fn) for ENV_NAME with the Himalaya terrain and snow applied. Built once and cached; the rollout tab reuses eval_env.""" if ENVS: return ENVS["env"], ENVS["eval_env"], ENVS["randomization_fn"] from mujoco_playground import registry from mujoco_playground._src import mjx_env STATE["status"] = "building env" # registry.load() clones mujoco_menagerie on demand; the snow-gait and get-up # envs are constructed directly, which skips that, so a fresh container fails # with "Error opening file ... left_hip_pitch_link.STL". mjx_env.ensure_menagerie_exists() if TASK == "getup": import g1_getup log("task: getup (fall recovery) -- flat ground, full-collision G1") env, eval_env = g1_getup.G1Getup(), g1_getup.G1Getup() randomization_fn = None # no terrain or snow: this task is about the body ENVS.update(env=env, eval_env=eval_env, randomization_fn=randomization_fn) return env, eval_env, randomization_fn if GAIT == "snow": import snow_gait log(f"gait: snow -- swing {FOOT_HEIGHT * 100:.0f} cm (street 15), " f"stance >= {FOOT_SEPARATION * 100:.0f} cm, {GAIT_FREQ[0]}-{GAIT_FREQ[1]} Hz " "(street 1.25-1.5), relaxed hips") cfg = snow_gait.snow_gait_config(foot_height=FOOT_HEIGHT, foot_separation=FOOT_SEPARATION, gait_freq=GAIT_FREQ) env = snow_gait.G1SnowGait(config=cfg) eval_env = snow_gait.G1SnowGait(config=cfg) else: log(f"loading {ENV_NAME} ...") env = registry.load(ENV_NAME) env_cfg = registry.get_default_config(ENV_NAME) eval_env = registry.load(ENV_NAME, config=env_cfg) randomization_fn = registry.get_domain_randomizer(ENV_NAME) if TERRAIN == "himalaya": from himalaya_terrain import apply_terrain, make_terrains, randomizer log(f"building {NUM_TERRAINS} Himalaya crops: {HIMALAYA_PATCH:.0f} m of Khumbu " f"-> 20 m arena, relief {HIMALAYA_RELIEF} m ...") grids = make_terrains(NUM_TERRAINS, seed=SEED, patch_m=HIMALAYA_PATCH, relief=HIMALAYA_RELIEF) apply_terrain(env, grids, HIMALAYA_RELIEF) apply_terrain(eval_env, grids, HIMALAYA_RELIEF) randomization_fn = randomizer(randomization_fn, grids) log("terrain: himalaya (per-env crops via domain randomization)") else: log("terrain: playground stock rough terrain") if SNOW: from himalaya_terrain import snow_randomizer randomization_fn = snow_randomizer(randomization_fn, SNOW_FRICTION, SNOW_DEPTH) log(f"snow: foot-floor friction U{SNOW_FRICTION}, depth U{SNOW_DEPTH} m " "(soft contact, per env)") ENVS.update(env=env, eval_env=eval_env, randomization_fn=randomization_fn) return env, eval_env, randomization_fn def ppo_config(): """Brax PPO config for the active task.""" if TASK == "getup": import g1_getup return g1_getup.brax_ppo_config() from mujoco_playground.config import locomotion_params return locomotion_params.brax_ppo_config(ENV_NAME) def net_config(): cfg = ppo_config().get("network_factory", None) return dict(cfg) if cfg is not None else None def train_worker() -> None: try: STATE["status"] = "importing" log("importing jax / playground / brax ...") import functools import jax log(f"jax {jax.__version__} devices: {jax.devices()}") install_jax_pmap_shims() import brax log(f"brax {brax.__version__}") if not any(d.platform == "gpu" for d in jax.devices()): log("WARNING: no GPU visible to JAX -- this will be very slow") from brax.training.agents.ppo import networks as ppo_networks from brax.training.agents.ppo import train as ppo from mujoco_playground import wrapper env, eval_env, randomization_fn = build_envs() ppo_params = ppo_config() ppo_params.num_timesteps = int(STATE["target"]) ppo_params.num_evals = NUM_EVALS # each eval = one log line + checkpoint upload log(f"num_envs={ppo_params.get('num_envs')} " f"batch_size={ppo_params.get('batch_size')} " f"timesteps={int(STATE['target']):,}") OUT.mkdir(parents=True, exist_ok=True) init_upload() t_start = [None] def progress(step, metrics): if STATE["stop"]: raise Stopped() import time STATE["step"] = int(step) rew = metrics.get("eval/episode_reward", float("nan")) ln = metrics.get("eval/avg_episode_length", float("nan")) if t_start[0] is None: t_start[0] = (time.time(), int(step)) sps = float("nan") else: t0, s0 = t_start[0] sps = (int(step) - s0) / max(time.time() - t0, 1e-9) log(f"step {int(step):>12,} reward {rew:8.2f} ep_len {ln:7.1f} " f"{sps:,.0f} steps/s") def policy_params_fn(step, make_policy, params): p = OUT / f"ckpt_{int(step)}.pkl" with open(p, "wb") as f: pickle.dump(params, f) _upload(p) ppo_kwargs = dict(ppo_params) net_cfg = ppo_kwargs.pop("network_factory", None) if net_cfg is not None: ppo_kwargs["network_factory"] = functools.partial( ppo_networks.make_ppo_networks, **dict(net_cfg)) STATE["status"] = "training" log("compiling (first step takes several minutes) ...") _, params, _ = ppo.train( **ppo_kwargs, environment=env, eval_env=eval_env, wrap_env_fn=wrapper.wrap_for_brax_training, **({"randomization_fn": randomization_fn} if randomization_fn else {}), progress_fn=progress, policy_params_fn=policy_params_fn, seed=SEED, ) final = OUT / "final.pkl" with open(final, "wb") as f: pickle.dump(params, f) _upload(final) STATE["status"] = "done" log("training complete") except Stopped: STATE["status"] = "stopped" log("stopped by user") except Exception as e: STATE["status"] = f"error: {type(e).__name__}" log(f"FAILED: {type(e).__name__}: {e}") log(traceback.format_exc()[-2000:]) def start(timesteps=None): """Budget is settable from the UI so a rerun does not need a Space restart (env vars are only injected at container start).""" t = STATE["thread"] if t is not None and t.is_alive(): return status_md() if timesteps: STATE["target"] = int(timesteps) STATE["step"] = 0 STATE["stop"] = False STATE["thread"] = threading.Thread(target=train_worker, daemon=True) STATE["thread"].start() return status_md() def stop(): STATE["stop"] = True log("stop requested -- will halt at the next eval boundary") return status_md() def status_md() -> str: alive = STATE["thread"] is not None and STATE["thread"].is_alive() target = int(STATE["target"]) pct = 100.0 * STATE["step"] / max(target, 1) return (f"**{ENV_NAME}** · terrain `{SURFACE}` — status: `{STATE['status']}`" f"{' (running)' if alive else ''} \n" f"step {STATE['step']:,} / {target:,} ({pct:.1f}%) \n" f"checkpoints: `{OUT}`" + (f" → pushing to `{HF_REPO}`" if HF_REPO else " \n⚠️ `HF_REPO` unset — checkpoints are not pushed to the Hub")) def logs() -> str: return "\n".join(LOG) or "(no output yet)" def list_ckpts(): files = sorted(OUT.glob("*.pkl"), key=lambda p: p.stat().st_mtime, reverse=True) return [f.name for f in files] def refresh_ckpts(): names = list_ckpts() return gr.Dropdown(choices=names, value=names[0] if names else None) def render(ckpt: str | None, vx: float, friction: float, depth: float, seconds: float): """Load a checkpoint and film the policy walking on a snowy Himalaya crop.""" from rollout import apply_snow, load_params, make_inference_fn, rollout, write_video if not ckpt: return None, "no checkpoint yet -- train first (or wait for ckpt_0.pkl)" path = OUT / ckpt if not path.exists(): return None, f"{ckpt} not found" try: _, eval_env, _ = build_envs() params = load_params(path) policy = make_inference_fn(eval_env, net_config())(params, deterministic=True) if SNOW and TASK != "getup": apply_snow(eval_env, float(friction), float(depth)) log(f"rollout {ckpt}: task={TASK} vx={vx} friction={friction} depth={depth} m, {seconds}s ...") qpos, info = rollout(eval_env, policy, seconds=float(seconds), command=(float(vx), 0.0, 0.0)) out = write_video(eval_env, qpos, OUT / "rollouts" / f"{path.stem}.mp4", fps=1.0 / eval_env.dt) if TASK == "getup": # Root height: ~0.76 m standing, ~0.1-0.3 m sprawled. msg = (f"{ckpt}: started at {info['start_height_m']:.2f} m, ended at " f"{info['end_height_m']:.2f} m (peak {info['peak_height_m']:.2f}; " f"standing is ~0.76 m)") else: verdict = (f"fell at {info['fell_at']:.1f}s" if info["fell_at"] is not None else f"stayed up for {info['seconds']:.1f}s") msg = f"{ckpt}: {verdict}, walked {info['distance_m']:.2f} m (commanded {vx} m/s)" log(msg) return str(out), msg except Exception as e: log(f"rollout FAILED: {type(e).__name__}: {e}") log(traceback.format_exc()[-1500:]) return None, f"rollout failed: {type(e).__name__}: {str(e)[:300]}" with gr.Blocks(title="G1 rough-terrain training") as demo: gr.Markdown(f"# Unitree G1 — {'fall recovery' if TASK == 'getup' else 'Himalayan terrain locomotion'} training" + (" (snow gait)" if TASK == "walk" and GAIT == "snow" else "")) st = gr.Markdown(status_md()) with gr.Row(): steps_in = gr.Number(value=NUM_TIMESTEPS, precision=0, label="timesteps", minimum=100_000, maximum=2_000_000_000) gr.Button("Start", variant="primary").click(start, inputs=steps_in, outputs=st) gr.Button("Stop", variant="stop").click(stop, outputs=st) out = gr.Textbox(label="training log", lines=26, max_lines=26, autoscroll=True, value=logs()) timer = gr.Timer(3.0) timer.tick(logs, outputs=out) timer.tick(status_md, outputs=st) gr.Markdown("## Watch a checkpoint walk") with gr.Row(): ckpt_dd = gr.Dropdown(choices=list_ckpts(), label="checkpoint", value=(list_ckpts() or [None])[0]) gr.Button("Refresh").click(refresh_ckpts, outputs=ckpt_dd) vx_in = gr.Slider(0.0, 1.0, value=0.5, step=0.1, label="forward speed (m/s)") mu_in = gr.Slider(0.2, 1.0, value=0.5, step=0.05, label="snow friction") depth_in = gr.Slider(0.0, 0.10, value=0.05, step=0.01, label="snow depth (m)") secs_in = gr.Slider(2, 20, value=8, step=1, label="seconds") verdict = gr.Markdown("") video = gr.Video(label="rollout", autoplay=True) gr.Button("Render rollout", variant="primary").click( render, inputs=[ckpt_dd, vx_in, mu_in, depth_in, secs_in], outputs=[video, verdict]) if AUTO_START: start() demo.queue().launch()