update comments

This commit is contained in:
Andrew
2024-04-19 06:44:59 +02:00
parent 3e12452bc9
commit ab7d8fe167
+10 -11
View File
@@ -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",