{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### Intro to JAX\n[JAX](https://github.com/google/jax) is a framework which is used for high-performance numerical computing and machine learning research developed at [Google Research](https://research.google/) teams. It allows you to build Python applications with a NumPy-consistent API that specializes in differentiating, vectorizing, parallelizing, and compiling to GPU/TPU Just-In-Time. JAX was designed with performance and speed as a first priority, and is natively compatible with common machine learning accelerators such as [GPUs](https://www.kaggle.com/docs/efficient-gpu-usage) and [TPUs](https://www.kaggle.com/docs/tpu). Large ML models can take ages to train -- you might be interested in using JAX for applications where speed and performance are particularly important!\n### When to use JAX vs TensorFlow?\n[TensorFlow](https://www.tensorflow.org/guide) is a fantastic product, with a rich and fully-featured ecosystem, capable of supporting most every use case a machine learning practitioner might have (e.g. [TFLite](https://www.tensorflow.org/lite) for on-device inference computing, [TFHub](https://tfhub.dev/) for sharing pre-trained models, and many additional specialized applications as well). This type of broad mandate both contrasts and compliments JAX's philosophy, which is more narrowly focused on speed and performance.  We recommend using JAX in situations where you do want to maximize speed and performance but you do not require any of the long tail of features and additional functionalities that only the [TensorFlow ecosystem](https://www.tensorflow.org/learn) can provide.\n### Intro to the FLAX\nJust like [JAX](https://jax.readthedocs.io/en/latest/notebooks/quickstart.html) focuses on speed, other members of the JAX ecosystem are encouraged to specialize as well.  For example, [Flax](https://flax.readthedocs.io/en/latest/) focuses on neural networks and [jgraph](https://github.com/deepmind/jraph) focuses on graph networks.  \n\n[Flax](https://flax.readthedocs.io/en/latest/) is a JAX-based neural network library that was initially developed by  Google Research's Brain Team (in close collaboration with the JAX team) but is now open source.  If you want to train machine learning models on GPUs and TPUs at an accelerated speed, or if you have an ML project that might benefit from bringing together both [Autograd](https://github.com/hips/autograd) and [XLA](https://www.tensorflow.org/xla), consider using [Flax](https://flax.readthedocs.io/en/latest/) for your next project! [Flax](https://flax.readthedocs.io/en/latest/) is especially well-suited for projects that use large language models, and is a popular choice for cutting-edge [machine learning research](https://arxiv.org/search/?query=JAX&searchtype=all&abstracts=show&order=-announced_date_first&size=50).","metadata":{}},{"cell_type":"markdown","source":"## Imports","metadata":{}},{"cell_type":"code","source":"# Importing all the libraries necessary for the project\nimport jax\nimport flax\nimport numpy as np\nimport jax.numpy as jnp\nimport tensorflow as tf\nimport pandas as pd\nimport os\nfrom flax import linen as nn # the Linen API\nfrom flax.training import train_state \nimport optax\nimport matplotlib.pyplot as plt\n%matplotlib inline","metadata":{"execution":{"iopub.status.busy":"2022-07-12T09:47:35.324310Z","iopub.execute_input":"2022-07-12T09:47:35.324884Z","iopub.status.idle":"2022-07-12T09:47:48.521908Z","shell.execute_reply.started":"2022-07-12T09:47:35.324777Z","shell.execute_reply":"2022-07-12T09:47:48.520529Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# List all the available devices\njax.local_devices()","metadata":{"execution":{"iopub.status.busy":"2022-07-12T09:47:48.524218Z","iopub.execute_input":"2022-07-12T09:47:48.524993Z","iopub.status.idle":"2022-07-12T09:47:48.565323Z","shell.execute_reply.started":"2022-07-12T09:47:48.524956Z","shell.execute_reply":"2022-07-12T09:47:48.563293Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Load and Pre-process the dataset\nIn this notebook, we'll be using dataset from very famous Kaggle competition [Digit Recognizer](https://www.kaggle.com/c/digit-recognizer)","metadata":{}},{"cell_type":"code","source":"mnist_train = pd.read_csv('../input/digit-recognizer/train.csv')","metadata":{"execution":{"iopub.status.busy":"2022-07-12T09:47:48.567861Z","iopub.execute_input":"2022-07-12T09:47:48.568743Z","iopub.status.idle":"2022-07-12T09:47:52.470474Z","shell.execute_reply.started":"2022-07-12T09:47:48.568698Z","shell.execute_reply":"2022-07-12T09:47:52.468824Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = mnist_train[\"label\"]\nfeatures = mnist_train.drop(labels = [\"label\"],axis = 1)\nfeatures = features/255.0\nfeatures = features.values.reshape(-1,28,28,1)\nlabels = labels.to_numpy()","metadata":{"execution":{"iopub.status.busy":"2022-07-12T09:47:52.474655Z","iopub.execute_input":"2022-07-12T09:47:52.475253Z","iopub.status.idle":"2022-07-12T09:47:52.737804Z","shell.execute_reply.started":"2022-07-12T09:47:52.475198Z","shell.execute_reply":"2022-07-12T09:47:52.736021Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels.shape","metadata":{"execution":{"iopub.status.busy":"2022-07-12T09:47:52.739429Z","iopub.execute_input":"2022-07-12T09:47:52.739827Z","iopub.status.idle":"2022-07-12T09:47:52.750023Z","shell.execute_reply.started":"2022-07-12T09:47:52.739795Z","shell.execute_reply":"2022-07-12T09:47:52.748445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"features.shape","metadata":{"execution":{"iopub.status.busy":"2022-07-12T09:47:52.752843Z","iopub.execute_input":"2022-07-12T09:47:52.753738Z","iopub.status.idle":"2022-07-12T09:47:52.763023Z","shell.execute_reply.started":"2022-07-12T09:47:52.753686Z","shell.execute_reply":"2022-07-12T09:47:52.761813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Creating a dictionary with images and labels\n","metadata":{}},{"cell_type":"code","source":"train_ds = {\n        'images': features,\n        'labels': labels\n       }","metadata":{"execution":{"iopub.status.busy":"2022-07-12T09:47:52.764485Z","iopub.execute_input":"2022-07-12T09:47:52.765268Z","iopub.status.idle":"2022-07-12T09:47:52.774614Z","shell.execute_reply.started":"2022-07-12T09:47:52.765234Z","shell.execute_reply":"2022-07-12T09:47:52.773070Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define model architecture\nWe'll be using [FLAX Linen package](https://flax.readthedocs.io/en/latest/flax.linen.html) for defining the model architecture from scratch.","metadata":{}},{"cell_type":"code","source":"class CNN(nn.Module):\n    @nn.compact\n    def __call__(self, x):\n        x = nn.Conv(features=32, kernel_size=(3, 3))(x)\n        x = nn.relu(x)\n        x = nn.max_pool(x, window_shape=(2, 2), strides=(2, 2))\n        x = nn.Conv(features=64, kernel_size=(3, 3))(x)\n        x = nn.relu(x)\n        x = nn.max_pool(x, window_shape=(2, 2), strides=(2, 2))\n        x = x.reshape((x.shape[0], -1))\n        x = nn.Dense(features=256)(x)\n        x = nn.relu(x)\n        x = nn.Dense(features=10)(x)   \n        x = nn.log_softmax(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-07-12T09:52:59.730585Z","iopub.execute_input":"2022-07-12T09:52:59.731041Z","iopub.status.idle":"2022-07-12T09:52:59.743483Z","shell.execute_reply.started":"2022-07-12T09:52:59.731005Z","shell.execute_reply":"2022-07-12T09:52:59.741726Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define loss and compute metrics\nNow, we will define the functions which calculates the training loss and compute the metrics using the given predicted values and labels","metadata":{}},{"cell_type":"code","source":"def cross_entropy_loss(*, logits, labels):\n    one_hot_labels = jax.nn.one_hot(labels, num_classes=10)\n    return -jnp.mean(jnp.sum(one_hot_labels * logits, axis=-1))\n\ndef compute_metrics(logits, labels):\n    loss = cross_entropy_loss(logits=logits, labels=labels)\n    accuracy = jnp.mean(jnp.argmax(logits, -1) == labels)\n    metrics = {\n      'loss': loss,\n      'accuracy': accuracy,\n    }\n    return metrics","metadata":{"execution":{"iopub.status.busy":"2022-07-12T09:53:01.493258Z","iopub.execute_input":"2022-07-12T09:53:01.493679Z","iopub.status.idle":"2022-07-12T09:53:01.503038Z","shell.execute_reply.started":"2022-07-12T09:53:01.493648Z","shell.execute_reply":"2022-07-12T09:53:01.501422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Defining train step\nWe will now define the function for training a single batch of data, which will take the current train state and the training data as input and return the updated train state along with the compute metrics.","metadata":{}},{"cell_type":"code","source":"@jax.jit\ndef train_step(state, batch):\n    def loss_fn(params):\n        logits = CNN().apply({'params': params}, batch['images'])\n        loss = cross_entropy_loss(logits=logits, labels=batch['labels'])\n        return loss, logits\n    grad_fn = jax.value_and_grad(loss_fn, has_aux=True)\n    (_, logits), grads = grad_fn(state.params)\n    state = state.apply_gradients(grads=grads)\n    metrics = compute_metrics(logits, batch['labels'])\n    return state, metrics","metadata":{"execution":{"iopub.status.busy":"2022-07-12T09:53:03.457450Z","iopub.execute_input":"2022-07-12T09:53:03.457844Z","iopub.status.idle":"2022-07-12T09:53:03.466462Z","shell.execute_reply.started":"2022-07-12T09:53:03.457814Z","shell.execute_reply":"2022-07-12T09:53:03.465250Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"per_device_batch_size = 32\n\ntotal_batch_size = per_device_batch_size * jax.local_device_count()\nprint(\"The overall batch size (both for training and eval) is\", total_batch_size)","metadata":{"execution":{"iopub.status.busy":"2022-07-12T09:53:03.871926Z","iopub.execute_input":"2022-07-12T09:53:03.872592Z","iopub.status.idle":"2022-07-12T09:53:03.879835Z","shell.execute_reply.started":"2022-07-12T09:53:03.872554Z","shell.execute_reply":"2022-07-12T09:53:03.878214Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train Function\nWe will now define the training function which takes the current train state, dataset, batch_size, and the rng as the input and returns the new train state along with the compute metrics","metadata":{}},{"cell_type":"code","source":"def train_epoch(state, train_ds, batch_size, epoch, rng):\n    train_ds_size = len(train_ds['images'])\n    steps_per_epoch = train_ds_size // batch_size \n\n    perms = jax.random.permutation(rng, train_ds_size)\n    perms = perms[:steps_per_epoch * batch_size]  # skip incomplete batch\n    perms = perms.reshape((steps_per_epoch, batch_size))\n    batch_metrics = []\n    for perm in perms:\n        batch = {k: v[perm, ...] for k, v in train_ds.items()}\n        state, metrics = train_step(state, batch)\n        batch_metrics.append(metrics)\n\n    training_batch_metrics = jax.device_get(batch_metrics)\n    training_epoch_metrics = {\n      k: np.mean([metrics[k] for metrics in training_batch_metrics])\n      for k in training_batch_metrics[0]}\n\n    print('Training - epoch: %d, loss: %.4f, accuracy: %.2f' % (epoch, training_epoch_metrics['loss'], training_epoch_metrics['accuracy'] * 100))\n\n    return state, training_epoch_metrics","metadata":{"execution":{"iopub.status.busy":"2022-07-12T09:53:05.365488Z","iopub.execute_input":"2022-07-12T09:53:05.365961Z","iopub.status.idle":"2022-07-12T09:53:05.380314Z","shell.execute_reply.started":"2022-07-12T09:53:05.365925Z","shell.execute_reply":"2022-07-12T09:53:05.378298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Create train state\nCreating the initial train state which we'll be passing to the neural network while training","metadata":{}},{"cell_type":"code","source":"def create_train_state(rng, learning_rate, momentum):\n    cnn = CNN()\n    params = cnn.init(rng, jnp.ones([1, 28,28,1]))['params']\n    tx = optax.sgd(learning_rate, momentum)\n    return train_state.TrainState.create(\n      apply_fn=cnn.apply, params=params, tx=tx)\n\nrng = jax.random.PRNGKey(0)\nrng, init_rng = jax.random.split(rng)\n\nlearning_rate = 2e-5\nmomentum = 0.9\nstate = create_train_state(init_rng, learning_rate, momentum)\ndel init_rng","metadata":{"execution":{"iopub.status.busy":"2022-07-12T09:53:06.714224Z","iopub.execute_input":"2022-07-12T09:53:06.714647Z","iopub.status.idle":"2022-07-12T09:53:06.813015Z","shell.execute_reply.started":"2022-07-12T09:53:06.714614Z","shell.execute_reply":"2022-07-12T09:53:06.811794Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training\nNext, we will train the neural network for 10 epochs andsee how well the model does.","metadata":{}},{"cell_type":"code","source":"num_epochs = 10\ntraining_accuracy = []\nfor epoch in range(1, num_epochs + 1):\n    rng, input_rng = jax.random.split(rng)\n    state, train_metrics = train_epoch(state, train_ds, total_batch_size, epoch, input_rng)\n    training_accuracy.append(train_metrics[\"accuracy\"])","metadata":{"execution":{"iopub.status.busy":"2022-07-12T09:53:07.882572Z","iopub.execute_input":"2022-07-12T09:53:07.883040Z","iopub.status.idle":"2022-07-12T10:00:47.557987Z","shell.execute_reply.started":"2022-07-12T09:53:07.883003Z","shell.execute_reply":"2022-07-12T10:00:47.556367Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plot the Accuracy \nplt.plot(training_accuracy)\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Accuracy\")\nplt.legend()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-07-12T09:52:57.817535Z","iopub.status.idle":"2022-07-12T09:52:57.817974Z","shell.execute_reply.started":"2022-07-12T09:52:57.817771Z","shell.execute_reply":"2022-07-12T09:52:57.817790Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Conclusion**\nHere in this notebook, we've illustrated how [JAX](https://github.com/google/jax) and [FLAX](https://flax.readthedocs.io/en/latest/) can be used to train the neural network for the toy image classification dataset, with the accuracy of more than 80%. To see more examples of how to use [JAX](https://github.com/google/jax) and [FLAX](https://flax.readthedocs.io/en/latest/) with different data formats, please see this discussion post.  \n\nNow, it's your turn to  create some amazing notebooks using [JAX](https://github.com/google/jax) and [FLAX](https://flax.readthedocs.io/en/latest/). \n\n### **Useful resources which helped me:**\n* https://flax.readthedocs.io/en/latest/notebooks/annotated_mnist.html\n* https://jax.readthedocs.io/en/latest/","metadata":{}}]}