Add Franka Panda env to MJX colab tutorial.

PiperOrigin-RevId: 682506877
Change-Id: Ifde31e993f63baa783e973c496e5bd7bd6465393
This commit is contained in:
Baruch Tabanpour
2024-10-04 17:35:19 -07:00
committed by Copybara-Service
parent 6be5decf6c
commit 0887f54307
+341
View File
@@ -415,6 +415,7 @@
"cell_type": "code",
"execution_count": 0,
"metadata": {
"cellView": "form",
"id": "mtGMYNLE3QJN"
},
"outputs": [],
@@ -1602,6 +1603,346 @@
" eval_env.render(rollout[::render_every], camera='track'),\n",
" fps=1.0 / eval_env.dt / render_every)"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "gluTlHURuC6i"
},
"source": [
"# Manipulation Environments and Policies\n",
"\n",
"By now, we have shown how MJX can be used to train policies for classic control and robotic locomotion. MJX can also be used for robotic manipulation!\n",
"\n",
"We demonstrate a task on the Franka Panda below, which trains a policy to pickup a cube and bring it to a mocap target position in about 3 minutes on an A100. We will be adding more support for manipulation environments in MJX (i.e. more performant collisions), so stay tuned!\n",
"\n"
]
},
{
"cell_type": "code",
"execution_count": 0,
"metadata": {
"id": "RCv16hZIu5Dm"
},
"outputs": [],
"source": [
"%%shell\n",
"if [ ! -d \"mujoco_menagerie\" ]; then\n",
" git clone https://github.com/google-deepmind/mujoco_menagerie\n",
"fi\n"
]
},
{
"cell_type": "code",
"execution_count": 0,
"metadata": {
"id": "AuU-9nUquEu4"
},
"outputs": [],
"source": [
"#@title Franka Panda BringToTarget Environment\n",
"\n",
"FRANKA_PANDA_ROOT_PATH = epath.Path('mujoco_menagerie/franka_emika_panda')\n",
"\n",
"\n",
"def default_config():\n",
" \"\"\"Returns reward config for the environment.\"\"\"\n",
"\n",
" return config_dict.create(\n",
" # Environment timestep. Should match the robot decision frequency.\n",
" dt=0.02,\n",
" # Lowers action magnitude for less-jerky motion. Also sometimes helps\n",
" # sample efficiency.\n",
" action_scale=0.04,\n",
" # The coefficients for all reward terms used for training.\n",
" reward_scales=config_dict.create(\n",
" # Gripper goes to the box.\n",
" gripper_box=4.0,\n",
" # Box goes to the target mocap.\n",
" box_target=8.0,\n",
" # Do not collide the gripper with the floor.\n",
" no_floor_collision=0.25,\n",
" # Arm stays close to target pose.\n",
" robot_target_qpos=0.3,\n",
" ),\n",
" )\n",
"\n",
"\n",
"def _load_sys(path: epath.Path) -> base.System:\n",
" \"\"\"Load a mujoco model from a path.\"\"\"\n",
" assets = {}\n",
" for f in path.parent.glob('*.xml'):\n",
" assets[f.name] = f.read_bytes()\n",
" for f in (path.parent / 'assets').glob('*'):\n",
" assets[f.name] = f.read_bytes()\n",
" xml = path.read_text()\n",
" model = mujoco.MjModel.from_xml_string(xml, assets)\n",
" return mjcf.load_model(model)\n",
"\n",
"\n",
"def _get_collision_info(\n",
" contact: Any, geom1: int, geom2: int) -> Tuple[jax.Array, jax.Array]:\n",
" if geom1 > geom2:\n",
" geom1, geom2 = geom2, geom1\n",
" mask = (jp.array([geom1, geom2]) == contact.geom).all(axis=1)\n",
" idx = jp.where(mask, contact.dist, 1e4).argmin()\n",
" dist = contact.dist[idx] * mask[idx]\n",
" normal = (dist < 0) * contact.frame[idx, 0, :3]\n",
" return dist, normal\n",
"\n",
"\n",
"def _geoms_colliding(\n",
" state: Optional[State], geom1: int, geom2: int\n",
") -> jax.Array:\n",
" return _get_collision_info(state.contact, geom1, geom2)[0] < 0\n",
"\n",
"\n",
"class PandaBringToTarget(PipelineEnv):\n",
" \"\"\"Environment for training franka panda to bring an object to target.\"\"\"\n",
"\n",
" def __init__(self, **kwargs):\n",
" global root_path\n",
" sys = _load_sys(FRANKA_PANDA_ROOT_PATH / 'mjx_single_cube.xml')\n",
" self._config = config = default_config()\n",
" nsteps = int(np.round(config.dt / sys.opt.timestep))\n",
" kwargs['backend'] = 'mjx'\n",
" kwargs['n_frames'] = nsteps\n",
" super().__init__(sys, **kwargs)\n",
"\n",
" # define constants\n",
" model = sys.mj_model\n",
" arm_joints = ['joint1', 'joint2', 'joint3', 'joint4', 'joint5',\n",
" 'joint6', 'joint7']\n",
" finger_joints = ['finger_joint1', 'finger_joint2']\n",
" all_joints = arm_joints + finger_joints\n",
" self._robot_arm_qposadr = np.array([\n",
" model.jnt_qposadr[model.joint(j).id] for j in arm_joints])\n",
" self._robot_qposadr = np.array([\n",
" model.jnt_qposadr[model.joint(j).id] for j in all_joints])\n",
" self._gripper_site = model.site('gripper').id\n",
" self._left_finger_geom = model.geom('left_finger_pad').id\n",
" self._right_finger_geom = model.geom('right_finger_pad').id\n",
" self._hand_geom = model.geom('hand_capsule').id\n",
" self._box_body = model.body('box').id\n",
" self._box_qposadr = model.jnt_qposadr[model.body('box').jntadr[0]]\n",
" # TODO(btaba): replace with mocap_pos once MJX version 3.2.3 is released.\n",
" self._target_id = model.body('mocap_target').id\n",
" self._floor_geom = model.geom('floor').id\n",
" self._init_q = sys.mj_model.keyframe('home').qpos\n",
" self._init_box_pos = jp.array(\n",
" self._init_q[self._box_qposadr : self._box_qposadr + 3],\n",
" dtype=jp.float32)\n",
" self._init_ctrl = sys.mj_model.keyframe('home').ctrl\n",
" self._lowers = model.actuator_ctrlrange[:, 0]\n",
" self._uppers = model.actuator_ctrlrange[:, 1]\n",
"\n",
" def reset(self, rng: jax.Array) -> State:\n",
" rng, rng_box, rng_target = jax.random.split(rng, 3)\n",
"\n",
" # intialize box position\n",
" box_pos = jax.random.uniform(\n",
" rng_box, (3,),\n",
" minval=jp.array([-0.2, -0.2, 0.0]),\n",
" maxval=jp.array([0.2, 0.2, 0.0])) + self._init_box_pos\n",
"\n",
" # initialize target position\n",
" target_pos = jax.random.uniform(\n",
" rng_target, (3,),\n",
" minval=jp.array([-0.2, -0.2, 0.2]),\n",
" maxval=jp.array([0.2, 0.2, 0.4])) + self._init_box_pos\n",
"\n",
" # initialize pipeline state\n",
" init_q = jp.array(self._init_q).at[\n",
" self._box_qposadr : self._box_qposadr + 3].set(box_pos)\n",
" pipeline_state = self.pipeline_init(\n",
" init_q, jp.zeros(self.sys.nv)\n",
" )\n",
" pipeline_state = pipeline_state.replace(ctrl=self._init_ctrl)\n",
" # set target mocap position\n",
" # TODO(btaba): replace with mocap_pos once MJX version 3.2.3 is released.\n",
" pipeline_state = pipeline_state.replace(\n",
" xpos=pipeline_state.xpos.at[self._target_id, :].set(target_pos))\n",
"\n",
" # initialize env state and info\n",
" metrics = {\n",
" 'out_of_bounds': jp.array(0.0),\n",
" **{k: 0.0 for k in self._config.reward_scales.keys()},\n",
" }\n",
" info = {'rng': rng, 'target_pos': target_pos, 'reached_box': 0.0}\n",
" obs = self._get_obs(pipeline_state, info)\n",
" reward, done = jp.zeros(2)\n",
" state = State(pipeline_state, obs, reward, done, metrics, info)\n",
" return state\n",
"\n",
" def step(self, state: State, action: jax.Array) -> State:\n",
" delta = action * self._config.action_scale\n",
" ctrl = state.pipeline_state.ctrl + delta\n",
" ctrl = jp.clip(ctrl, self._lowers, self._uppers)\n",
"\n",
" # step the physics\n",
" data = self.pipeline_step(state.pipeline_state, ctrl)\n",
"\n",
" # compute reward terms\n",
" target_pos = state.info['target_pos']\n",
" box_pos = data.xpos[self._box_body]\n",
" gripper_pos = data.site_xpos[self._gripper_site]\n",
" box_target = 1 - jp.tanh(5 * jp.linalg.norm(target_pos - box_pos))\n",
" gripper_box = 1 - jp.tanh(5 * jp.linalg.norm(box_pos - gripper_pos))\n",
" robot_target_qpos = 1 - jp.tanh(\n",
" jp.linalg.norm(\n",
" state.pipeline_state.qpos[self._robot_arm_qposadr]\n",
" - self._init_q[self._robot_arm_qposadr]\n",
" )\n",
" )\n",
"\n",
" hand_floor_collision = [\n",
" _geoms_colliding(state.pipeline_state, self._floor_geom, g)\n",
" for g in [\n",
" self._left_finger_geom,\n",
" self._right_finger_geom,\n",
" self._hand_geom,\n",
" ]\n",
" ]\n",
" floor_collision = sum(hand_floor_collision) > 0\n",
" no_floor_collision = 1 - floor_collision\n",
"\n",
" state.info['reached_box'] = 1.0 * jp.maximum(\n",
" state.info['reached_box'],\n",
" (jp.linalg.norm(box_pos - gripper_pos) < 0.012),\n",
" )\n",
"\n",
" rewards = {\n",
" 'box_target': box_target * state.info['reached_box'],\n",
" 'gripper_box': gripper_box,\n",
" 'no_floor_collision': no_floor_collision,\n",
" 'robot_target_qpos': robot_target_qpos,\n",
" }\n",
" rewards = {k: v * self._config.reward_scales[k] for k, v in rewards.items()}\n",
" reward = jp.clip(sum(rewards.values()), -1e4, 1e4)\n",
"\n",
" out_of_bounds = jp.any(jp.abs(box_pos) > 1.0)\n",
" out_of_bounds |= box_pos[2] < 0.0\n",
" state.metrics.update(\n",
" out_of_bounds=out_of_bounds.astype(float),\n",
" **rewards)\n",
"\n",
" obs = self._get_obs(data, state.info)\n",
" done = out_of_bounds | jp.isnan(data.qpos).any() | jp.isnan(data.qvel).any()\n",
" done = done.astype(float)\n",
" state = State(data, obs, reward, done, state.metrics, state.info)\n",
"\n",
" return state\n",
"\n",
" def _get_obs(self, data: PipelineState, info: dict[str, Any]) -> jax.Array:\n",
" gripper_pos = data.site_xpos[self._gripper_site]\n",
" gripper_mat = data.site_xmat[self._gripper_site].ravel()\n",
" obs = jp.concatenate([\n",
" data.qpos,\n",
" data.qvel,\n",
" gripper_pos,\n",
" gripper_mat[3:],\n",
" data.xmat[self._box_body].ravel()[3:],\n",
" data.xpos[self._box_body] - data.site_xpos[self._gripper_site],\n",
" info['target_pos'] - data.xpos[self._box_body],\n",
" data.ctrl - data.qpos[self._robot_qposadr[:-1]],\n",
" ])\n",
"\n",
" return obs\n",
"\n",
"envs.register_environment('PandaBringToTarget', PandaBringToTarget)"
]
},
{
"cell_type": "code",
"execution_count": 0,
"metadata": {
"id": "76g9uILMQVkc"
},
"outputs": [],
"source": [
"# instantiate the environment\n",
"env_name = 'PandaBringToTarget'\n",
"env = envs.get_environment(env_name)\n",
"\n",
"# define the jit reset/step functions\n",
"jit_reset = jax.jit(env.reset)\n",
"jit_step = jax.jit(env.step)"
]
},
{
"cell_type": "code",
"execution_count": 0,
"metadata": {
"id": "10_vs9IDvnke"
},
"outputs": [],
"source": [
"#@title Train Pick-up-cube Policy\n",
"\n",
"make_networks_factory = functools.partial(\n",
" ppo_networks.make_ppo_networks,\n",
" policy_hidden_layer_sizes=(32, 32, 32, 32))\n",
"\n",
"train_fn = functools.partial(\n",
" ppo.train, num_timesteps=20_000_000, num_evals=4, reward_scaling=0.1,\n",
" episode_length=150, normalize_observations=True, action_repeat=1,\n",
" unroll_length=10, num_minibatches=32, num_updates_per_batch=8,\n",
" discounting=0.97, learning_rate=1e-3, entropy_cost=2e-2, num_envs=2048,\n",
" batch_size=512, num_resets_per_eval=1,\n",
" network_factory=make_networks_factory, seed=0)\n",
"\n",
"\n",
"x_data, y_data, y_dataerr = [], [], []\n",
"times = [datetime.now()]\n",
"def progress(num_steps, metrics):\n",
" times.append(datetime.now())\n",
" x_data.append(num_steps)\n",
" y_data.append(metrics['eval/episode_reward'])\n",
" y_dataerr.append(metrics['eval/episode_reward_std'])\n",
"\n",
" plt.xlim([0, train_fn.keywords['num_timesteps'] * 1.25])\n",
" plt.ylim([0, 2000])\n",
" plt.xlabel('# environment steps')\n",
" plt.ylabel('reward per episode')\n",
" plt.title(f'y={y_data[-1]:.3f}')\n",
" plt.errorbar(x_data, y_data, yerr=y_dataerr)\n",
" plt.show()\n",
"\n",
"make_inference_fn, params, _= train_fn(environment=env, progress_fn=progress)\n",
"jit_inference_fn = jax.jit(make_inference_fn(params, deterministic=True))\n",
"\n",
"print(f'time to jit: {times[1] - times[0]}')\n",
"print(f'time to train: {times[-1] - times[1]}')\n"
]
},
{
"cell_type": "code",
"execution_count": 0,
"metadata": {
"id": "jDJLcI0Bv5lD"
},
"outputs": [],
"source": [
"# initialize the state\n",
"rng = jax.random.PRNGKey(0)\n",
"state = jit_reset(rng)\n",
"rollout = [state.pipeline_state]\n",
"\n",
"# grab a trajectory\n",
"n_steps = 150\n",
"render_every = 2\n",
"\n",
"for i in range(n_steps):\n",
" act_rng, rng = jax.random.split(rng)\n",
" ctrl, _ = jit_inference_fn(state.obs, act_rng)\n",
" state = jit_step(state, ctrl)\n",
" rollout.append(state.pipeline_state)\n",
"\n",
" if state.done:\n",
" break\n",
"\n",
"media.show_video(env.render(rollout[::render_every]), fps=1.0 / env.dt / render_every)"
]
}
],
"metadata": {