G1 rough-terrain training app (MuJoCo Playground + Brax PPO)

#1
Files changed (3) hide show
  1. README.md +20 -4
  2. app.py +178 -0
  3. requirements.txt +6 -0
README.md CHANGED
@@ -1,13 +1,29 @@
1
  ---
2
- title: Armins
3
- emoji: 🌍
4
  colorFrom: red
5
  colorTo: green
6
  sdk: gradio
7
  sdk_version: 6.26.0
8
- python_version: '3.13'
9
  app_file: app.py
10
  pinned: false
 
11
  ---
12
 
13
- Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
+ title: G1 Rough Terrain Training
3
+ emoji: 🦿
4
  colorFrom: red
5
  colorTo: green
6
  sdk: gradio
7
  sdk_version: 6.26.0
8
+ python_version: '3.11'
9
  app_file: app.py
10
  pinned: false
11
+ short_description: GPU training for Unitree G1 rough-terrain locomotion
12
  ---
13
 
14
+ # Unitree G1 — rough-terrain locomotion training
15
+
16
+ Trains `G1JoystickRoughTerrain` from [MuJoCo Playground](https://playground.mujoco.org/)
17
+ with Brax PPO on the Space's GPU. Thousands of MJX environments step in parallel,
18
+ versus the 16 CPU processes a laptop manages.
19
+
20
+ Training starts automatically on boot. Checkpoints are written to `/data` when
21
+ persistent storage is attached, and pushed to `HF_REPO` if that variable and a
22
+ write-scoped `HF_TOKEN` secret are set.
23
+
24
+ **Settings that matter**
25
+ - *Sleep time* → **never**. The Space idles on HTTP traffic, not CPU/GPU load, so
26
+ a background training thread will not keep it awake.
27
+ - *Persistent storage* — without it, checkpoints are lost on restart. Set
28
+ `HF_REPO` so they are pushed to the Hub instead.
29
+ - `NUM_TIMESTEPS` (default 200M), `ENV_NAME`, `SEED` are configurable variables.
app.py ADDED
@@ -0,0 +1,178 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Unitree G1 rough-terrain locomotion training, on the Space's GPU.
2
+
3
+ Training runs in a background thread so the Gradio server stays responsive; the
4
+ UI is a log tail plus start/stop. Brax PPO cannot be interrupted from outside,
5
+ so the stop button sets a flag that the progress callback checks and raises on.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import os
11
+ import pickle
12
+ import threading
13
+ import traceback
14
+ from collections import deque
15
+ from datetime import datetime
16
+ from pathlib import Path
17
+
18
+ import gradio as gr
19
+
20
+ ENV_NAME = os.environ.get("ENV_NAME", "G1JoystickRoughTerrain")
21
+ NUM_TIMESTEPS = int(os.environ.get("NUM_TIMESTEPS", 200_000_000))
22
+ SEED = int(os.environ.get("SEED", 0))
23
+ HF_REPO = os.environ.get("HF_REPO", "").strip()
24
+ AUTO_START = os.environ.get("AUTO_START", "1") == "1"
25
+
26
+ # /data exists only when persistent storage is attached; fall back to /tmp.
27
+ OUT = Path("/data/ckpt") if Path("/data").is_dir() else Path("/tmp/ckpt")
28
+
29
+ LOG: deque[str] = deque(maxlen=800)
30
+ STATE = {"thread": None, "status": "idle", "stop": False, "step": 0}
31
+
32
+
33
+ def log(msg: str) -> None:
34
+ LOG.append(f"[{datetime.now():%H:%M:%S}] {msg}")
35
+ print(msg, flush=True)
36
+
37
+
38
+ class Stopped(Exception):
39
+ pass
40
+
41
+
42
+ def _upload(path: Path) -> None:
43
+ if not HF_REPO:
44
+ return
45
+ try:
46
+ from huggingface_hub import HfApi
47
+ api = HfApi()
48
+ api.create_repo(HF_REPO, repo_type="model", exist_ok=True)
49
+ api.upload_file(path_or_fileobj=str(path), path_in_repo=path.name,
50
+ repo_id=HF_REPO, repo_type="model")
51
+ log(f"uploaded {path.name} -> {HF_REPO}")
52
+ except Exception as e: # never let upload kill training
53
+ log(f"upload failed ({type(e).__name__}): {e}")
54
+
55
+
56
+ def train_worker() -> None:
57
+ try:
58
+ STATE["status"] = "importing"
59
+ log("importing jax / playground / brax ...")
60
+ import functools
61
+ import jax
62
+
63
+ log(f"jax devices: {jax.devices()}")
64
+ if not any(d.platform == "gpu" for d in jax.devices()):
65
+ log("WARNING: no GPU visible to JAX -- this will be very slow")
66
+
67
+ from brax.training.agents.ppo import networks as ppo_networks
68
+ from brax.training.agents.ppo import train as ppo
69
+ from mujoco_playground import registry, wrapper
70
+ from mujoco_playground.config import locomotion_params
71
+
72
+ STATE["status"] = "building env"
73
+ log(f"loading {ENV_NAME} ...")
74
+ env = registry.load(ENV_NAME)
75
+ env_cfg = registry.get_default_config(ENV_NAME)
76
+ ppo_params = locomotion_params.brax_ppo_config(ENV_NAME)
77
+ ppo_params.num_timesteps = NUM_TIMESTEPS
78
+ log(f"num_envs={ppo_params.get('num_envs')} "
79
+ f"batch_size={ppo_params.get('batch_size')} "
80
+ f"timesteps={NUM_TIMESTEPS:,}")
81
+
82
+ OUT.mkdir(parents=True, exist_ok=True)
83
+
84
+ def progress(step, metrics):
85
+ if STATE["stop"]:
86
+ raise Stopped()
87
+ STATE["step"] = int(step)
88
+ rew = metrics.get("eval/episode_reward", float("nan"))
89
+ ln = metrics.get("eval/avg_episode_length", float("nan"))
90
+ log(f"step {int(step):>12,} reward {rew:8.2f} ep_len {ln:7.1f}")
91
+
92
+ def policy_params_fn(step, make_policy, params):
93
+ p = OUT / f"ckpt_{int(step)}.pkl"
94
+ with open(p, "wb") as f:
95
+ pickle.dump(params, f)
96
+ _upload(p)
97
+
98
+ ppo_kwargs = dict(ppo_params)
99
+ net_cfg = ppo_kwargs.pop("network_factory", None)
100
+ if net_cfg is not None:
101
+ ppo_kwargs["network_factory"] = functools.partial(
102
+ ppo_networks.make_ppo_networks, **dict(net_cfg))
103
+
104
+ STATE["status"] = "training"
105
+ log("compiling (first step takes several minutes) ...")
106
+ _, params, _ = ppo.train(
107
+ **ppo_kwargs,
108
+ environment=env,
109
+ eval_env=registry.load(ENV_NAME, config=env_cfg),
110
+ wrap_env_fn=wrapper.wrap_for_brax_training,
111
+ randomization_fn=registry.get_domain_randomizer(ENV_NAME),
112
+ progress_fn=progress,
113
+ policy_params_fn=policy_params_fn,
114
+ seed=SEED,
115
+ )
116
+ final = OUT / "final.pkl"
117
+ with open(final, "wb") as f:
118
+ pickle.dump(params, f)
119
+ _upload(final)
120
+ STATE["status"] = "done"
121
+ log("training complete")
122
+
123
+ except Stopped:
124
+ STATE["status"] = "stopped"
125
+ log("stopped by user")
126
+ except Exception as e:
127
+ STATE["status"] = f"error: {type(e).__name__}"
128
+ log(f"FAILED: {type(e).__name__}: {e}")
129
+ log(traceback.format_exc()[-2000:])
130
+
131
+
132
+ def start():
133
+ t = STATE["thread"]
134
+ if t is not None and t.is_alive():
135
+ return status_md()
136
+ STATE["stop"] = False
137
+ STATE["thread"] = threading.Thread(target=train_worker, daemon=True)
138
+ STATE["thread"].start()
139
+ return status_md()
140
+
141
+
142
+ def stop():
143
+ STATE["stop"] = True
144
+ log("stop requested -- will halt at the next eval boundary")
145
+ return status_md()
146
+
147
+
148
+ def status_md() -> str:
149
+ alive = STATE["thread"] is not None and STATE["thread"].is_alive()
150
+ pct = 100.0 * STATE["step"] / max(NUM_TIMESTEPS, 1)
151
+ return (f"**{ENV_NAME}** — status: `{STATE['status']}`"
152
+ f"{' (running)' if alive else ''} \n"
153
+ f"step {STATE['step']:,} / {NUM_TIMESTEPS:,} ({pct:.1f}%) \n"
154
+ f"checkpoints: `{OUT}`"
155
+ + (f" → pushing to `{HF_REPO}`" if HF_REPO else
156
+ " \n⚠️ `HF_REPO` unset — checkpoints are not pushed to the Hub"))
157
+
158
+
159
+ def logs() -> str:
160
+ return "\n".join(LOG) or "(no output yet)"
161
+
162
+
163
+ with gr.Blocks(title="G1 rough-terrain training") as demo:
164
+ gr.Markdown("# Unitree G1 — rough-terrain locomotion training")
165
+ st = gr.Markdown(status_md())
166
+ with gr.Row():
167
+ gr.Button("Start", variant="primary").click(start, outputs=st)
168
+ gr.Button("Stop", variant="stop").click(stop, outputs=st)
169
+ out = gr.Textbox(label="training log", lines=26, max_lines=26,
170
+ autoscroll=True, value=logs())
171
+ timer = gr.Timer(3.0)
172
+ timer.tick(logs, outputs=out)
173
+ timer.tick(status_md, outputs=st)
174
+
175
+ if AUTO_START:
176
+ start()
177
+
178
+ demo.queue().launch()
requirements.txt ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ jax[cuda12]>=0.4.35
2
+ mujoco>=3.2.7
3
+ mujoco-mjx>=3.2.7
4
+ playground==0.2.0
5
+ brax>=0.12.1
6
+ huggingface_hub>=0.35