Add snow to the Himalayan terrain: per-env foot-floor friction and soft-contact depth

#7
Files changed (3) hide show
  1. README.md +9 -1
  2. app.py +19 -23
  3. himalaya_terrain.py +57 -0
README.md CHANGED
@@ -8,7 +8,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 on real Himalayan terrain
12
  ---
13
 
14
  # Unitree G1 — Himalayan terrain locomotion training
@@ -21,6 +21,14 @@ every environment its own crop through the domain randomizer. Set `TERRAIN=playg
21
  to fall back to the stock 5 cm noise field; `HIMALAYA_RELIEF`, `HIMALAYA_PATCH`,
22
  `NUM_TERRAINS` tune the terrain.
23
 
 
 
 
 
 
 
 
 
24
 
25
  Trains `G1JoystickRoughTerrain` from [MuJoCo Playground](https://playground.mujoco.org/)
26
  with Brax PPO on the Space's GPU. Thousands of MJX environments step in parallel,
 
8
  python_version: '3.11'
9
  app_file: app.py
10
  pinned: false
11
+ short_description: GPU training for Unitree G1 on snowy Himalayan terrain
12
  ---
13
 
14
  # Unitree G1 — Himalayan terrain locomotion training
 
21
  to fall back to the stock 5 cm noise field; `HIMALAYA_RELIEF`, `HIMALAYA_PATCH`,
22
  `NUM_TERRAINS` tune the terrain.
23
 
24
+ **Snow** (`SNOW=1`, default): every environment also draws its own snow —
25
+ foot–floor friction from `SNOW_FRICTION` (default `0.3,0.7`; packed snow ≈ 0.5,
26
+ rock ≈ 1.0) and a soft-contact snow depth from `SNOW_DEPTH` (default `0.0,0.08` m,
27
+ 0–3 in: feet settle about half the depth at rest and up to twice it on a
28
+ landing). Randomizing both is what makes the policy robust across snow
29
+ conditions instead of tuned to one. `SNOW=0` gives bare terrain. Checkpoints go
30
+ under `himalaya-snow/` in `HF_REPO`.
31
+
32
 
33
  Trains `G1JoystickRoughTerrain` from [MuJoCo Playground](https://playground.mujoco.org/)
34
  with Brax PPO on the Space's GPU. Thousands of MJX environments step in parallel,
app.py CHANGED
@@ -1,18 +1,8 @@
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.
5
-
6
- Start/stop are deliberately NOT exposed as named API endpoints (api_name=False).
7
- On a public Space, auto-named endpoints get probed -- a caller sent the string
8
- "ping" to /start and then hit /stop, which killed a run mid-flight. The read-only
9
- log and status endpoints stay exposed so the run can be monitored remotely. Brax PPO cannot be interrupted from outside,
10
  so the stop button sets a flag that the progress callback checks and raises on.
11
-
12
- Start/stop are deliberately NOT exposed as named API endpoints (api_name=False).
13
- On a public Space, auto-named endpoints get probed -- a caller sent the string
14
- "ping" to /start and then hit /stop, which killed a run mid-flight. The read-only
15
- log and status endpoints stay exposed so the run can be monitored remotely.
16
  """
17
 
18
  from __future__ import annotations
@@ -41,6 +31,10 @@ TERRAIN = os.environ.get("TERRAIN", "himalaya")
41
  NUM_TERRAINS = int(os.environ.get("NUM_TERRAINS", 64))
42
  HIMALAYA_RELIEF = float(os.environ.get("HIMALAYA_RELIEF", 0.3)) # metres over the 20 m arena
43
  HIMALAYA_PATCH = float(os.environ.get("HIMALAYA_PATCH", 600.0)) # metres of real ground per arena
 
 
 
 
44
 
45
  # /data exists only when persistent storage is attached; fall back to /tmp.
46
  OUT = Path("/data/ckpt") if Path("/data").is_dir() else Path("/tmp/ckpt")
@@ -76,7 +70,7 @@ def init_upload() -> None:
76
  who = api.whoami().get("name")
77
  api.create_repo(HF_REPO, repo_type="model", exist_ok=True)
78
  UPLOAD["enabled"] = True
79
- sub = f"/tree/main/{TERRAIN}" if TERRAIN != "playground" else ""
80
  log(f"checkpoints -> https://huggingface.co/{HF_REPO}{sub} (as {who})")
81
  except Exception as e:
82
  log(f"cannot write to HF_REPO={HF_REPO}: {type(e).__name__}: "
@@ -91,7 +85,7 @@ def _upload(path: Path) -> None:
91
  return
92
  try:
93
  from huggingface_hub import HfApi
94
- prefix = f"{TERRAIN}/" if TERRAIN != "playground" else ""
95
  HfApi(token=HF_TOKEN).upload_file(
96
  path_or_fileobj=str(path), path_in_repo=prefix + path.name,
97
  repo_id=HF_REPO, repo_type="model")
@@ -187,6 +181,11 @@ def train_worker() -> None:
187
  log("terrain: himalaya (per-env crops via domain randomization)")
188
  else:
189
  log("terrain: playground stock rough terrain")
 
 
 
 
 
190
  ppo_params = locomotion_params.brax_ppo_config(ENV_NAME)
191
  ppo_params.num_timesteps = int(STATE["target"])
192
  log(f"num_envs={ppo_params.get('num_envs')} "
@@ -259,11 +258,8 @@ def start(timesteps=None):
259
  t = STATE["thread"]
260
  if t is not None and t.is_alive():
261
  return status_md()
262
- try:
263
- if timesteps is not None:
264
- STATE["target"] = max(100_000, int(float(timesteps)))
265
- except (TypeError, ValueError):
266
- log(f"ignoring non-numeric timesteps input: {timesteps!r}")
267
  STATE["step"] = 0
268
  STATE["stop"] = False
269
  STATE["thread"] = threading.Thread(target=train_worker, daemon=True)
@@ -281,7 +277,7 @@ def status_md() -> str:
281
  alive = STATE["thread"] is not None and STATE["thread"].is_alive()
282
  target = int(STATE["target"])
283
  pct = 100.0 * STATE["step"] / max(target, 1)
284
- return (f"**{ENV_NAME}** · terrain `{TERRAIN}` — status: `{STATE['status']}`"
285
  f"{' (running)' if alive else ''} \n"
286
  f"step {STATE['step']:,} / {target:,} ({pct:.1f}%) \n"
287
  f"checkpoints: `{OUT}`"
@@ -297,10 +293,10 @@ with gr.Blocks(title="G1 rough-terrain training") as demo:
297
  gr.Markdown("# Unitree G1 — Himalayan terrain locomotion training")
298
  st = gr.Markdown(status_md())
299
  with gr.Row():
300
- steps_in = gr.Number(value=NUM_TIMESTEPS, precision=0, label="timesteps")
301
- gr.Button("Start", variant="primary").click(
302
- start, inputs=steps_in, outputs=st, api_name=False)
303
- gr.Button("Stop", variant="stop").click(stop, outputs=st, api_name=False)
304
  out = gr.Textbox(label="training log", lines=26, max_lines=26,
305
  autoscroll=True, value=logs())
306
  timer = gr.Timer(3.0)
 
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
 
31
  NUM_TERRAINS = int(os.environ.get("NUM_TERRAINS", 64))
32
  HIMALAYA_RELIEF = float(os.environ.get("HIMALAYA_RELIEF", 0.3)) # metres over the 20 m arena
33
  HIMALAYA_PATCH = float(os.environ.get("HIMALAYA_PATCH", 600.0)) # metres of real ground per arena
34
+ SNOW = os.environ.get("SNOW", "1") == "1" # snow on the terrain
35
+ SNOW_FRICTION = tuple(float(v) for v in os.environ.get("SNOW_FRICTION", "0.3,0.7").split(","))
36
+ SNOW_DEPTH = tuple(float(v) for v in os.environ.get("SNOW_DEPTH", "0.0,0.08").split(","))
37
+ SURFACE = TERRAIN + ("-snow" if SNOW else "") # checkpoint prefix
38
 
39
  # /data exists only when persistent storage is attached; fall back to /tmp.
40
  OUT = Path("/data/ckpt") if Path("/data").is_dir() else Path("/tmp/ckpt")
 
70
  who = api.whoami().get("name")
71
  api.create_repo(HF_REPO, repo_type="model", exist_ok=True)
72
  UPLOAD["enabled"] = True
73
+ sub = f"/tree/main/{SURFACE}" if SURFACE != "playground" else ""
74
  log(f"checkpoints -> https://huggingface.co/{HF_REPO}{sub} (as {who})")
75
  except Exception as e:
76
  log(f"cannot write to HF_REPO={HF_REPO}: {type(e).__name__}: "
 
85
  return
86
  try:
87
  from huggingface_hub import HfApi
88
+ prefix = f"{SURFACE}/" if SURFACE != "playground" else ""
89
  HfApi(token=HF_TOKEN).upload_file(
90
  path_or_fileobj=str(path), path_in_repo=prefix + path.name,
91
  repo_id=HF_REPO, repo_type="model")
 
181
  log("terrain: himalaya (per-env crops via domain randomization)")
182
  else:
183
  log("terrain: playground stock rough terrain")
184
+ if SNOW:
185
+ from himalaya_terrain import snow_randomizer
186
+ randomization_fn = snow_randomizer(randomization_fn, SNOW_FRICTION, SNOW_DEPTH)
187
+ log(f"snow: foot-floor friction U{SNOW_FRICTION}, depth U{SNOW_DEPTH} m "
188
+ "(soft contact, per env)")
189
  ppo_params = locomotion_params.brax_ppo_config(ENV_NAME)
190
  ppo_params.num_timesteps = int(STATE["target"])
191
  log(f"num_envs={ppo_params.get('num_envs')} "
 
258
  t = STATE["thread"]
259
  if t is not None and t.is_alive():
260
  return status_md()
261
+ if timesteps:
262
+ STATE["target"] = int(timesteps)
 
 
 
263
  STATE["step"] = 0
264
  STATE["stop"] = False
265
  STATE["thread"] = threading.Thread(target=train_worker, daemon=True)
 
277
  alive = STATE["thread"] is not None and STATE["thread"].is_alive()
278
  target = int(STATE["target"])
279
  pct = 100.0 * STATE["step"] / max(target, 1)
280
+ return (f"**{ENV_NAME}** · terrain `{SURFACE}` — status: `{STATE['status']}`"
281
  f"{' (running)' if alive else ''} \n"
282
  f"step {STATE['step']:,} / {target:,} ({pct:.1f}%) \n"
283
  f"checkpoints: `{OUT}`"
 
293
  gr.Markdown("# Unitree G1 — Himalayan terrain locomotion training")
294
  st = gr.Markdown(status_md())
295
  with gr.Row():
296
+ steps_in = gr.Number(value=NUM_TIMESTEPS, precision=0, label="timesteps",
297
+ minimum=100_000, maximum=2_000_000_000)
298
+ gr.Button("Start", variant="primary").click(start, inputs=steps_in, outputs=st)
299
+ gr.Button("Stop", variant="stop").click(stop, outputs=st)
300
  out = gr.Textbox(label="training log", lines=26, max_lines=26,
301
  autoscroll=True, value=logs())
302
  timer = gr.Timer(3.0)
himalaya_terrain.py CHANGED
@@ -109,3 +109,60 @@ def randomizer(base_fn, grids: np.ndarray):
109
  return model, in_axes
110
 
111
  return fn
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
109
  return model, in_axes
110
 
111
  return fn
112
+
113
+
114
+ # --------------------------------------------------------------------------
115
+ # Snow
116
+ # --------------------------------------------------------------------------
117
+ # Playground's G1 touches the floor only through the two explicit foot-floor
118
+ # contact pairs (pair 0 = left foot, 1 = right foot), so snow is expressed on
119
+ # `pair_friction` / `pair_solref` / `pair_solimp` and overrides the geom params.
120
+ #
121
+ # Depth is a soft-contact calibration, not a separate layer: with a 35 kg block
122
+ # on the G1's two-foot footprint, `depth` metres of snow settles ~0.5*depth at
123
+ # rest and 2*depth under a landing impact (solimp d0=0.3 keeps the surface firm
124
+ # enough that a foot can't drop straight through; the hfield base is 1 m thick
125
+ # anyway). depth=0 reproduces the stock rigid contact (solref 0.02).
126
+ SNOW_FRICTION = (0.3, 0.7) # packed snow ~0.5; stock randomizer is (0.4, 1.0)
127
+ SNOW_DEPTH = (0.0, 0.08) # metres, ~0 to 3 in
128
+
129
+
130
+ def snow_params(depth):
131
+ """(solref[2], solimp[5]) for `depth` metres of snow (jnp or float)."""
132
+ import jax.numpy as jnp
133
+ tc = 0.02 + 1.5 * depth
134
+ width = jnp.maximum(2.0 * depth, 1e-3)
135
+ solref = jnp.stack([tc, jnp.ones_like(tc)], axis=-1)
136
+ solimp = jnp.stack([0.3 * jnp.ones_like(tc), 0.95 * jnp.ones_like(tc), width,
137
+ 0.5 * jnp.ones_like(tc), 2.0 * jnp.ones_like(tc)], axis=-1)
138
+ return solref, solimp
139
+
140
+
141
+ def snow_randomizer(base_fn, friction=SNOW_FRICTION, depth=SNOW_DEPTH):
142
+ """Wrap a Playground domain randomizer so every env also gets its own snow:
143
+ foot-floor friction ~U(*friction) and a soft-contact depth ~U(*depth)."""
144
+ import jax
145
+ import jax.numpy as jnp
146
+
147
+ def fn(model, rng):
148
+ model, in_axes = base_fn(model, rng)
149
+ n = rng.shape[0]
150
+ k_f, k_d = jax.random.split(jax.random.fold_in(rng[0], 7))
151
+ mu = jax.random.uniform(k_f, (n,), minval=friction[0], maxval=friction[1])
152
+ dp = jax.random.uniform(k_d, (n,), minval=depth[0], maxval=depth[1])
153
+
154
+ def batched(name):
155
+ x = getattr(model, name)
156
+ return x if getattr(in_axes, name) is not None else jnp.broadcast_to(x, (n,) + x.shape)
157
+
158
+ pf = batched("pair_friction")
159
+ pf = pf.at[:, 0:2, 0:2].set(mu[:, None, None])
160
+ sr, si = snow_params(dp) # (n,2), (n,5)
161
+ psr = batched("pair_solref").at[:, 0:2, :].set(sr[:, None, :])
162
+ psi = batched("pair_solimp").at[:, 0:2, :].set(si[:, None, :])
163
+
164
+ model = model.tree_replace({"pair_friction": pf, "pair_solref": psr, "pair_solimp": psi})
165
+ in_axes = in_axes.tree_replace({"pair_friction": 0, "pair_solref": 0, "pair_solimp": 0})
166
+ return model, in_axes
167
+
168
+ return fn