Spaces:
Sleeping
Sleeping
Add snow to the Himalayan terrain: per-env foot-floor friction and soft-contact depth
#7
by arminfg - opened
- README.md +9 -1
- app.py +19 -23
- 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
|
| 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/{
|
| 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"{
|
| 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 |
-
|
| 263 |
-
|
| 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 `{
|
| 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 |
-
|
| 302 |
-
|
| 303 |
-
gr.Button("Stop", variant="stop").click(stop, outputs=st
|
| 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
|