Skip to content

Commit 2ebfae6

Browse files
authored
Merge pull request #345 from 136dxj/main
feat:improve handstand performance
2 parents 61ad9a6 + b3c484d commit 2ebfae6

2 files changed

Lines changed: 125 additions & 20 deletions

File tree

‎conf/ppo/task/go2_handstand/mujoco.yaml‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,7 @@ env:
66
sim_dt: 0.005
77
algo:
88
num_envs: 1024
9-
max_iterations: 1510
9+
max_iterations: 3000
1010
obs_groups:
1111
actor:
1212
- actor
@@ -25,6 +25,8 @@ reward:
2525
penalty_contact: -0.2
2626
action_rate: -0.01
2727
tar: 0.3
28+
feet_air_time: 1
29+
world_z_vel_penalty: -1
2830
tracking_sigma: 0.25
2931
base_height_target: 0.3
3032

‎src/unilab/envs/locomotion/go2/handstand.py‎

Lines changed: 122 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -98,8 +98,21 @@ def _compute_reset_obs(
9898
dof_vel: Any,
9999
) -> dict[str, np.ndarray]:
100100
height = env.torso_height[env_ids].reshape(-1, 1)
101+
env.feet_phase[env_ids, :] = 0
102+
env.feet_phase[:, 2] = 0.0 # RL starts at 0
103+
env.feet_phase[:, 3] = 0.5
104+
# Reset air time tracking
105+
env._feet_air_time[env_ids, :] = 0.0
106+
env._last_contacts[env_ids, :] = False
107+
101108
return env._compute_obs( # type: ignore[no-any-return]
102-
info_updates, linvel, gyro, gravity, dof_pos, dof_vel, height
109+
info_updates,
110+
linvel,
111+
gyro,
112+
gravity,
113+
dof_pos,
114+
dof_vel,
115+
height, # , env.feet_phase[env_ids, 2:]
103116
)
104117

105118

@@ -128,8 +141,14 @@ def __init__(self, cfg: Go2HandStandCfg, num_envs=1, backend_type="mujoco"):
128141
self._init_domain_randomization(Go2HandStandDomainRandomizationProvider())
129142
self.phase = np.zeros((num_envs,), dtype=np.float32)
130143
self.feet_phase = np.zeros((num_envs, len(cfg.sensor.feet_force)), dtype=np.float32)
131-
self.gait_frequency = 2
144+
# Initialize rear leg (RL, RR) phases with alternating values for gait
145+
self.feet_phase[:, 2] = 0.0 # RL starts at 0
146+
self.feet_phase[:, 3] = 0.5 # RR starts at 0.5 (alternating)
147+
self.gait_frequency = 2 # Slower stepping: 2 seconds per cycle
132148
self.feet_force = np.zeros((num_envs, len(cfg.sensor.feet_force), 1), dtype=np.float32)
149+
# Track air time for stepping reward
150+
self._feet_air_time = np.zeros((num_envs, len(cfg.sensor.feet_force)), dtype=np.float32)
151+
self._last_contacts = np.zeros((num_envs, 2), dtype=bool)
133152
self.feet_pos = np.zeros((num_envs, len(cfg.sensor.feet_pos), 3), dtype=np.float32)
134153
self.torso_height = np.zeros((num_envs,), dtype=np.float32)
135154
self._z_des = 0.55
@@ -153,6 +172,8 @@ def _init_reward_functions(self):
153172
"penalty_contact": self._reward_penalty_contact,
154173
"action_rate": rewards.action_rate,
155174
"tar": self._reward_tar,
175+
"feet_air_time": self._reward_feet_air_time,
176+
"world_z_vel_penalty": self._reward_world_z_vel_penalty,
156177
}
157178

158179
def update_state(self, state: NpEnvState) -> NpEnvState:
@@ -172,18 +193,58 @@ def update_state(self, state: NpEnvState) -> NpEnvState:
172193
arr = self._backend.get_sensor_data(name)
173194
contact_arrays.append(arr)
174195
result = np.concatenate(contact_arrays, axis=1)
196+
# print(linvel)
197+
# Update phase for stepping pattern (only when high enough)
198+
height_mask = self.torso_height >= self._z_des * 0.8
199+
dt = self._cfg.ctrl_dt
200+
phase_increment = self.gait_frequency * dt
201+
self.feet_phase = (self.feet_phase + phase_increment) % 1.0
202+
# Only allow stepping for rear legs (2=RL, 3=RR) when height is sufficient
203+
self.feet_phase[:, :2] = 0.0 # Front legs phase locked at 0 (should be in stance)
204+
# Reset rear leg phase when height is too low
205+
for i in [2, 3]:
206+
self.feet_phase[~height_mask, i] = 0.0
207+
208+
# Update feet air time for stepping reward (only rear legs)
209+
contact = self.feet_force[:, [2, 3], 0] > 1.0
210+
contact_filt = np.logical_or(contact, self._last_contacts)
211+
self._last_contacts = contact
212+
# Increment air time
213+
self._feet_air_time[:, [2, 3]] += self._cfg.ctrl_dt
214+
# Reset air time for feet in contact
215+
self._feet_air_time[:, [2, 3]] *= ~contact_filt
175216

176217
terminated_z = gravity[:, 2] <= -0.25
177218
terminated_contact = np.any(result, axis=1)
219+
# After 100 steps, terminate if height is too low (failed to maintain target)
220+
# step_count = state.info.get("steps", np.zeros((self._num_envs,), dtype=np.uint32))
221+
# terminated_height = (step_count >= 100) & (self.torso_height < self._z_des * 0.8)
222+
# terminated = np.logical_or(
223+
# np.logical_or(terminated_contact, terminated_z),
224+
# terminated_height
225+
# )
178226
terminated = np.logical_or(terminated_contact, terminated_z)
179227
reward = self._compute_reward(state.info, linvel, gyro, dof_pos)
180228
obs = self._compute_obs(
181-
state.info, linvel, gyro, gravity, dof_pos, dof_vel, self.torso_height.reshape(-1, 1)
229+
state.info,
230+
linvel,
231+
gyro,
232+
gravity,
233+
dof_pos,
234+
dof_vel,
235+
self.torso_height.reshape(-1, 1), # , self.feet_phase[:,[2,3]]
182236
)
183237
return state.replace(obs=obs, reward=reward, terminated=terminated)
184238

185239
def _compute_obs(
186-
self, info: dict, linvel, gyro, gravity, dof_pos, dof_vel, height
240+
self,
241+
info: dict,
242+
linvel,
243+
gyro,
244+
gravity,
245+
dof_pos,
246+
dof_vel,
247+
height, # , feet_phase
187248
) -> dict[str, np.ndarray]:
188249
noise_cfg = self._cfg.noise_config
189250
diff = dof_pos - self.default_angles
@@ -257,14 +318,6 @@ def _compute_reward(self, info: dict, linvel, gyro, dof_pos) -> np.ndarray:
257318

258319
# ── reward functions (robot-specific) ────────────────────────────
259320

260-
def _reward_swing_feet_z(self, ctx: RewardContext) -> np.ndarray:
261-
is_swing = self.feet_phase >= 0.6
262-
target_height = 0.1
263-
height_error = np.square(self.feet_pos[:, :, 2] - target_height)
264-
swing_rew = np.exp(-height_error / 0.01) * is_swing
265-
reward: np.ndarray = np.sum(swing_rew, axis=1) / len(self._cfg.sensor.feet_pos)
266-
return reward
267-
268321
def _reward_foot_drag(self, ctx: RewardContext) -> np.ndarray:
269322
foot_pos = self.get_foot_pos()
270323
foot_heights = foot_pos[..., 2]
@@ -284,18 +337,45 @@ def _reward_penalty_contact(self, ctx: RewardContext) -> np.ndarray:
284337
result = np.concatenate(contact_arrays, axis=1)
285338
return np.asarray(np.any(result, axis=1))
286339

287-
def _reward_contact(self, ctx: RewardContext) -> np.ndarray:
288-
contact = self.feet_force[:, :, 2] > 0.1
340+
def _reward_stand_contact(self, ctx: RewardContext) -> np.ndarray:
341+
# res = np.zeros(self._num_envs, dtype=np.float32)
342+
# # When height is above 0.8 * target, encourage rear leg (RL, RR) stepping
343+
# height_mask = self.torso_height >= self._z_des * 0.8
344+
# for i in [2, 3]: # Rear legs only
345+
# is_stance = self.feet_phase[:, i] < 0.6
346+
# target_height = np.where(is_stance, 0.0, 0.1)
347+
# foot_height = self.feet_pos[:, i, 2]
348+
# height_error = np.abs(foot_height - target_height)
349+
# # Exponential reward for matching target height
350+
# foot_rew = np.exp(-height_error / 0.02)
351+
# res += foot_rew * height_mask.astype(np.float32)
352+
# return res / 2 # Average over 2 rear legs
353+
contact = self.feet_force[:, :, 0] > 0.1
354+
height_mask = self.torso_height >= self._z_des * 0.8
289355
res = np.zeros(self._num_envs, dtype=np.float32)
290-
for i in range(len(self._cfg.sensor.feet_force)):
356+
for i in [2, 3]:
291357
is_contact = (self.feet_phase[:, i] < 0.6) | (self.gait_frequency < 1.0e-8)
292-
res += (contact[:, i] == is_contact).astype(np.float32)
358+
res += (contact[:, i] == is_contact).astype(np.float32) * height_mask.astype(np.float32)
293359
return res / len(self._cfg.sensor.feet_force)
294360

361+
def _reward_swing_feet_z(self, ctx: RewardContext) -> np.ndarray:
362+
is_swing = self.feet_phase >= 0.6
363+
height_mask = self.torso_height >= self._z_des * 0.8
364+
target_height = 0.1
365+
height_error = np.square(self.feet_pos[:, [2, 3], 2] - target_height)
366+
swing_rew = np.exp(-height_error / 0.01) * is_swing[:, 2:]
367+
reward: np.ndarray = (
368+
np.sum(swing_rew, axis=1)
369+
/ len(self._cfg.sensor.feet_pos)
370+
* height_mask.astype(np.float32)
371+
)
372+
return reward
373+
295374
def _reward_height(self, ctx: RewardContext) -> np.ndarray:
296-
height = np.minimum(self.torso_height, self._z_des)
297-
error = self._z_des - height
298-
return np.exp(-error / 0.25)
375+
# height = np.minimum(self.torso_height, self._z_des)
376+
height = self.torso_height
377+
error = np.abs(self._z_des - height)
378+
return np.exp(-error / 0.1)
299379

300380
def _reward_orientation(self, ctx: RewardContext) -> np.ndarray:
301381
gravity = -1 * self._backend.get_sensor_data("upvector")
@@ -325,5 +405,28 @@ def _reward_tar(self, ctx: RewardContext) -> np.ndarray:
325405

326406
return cast(np.ndarray, np.exp(-error / 1) * mask)
327407

408+
def _reward_feet_air_time(self, ctx: RewardContext) -> np.ndarray:
409+
"""Reward rear legs (RL, RR) for long steps - reward on first contact."""
410+
# Only apply when robot is high enough
411+
height_mask = self.torso_height >= self._z_des * 0.8
412+
# Target air time - reward for staying in air longer than this
413+
target_air_time = 0.2
414+
# First contact detection: air_time > 0 and currently in contact
415+
contact = self.feet_force[:, [2, 3], 0] > 1.0
416+
first_contact = (self._feet_air_time[:, [2, 3]] > 0.0) & contact
417+
# Reward: (actual_air_time - target) for feet making first contact
418+
rew = (self._feet_air_time[:, [2, 3]] - target_air_time) * first_contact
419+
# Sum over rear legs and apply height mask
420+
return np.sum(rew, axis=1) * height_mask
421+
422+
def _reward_world_z_vel_penalty(self, ctx: RewardContext) -> np.ndarray:
423+
"""Penalize vertical velocity after standing up to prevent bouncing."""
424+
# Only apply when robot is high enough
425+
height_mask = self.torso_height >= self._z_des * 0.8
426+
# Get world frame z velocity
427+
world_z_vel = self._backend.get_base_lin_vel()[:, 2]
428+
# Penalize absolute z velocity (both up and down)
429+
return np.abs(world_z_vel) * height_mask
430+
328431
# def _cost_pose(self, qpos: jax.Array) -> jax.Array:
329432
# return jp.sum(jp.square(qpos[self._joint_ids] - self._joint_pose))

0 commit comments

Comments
 (0)