Spaces:
Sleeping
Sleeping
Download app.py from iteratehack/armins: direct link, hf CLI and curl.
- Browser
- Download file 18.7 kB
-
https://huggingface.co/spaces/iteratehack/armins/resolve/main/app.py
- Command line
-
hf download hf://spaces/iteratehack/armins/app.py
-
curl -L -o app.py https://huggingface.co/spaces/iteratehack/armins/resolve/main/app.py
18.7 kB
| """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. <your-user>/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() | |