Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
137 changes: 70 additions & 67 deletions examples/gradient_accumulation_and_microbatching.ipynb
Original file line number Diff line number Diff line change
@@ -1,21 +1,23 @@
{
"cells": [
{
"cell_type": "markdown",
"metadata": {
"id": "-2Bmeh2sl_fp"
},
"cell_type": "markdown",
"source": [
"# Summary\n",
"\n",
"The purpose of this notebook is to demonstrate example usages of utilities defined by the optax microbatching API. microbatching is a general purpose function transformation that lifts a function that operates over a batch to one that operates over a potentially much larger batch, by splitting up the work into smaller chunks and accumulating the results. Like other jax transformations, it's designed to be quite general - any function that can normally be traced by other jax transformations should work here. This notebook is broken up into multiple sections to illustrate usages of different functions in the API."
]
},
{
"cell_type": "code",
"execution_count": 1,
"metadata": {
"id": "5LUEGvWml-mH"
},
"cell_type": "code",
"outputs": [],
"source": [
"# ! pip install -q \"optax @ git+https://github.com/google-deepmind/optax\"\n",
"import jax\n",
Expand All @@ -26,15 +28,13 @@
"import time\n",
"from optax import microbatching\n",
"import gc"
],
"outputs": [],
"execution_count": 1
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "T5Ie-_tfqgAV"
},
"cell_type": "markdown",
"source": [
"# Setup\n",
"\n",
Expand All @@ -43,10 +43,12 @@
]
},
{
"cell_type": "code",
"execution_count": 13,
"metadata": {
"id": "nInGNjjeqcvU"
},
"cell_type": "code",
"outputs": [],
"source": [
"class TransformerBlock(nnx.Module):\n",
" def __init__(self, hidden_size: int, num_heads: int, *, rngs: nnx.Rngs):\n",
Expand Down Expand Up @@ -79,15 +81,15 @@
"\n",
" logits = self.final_layer(x)\n",
" return logits"
],
"outputs": [],
"execution_count": 13
]
},
{
"cell_type": "code",
"execution_count": 14,
"metadata": {
"id": "QEcf3p7fqoJ0"
},
"cell_type": "code",
"outputs": [],
"source": [
"hidden_size = 512\n",
"num_heads = 8\n",
Expand All @@ -107,15 +109,15 @@
"\n",
"adamw = optax.adamw(0.01)\n",
"opt_state = adamw.init(params)"
],
"outputs": [],
"execution_count": 14
]
},
{
"cell_type": "code",
"execution_count": 15,
"metadata": {
"id": "iR9qzYp6qjDS"
},
"cell_type": "code",
"outputs": [],
"source": [
"def loss_fn(params, batch):\n",
" model = nnx.merge(graphdef, params)\n",
Expand All @@ -133,26 +135,24 @@
" return params, opt_state\n",
"\n",
"update_fn.lower(params, opt_state, batch).compile()"
],
"outputs": [],
"execution_count": 15
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "KlRFCRc5nHLP"
},
"cell_type": "markdown",
"source": [
"# Part 1: Microbatching for Gradient Accumulation\n",
"\n",
"This section compares three methdos for performing gradient accumulation in jax/optax, one of which is through the microbatching API."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "COxGVTPUq9bZ"
},
"cell_type": "markdown",
"source": [
"### Option 1: Manual Gradient Accumulation\n",
"\n",
Expand All @@ -164,10 +164,12 @@
]
},
{
"cell_type": "code",
"execution_count": 8,
"metadata": {
"id": "SEf4ohj-q56F"
},
"cell_type": "code",
"outputs": [],
"source": [
"@functools.partial(jax.jit, donate_argnums=(2,))\n",
"def add_gradient(params, batch, accumulated_gradients):\n",
Expand All @@ -184,15 +186,15 @@
"accumulated_gradients = optax.tree.zeros_like(params)\n",
"add_gradient.lower(params, batch, accumulated_gradients).compile()\n",
"update_params.lower(params, opt_state, accumulated_gradients).compile()"
],
"outputs": [],
"execution_count": 8
]
},
{
"cell_type": "code",
"execution_count": 9,
"metadata": {
"id": "FjIvz_kRrAXH"
},
"cell_type": "code",
"outputs": [],
"source": [
"start_time = time.perf_counter()\n",
"for i in range(accumulation_steps):\n",
Expand All @@ -203,26 +205,26 @@
")\n",
"end_time = time.perf_counter()\n",
"print('Total Time', end_time - start_time)"
],
"outputs": [],
"execution_count": 9
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "gPfNKr5RrHmq"
},
"cell_type": "markdown",
"source": [
"### Option 2: optax.MultiSteps\n",
"\n",
"By wrapping our optimizer with optax.MultiSteps, we can have optax handle the gradient accumulation for us. Now we only have to define and compile a single update_fn, which is slightly simpler. The opt_state now keeps track of the accumulated gradients for us. This is more convenient as we have only a single jitted function now."
]
},
{
"cell_type": "code",
"execution_count": 11,
"metadata": {
"id": "vTMm84VyrDZV"
},
"cell_type": "code",
"outputs": [],
"source": [
"multi_adam = optax.MultiSteps(adamw, accumulation_steps)\n",
"\n",
Expand All @@ -235,15 +237,15 @@
"\n",
"multi_opt_state = multi_adam.init(params)\n",
"update_fn_v2.lower(params, batch, multi_opt_state).compile()"
],
"outputs": [],
"execution_count": 11
]
},
{
"cell_type": "code",
"execution_count": 12,
"metadata": {
"id": "qi-zg9-5rGsg"
},
"cell_type": "code",
"outputs": [],
"source": [
"start_time = time.perf_counter()\n",
"for i in range(accumulation_steps):\n",
Expand All @@ -252,33 +254,33 @@
"jax.block_until_ready((params, multi_opt_state))\n",
"end_time = time.perf_counter()\n",
"print('Total Time', end_time - start_time)"
],
"outputs": [],
"execution_count": 12
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "SBGm3wY5rUIH"
},
"cell_type": "markdown",
"source": [
"### Option 3: `microbatching.microbatch`\n",
"\n",
"microbatching differs from the approach above in that it transfers the entire batch of data to device memory, then splits it up perfoming the forward-backward pass on smaller batches and accumulating them using jax.lax.scan. Like Option 2 above, the full train step can be written as a single jitted function, however now the train step is doing 16X as much work."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "I5HrKPuVrQiA"
},
"cell_type": "markdown",
"source": []
},
{
"cell_type": "code",
"execution_count": 14,
"metadata": {
"id": "aBatfdC7rO--"
},
"cell_type": "code",
"outputs": [],
"source": [
"@functools.partial(jax.jit, donate_argnums=(0, 2))\n",
"def update_fn_v3(params, batch, opt_state):\n",
Expand All @@ -293,40 +295,40 @@
"\n",
"full_batch = jnp.vstack([batch]*accumulation_steps)\n",
"update_fn_v3.lower(params, full_batch, opt_state).compile()"
],
"outputs": [],
"execution_count": 14
]
},
{
"cell_type": "code",
"execution_count": 17,
"metadata": {
"id": "JKxDiYD9rXMF"
},
"cell_type": "code",
"outputs": [],
"source": [
"start_time = time.perf_counter()\n",
"params, opt_state = jax.block_until_ready(update_fn_v3(params, full_batch, opt_state))\n",
"end_time = time.perf_counter()\n",
"print('Total Time', end_time - start_time)"
],
"outputs": [],
"execution_count": 17
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "MZ2Xwu9znYYl"
},
"cell_type": "markdown",
"source": [
"# Part 2: `microbatching.micro_vmap`\n",
"\n",
"micro_vmap combines microbatching with jax.vmap, providing a new transformation with a similar API as jax.vmap, but that works with much larger batches than jax.vmap. It is especially useful when the function being vmapped requries more memory than that of the inputs/outputs for intermediates, or if you want to aggregate across the vmapped dimension."
]
},
{
"cell_type": "code",
"execution_count": 12,
"metadata": {
"id": "fd4XmvYjnoh3"
},
"cell_type": "code",
"outputs": [],
"source": [
"def expensive_function(x):\n",
" return jax.nn.softmax(jnp.sin(jnp.outer(x, x))).sum(axis=0)\n",
Expand All @@ -341,41 +343,41 @@
"gc.collect()\n",
"print('Processed small batch', result.shape)\n",
"print(result)"
],
"outputs": [],
"execution_count": 12
]
},
{
"cell_type": "code",
"execution_count": 11,
"metadata": {
"id": "nFGg2x5erilE"
},
"cell_type": "code",
"outputs": [],
"source": [
"result = microbatching.micro_vmap(expensive_function, microbatch_size=32)(X)\n",
"result = jax.block_until_ready(result)\n",
"gc.collect()\n",
"print('Processed Full Batch', result.shape)\n",
"print(result)"
],
"outputs": [],
"execution_count": 11
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "PWuT3BU1xJ2S"
},
"cell_type": "markdown",
"source": [
"# Part 3: `microbatching.micro_grad`\n",
"\n",
"micro_grad provides a simple and performant way to compute a sum or average of transformed per-example grads. While normally computing per-example gradients with jax is more expensive than computing normal gradients, and fail to run for the same batch sizes, the microbatching provides a sound mechanism to bypass this issue that we surface through the convenient and familiar API. Below we use the API to collect metrics about the per-example gradients, which can be useful for understanding and debugging the behavior of training runs."
]
},
{
"cell_type": "code",
"execution_count": 25,
"metadata": {
"id": "3hPNgICJuskf"
},
"cell_type": "code",
"outputs": [],
"source": [
"def metrics_fn(per_example_grad):\n",
" leaf_norms = jax.tree.map(jnp.linalg.norm, per_example_grad)\n",
Expand All @@ -384,34 +386,35 @@
"grad_fn = microbatching.micro_grad(loss_fn, metrics_fn=metrics_fn, microbatch_size=8)\n",
"\n",
"grad, aux = jax.jit(grad_fn)(params, batch)"
],
"outputs": [],
"execution_count": 25
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "ph9T_WAWykls"
},
"cell_type": "code",
"outputs": [],
"source": [
"# This shows the norm for the gradient of the embedding layer per example.\n",
"# High uniformity of the norm values is encouraging.\n",
"aux.metrics['embedding']['embedding'].get_value()"
],
"outputs": [],
"execution_count": 29
"aux.metrics['embedding']['embedding'].value"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "XaMThtUpnakg"
},
"cell_type": "markdown",
"source": []
}
],
"metadata": {
"colab": {
"private_outputs": true
},
"language_info": {
"name": "python"
}
},
"nbformat": 4,
Expand Down
Loading