From ab7d8fe167e7bd0b9cdbb3394af97e3911095e65 Mon Sep 17 00:00:00 2001 From: Andrew Date: Fri, 19 Apr 2024 06:44:59 +0200 Subject: [PATCH] update comments --- mjx/training_apg.ipynb | 21 ++++++++++----------- 1 file changed, 10 insertions(+), 11 deletions(-) diff --git a/mjx/training_apg.ipynb b/mjx/training_apg.ipynb index 4f1e4a7a..c11f931c 100644 --- a/mjx/training_apg.ipynb +++ b/mjx/training_apg.ipynb @@ -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",