{"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":"# Transfer learning using JAX/FLAX","metadata":{}},{"cell_type":"markdown","source":"# What this notebook covers\n\nThis notebook describes how to do transfer learning using JAX/FLAX with the help of jax-resnet library.\n","metadata":{}},{"cell_type":"markdown","source":"# About Dataset\n\nThe [CIFAR-10 dataset](https://www.tensorflow.org/datasets/catalog/cifar10) consists of 60000 32x32 colour images in 10 classes, with 6000 images per class. There are 50000 training images and 10000 test images.\nClasses of the dataset are - airplane, automobile, bird, cat, deer, dog, frog, horse, ship, truck","metadata":{}},{"cell_type":"markdown","source":"# How to use this notebook\n\nYou can copy & edit this notebook , which you can find on top-right corner of kaggle notebook viewer.\n","metadata":{}},{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"# install jax-resnet library\n!pip install jax-resnet","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-05-11T08:00:45.734781Z","iopub.execute_input":"2022-05-11T08:00:45.735054Z","iopub.status.idle":"2022-05-11T08:00:53.36085Z","shell.execute_reply.started":"2022-05-11T08:00:45.735022Z","shell.execute_reply":"2022-05-11T08:00:53.359837Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Import Libraries\nimport os\nimport cv2\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nfrom typing import Callable\nfrom tqdm.notebook import tqdm\n\n\nfrom sklearn import preprocessing\nfrom sklearn.model_selection import train_test_split\nimport tensorflow.keras as keras\nimport tensorflow as tf\nimport tensorflow_datasets as tfds \n\nimport jax\nimport optax\nimport flax\nimport jax.numpy as jnp\nfrom jax import jit\nfrom jax import lax\nfrom jax_resnet import pretrained_resnet, slice_variables, Sequential\nfrom flax.jax_utils import replicate, unreplicate\nfrom flax.training import train_state\nfrom flax import linen as nn\nfrom flax.core import FrozenDict,frozen_dict\nfrom flax.training.common_utils import shard\n\nimport warnings\nimport logging\nfrom functools import partial\n\nwarnings.simplefilter('ignore')\nos.environ['TF_CPP_MIN_LOG_LEVEL'] = '3' ","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-05-11T08:05:09.07438Z","iopub.execute_input":"2022-05-11T08:05:09.074704Z","iopub.status.idle":"2022-05-11T08:05:09.083329Z","shell.execute_reply.started":"2022-05-11T08:05:09.074665Z","shell.execute_reply":"2022-05-11T08:05:09.082676Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configuration","metadata":{}},{"cell_type":"code","source":"\nConfig = {\n    'NUM_LABELS': 10,\n    'N_SPLITS': 5,\n    'BATCH_SIZE': 32,\n    'N_EPOCHS': 10,\n    'LR': 0.001,\n    'WIDTH': 32,\n    'HEIGHT': 32,\n    'IMAGE_SIZE': 128,\n    'WEIGHT_DECAY': 1e-5,\n    'FREEZE_BACKBONE': True\n}","metadata":{"execution":{"iopub.status.busy":"2022-05-11T08:23:23.634891Z","iopub.execute_input":"2022-05-11T08:23:23.635171Z","iopub.status.idle":"2022-05-11T08:23:23.639982Z","shell.execute_reply.started":"2022-05-11T08:23:23.635139Z","shell.execute_reply":"2022-05-11T08:23:23.639152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Preprocessing ","metadata":{}},{"cell_type":"code","source":"\ndef transform_images(row, size):\n    '''\n    Resize image \n    INPUT row , size\n    RETURNS resized image and its label\n    '''\n    x_train = tf.image.resize(row['image'], (size, size))\n    return x_train, row['label']\n\ndef load_datasets():\n    '''\n    load and transform dataset from tfds\n    RETURNS train and test dataset\n    \n    '''\n    \n    # Construct a tf.data.Dataset\n    train_ds,test_ds = tfds.load('cifar10', split=['train','test'], shuffle_files=True)\n\n    train_ds = train_ds.map(lambda row:transform_images(row,Config[\"IMAGE_SIZE\"]))\n    test_ds = test_ds.map(lambda row:transform_images(row,Config[\"IMAGE_SIZE\"]))\n    \n    # Build your input pipeline\n    train_dataset = train_ds.batch(Config[\"BATCH_SIZE\"]).prefetch(tf.data.AUTOTUNE)\n    test_dataset = test_ds.batch(Config[\"BATCH_SIZE\"]).prefetch(tf.data.AUTOTUNE)\n    \n    return train_dataset,test_dataset","metadata":{"execution":{"iopub.status.busy":"2022-05-11T08:00:53.388596Z","iopub.execute_input":"2022-05-11T08:00:53.388954Z","iopub.status.idle":"2022-05-11T08:00:53.397795Z","shell.execute_reply.started":"2022-05-11T08:00:53.388851Z","shell.execute_reply":"2022-05-11T08:00:53.396684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#  Model","metadata":{}},{"cell_type":"markdown","source":"Here I am using pretrained ```resnet18``` model using ```jax-resnet``` library. For more information about ```jax-resnet``` you can refer to - https://github.com/n2cholas/jax-resnet.\n\nBelow implementation of transfer learning using ```jax-resnet``` is taken from the [happywhale flax jax tpu gpu resnet baseline](https://www.kaggle.com/code/alexlwh/happywhale-flax-jax-tpu-gpu-resnet-baseline) by [@alexlwh](https://www.kaggle.com/alexlwh).\n\nTo implement transfer learning - \n\n1. Remove the last 2 layers of the pre-trained ```ResNet18``` model and replace them with our layers ( ```Head model``` )\n2. Weights of ```ResNet18``` pre-trained model wil be used as feature extractor and are frozen ( it will not updated during training)\n\nIf you want to learn more about transfer learning , you can refer to the following resources -\n* [Transfer learning - wiki](https://en.wikipedia.org/wiki/Transfer_learning)\n* [Transfer learning in tensorflow](https://www.tensorflow.org/tutorials/images/transfer_learning)\n* [A Gentle Introduction to Transfer Learning for Deep Learning](https://machinelearningmastery.com/transfer-learning-for-deep-learning/)\n","metadata":{}},{"cell_type":"code","source":"\"\"\"\nreference - https://www.kaggle.com/code/alexlwh/happywhale-flax-jax-tpu-gpu-resnet-baseline\n\"\"\"\nclass MarginLayer(nn.Module):\n    @nn.compact\n    def __call__(self, inputs):\n        raise NotImplementedError\n\nclass Head(nn.Module):\n    '''head model'''\n    batch_norm_cls: partial = partial(nn.BatchNorm, momentum=0.9)\n    @nn.compact\n    def __call__(self, inputs, train: bool):\n        output_n = inputs.shape[-1]\n        x = self.batch_norm_cls(use_running_average=not train)(inputs)\n        x = nn.Dropout(rate=0.25)(x, deterministic=not train)\n        x = nn.Dense(features=output_n)(x)\n        x = nn.relu(x)\n        x = self.batch_norm_cls(use_running_average=not train)(x)\n        x = nn.Dropout(rate=0.5)(x, deterministic=not train)\n        x = nn.Dense(features=Config[\"NUM_LABELS\"])(x)\n        return x\n\nclass Model(nn.Module):\n    '''Combines backbone and head model'''\n    backbone: Sequential\n    head: Head\n        \n    def __call__(self, inputs, train: bool):\n        x = self.backbone(inputs)\n        # average pool layer\n        x = jnp.mean(x, axis=(1, 2))\n        x = self.head(x, train)\n        return x\n\n    \ndef _get_backbone_and_params(model_arch: str):\n    '''\n    Get backbone and params\n    1. Loads pretrained model (resnet18)\n    2. Get model and param structure except last 2 layers\n    3. Extract the corresponding subset of the variables dict\n    INPUT : model_arch\n    RETURNS backbone , backbone_params\n    '''\n    if model_arch == 'resnet18':\n        resnet_tmpl, params = pretrained_resnet(18)\n        model = resnet_tmpl()\n    else:\n        raise NotImplementedError\n        \n    # get model & param structure for backbone\n    start, end = 0, len(model.layers) - 2\n    backbone = Sequential(model.layers[start:end])\n    backbone_params = slice_variables(params, start, end)\n    return backbone, backbone_params\n\n\ndef get_model_and_variables(model_arch: str, head_init_key: int):\n    '''\n    Get model and variables \n    1. Initialise inputs(shape=(1,image_size,image_size,3))\n    2. Get backbone and params\n    3. Apply backbone model and get outputs\n    4. Initialise head\n    5. Create final model using backbone and head\n    6. Combine params from backbone and head\n    \n    INPUT model_arch, head_init_key\n    RETURNS  model, variables \n    '''\n    \n    #backbone\n    inputs = jnp.ones((1, Config['IMAGE_SIZE'],Config['IMAGE_SIZE'], 3), jnp.float32)\n    backbone, backbone_params = _get_backbone_and_params(model_arch)\n    key = jax.random.PRNGKey(head_init_key)\n    backbone_output = backbone.apply(backbone_params, inputs, mutable=False)\n    \n    #head\n    head_inputs = jnp.ones((1, backbone_output.shape[-1]), jnp.float32)\n    head = Head()\n    head_params = head.init(key, head_inputs, train=False)\n    \n    #final model\n    model = Model(backbone, head)\n    variables = FrozenDict({\n        'params': {\n            'backbone': backbone_params['params'],\n            'head': head_params['params']\n        },\n        'batch_stats': {\n            'backbone': backbone_params['batch_stats'],\n            'head': head_params['batch_stats']\n        }\n    })\n    return model, variables\n","metadata":{"execution":{"iopub.status.busy":"2022-05-16T15:46:50.328068Z","iopub.execute_input":"2022-05-16T15:46:50.328461Z","iopub.status.idle":"2022-05-16T15:46:50.444192Z","shell.execute_reply.started":"2022-05-16T15:46:50.328339Z","shell.execute_reply":"2022-05-16T15:46:50.442472Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model, variables = get_model_and_variables('resnet18', 0)\ninputs = jnp.ones((1,Config['IMAGE_SIZE'], Config['IMAGE_SIZE'],3), jnp.float32)\nkey = jax.random.PRNGKey(0)\no = model.apply(variables, inputs, train=False, mutable=False)","metadata":{"execution":{"iopub.status.busy":"2022-05-11T08:00:53.426803Z","iopub.execute_input":"2022-05-11T08:00:53.427419Z","iopub.status.idle":"2022-05-11T08:00:54.898234Z","shell.execute_reply.started":"2022-05-11T08:00:53.427356Z","shell.execute_reply":"2022-05-11T08:00:54.897524Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"o","metadata":{"execution":{"iopub.status.busy":"2022-05-11T08:00:54.899517Z","iopub.execute_input":"2022-05-11T08:00:54.89978Z","iopub.status.idle":"2022-05-11T08:00:54.909047Z","shell.execute_reply.started":"2022-05-11T08:00:54.899745Z","shell.execute_reply":"2022-05-11T08:00:54.908351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset,test_dataset=load_datasets()","metadata":{"execution":{"iopub.status.busy":"2022-05-11T08:00:54.910335Z","iopub.execute_input":"2022-05-11T08:00:54.91073Z","iopub.status.idle":"2022-05-11T08:01:59.690919Z","shell.execute_reply.started":"2022-05-11T08:00:54.910692Z","shell.execute_reply":"2022-05-11T08:01:59.690172Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_dataset)","metadata":{"execution":{"iopub.status.busy":"2022-05-11T08:01:59.692229Z","iopub.execute_input":"2022-05-11T08:01:59.692641Z","iopub.status.idle":"2022-05-11T08:01:59.698179Z","shell.execute_reply.started":"2022-05-11T08:01:59.692604Z","shell.execute_reply":"2022-05-11T08:01:59.697039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset","metadata":{"execution":{"iopub.status.busy":"2022-05-11T08:01:59.699624Z","iopub.execute_input":"2022-05-11T08:01:59.700029Z","iopub.status.idle":"2022-05-11T08:01:59.711777Z","shell.execute_reply.started":"2022-05-11T08:01:59.699994Z","shell.execute_reply":"2022-05-11T08:01:59.711079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"total_batch_size = Config[\"BATCH_SIZE\"] * jax.local_device_count()\nnum_train_steps = len(train_dataset)","metadata":{"execution":{"iopub.status.busy":"2022-05-11T08:01:59.712833Z","iopub.execute_input":"2022-05-11T08:01:59.713015Z","iopub.status.idle":"2022-05-11T08:01:59.720266Z","shell.execute_reply.started":"2022-05-11T08:01:59.712993Z","shell.execute_reply":"2022-05-11T08:01:59.719564Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_train_steps","metadata":{"execution":{"iopub.status.busy":"2022-05-11T08:01:59.721618Z","iopub.execute_input":"2022-05-11T08:01:59.721911Z","iopub.status.idle":"2022-05-11T08:01:59.732875Z","shell.execute_reply.started":"2022-05-11T08:01:59.721877Z","shell.execute_reply":"2022-05-11T08:01:59.732105Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Flax provides a class ```flax.training.train_state.TrainState``` , which stores the model parameters, the loss function, the optimizer, and exposes an ```apply_gradients``` function to update the model's weight parameters.","metadata":{}},{"cell_type":"code","source":"class TrainState(train_state.TrainState):\n    batch_stats: FrozenDict\n    loss_fn: Callable = flax.struct.field(pytree_node=False)\n    eval_fn: Callable = flax.struct.field(pytree_node=False)","metadata":{"execution":{"iopub.status.busy":"2022-05-11T08:01:59.734098Z","iopub.execute_input":"2022-05-11T08:01:59.734492Z","iopub.status.idle":"2022-05-11T08:01:59.742346Z","shell.execute_reply.started":"2022-05-11T08:01:59.734457Z","shell.execute_reply":"2022-05-11T08:01:59.741623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\nreference - https://github.com/deepmind/optax/issues/159#issuecomment-896459491\n\"\"\"\ndef zero_grads():\n    '''\n    Zero out the previous gradient computation\n    '''\n    def init_fn(_): \n        return ()\n    def update_fn(updates, state, params=None):\n        return jax.tree_map(jnp.zeros_like, updates), ()\n    return optax.GradientTransformation(init_fn, update_fn)","metadata":{"execution":{"iopub.status.busy":"2022-05-11T08:01:59.745792Z","iopub.execute_input":"2022-05-11T08:01:59.746115Z","iopub.status.idle":"2022-05-11T08:01:59.75289Z","shell.execute_reply.started":"2022-05-11T08:01:59.746083Z","shell.execute_reply":"2022-05-11T08:01:59.752098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\nreference - https://colab.research.google.com/drive/1g_pt2Rc3bv6H6qchvGHD-BpgF-Pt4vrC#scrollTo=TqDvTL_tIQCH&line=2&uniqifier=1\n\"\"\"\ndef create_mask(params, label_fn):\n    def _map(params, mask, label_fn):\n        for k in params:\n            if label_fn(k):\n                mask[k] = 'zero'\n            else:\n                if isinstance(params[k], FrozenDict):\n                    mask[k] = {}\n                    _map(params[k], mask[k], label_fn)\n                else:\n                    mask[k] = 'adam'\n    mask = {}\n    _map(params, mask, label_fn)\n    return frozen_dict.freeze(mask)","metadata":{"execution":{"iopub.status.busy":"2022-05-11T08:01:59.754234Z","iopub.execute_input":"2022-05-11T08:01:59.754495Z","iopub.status.idle":"2022-05-11T08:01:59.76321Z","shell.execute_reply.started":"2022-05-11T08:01:59.754462Z","shell.execute_reply":"2022-05-11T08:01:59.762545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Here we are using ```Adam optimizer``` with ```weight decay```, and again we are using optax library.\n\nHere is the interesting article on adam optimizer - https://www.fast.ai/2018/07/02/adam-weight-decay/","metadata":{}},{"cell_type":"code","source":"adamw = optax.adamw(\n    learning_rate=Config['LR'],\n    b1=0.9, b2=0.999, \n    eps=1e-6, weight_decay=1e-2\n)\n\noptimizer = optax.multi_transform(\n    {'adam': adamw, 'zero': zero_grads()},\n    create_mask(variables['params'], lambda s: s.startswith('backbone'))\n)","metadata":{"execution":{"iopub.status.busy":"2022-05-11T08:05:21.313992Z","iopub.execute_input":"2022-05-11T08:05:21.31455Z","iopub.status.idle":"2022-05-11T08:05:21.320321Z","shell.execute_reply.started":"2022-05-11T08:05:21.31451Z","shell.execute_reply":"2022-05-11T08:05:21.319415Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Here i am using accuracy metric. (https://developers.google.com/machine-learning/crash-course/classification/accuracy)","metadata":{}},{"cell_type":"code","source":"def accuracy(logits,labels):\n    '''\n    calculates accuracy based on logits and labels\n    INPUT logits , labels\n    RETURNS accuracy\n    '''\n    return [jnp.mean(jnp.argmax(logits, -1) == labels)]","metadata":{"execution":{"iopub.status.busy":"2022-05-11T08:05:21.333817Z","iopub.execute_input":"2022-05-11T08:05:21.334018Z","iopub.status.idle":"2022-05-11T08:05:21.338709Z","shell.execute_reply.started":"2022-05-11T08:05:21.333994Z","shell.execute_reply":"2022-05-11T08:05:21.337872Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Now , computing the ```softmax cross entropy ``` between sets of ```logits``` and ```labels``` using the [optax](https://optax.readthedocs.io/en/latest/) library.","metadata":{}},{"cell_type":"code","source":"#loss function and evaluation function \nloss_fn = optax.softmax_cross_entropy\neval_fn = accuracy","metadata":{"execution":{"iopub.status.busy":"2022-05-11T08:05:21.340539Z","iopub.execute_input":"2022-05-11T08:05:21.340811Z","iopub.status.idle":"2022-05-11T08:05:21.348662Z","shell.execute_reply.started":"2022-05-11T08:05:21.340777Z","shell.execute_reply":"2022-05-11T08:05:21.347676Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Instantiate a TrainState.\nstate = TrainState.create(\n    apply_fn = model.apply,\n    params = variables['params'],\n    tx = optimizer,\n    batch_stats = variables['batch_stats'],\n    loss_fn = loss_fn,\n    eval_fn = eval_fn\n)","metadata":{"execution":{"iopub.status.busy":"2022-05-11T08:05:21.350722Z","iopub.execute_input":"2022-05-11T08:05:21.351005Z","iopub.status.idle":"2022-05-11T08:05:21.398207Z","shell.execute_reply.started":"2022-05-11T08:05:21.350969Z","shell.execute_reply":"2022-05-11T08:05:21.397569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_step(state: TrainState, batch, labels, dropout_rng):\n    dropout_rng, new_dropout_rng = jax.random.split(dropout_rng)\n    \n    # params as input because we differentiate wrt it \n    def loss_function(params):\n        # if you set state.params, then params can't be backpropagated through!\n        variables = {'params': params, 'batch_stats': state.batch_stats}\n        \n        # return mutated states if mutable is specified\n        logits, new_batch_stats = state.apply_fn(\n            variables, batch, train=True, \n            mutable=['batch_stats'],\n            rngs={'dropout': dropout_rng}\n        )\n        # logits: (BS, OUTPUT_N), one_hot: (BS, OUTPUT_N)\n        one_hot = jax.nn.one_hot(labels,Config[\"NUM_LABELS\"])\n        loss = state.loss_fn(logits, one_hot).mean()\n        return loss, (logits, new_batch_stats)\n    \n    \n    # backpropagation and update params & batch_stats \n    grad_fn = jax.value_and_grad(loss_function, has_aux=True) #differentiate the loss function\n    (loss, aux), grads = grad_fn(state.params)\n    logits, new_batch_stats = aux\n    grads = lax.pmean(grads, axis_name='batch') #compute the mean gradient over all devices\n    new_state = state.apply_gradients(\n        grads=grads, batch_stats=new_batch_stats['batch_stats'] #applies the gradients to the weights.\n    )\n    \n    # evaluation metrics\n    accuracy = state.eval_fn(logits, labels)\n    \n    # store metadata\n    metadata = jax.lax.pmean(\n        {'loss': loss, 'accuracy': accuracy},\n        axis_name='batch'\n    )\n    return new_state, metadata, new_dropout_rng\n\n\ndef val_step(state: TrainState, batch, labels):\n    variables = {'params': state.params, 'batch_stats': state.batch_stats}\n    logits = state.apply_fn(variables, batch, train=False) # stack the model's forward pass with the logits function\n    return state.eval_fn(logits, labels)\n\ndef test_step(state: TrainState, batch):\n    variables = {'params': state.params, 'batch_stats': state.batch_stats}\n    logits = state.apply_fn(variables, batch, train=False) # stack the model's forward pass with the logits function\n    return logits\n\nparallel_train_step = jax.pmap(train_step, axis_name='batch', donate_argnums=(0,))\nparallel_val_step = jax.pmap(val_step, axis_name='batch', donate_argnums=(0,))\nparallel_test_step = jax.pmap(test_step, axis_name='batch', donate_argnums=(0,))\n\n# required for parallelism\nstate = replicate(state)\n\n# control randomness on dropout and update inside train_step\nrng = jax.random.PRNGKey(0)\ndropout_rng = jax.random.split(rng, jax.local_device_count())  # for parallelism\n","metadata":{"execution":{"iopub.status.busy":"2022-05-11T08:05:21.399535Z","iopub.execute_input":"2022-05-11T08:05:21.399777Z","iopub.status.idle":"2022-05-11T08:05:21.582262Z","shell.execute_reply.started":"2022-05-11T08:05:21.399745Z","shell.execute_reply":"2022-05-11T08:05:21.581498Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for epoch_i in tqdm(range(Config['N_EPOCHS']), desc=f\"{Config['N_EPOCHS']} epochs\", position=0, leave=True):\n    # training set\n    train_loss, train_accuracy = [], []\n    iter_n = len(train_dataset)\n    \n    with tqdm(total=iter_n, desc=f\"{iter_n} iterations\", leave=False) as progress_bar:\n        for _batch in train_dataset:\n            batch=_batch[0]  # train_dataset is tuple containing (image,labels)\n            labels=_batch[1]\n\n            batch = jnp.array(batch, dtype=jnp.float32)\n            labels = jnp.array(labels, dtype=jnp.float32)\n            \n            batch, labels = shard(batch), shard(labels)\n           \n            # backprop and update param & batch statsp\n            \n            state, train_metadata, dropout_rng = parallel_train_step(state, batch, labels, dropout_rng)\n            train_metadata = unreplicate(train_metadata)\n            \n            # update train statistics\n            _train_loss, _train_top1_acc = map(float, [train_metadata['loss'], *train_metadata['accuracy']])\n            train_loss.append(_train_loss)\n            train_accuracy.append(_train_top1_acc)\n            progress_bar.update(1)\n            \n    avg_train_loss = sum(train_loss)/len(train_loss)\n    avg_train_acc = sum(train_accuracy)/len(train_accuracy)\n    print(f\"[{epoch_i+1}/{Config['N_EPOCHS']}] Train Loss: {avg_train_loss:.03} | Train Accuracy: {avg_train_acc:.03}\")\n    \n    # validation set\n    \n    valid_accuracy = []\n    iter_n = len(test_dataset)\n    with tqdm(total=iter_n, desc=f\"{iter_n} iterations\", leave=False) as progress_bar:\n        for _batch in test_dataset:\n            batch = _batch[0]\n            labels = _batch[1]\n\n            batch = jnp.array(batch, dtype=jnp.float32)\n            labels = jnp.array(labels, dtype=jnp.float32)\n\n            batch, labels = shard(batch), shard(labels)\n            metric = parallel_val_step(state, batch, labels)[0]\n            valid_accuracy.append(metric)\n            progress_bar.update(1)\n\n\n    avg_valid_acc = sum(valid_accuracy)/len(valid_accuracy)\n    avg_valid_acc = np.array(avg_valid_acc)[0]\n    print(f\"[{epoch_i+1}/{Config['N_EPOCHS']}] Valid Accuracy: {avg_valid_acc:.03}\")\n\n","metadata":{"execution":{"iopub.status.busy":"2022-05-11T08:23:29.122291Z","iopub.execute_input":"2022-05-11T08:23:29.122616Z","iopub.status.idle":"2022-05-11T08:29:47.681311Z","shell.execute_reply.started":"2022-05-11T08:23:29.122579Z","shell.execute_reply":"2022-05-11T08:29:47.680594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Conclusion\n\nSo we learned how to do transfer learning using ```jax/flax``` with help of ```jax-resnet``` library by removing the first two layers of ```resnet18``` and adding our layers , trained this model on a very famous ```cifar-10``` dataset and achieved good results. This notebook is minimal and reproducible , can be used for other tasks with different dataset by just ```copy & edit``` this notebook.","metadata":{}},{"cell_type":"markdown","source":"# References \n\n* https://www.kaggle.com/code/alexlwh/happywhale-flax-jax-tpu-gpu-resnet-baseline\n* https://www.kaggle.com/code/heyytanay/herbarium-jax-flax-training-kfolds-w-b\n* https://flax.readthedocs.io/en/latest/notebooks/annotated_mnist.html\n* https://optax.readthedocs.io/\n* https://flax.readthedocs.io/\n* https://jax.readthedocs.io/\n* https://github.com/google/flax/tree/main/examples\n* https://github.com/n2cholas/jax-resnet\n* https://www.tensorflow.org/datasets/catalog/cifar10\n* https://www.tensorflow.org/tutorials\n","metadata":{}}]}