Update MJX notebook to work with brax 0.10.0.
PiperOrigin-RevId: 605105172 Change-Id: Icd209ada7977439ed7d5213588412d37dd54ffc0
This commit is contained in:
committed by
Copybara-Service
parent
e4fada1963
commit
3baab3784a
+40
-37
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user