Add a snow gait: higher swing, wider stance, slower cadence, foot-separation cost

#11
Files changed (3) hide show
  1. README.md +17 -0
  2. app.py +26 -6
  3. snow_gait.py +109 -0
README.md CHANGED
@@ -41,6 +41,23 @@ crop (`rollout.py`; MJX + `MUJOCO_GL=egl`) and reports how long it stayed up and
41
  how far it walked. `NUM_EVALS` (default 40) sets how many evals — and therefore
42
  log lines and checkpoint uploads — a run makes.
43
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
44
  ## Fall recovery (`TASK=getup`)
45
 
46
  Walking and getting up are different tasks, and the joystick task cannot learn
 
41
  how far it walked. `NUM_EVALS` (default 40) sets how many evals — and therefore
42
  log lines and checkpoint uploads — a run makes.
43
 
44
+ ## Snow gait (`GAIT=snow`)
45
+
46
+ Playground's joystick task is tuned for pavement: a 15 cm swing-height reference
47
+ at 1.25–1.5 Hz with the hips pinned to a narrow stance. `GAIT=snow` retunes it
48
+ for deep snow via `snow_gait.py` — a 22 cm swing (`FOOT_HEIGHT`), a slower
49
+ 1.0–1.3 Hz cadence (`GAIT_FREQ`) with more air-time reward for longer strides, a
50
+ relaxed hip-deviation cost so the stance can widen, and a new cost for letting
51
+ the feet come within `FOOT_SEPARATION` (20 cm) of each other *laterally*.
52
+
53
+ That last term is also a fix, not just styling. The joystick task ends an episode
54
+ when the two feet touch, and the 200M-step run plateaued because the gait crossed
55
+ its feet after ~45 steps and took the −100 termination every episode while still
56
+ upright. The separation cost penalises exactly that, and is measured along the
57
+ pelvis's y-axis so feet may still pass fore-aft as walking requires.
58
+
59
+ Checkpoints go under a `-snowgait` suffix.
60
+
61
  ## Fall recovery (`TASK=getup`)
62
 
63
  Walking and getting up are different tasks, and the joystick task cannot learn
app.py CHANGED
@@ -38,7 +38,15 @@ SNOW_DEPTH = tuple(float(v) for v in os.environ.get("SNOW_DEPTH", "0.0,0.08").sp
38
  # "getup" = fall recovery (g1_getup.py): most episodes start fallen, nothing
39
  # terminates on being down, and the model gains torso/pelvis collision geoms.
40
  TASK = os.environ.get("TASK", "walk").strip().lower()
41
- SURFACE = TERRAIN + ("-snow" if SNOW else "") + ("-getup" if TASK == "getup" else "")
 
 
 
 
 
 
 
 
42
  NUM_EVALS = int(os.environ.get("NUM_EVALS", 40)) # evals (and checkpoints) per run
43
 
44
  # /data exists only when persistent storage is attached; fall back to /tmp.
@@ -167,10 +175,21 @@ def build_envs():
167
  randomization_fn = None # no terrain or snow: this task is about the body
168
  ENVS.update(env=env, eval_env=eval_env, randomization_fn=randomization_fn)
169
  return env, eval_env, randomization_fn
170
- log(f"loading {ENV_NAME} ...")
171
- env = registry.load(ENV_NAME)
172
- env_cfg = registry.get_default_config(ENV_NAME)
173
- eval_env = registry.load(ENV_NAME, config=env_cfg)
 
 
 
 
 
 
 
 
 
 
 
174
  randomization_fn = registry.get_domain_randomizer(ENV_NAME)
175
  if TERRAIN == "himalaya":
176
  from himalaya_terrain import apply_terrain, make_terrains, randomizer
@@ -376,7 +395,8 @@ def render(ckpt: str | None, vx: float, friction: float, depth: float, seconds:
376
 
377
 
378
  with gr.Blocks(title="G1 rough-terrain training") as demo:
379
- gr.Markdown(f"# Unitree G1 — {'fall recovery' if TASK == 'getup' else 'Himalayan terrain locomotion'} training")
 
380
  st = gr.Markdown(status_md())
381
  with gr.Row():
382
  steps_in = gr.Number(value=NUM_TIMESTEPS, precision=0, label="timesteps",
 
38
  # "getup" = fall recovery (g1_getup.py): most episodes start fallen, nothing
39
  # terminates on being down, and the model gains torso/pelvis collision geoms.
40
  TASK = os.environ.get("TASK", "walk").strip().lower()
41
+ # "snow" retunes the walking gait for deep snow: higher swing, wider stance,
42
+ # slower cadence, and a lateral foot-separation cost (see snow_gait.py).
43
+ GAIT = os.environ.get("GAIT", "street").strip().lower()
44
+ FOOT_HEIGHT = float(os.environ.get("FOOT_HEIGHT", 0.22)) # swing reference, m
45
+ FOOT_SEPARATION = float(os.environ.get("FOOT_SEPARATION", 0.20)) # min lateral gap, m
46
+ GAIT_FREQ = tuple(float(v) for v in os.environ.get("GAIT_FREQ", "1.0,1.3").split(","))
47
+ SURFACE = (TERRAIN + ("-snow" if SNOW else "")
48
+ + ("-getup" if TASK == "getup" else "")
49
+ + ("-snowgait" if TASK == "walk" and GAIT == "snow" else ""))
50
  NUM_EVALS = int(os.environ.get("NUM_EVALS", 40)) # evals (and checkpoints) per run
51
 
52
  # /data exists only when persistent storage is attached; fall back to /tmp.
 
175
  randomization_fn = None # no terrain or snow: this task is about the body
176
  ENVS.update(env=env, eval_env=eval_env, randomization_fn=randomization_fn)
177
  return env, eval_env, randomization_fn
178
+ if GAIT == "snow":
179
+ import snow_gait
180
+ log(f"gait: snow -- swing {FOOT_HEIGHT * 100:.0f} cm (street 15), "
181
+ f"stance >= {FOOT_SEPARATION * 100:.0f} cm, {GAIT_FREQ[0]}-{GAIT_FREQ[1]} Hz "
182
+ "(street 1.25-1.5), relaxed hips")
183
+ cfg = snow_gait.snow_gait_config(foot_height=FOOT_HEIGHT,
184
+ foot_separation=FOOT_SEPARATION,
185
+ gait_freq=GAIT_FREQ)
186
+ env = snow_gait.G1SnowGait(config=cfg)
187
+ eval_env = snow_gait.G1SnowGait(config=cfg)
188
+ else:
189
+ log(f"loading {ENV_NAME} ...")
190
+ env = registry.load(ENV_NAME)
191
+ env_cfg = registry.get_default_config(ENV_NAME)
192
+ eval_env = registry.load(ENV_NAME, config=env_cfg)
193
  randomization_fn = registry.get_domain_randomizer(ENV_NAME)
194
  if TERRAIN == "himalaya":
195
  from himalaya_terrain import apply_terrain, make_terrains, randomizer
 
395
 
396
 
397
  with gr.Blocks(title="G1 rough-terrain training") as demo:
398
+ gr.Markdown(f"# Unitree G1 — {'fall recovery' if TASK == 'getup' else 'Himalayan terrain locomotion'} training"
399
+ + (" (snow gait)" if TASK == "walk" and GAIT == "snow" else ""))
400
  st = gr.Markdown(status_md())
401
  with gr.Row():
402
  steps_in = gr.Number(value=NUM_TIMESTEPS, precision=0, label="timesteps",
snow_gait.py ADDED
@@ -0,0 +1,109 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """A snow gait for the G1: high, wide, deliberate steps instead of a street walk.
2
+
3
+ Playground's joystick task rewards tracking a swing-height reference of 15 cm at
4
+ 1.25-1.5 Hz with the hips held near the nominal (narrow) pose. That is a gait for
5
+ pavement. Walking through snow, people lift their feet clear of the surface,
6
+ plant them wider apart, and take slower, longer strides.
7
+
8
+ Four changes, three of them just numbers the joystick task already exposes:
9
+
10
+ * `max_foot_height` 0.15 -> 0.22 m. This is the reference trajectory the
11
+ `feet_phase` reward tracks (`gait.get_rz`), so raising it literally asks for
12
+ a higher swing -- the direct "pick your feet up" knob.
13
+ * `joint_deviation_hip` -0.25 -> -0.05. That cost pins the hips (including
14
+ hip *roll*) to the nominal narrow stance; relaxing it permits a wide one.
15
+ * gait frequency 1.25-1.5 Hz -> 1.0-1.3 Hz, and `feet_air_time` 2.0 -> 3.0.
16
+ Slower cadence at the same commanded speed means a longer stride.
17
+ * a new `feet_separation` cost, which the joystick task has no equivalent of.
18
+
19
+ That last one is not cosmetic. The joystick task ends an episode when the two
20
+ feet touch each other, and the 200M-step run plateaued because the learned gait
21
+ crossed its feet after ~45 steps and ate the -100 termination every episode
22
+ while still upright. Penalising *lateral* separation below a threshold attacks
23
+ that directly, and a wide stance is what you want in snow anyway. The cost is
24
+ lateral-only, measured along the pelvis's y-axis, so the feet may still pass each
25
+ other fore-aft as any walking gait requires.
26
+ """
27
+
28
+ from __future__ import annotations
29
+
30
+ from typing import Any, Dict, Optional, Union
31
+
32
+ import jax
33
+ import jax.numpy as jp
34
+ from ml_collections import config_dict
35
+ from mujoco import mjx
36
+
37
+ from mujoco_playground._src.locomotion.g1 import joystick as g1_joystick
38
+
39
+ # Defaults, all overridable from the Space (see app.py).
40
+ FOOT_HEIGHT = 0.22 # m, swing-height reference (street gait: 0.15)
41
+ FOOT_SEPARATION = 0.20 # m, minimum lateral gap between the feet
42
+ SEPARATION_COST = -2.0 # scale for falling below that gap
43
+ HIP_DEVIATION_COST = -0.05 # street gait: -0.25
44
+ AIR_TIME_REWARD = 3.0 # street gait: 2.0
45
+ GAIT_FREQ = (1.0, 1.3) # Hz, street gait: (1.25, 1.5)
46
+
47
+
48
+ def snow_gait_config(
49
+ foot_height: float = FOOT_HEIGHT,
50
+ foot_separation: float = FOOT_SEPARATION,
51
+ separation_cost: float = SEPARATION_COST,
52
+ hip_deviation_cost: float = HIP_DEVIATION_COST,
53
+ air_time_reward: float = AIR_TIME_REWARD,
54
+ gait_freq: tuple[float, float] = GAIT_FREQ,
55
+ ) -> config_dict.ConfigDict:
56
+ cfg = g1_joystick.default_config()
57
+ rc = cfg.reward_config
58
+ rc.max_foot_height = foot_height
59
+ rc.min_foot_separation = foot_separation # new
60
+ rc.scales.feet_separation = separation_cost # new
61
+ rc.scales.joint_deviation_hip = hip_deviation_cost
62
+ rc.scales.feet_air_time = air_time_reward
63
+ cfg.gait_freq = list(gait_freq) # new
64
+ return cfg
65
+
66
+
67
+ class G1SnowGait(g1_joystick.Joystick):
68
+ """The joystick task, retuned for deep snow."""
69
+
70
+ def __init__(
71
+ self,
72
+ task: str = "rough_terrain",
73
+ config: Optional[config_dict.ConfigDict] = None,
74
+ config_overrides: Optional[Dict[str, Union[str, int, list[Any]]]] = None,
75
+ ) -> None:
76
+ super().__init__(task=task, config=config or snow_gait_config(),
77
+ config_overrides=config_overrides)
78
+ self._pelvis_body_id = self._mj_model.body("pelvis").id
79
+
80
+ def reset(self, rng: jax.Array) -> Any:
81
+ # The base class samples gait frequency internally from a hard-coded
82
+ # 1.25-1.5 Hz; phase_dt is the only thing it derives from it, so a slower
83
+ # cadence is a matter of rewriting that one entry after the fact.
84
+ rng, freq_rng = jax.random.split(rng)
85
+ state = super().reset(rng)
86
+ lo, hi = self._config.gait_freq
87
+ gait_freq = jax.random.uniform(freq_rng, (1,), minval=lo, maxval=hi)
88
+ state.info["phase_dt"] = 2 * jp.pi * self.dt * gait_freq
89
+ return state
90
+
91
+ def _get_reward(self, data: mjx.Data, action: jax.Array, info: dict[str, Any],
92
+ metrics: dict[str, Any], done: jax.Array,
93
+ first_contact: jax.Array, contact: jax.Array) -> dict[str, jax.Array]:
94
+ rewards = super()._get_reward(data, action, info, metrics, done,
95
+ first_contact, contact)
96
+ rewards["feet_separation"] = self._cost_feet_separation(data)
97
+ return rewards
98
+
99
+ def _cost_feet_separation(self, data: mjx.Data) -> jax.Array:
100
+ """How far the feet are inside the minimum lateral gap, in metres.
101
+
102
+ Measured along the pelvis's y-axis so that fore-aft passing during swing
103
+ costs nothing -- only the feet converging sideways does.
104
+ """
105
+ feet = data.site_xpos[self._feet_site_id] # (2, 3), world frame
106
+ lateral_axis = data.xmat[self._pelvis_body_id].reshape(3, 3)[:, 1]
107
+ lateral_gap = jp.abs(jp.dot(feet[0] - feet[1], lateral_axis))
108
+ return jp.clip(self._config.reward_config.min_foot_separation - lateral_gap,
109
+ 0.0, None)