{
"cells": [
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Policy learning and Policy Gradients\n",
"\n",
"This is a recap of policy learning contextualizes how we can use MJX's differentiability for policy learning. If the below concepts are unfamiliar, there are many great resources [online](https://spinningup.openai.com/en/latest/spinningup/rl_intro.html)! \n",
"\n",
"The goal of policy learning is to find a control policy $\\pi$ which outputs actions $a_t \\sim \\pi(\\cdot| x_t, \\theta)$ maximizing the total rewards $\\sum r_t$ over some time period, where $r_t$ is shorthand for a reward function evaluated at the state and action of time t:\n",
"\n",
"$$r_t = r(x_t, a_t)$$\n",
"\n",
"$\\theta$ are the parameters of the policy; weights in the common case that the policy is a neural network. **Policy gradient methods** involve estimating the gradient of the rewards with respect to the weights, and using this value in a first-order optimization algorithm such as Gradient Descent or [Adam](https://arxiv.org/abs/1412.6980). How we estimate the policy gradient depends on what state transition model we assume. \n",
"\n",
"#### Zeroth-Order Policy Gradients (ZoPG)\n",
"\n",
"Referring to `mjx.step` as the simulation function f, we borrow some [terminology](https://arxiv.org/abs/2202.00817) to differentiate between zeroth-order gradients, which only depend on values of f, and first-order gradients, which depend on its jacobian.\n",
"\n",
"Reinforcement learning (RL) algorithms such as the standard [PPO](https://github.com/google/brax/blob/main/brax/training/agents/ppo/train.py) assume the stochastic state transition model $x_{t+1} \\sim P(\\cdot | x_t, a_t)$. This leads to a ZoPG of the form:\n",
"\n",
"$$\n",
"\\nabla_\\theta J(\\pi_\\theta) = \\mathbb{E}_{\\tau \\sim \\pi_\\theta}\\left[ \\sum \\nabla_\\theta \\log\\pi_\\theta (a_t | s_t) R(\\tau) \\right]\n",
"$$\n",
"\n",
"Where $R(\\tau)$ is some function depending on the rollout $\\tau = \\{x_t, a_t\\}_{t=0}^{T}$. Despite this method's popularity and extensive research into its refinement, a fundamental property is that the gradient has high variance. This allows the optimizer to thoroughly explore the space of policies, leading to the robust and often surprisingly good policies that have been achieved. However, the variance comes at the cost of requiring many samples $(x_t, a_t)$ to converge.\n",
"\n",
"#### First-Order Policy Gradients (FoPG)\n",
"On the other hand, if you assume a deterministic state transition model $x_{t+1} = f(x_t, a_t)$, you end up with the first-order policy gradient. Other common names include Analytical Policy Gradients (APG) and Backpropogation through Time (BPTT). Unlike ZoPG methods, which model the state evolution as a probabilistic black box, the FoPG explicitly contains the jacobians of the simulation function f. For example, let's look at the gradient of the reward $r_t$, in the case that it only depends on state.\n",
"$$\n",
"\\frac{\\partial r_t}{\\partial \\theta} = \\frac{\\partial r_t}{\\partial x_t}\\frac{\\partial x_t}{\\partial \\theta} \n",
"$$\n",
"\n",
"$$\n",
"\\frac{\\partial x_t}{\\partial \\theta} = \\textcolor{Navy}{\\frac{\\partial f(x_{t-1}, a_{t-1})}{\\partial x_{t-1}}}\\frac{\\partial x_{t-1}}{\\partial \\theta} + \\textcolor{Navy}{\\frac{\\partial f(x_{t-1}, a_{t-1})}{\\partial a_{t-1}}} \\frac{\\partial a_{t-1}}{\\partial \\theta}\n",
"$$\n",
"\n",
"The navy-colored terms in the above expression are enabled by MJX's differentiability and are the key difference between FoPG's and ZoPG's. An important consideration is what these jacobians look like near contact points. To see why certain gradients within the jacobian can be pathological, imagine a hard sphere falling toward a block of marble. How does its velocity change with respect to distance ($\\frac{\\partial \\dot{z}_t}{\\partial z_t}$), the instant before it touches the ground? This is the case of an **uninformative gradient**, due to [hard contact](https://arxiv.org/html/2404.02887v1). Fortunately, the default contact settings in Mujoco are sufficiently [soft](https://mujoco.readthedocs.io/en/stable/computation/index.html#soft-contact-model) for learning via FoPG's. With soft contacts, the ground applies an increasing force on the ball as it penetrates it, unlike rigid contacts, which instantly provide enough force for deflection.\n",
"\n",
"A helpful way to think about FoPG's is via the chain rule and computation graphs, as illustrated below for how $r_2$ influences the policy gradient, again for the case that the reward does not depend on action:\n",
"\n",
"
\n",
"\n",
"Note that there three distinct gradient chains in this example. The red pathway considers how the immediately prior action affected the state. The blue path explains the name *Backpropogation through Time*, capturing how actions affect downstream rewards. The least intuitive may be the green chain, which shows how the reward depends on how actions depend on previous actions.Experience shows that blocking *any* of these three pathways via jax.lax.stop_grad can badly hinder policy learning. As the length of $x_t$ backbone increases, [gradient explosion](https://arxiv.org/abs/2111.05803) becomes a crucial consideration. In practice, this can be resolved via decaying downstream gradients or periodically truncating the gradient.\n",
"\n",
"**The Sharp Bits of FoPG's**\n",
"\n",
"While FoPG's have been shown to be very sample efficient, especially as the [dimension of the state space increases](https://arxiv.org/abs/2204.07137), one fundamental shortcoming is that due to the lower gradient variance, FoPG's also have less exploration power than ZoPG's and benefit from the practioner being more explicit in the problem formulation.\n",
"\n",
"Additionally, discontinuous reward formulations are ubiquitious in RL, for instance, a large penalty when the robot falls. It can be significantly more [challenging](https://arxiv.org/abs/2403.14864) to design robust policies with FoPG's, since they cannot backprop through such penalties.\n",
"\n",
"Last, despite the sample efficiency, FoPG methods can still struggle with wall-clock time. Because the gradients have low variance, they do not benefit significantly from massive parallelization of data collection - unlike [RL](https://arxiv.org/abs/2109.11978). Additionally, the policy gradient is typically calculated via autodifferentiation. This can be 3-5x slower than unrolling the simulation forward, and memory intensive, with memory requirements scaling with $O(m \\cdot (m+n) \\cdot T)$, where m and n are the state and control dimensions, $m \\cdot (m+n)$ is the jacobian dimension, and T is the number of steps propogated through.\n",
"\n",
"Note that with certain models, using autodifferentiation through mjx.step currently causes [nan gradients](https://github.com/google-deepmind/mujoco/issues/1517). For now, we address this issue by using double-precision floats, at the cost of doubling the memory requirements and training time."
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"---\n",
"\n",
"**Coming Up**\n",
"\n",
"In this tutorial, we demonstrate two ways to use FoPG's, using Brax's simple APG [algorithm](https://github.com/google/brax/tree/main/brax/training/agents/apg). This algorithm essentially uses FoPG's to perform live gradient descent on the policy, unrolling it for a short window, using the data to do a policy update, then continuing where it left off."
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Setup: Imports and installations"
]
},
{
"cell_type": "code",
"execution_count": 1,
"metadata": {},
"outputs": [],
"source": [
"import os\n",
"os.environ[\"XLA_PYTHON_CLIENT_MEM_FRACTION\"] = \"0.8\" # 0.9 causes too much lag. \n",
"from datetime import datetime\n",
"import functools\n",
"\n",
"# Math\n",
"import jax.numpy as jp\n",
"import numpy as np\n",
"import jax\n",
"from jax import config # Analytical gradients work much better with double precision.\n",
"config.update(\"jax_debug_nans\", True)\n",
"config.update(\"jax_enable_x64\", True)\n",
"config.update('jax_default_matmul_precision', jax.lax.Precision.HIGH)\n",
"from brax import math\n",
"\n",
"# Sim\n",
"import mujoco\n",
"import mujoco.mjx as mjx\n",
"\n",
"# Brax\n",
"from brax import envs\n",
"from brax.base import Motion, Transform\n",
"from brax.io import mjcf\n",
"from brax.envs.base import PipelineEnv, State\n",
"from brax.mjx.pipeline import _reformat_contact\n",
"from brax.training.acme import running_statistics\n",
"from brax.io import model\n",
"\n",
"# Algorithms\n",
"from brax.training.agents.apg import train as apg\n",
"from brax.training.agents.apg import networks as apg_networks\n",
"from brax.training.agents.ppo import train as ppo\n",
"\n",
"# Supporting\n",
"from etils import epath\n",
"import mediapy as media\n",
"import matplotlib.pyplot as plt\n",
"from ml_collections import config_dict\n",
"from typing import Any, Dict\n"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"#### Quadruped Env"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"!git clone https://github.com/google-deepmind/mujoco_menagerie.git"
]
},
{
"cell_type": "code",
"execution_count": 3,
"metadata": {},
"outputs": [
{
"data": {
"text/html": [
"