update comments
This commit is contained in:
+10
-11
@@ -521,17 +521,17 @@
|
||||
" pipeline_state=data, obs=obs, reward=reward,\n",
|
||||
" done=done)\n",
|
||||
" \n",
|
||||
" #### Demonstration Replay\n",
|
||||
" #### Reset state to reference if it gets too far\n",
|
||||
" error = (((x.pos - ref_x.pos) ** 2).sum(-1)**0.5).mean()\n",
|
||||
" replay = jp.where(error > self.err_threshold, 1.0, 0.0)\n",
|
||||
" to_reference = jp.where(error > self.err_threshold, 1.0, 0.0)\n",
|
||||
"\n",
|
||||
" replay = jp.array(replay, dtype=int) # keeps output types same as input. \n",
|
||||
" to_reference = jp.array(to_reference, dtype=int) # keeps output types same as input. \n",
|
||||
" ref_data = self.mjx_to_brax(ref_data)\n",
|
||||
"\n",
|
||||
" data = jax.tree_util.tree_map(lambda x, y: \n",
|
||||
" jp.array((1-replay)*x + replay*y, x.dtype), data, ref_data)\n",
|
||||
" jp.array((1-to_reference)*x + to_reference*y, x.dtype), data, ref_data)\n",
|
||||
" \n",
|
||||
" x, xd = data.x, data.xd # Data may or may not have changed.\n",
|
||||
" x, xd = data.x, data.xd # Data may have changed.\n",
|
||||
" obs = self._get_obs(data.qpos, x, xd, state.info)\n",
|
||||
" \n",
|
||||
" return state.replace(pipeline_state=data, obs=obs)\n",
|
||||
@@ -562,6 +562,9 @@
|
||||
" return obs\n",
|
||||
" \n",
|
||||
" def mjx_to_brax(self, data):\n",
|
||||
" \"\"\" \n",
|
||||
" Apply the brax wrapper on the core MJX data structure.\n",
|
||||
" \"\"\"\n",
|
||||
" q, qd = data.qpos, data.qvel\n",
|
||||
" x = Transform(pos=data.xpos[1:], rot=data.xquat[1:])\n",
|
||||
" cvel = Motion(vel=data.cvel[1:, 3:], ang=data.cvel[1:, :3])\n",
|
||||
@@ -609,7 +612,7 @@
|
||||
" return pos_err + vel_err\n",
|
||||
"\n",
|
||||
" def _reward_feet_height(self, feet_pos, feet_pos_ref):\n",
|
||||
" return jp.sum(jp.abs(feet_pos - feet_pos_ref)) # try to drive it to 0 with l1 norm\n",
|
||||
" return jp.sum(jp.abs(feet_pos - feet_pos_ref)) # try to drive it to 0 using the l1 norm.\n",
|
||||
"\n",
|
||||
"envs.register_environment('trotting_anymal', TrotAnymal)"
|
||||
]
|
||||
@@ -1140,9 +1143,7 @@
|
||||
" # Tracking of linear velocity commands (xy axes)\n",
|
||||
" local_vel = math.rotate(xd.vel[0], math.quat_inv(x.rot[0]))\n",
|
||||
" lin_vel_error = jp.sum(jp.square(commands[:2] - local_vel[:2]))\n",
|
||||
" lin_vel_reward = jp.exp(\n",
|
||||
" -lin_vel_error\n",
|
||||
" )\n",
|
||||
" lin_vel_reward = jp.exp(-lin_vel_error)\n",
|
||||
" return lin_vel_reward\n",
|
||||
" def _reward_orientation(self, x: Transform) -> jax.Array:\n",
|
||||
" # Penalize non flat base orientation\n",
|
||||
@@ -1168,7 +1169,6 @@
|
||||
" rews = jp.sum(jp.square(feet_pos - xyt), axis=1) \n",
|
||||
" rews = jp.clip(rews, 0, 10)\n",
|
||||
" return jp.sum(rews)\n",
|
||||
" \n",
|
||||
" def _reward_feet_height(self, data, state_info):\n",
|
||||
" \"\"\" \n",
|
||||
" Feet height tracks rectified sine waves \n",
|
||||
@@ -1205,7 +1205,6 @@
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Reconstruct the trotting inference function\n",
|
||||
"\n",
|
||||
"make_networks_factory = functools.partial(\n",
|
||||
" apg_networks.make_apg_networks,\n",
|
||||
" hidden_layer_sizes=(256, 128)\n",
|
||||
|
||||
Reference in New Issue
Block a user