armins / app.py
arminfg's picture
Fix: fetch mujoco_menagerie for directly-built envs (snow gait crashed on a fresh container) (#13)
f33b375
Raw History Blame Contribute Delete
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()