From 3baab3784ac4486e9adc8884ba3edaf26e26327d Mon Sep 17 00:00:00 2001 From: Baruch Tabanpour Date: Wed, 7 Feb 2024 14:47:57 -0800 Subject: [PATCH] Update MJX notebook to work with brax 0.10.0. PiperOrigin-RevId: 605105172 Change-Id: Icd209ada7977439ed7d5213588412d37dd54ffc0 --- mjx/tutorial.ipynb | 77 ++++++++++++++++++++++++---------------------- 1 file changed, 40 insertions(+), 37 deletions(-) diff --git a/mjx/tutorial.ipynb b/mjx/tutorial.ipynb index ade2a83b..c6056318 100644 --- a/mjx/tutorial.ipynb +++ b/mjx/tutorial.ipynb @@ -155,7 +155,6 @@ "cell_type": "code", "execution_count": null, "metadata": { - "cellView": "form", "id": "ObF1UXrkb0Nd" }, "outputs": [], @@ -175,7 +174,7 @@ "from brax import envs\n", "from brax import math\n", "from brax.base import Base, Motion, Transform\n", - "from brax.envs.base import Env, MjxEnv, State\n", + "from brax.envs.base import Env, PipelineEnv, State\n", "from brax.mjx.base import State as MjxState\n", "from brax.training.agents.ppo import train as ppo\n", "from brax.training.agents.ppo import networks as ppo_networks\n", @@ -418,7 +417,7 @@ "source": [ "#@title Humanoid Env\n", "\n", - "class Humanoid(MjxEnv):\n", + "class Humanoid(PipelineEnv):\n", "\n", " def __init__(\n", " self,\n", @@ -432,7 +431,7 @@ " **kwargs,\n", " ):\n", " path = epath.Path(epath.resource_path('mujoco')) / (\n", - " 'mjx/benchmark/model/humanoid'\n", + " 'mjx/test_data/humanoid'\n", " )\n", " mj_model = mujoco.MjModel.from_xml_path(\n", " (path / 'humanoid.xml').as_posix())\n", @@ -440,11 +439,14 @@ " mj_model.opt.iterations = 6\n", " mj_model.opt.ls_iterations = 6\n", "\n", + " sys = mjcf.load_model(mj_model)\n", + "\n", " physics_steps_per_control_step = 5\n", " kwargs['n_frames'] = kwargs.get(\n", " 'n_frames', physics_steps_per_control_step)\n", + " kwargs['backend'] = 'mjx'\n", "\n", - " super().__init__(model=mj_model, **kwargs)\n", + " super().__init__(sys, **kwargs)\n", "\n", " self._forward_reward_weight = forward_reward_weight\n", " self._ctrl_cost_weight = ctrl_cost_weight\n", @@ -470,7 +472,7 @@ "\n", " data = self.pipeline_init(qpos, qvel)\n", "\n", - " obs = self._get_obs(data.data, jp.zeros(self.sys.nu))\n", + " obs = self._get_obs(data, jp.zeros(self.sys.nu))\n", " reward, done, zero = jp.zeros(3)\n", " metrics = {\n", " 'forward_reward': zero,\n", @@ -490,8 +492,8 @@ " data0 = state.pipeline_state\n", " data = self.pipeline_step(data0, action)\n", "\n", - " com_before = data0.data.subtree_com[1]\n", - " com_after = data.data.subtree_com[1]\n", + " com_before = data0.subtree_com[1]\n", + " com_after = data.subtree_com[1]\n", " velocity = (com_after - com_before) / self.dt\n", " forward_reward = self._forward_reward_weight * velocity[0]\n", "\n", @@ -505,7 +507,7 @@ "\n", " ctrl_cost = self._ctrl_cost_weight * jp.sum(jp.square(action))\n", "\n", - " obs = self._get_obs(data.data, action)\n", + " obs = self._get_obs(data, action)\n", " reward = forward_reward + healthy_reward - ctrl_cost\n", " done = 1.0 - is_healthy if self._terminate_when_unhealthy else 0.0\n", " state.metrics.update(\n", @@ -604,7 +606,7 @@ "source": [ "## Train Humanoid Policy\n", "\n", - "Let's now train a policy with PPO to make the Humanoid run forwards. Training takes about 9-10 minutes on a Tesla A100 GPU." + "Let's now train a policy with PPO to make the Humanoid run forwards. Training takes about 8-9 minutes on a Tesla A100 GPU." ] }, { @@ -764,7 +766,7 @@ }, "outputs": [], "source": [ - "mj_model = eval_env._model\n", + "mj_model = eval_env.sys.mj_model\n", "mj_data = mujoco.MjData(mj_model)\n", "\n", "renderer = mujoco.Renderer(mj_model)\n", @@ -962,7 +964,7 @@ " return default_config\n", "\n", "\n", - "class BarkourEnv(MjxEnv):\n", + "class BarkourEnv(PipelineEnv):\n", " \"\"\"Environment for training the barkour quadruped joystick policy in MJX.\"\"\"\n", "\n", " def __init__(\n", @@ -973,18 +975,19 @@ " **kwargs,\n", " ):\n", " path = epath.Path('mujoco_menagerie/google_barkour_vb/scene_mjx.xml')\n", + " sys = mjcf.load(path.as_posix())\n", " self._dt = 0.02 # this environment is 50 fps\n", - " self.brax_sys = mjcf.load(path).replace(dt=self._dt)\n", - " model = self.brax_sys.get_model()\n", - " model.opt.timestep = 0.004\n", + " sys = sys.tree_replace({'opt.timestep': 0.004, 'dt': 0.004})\n", "\n", " # override menagerie params for smoother policy\n", - " model.dof_damping[6:] = 0.5239\n", - " model.actuator_gainprm[:, 0] = 35.0\n", - " model.actuator_biasprm[:, 1] = -35.0\n", + " sys = sys.replace(\n", + " dof_damping=sys.dof_damping.at[6:].set(0.5239),\n", + " actuator_gainprm=sys.actuator_gainprm.at[:, 0].set(35.0),\n", + " actuator_biasprm=sys.actuator_biasprm.at[:, 1].set(-35.0),\n", + " )\n", "\n", - " n_frames = kwargs.pop('n_frames', int(self._dt / model.opt.timestep))\n", - " super().__init__(model=model, n_frames=n_frames)\n", + " n_frames = kwargs.pop('n_frames', int(self._dt / sys.opt.timestep))\n", + " super().__init__(sys, backend='mjx', n_frames=n_frames)\n", "\n", " self.reward_config = get_config()\n", " # set custom from kwargs\n", @@ -993,13 +996,13 @@ " self.reward_config.rewards.scales[k[:-6]] = v\n", "\n", " self._torso_idx = mujoco.mj_name2id(\n", - " model, mujoco.mjtObj.mjOBJ_BODY.value, 'torso'\n", + " sys.mj_model, mujoco.mjtObj.mjOBJ_BODY.value, 'torso'\n", " )\n", " self._action_scale = action_scale\n", " self._obs_noise = obs_noise\n", " self._kick_vel = kick_vel\n", - " self._init_q = jp.array(model.keyframe('home').qpos)\n", - " self._default_pose = model.keyframe('home').qpos[7:]\n", + " self._init_q = jp.array(sys.mj_model.keyframe('home').qpos)\n", + " self._default_pose = sys.mj_model.keyframe('home').qpos[7:]\n", " self.lowers = jp.array([-0.7, -1.0, 0.05] * 4)\n", " self.uppers = jp.array([0.52, 2.1, 2.1] * 4)\n", " feet_site = [\n", @@ -1009,7 +1012,7 @@ " 'foot_hind_right',\n", " ]\n", " feet_site_id = [\n", - " mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_SITE.value, f)\n", + " mujoco.mj_name2id(sys.mj_model, mujoco.mjtObj.mjOBJ_SITE.value, f)\n", " for f in feet_site\n", " ]\n", " assert not any(id_ == -1 for id_ in feet_site_id), 'Site not found.'\n", @@ -1021,13 +1024,13 @@ " 'lower_leg_hind_right',\n", " ]\n", " lower_leg_body_id = [\n", - " mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY.value, l)\n", + " mujoco.mj_name2id(sys.mj_model, mujoco.mjtObj.mjOBJ_BODY.value, l)\n", " for l in lower_leg_body\n", " ]\n", " assert not any(id_ == -1 for id_ in lower_leg_body_id), 'Body not found.'\n", " self._lower_leg_body_id = np.array(lower_leg_body_id)\n", " self._foot_radius = 0.0175\n", - " self._nv = model.nv\n", + " self._nv = sys.nv\n", "\n", " def sample_command(self, rng: jax.Array) -\u003e jax.Array:\n", " lin_vel_x = [-0.6, 1.5] # min max [m/s]\n", @@ -1081,9 +1084,9 @@ " kick_theta = jax.random.uniform(kick_noise_2, maxval=2 * jp.pi)\n", " kick = jp.array([jp.cos(kick_theta), jp.sin(kick_theta)])\n", " kick *= jp.mod(state.info['step'], push_interval) == 0\n", - " qvel = state.pipeline_state.data.qvel # pytype: disable=attribute-error\n", + " qvel = state.pipeline_state.qvel # pytype: disable=attribute-error\n", " qvel = qvel.at[:2].set(kick * self._kick_vel + qvel[:2])\n", - " state = state.tree_replace({'pipeline_state.data.qvel': qvel})\n", + " state = state.tree_replace({'pipeline_state.qvel': qvel})\n", "\n", " # physics step\n", " motor_targets = self._default_pose + action * self._action_scale\n", @@ -1097,7 +1100,7 @@ " joint_vel = pipeline_state.qd[6:]\n", "\n", " # foot contact data based on z-position\n", - " foot_pos = pipeline_state.data.site_xpos[self._feet_site_id] # pytype: disable=attribute-error\n", + " foot_pos = pipeline_state.site_xpos[self._feet_site_id] # pytype: disable=attribute-error\n", " foot_contact_z = foot_pos[:, 2] - self._foot_radius\n", " contact = foot_contact_z \u003c 1e-3 # a mm or less off the floor\n", " contact_filt_mm = contact | state.info['last_contact']\n", @@ -1123,7 +1126,7 @@ " 'lin_vel_z': self._reward_lin_vel_z(xd),\n", " 'ang_vel_xy': self._reward_ang_vel_xy(xd),\n", " 'orientation': self._reward_orientation(x),\n", - " 'torques': self._reward_torques(pipeline_state.data.qfrc_actuator), # pytype: disable=attribute-error\n", + " 'torques': self._reward_torques(pipeline_state.qfrc_actuator), # pytype: disable=attribute-error\n", " 'action_rate': self._reward_action_rate(action, state.info['last_act']),\n", " 'stand_still': self._reward_stand_still(\n", " state.info['command'], joint_angles,\n", @@ -1267,8 +1270,8 @@ " ) -\u003e jax.Array:\n", " # get velocities at feet which are offset from lower legs\n", " # pytype: disable=attribute-error\n", - " pos = pipeline_state.data.site_xpos[self._feet_site_id] # feet position\n", - " feet_offset = pos - pipeline_state.data.xpos[self._lower_leg_body_id]\n", + " pos = pipeline_state.site_xpos[self._feet_site_id] # feet position\n", + " feet_offset = pos - pipeline_state.xpos[self._lower_leg_body_id]\n", " # pytype: enable=attribute-error\n", " offset = base.Transform.create(pos=feet_offset)\n", " foot_indices = self._lower_leg_body_id - 1 # we got rid of the world body\n", @@ -1284,7 +1287,7 @@ " self, trajectory: List[base.State], camera: str | None = None\n", " ) -\u003e Sequence[np.ndarray]:\n", " camera = camera or 'track'\n", - " return super().render(trajectory, camera)\n", + " return super().render(trajectory, camera=camera)\n", "\n", "envs.register_environment('barkour', BarkourEnv)" ] @@ -1445,7 +1448,7 @@ }, "outputs": [], "source": [ - "HTML(html.render(eval_env.brax_sys, rollout))" + "HTML(html.render(eval_env.sys, rollout))" ] } ], @@ -1453,13 +1456,13 @@ "accelerator": "GPU", "colab": { "gpuClass": "premium", - "gpuType": "A100", + "gpuType": "V100", "machine_shape": "hm", "private_outputs": true, "provenance": [ { - "file_id": "11cFRVCJ8Kn71tlQFbFcw4JzQZ00F8BRG", - "timestamp": 1704355889284 + "file_id": "1A58SK07tnOzix53E68D0TQ2ePCTZA61f", + "timestamp": 1707342610876 } ], "toc_visible": true