Update MJX notebook to work with brax 0.10.0.

PiperOrigin-RevId: 605105172
Change-Id: Icd209ada7977439ed7d5213588412d37dd54ffc0
This commit is contained in:
Baruch Tabanpour
2024-02-07 14:47:57 -08:00
committed by Copybara-Service
parent e4fada1963
commit 3baab3784a
+40 -37
View File
@@ -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