{"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":"# DR Grade Classification with ViT & Grad CAM","metadata":{}},{"cell_type":"markdown","source":"# Introduction\n\nDiabetic Retinopathy (DR) is a complication of diabetes, caused by high blood sugar levels damaging the back of the eye (retina). It can cause blindness if left undiagnosed and untreated.\n\nDR is split:\n- Grade 0 and 1: considered as No \"referable\" DR, because 1 is difficult to diagnose (early DR)\n- Grade 2: **background retinopathy** – tiny bulges develop in the blood vessels, which may bleed slightly but do not usually affect your vision\n- Grade 3:**pre-proliferative retinopathy** – more severe and widespread changes affect the blood vessels, including more significant bleeding into the eye\n- Grade 4: **proliferative retinopathy** – scar tissue and new blood vessels, which are weak and bleed easily, develop on the retina; this can result in some loss of vision\n\n*Source: [NHS UK](https://www.nhs.uk/conditions/diabetic-retinopathy/)*\n\nGlobally, the number of people with DR will grow from 126.6 million in 2010 to 191.0 million by 2030.  \n*Source: [10.4103/0301-4738.100542](https://www.ncbi.nlm.nih.gov/pmc/articles/PMC3491270/)*\n\n![dr_grades.png](https://www.ophthalytics.com/wp-content/uploads/2021/08/Copy-of-WEBSITE-CONTENT-1536x878.png)\n\n**Source:** Ophthalytics.  \n**Link:** https://www.ophthalytics.com/our-technology/diabetic-retinopathy/\n\nNote we will refer to Diabetic Retinoptahy as DR in the following.\n\nWith:\n- 0 - No DR\n- 1 - Mild\n- 2 - Moderate\n- 3 - Severe\n- 4 - Proliferative DR\n\n## Credits\n\nThis [implementation](https://www.kaggle.com/code/basu369victor/covid19-detection-with-vit-and-heatmap) was taken from [VICTOR BASU](https://www.kaggle.com/basu369victor).\n\nThat was inspired from these two official keras examples:\n 1. [**Image classification with Vision Transformer**](https://keras.io/examples/vision/image_classification_with_vision_transformer/)\n 2. [**Grad-CAM class activation visualization**](https://keras.io/examples/vision/grad_cam/)\n\n## Disclaimer\nThis notebook implements Vision Transformer (ViT) model by Alexey Dosovitskiy et al for image classification, and demonstrates it on the APTOS 2019 Diabetic Retionapathy Classification dataset.","metadata":{}},{"cell_type":"markdown","source":"## About the model\nThis example implements the [Vision Transformer (ViT)](https://arxiv.org/abs/2010.11929) model by Alexey Dosovitskiy et al. for image classification, and demonstrates it on the CIFAR-100 dataset. The ViT model applies the Transformer architecture with self-attention to sequences of image patches, without using convolution layers.\n\n![vit](https://neurohive.io/wp-content/uploads/2020/10/rsz_cov.png)","metadata":{}},{"cell_type":"code","source":"import numpy as np \nimport pandas as pd \nimport os\nimport gc\nfrom sklearn.model_selection import train_test_split\n\nimport tensorflow as tf\nfrom tensorflow import keras\nfrom tensorflow.keras import layers\nimport tensorflow_addons as tfa\n\n# Display\nfrom IPython.display import Image, display\nimport matplotlib.pyplot as plt\nimport matplotlib.cm as cm\n\n\nAUTOTUNE = tf.data.experimental.AUTOTUNE","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-10-23T15:20:05.594489Z","iopub.execute_input":"2022-10-23T15:20:05.594834Z","iopub.status.idle":"2022-10-23T15:20:05.600926Z","shell.execute_reply.started":"2022-10-23T15:20:05.594800Z","shell.execute_reply":"2022-10-23T15:20:05.599860Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Reshape images to 512x512","metadata":{}},{"cell_type":"code","source":"TRAIN_PATH = '../input/aptos2019-blindness-detection/train_images/'\nDF_TRAIN = pd.read_csv('../input/aptos2019-blindness-detection/train.csv', dtype='str')\nDF_TRAIN['image_path'] = TRAIN_PATH + DF_TRAIN[\"id_code\"] + \".png\" \nDF_TRAIN.head()","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:20:05.618005Z","iopub.execute_input":"2022-10-23T15:20:05.618252Z","iopub.status.idle":"2022-10-23T15:20:05.639296Z","shell.execute_reply.started":"2022-10-23T15:20:05.618228Z","shell.execute_reply":"2022-10-23T15:20:05.638463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"classes = {0 : \"No DR\",\n           1 : \"Mild\",\n           2 : \"Moderate\",\n           3 : \"Severe\",\n           4 : \"Proliferative\"}","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:20:05.642523Z","iopub.execute_input":"2022-10-23T15:20:05.642799Z","iopub.status.idle":"2022-10-23T15:20:05.646393Z","shell.execute_reply.started":"2022-10-23T15:20:05.642775Z","shell.execute_reply":"2022-10-23T15:20:05.645528Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir ./train_imgs_reshaped","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:20:05.700133Z","iopub.execute_input":"2022-10-23T15:20:05.700416Z","iopub.status.idle":"2022-10-23T15:20:06.686979Z","shell.execute_reply.started":"2022-10-23T15:20:05.700392Z","shell.execute_reply":"2022-10-23T15:20:06.685827Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DF_TRAIN['image_path'][0]","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:20:06.688869Z","iopub.execute_input":"2022-10-23T15:20:06.689231Z","iopub.status.idle":"2022-10-23T15:20:06.696813Z","shell.execute_reply.started":"2022-10-23T15:20:06.689187Z","shell.execute_reply":"2022-10-23T15:20:06.695854Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'../input/aptos2019-blindness-detection/train_images/000c1434d8d7.png'.split('/')[-1]","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:20:06.698717Z","iopub.execute_input":"2022-10-23T15:20:06.699274Z","iopub.status.idle":"2022-10-23T15:20:06.707378Z","shell.execute_reply.started":"2022-10-23T15:20:06.699239Z","shell.execute_reply":"2022-10-23T15:20:06.706549Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from PIL import Image\n\nfor img in DF_TRAIN['image_path']:\n    img_outfpath = \"./train_imgs_reshaped/\" + img.split('/')[-1]\n    image = Image.open(img)\n    image = image.resize((512,512),Image.ANTIALIAS)\n    image.save(fp=img_outfpath)","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:20:06.709436Z","iopub.execute_input":"2022-10-23T15:20:06.709944Z","iopub.status.idle":"2022-10-23T15:42:11.729174Z","shell.execute_reply.started":"2022-10-23T15:20:06.709831Z","shell.execute_reply":"2022-10-23T15:42:11.728251Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define DFs for reshaped images","metadata":{}},{"cell_type":"markdown","source":"One hot encode diagnosis","metadata":{}},{"cell_type":"code","source":"TRAIN_PATH_RS = './train_imgs_reshaped/'\nDF_TRAIN_RS = pd.read_csv('../input/aptos2019-blindness-detection/train.csv', dtype='str')\nDF_TRAIN_RS['image_path'] = TRAIN_PATH_RS + DF_TRAIN_RS[\"id_code\"] + \".png\" \nDF_TRAIN_RS.head()","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:42:11.730662Z","iopub.execute_input":"2022-10-23T15:42:11.731019Z","iopub.status.idle":"2022-10-23T15:42:11.751649Z","shell.execute_reply.started":"2022-10-23T15:42:11.730977Z","shell.execute_reply":"2022-10-23T15:42:11.750786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Source: https://stackoverflow.com/questions/37292872/how-can-i-one-hot-encode-in-python\ndef encode_and_bind(original_dataframe, feature_to_encode):\n    dummies = pd.get_dummies(original_dataframe[[feature_to_encode]])\n    res = pd.concat([original_dataframe, dummies], axis=1)\n    res = res.drop([feature_to_encode], axis=1)\n    return(res) ","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:42:11.752765Z","iopub.execute_input":"2022-10-23T15:42:11.753096Z","iopub.status.idle":"2022-10-23T15:42:11.757675Z","shell.execute_reply.started":"2022-10-23T15:42:11.753065Z","shell.execute_reply":"2022-10-23T15:42:11.756678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"res = encode_and_bind(DF_TRAIN_RS, 'diagnosis')","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:42:11.758960Z","iopub.execute_input":"2022-10-23T15:42:11.759301Z","iopub.status.idle":"2022-10-23T15:42:11.772697Z","shell.execute_reply.started":"2022-10-23T15:42:11.759268Z","shell.execute_reply":"2022-10-23T15:42:11.771738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"res.head()","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:42:11.775441Z","iopub.execute_input":"2022-10-23T15:42:11.775849Z","iopub.status.idle":"2022-10-23T15:42:11.786767Z","shell.execute_reply.started":"2022-10-23T15:42:11.775823Z","shell.execute_reply":"2022-10-23T15:42:11.786036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"classes = {\"diagnosis_0\" : \"No DR\",\n           \"diagnosis_1\" : \"Mild\",\n           \"diagnosis_2\" : \"Moderate\",\n           \"diagnosis_3\" : \"Severe\",\n           \"diagnosis_4\" : \"Proliferative\"}","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:42:11.788518Z","iopub.execute_input":"2022-10-23T15:42:11.788777Z","iopub.status.idle":"2022-10-23T15:42:11.795647Z","shell.execute_reply.started":"2022-10-23T15:42:11.788753Z","shell.execute_reply":"2022-10-23T15:42:11.794695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"res.rename(columns=classes, inplace=True)\nres.head()","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:42:11.796996Z","iopub.execute_input":"2022-10-23T15:42:11.797542Z","iopub.status.idle":"2022-10-23T15:42:11.813127Z","shell.execute_reply.started":"2022-10-23T15:42:11.797506Z","shell.execute_reply":"2022-10-23T15:42:11.812259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"res.columns","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:42:11.814435Z","iopub.execute_input":"2022-10-23T15:42:11.814827Z","iopub.status.idle":"2022-10-23T15:42:11.824444Z","shell.execute_reply.started":"2022-10-23T15:42:11.814793Z","shell.execute_reply":"2022-10-23T15:42:11.823581Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"target = np.array(res[['No DR', 'Mild', 'Moderate', 'Severe', 'Proliferative']])","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:42:11.825588Z","iopub.execute_input":"2022-10-23T15:42:11.825950Z","iopub.status.idle":"2022-10-23T15:42:11.834221Z","shell.execute_reply.started":"2022-10-23T15:42:11.825917Z","shell.execute_reply":"2022-10-23T15:42:11.833355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_train, X_test, y_train, y_test  = train_test_split(res['image_path'], target, test_size=0.33, random_state=42)\nprint(f\"train shape: {X_train.shape}- y_train shape: {y_train.shape}\")\nprint(f\"test shape: {X_test.shape}- y_test shape: {y_test.shape}\")","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:42:11.835601Z","iopub.execute_input":"2022-10-23T15:42:11.835929Z","iopub.status.idle":"2022-10-23T15:42:11.847200Z","shell.execute_reply.started":"2022-10-23T15:42:11.835903Z","shell.execute_reply":"2022-10-23T15:42:11.846299Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from PIL import Image\nimage_path = Image.open('./train_imgs_reshaped/000c1434d8d7.png')\nplt.imshow(image_path)\nplt.axis(\"off\")","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:42:11.848582Z","iopub.execute_input":"2022-10-23T15:42:11.848943Z","iopub.status.idle":"2022-10-23T15:42:12.221435Z","shell.execute_reply.started":"2022-10-23T15:42:11.848905Z","shell.execute_reply":"2022-10-23T15:42:12.220523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Configure the hyperparameters","metadata":{}},{"cell_type":"code","source":"num_classes = 5\ninput_shape = (512, 512, 3)\nlearning_rate = 1e-4 #0.001\nweight_decay = 0.0001\nbatch_size = 16 #256\nnum_epochs = 100\n# We'll resize input images to this size\nimage_size =  256 \n# Size of the patches to be extract from the input images\npatch_size = 7 \nnum_patches = (image_size // patch_size) ** 2\nprojection_dim = 64\nnum_heads = 4\n# Size of the transformer layers\ntransformer_units = [\n    projection_dim * 2,\n    projection_dim,\n]  \ntransformer_layers = 8\n# Size of the dense layers of the final classifier\nmlp_head_units = [56, 28] #[1024, 512]  ","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:42:12.222748Z","iopub.execute_input":"2022-10-23T15:42:12.223092Z","iopub.status.idle":"2022-10-23T15:42:12.228030Z","shell.execute_reply.started":"2022-10-23T15:42:12.223060Z","shell.execute_reply":"2022-10-23T15:42:12.226982Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_patches","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:42:12.229344Z","iopub.execute_input":"2022-10-23T15:42:12.229720Z","iopub.status.idle":"2022-10-23T15:42:12.242738Z","shell.execute_reply.started":"2022-10-23T15:42:12.229652Z","shell.execute_reply":"2022-10-23T15:42:12.241790Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@tf.function\ndef load(image_file, target):\n    image = tf.io.read_file(image_file)\n    image = tf.image.decode_png(image)\n\n    image_ = tf.cast(image, tf.uint8)\n    return image, target","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:42:12.244061Z","iopub.execute_input":"2022-10-23T15:42:12.244473Z","iopub.status.idle":"2022-10-23T15:42:12.251915Z","shell.execute_reply.started":"2022-10-23T15:42:12.244445Z","shell.execute_reply":"2022-10-23T15:42:12.251138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader = (\n    tf.data.Dataset\n    .from_tensor_slices((X_train,y_train))\n    .map(load, num_parallel_calls=AUTOTUNE)\n    .shuffle(7)\n    .batch(batch_size)\n)\ntest_loader = (\n    tf.data.Dataset\n    .from_tensor_slices((X_test,y_test))\n    .map(load, num_parallel_calls=AUTOTUNE)\n    .shuffle(7)\n    .batch(batch_size)\n)","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:42:12.253816Z","iopub.execute_input":"2022-10-23T15:42:12.254136Z","iopub.status.idle":"2022-10-23T15:42:14.569924Z","shell.execute_reply.started":"2022-10-23T15:42:12.254111Z","shell.execute_reply":"2022-10-23T15:42:14.569022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_batch = (\n    tf.data.Dataset\n    .from_tensor_slices((X_train,y_train))\n    .map(load, num_parallel_calls=AUTOTUNE)\n    .shuffle(7)\n    .batch(X_train.shape[0]-100)#X_train.shape[0]-100\n)\n#next(iter(train_batch))[0].shape","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:42:14.571351Z","iopub.execute_input":"2022-10-23T15:42:14.571935Z","iopub.status.idle":"2022-10-23T15:42:14.585428Z","shell.execute_reply.started":"2022-10-23T15:42:14.571897Z","shell.execute_reply":"2022-10-23T15:42:14.584530Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data Augmentation","metadata":{}},{"cell_type":"code","source":"data_augmentation = keras.Sequential(\n    [\n        layers.experimental.preprocessing.Normalization(),\n        layers.experimental.preprocessing.Resizing(image_size, image_size),\n        layers.experimental.preprocessing.RandomFlip(\"horizontal\"),\n        layers.experimental.preprocessing.RandomRotation(factor=0.02),\n        layers.experimental.preprocessing.RandomZoom(\n            height_factor = 0.2, width_factor = 0.2\n        ),\n    ],\n     name=\"data_augmentation\",\n)\n# Compute the mean and the variance of the training data for normalization.\nCompleteBatchData  =next(iter(train_batch))[0]\ndata_augmentation.layers[0].adapt(CompleteBatchData)","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:42:14.587117Z","iopub.execute_input":"2022-10-23T15:42:14.587512Z","iopub.status.idle":"2022-10-23T15:43:29.337717Z","shell.execute_reply.started":"2022-10-23T15:42:14.587470Z","shell.execute_reply":"2022-10-23T15:43:29.336457Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del CompleteBatchData\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:43:29.339489Z","iopub.execute_input":"2022-10-23T15:43:29.339867Z","iopub.status.idle":"2022-10-23T15:43:29.606355Z","shell.execute_reply.started":"2022-10-23T15:43:29.339827Z","shell.execute_reply":"2022-10-23T15:43:29.605552Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Implementing multilayer perceptron (MLP)","metadata":{}},{"cell_type":"code","source":"def mlp(x, hidden_units, dropout_rate):\n    for units in hidden_units:\n        x = layers.Dense(units, activation = tf.nn.gelu)(x)\n        x = layers.Dropout(dropout_rate)(x)\n    return x","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:43:29.607580Z","iopub.execute_input":"2022-10-23T15:43:29.607945Z","iopub.status.idle":"2022-10-23T15:43:29.612924Z","shell.execute_reply.started":"2022-10-23T15:43:29.607906Z","shell.execute_reply":"2022-10-23T15:43:29.611840Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Implement patch creation as a layer","metadata":{}},{"cell_type":"code","source":"class Patches(layers.Layer):\n    def __init__(self, patch_size):\n        super(Patches, self).__init__()\n        self.patch_size = patch_size\n        \n    def call(self, images):\n        batch_size = tf.shape(images)[0]\n        patches = tf.image.extract_patches(\n            images = images,\n            sizes = [1, self.patch_size, self.patch_size, 1],\n            strides=[1, self.patch_size, self.patch_size, 1],\n            rates=[1, 1, 1, 1],\n            padding=\"VALID\",\n        )\n        patch_dims = patches.shape[-1]\n        #print(patches.shape)\n        patches = tf.reshape(patches, [batch_size, -1, patch_dims])\n        return patches","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:43:29.616677Z","iopub.execute_input":"2022-10-23T15:43:29.617175Z","iopub.status.idle":"2022-10-23T15:43:29.627354Z","shell.execute_reply.started":"2022-10-23T15:43:29.617138Z","shell.execute_reply":"2022-10-23T15:43:29.626496Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(8, 8))\nimage = next(iter(train_loader))[0][5]\n\nplt.imshow(image)\nplt.axis(\"off\")\n\nresized_image = tf.image.resize(\n    tf.convert_to_tensor([image]), size=(image_size, image_size)\n)\n\nprint(resized_image.shape)\npatches = Patches(patch_size)(resized_image)\nprint(f\"Image size: {image_size} X {image_size}\")\nprint(f\"Patch size: {patch_size} X {patch_size}\")\nprint(f\"Patches per image: {patches.shape[1]}\")\nprint(f\"Elements per patch: {patches.shape[-1]}\")\n\nn = int(np.sqrt(patches.shape[1]))\n#print(n)\n\nplt.figure(figsize=(8, 8))\nfor i, patch in enumerate(patches[0]):\n    ax = plt.subplot(n, n, i + 1)\n    patch_img = tf.reshape(patch, (patch_size, patch_size, 3))\n    plt.imshow(patch_img.numpy().astype('uint8'))\n    plt.axis(\"off\")","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:43:29.628926Z","iopub.execute_input":"2022-10-23T15:43:29.629320Z","iopub.status.idle":"2022-10-23T15:44:34.954908Z","shell.execute_reply.started":"2022-10-23T15:43:29.629285Z","shell.execute_reply":"2022-10-23T15:44:34.953932Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(8, 8))\nimage = next(iter(train_loader))[0][5]\n\nplt.imshow(image)\nplt.axis(\"off\")\n\nresized_image = tf.image.resize(\n    tf.convert_to_tensor([image]), size=(image_size, image_size)\n)\n\nprint(resized_image.shape)\npatches = Patches(patch_size)(resized_image)\nprint(f\"Image size: {image_size} X {image_size}\")\nprint(f\"Patch size: {patch_size} X {patch_size}\")\nprint(f\"Patches per image: {patches.shape[1]}\")\nprint(f\"Elements per patch: {patches.shape[-1]}\")\n\nn = int(np.sqrt(patches.shape[1]))\n#print(n)\n\nplt.figure(figsize=(8, 8))\nfor i, patch in enumerate(patches[0]):\n    ax = plt.subplot(n, n, i + 1)\n    patch_img = tf.reshape(patch, (patch_size, patch_size, 3))\n    plt.imshow(patch_img.numpy().astype('uint8'))\n    plt.axis(\"off\")","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:44:34.956380Z","iopub.execute_input":"2022-10-23T15:44:34.956763Z","iopub.status.idle":"2022-10-23T15:45:40.606536Z","shell.execute_reply.started":"2022-10-23T15:44:34.956725Z","shell.execute_reply":"2022-10-23T15:45:40.605744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## The patch encoding layer\n\nThe **PatchEncoder** layer will linearly transform a **patch** by projecting it into a vector of size **projection_dim**. In addition, it adds a learnable position embedding to the projected vector.","metadata":{}},{"cell_type":"code","source":"class PatchEncoder(layers.Layer):\n    def __init__(self, num_of_patches, projection_dim):\n        super(PatchEncoder, self).__init__()\n        self.num_patches = num_patches\n        self.projection = layers.Dense(units = projection_dim)\n        self.position_embedding = layers.Embedding(\n            input_dim = num_patches, output_dim = projection_dim\n        )\n        \n    def call(self, patch):\n        positions = tf.range(start=0, limit=self.num_patches, delta=1)\n        encode = self.projection(patch) + self.position_embedding(positions)\n        return encode","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:45:40.607775Z","iopub.execute_input":"2022-10-23T15:45:40.608193Z","iopub.status.idle":"2022-10-23T15:45:40.615942Z","shell.execute_reply.started":"2022-10-23T15:45:40.608161Z","shell.execute_reply":"2022-10-23T15:45:40.614721Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"##  The ViT model\n\nThe ViT model consists of multiple Transformer blocks, which use the **layers.MultiHeadAttention layer** as a self-attention mechanism applied to the sequence of patches. The Transformer blocks produce a **[batch_size, num_patches, projection_dim]** tensor, which is processed via an classifier head with softmax to produce the final class probabilities output.<br>\nUnlike the technique described in the [paper](https://arxiv.org/abs/2010.11929), which prepends a learnable embedding to the sequence of encoded patches to serve as the image representation, all the outputs of the final Transformer block are reshaped with **layers.Flatten()** and used as the image representation input to the classifier head. Note that the **layers.GlobalAveragePooling1D** layer could also be used instead to aggregate the outputs of the Transformer block, especially when the number of patches and the projection dimensions are large.","metadata":{}},{"cell_type":"code","source":"def vit_model():\n    inputs = layers.Input(shape=input_shape)\n    # Augment data.\n    augmented = data_augmentation(inputs)\n    # Create patches.\n    patches = Patches(patch_size)(augmented)\n    # Encode patches.\n    encoded_patches = PatchEncoder(num_patches, projection_dim)(patches)\n    \n    # Create multiple layers of the Transformer block.\n    for _ in range(transformer_layers):\n        # Layer normalization 1.\n        x1 = layers.BatchNormalization()(encoded_patches)\n        # create a multi-head attention layer\n        attention_output = layers.MultiHeadAttention(\n            num_heads=num_heads, key_dim=projection_dim, dropout=0.1\n        )(x1, x1)\n        # Skip connection 1.\n        x2 = layers.Add()([attention_output, encoded_patches])\n        # Layer normalization 2.\n        x3 = layers.BatchNormalization()(x2)\n        # MLP.\n        x3 = mlp(x3, hidden_units=transformer_units, dropout_rate=0.1)\n        # Skip connection 2.\n        encoded_patches = layers.Add()([x3, x2])\n        \n    # Create a [batch_size, projection_dim] tensor.\n    representation = layers.LayerNormalization()(encoded_patches)\n    representation = layers.Flatten()(representation)\n    representation = layers.Dropout(0.5)(representation)\n    # Add MLP\n    features = mlp(representation, hidden_units = mlp_head_units, dropout_rate=0.5)\n    # Classify outputs.\n    logits = layers.Dense(num_classes, activation='softmax')(features)\n    # create keras model\n    model = keras.Model(inputs=inputs, outputs=logits)\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:45:40.617676Z","iopub.execute_input":"2022-10-23T15:45:40.618083Z","iopub.status.idle":"2022-10-23T15:45:40.627075Z","shell.execute_reply.started":"2022-10-23T15:45:40.618048Z","shell.execute_reply":"2022-10-23T15:45:40.626062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def experiment(model):\n    optimizer = tfa.optimizers.AdamW(\n        learning_rate=learning_rate, weight_decay=weight_decay\n    )\n    \n    model.compile(\n        optimizer=optimizer,\n        loss=keras.losses.CategoricalCrossentropy(from_logits=True),\n        metrics=[\n            keras.metrics.CategoricalAccuracy(name=\"accuracy\"),\n            keras.metrics.AUC( name=\"AUC\"),\n        ],\n     )\n    checkpoint_filepath = \"./tmp/checkpoint\"\n    checkpoint_callback = keras.callbacks.ModelCheckpoint(\n        checkpoint_filepath,\n        monitor=\"val_accuracy\",\n        save_best_only=True,\n        save_weights_only=True,\n    )\n\n    history = model.fit(train_loader ,\n                        batch_size=batch_size,\n                        epochs=num_epochs,\n                        validation_data=test_loader,\n                        callbacks=[checkpoint_callback],)\n    model.load_weights(checkpoint_filepath)\n    _, accuracy, auc = model.evaluate(test_loader)\n    print(f\"Test accuracy: {round(accuracy * 100, 2)}%\")\n    print(f\"Test AUC: {round(auc * 100, 2)}%\")\n\n    return history","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:45:40.628419Z","iopub.execute_input":"2022-10-23T15:45:40.628805Z","iopub.status.idle":"2022-10-23T15:45:40.640614Z","shell.execute_reply.started":"2022-10-23T15:45:40.628770Z","shell.execute_reply":"2022-10-23T15:45:40.639704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"vit_classifier = vit_model()\nvit_classifier.summary()","metadata":{"execution":{"iopub.status.busy":"2022-10-23T15:45:40.641947Z","iopub.execute_input":"2022-10-23T15:45:40.642303Z","iopub.status.idle":"2022-10-23T15:45:41.804633Z","shell.execute_reply.started":"2022-10-23T15:45:40.642266Z","shell.execute_reply":"2022-10-23T15:45:41.803809Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = experiment(vit_classifier)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-10-23T15:45:41.806457Z","iopub.execute_input":"2022-10-23T15:45:41.806816Z","iopub.status.idle":"2022-10-23T18:28:49.594160Z","shell.execute_reply.started":"2022-10-23T15:45:41.806778Z","shell.execute_reply":"2022-10-23T18:28:49.593308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model Performance Visulization","metadata":{}},{"cell_type":"code","source":"# list all data in history\nprint(history.history.keys())\n# summarize history for accuracy\nplt.figure(figsize=(12,10))\nplt.plot(history.history['accuracy'])\nplt.plot(history.history['val_accuracy'])\nplt.title('model accuracy')\nplt.ylabel('accuracy')\nplt.xlabel('epoch')\nplt.legend(['train', 'test'], loc='upper left')\nplt.show()\n# summarize history for loss\nplt.figure(figsize=(12,10))\nplt.plot(history.history['loss'])\nplt.plot(history.history['val_loss'])\nplt.title('model loss')\nplt.ylabel('loss')\nplt.xlabel('epoch')\nplt.legend(['train', 'test'], loc='upper left')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-10-23T18:28:49.595521Z","iopub.execute_input":"2022-10-23T18:28:49.595884Z","iopub.status.idle":"2022-10-23T18:28:49.943554Z","shell.execute_reply.started":"2022-10-23T18:28:49.595838Z","shell.execute_reply":"2022-10-23T18:28:49.942776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# summarize history for loss\nplt.figure(figsize=(12,10))\nplt.plot(history.history['AUC'])\nplt.plot(history.history['val_AUC'])\nplt.title('model AUC')\nplt.ylabel('AUC')\nplt.xlabel('epoch')\nplt.legend(['train', 'test'], loc='upper left')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-10-23T18:28:52.438985Z","iopub.execute_input":"2022-10-23T18:28:52.439371Z","iopub.status.idle":"2022-10-23T18:28:52.602759Z","shell.execute_reply.started":"2022-10-23T18:28:52.439337Z","shell.execute_reply":"2022-10-23T18:28:52.601798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"vit_classifier.load_weights(\"./tmp/checkpoint\")","metadata":{"execution":{"iopub.status.busy":"2022-10-23T18:32:56.403134Z","iopub.execute_input":"2022-10-23T18:32:56.403466Z","iopub.status.idle":"2022-10-23T18:32:57.095208Z","shell.execute_reply.started":"2022-10-23T18:32:56.403434Z","shell.execute_reply":"2022-10-23T18:32:57.094204Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_img_array(img):\n    \n    # `array` is a float32 Numpy array of shape (299, 299, 3)\n    array = keras.preprocessing.image.img_to_array(img)\n    # We add a dimension to transform our array into a \"batch\"\n    # of size (1, 299, 299, 3)\n    array = np.expand_dims(array, axis=0)\n    return array","metadata":{"execution":{"iopub.status.busy":"2022-10-23T18:33:03.828572Z","iopub.execute_input":"2022-10-23T18:33:03.828924Z","iopub.status.idle":"2022-10-23T18:33:03.833777Z","shell.execute_reply.started":"2022-10-23T18:33:03.828889Z","shell.execute_reply":"2022-10-23T18:33:03.832723Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## The Grad-CAM algorithm","metadata":{}},{"cell_type":"code","source":"def gradcam_heatmap(img_array, model, last_conv_layer_name, pred_index=None):\n    # First, we create a model that maps the input image to the activations\n    # of the last conv layer as well as the output predictions\n    grad_model = tf.keras.models.Model(\n        [model.input], [model.get_layer(last_conv_layer_name).output,  model.output]\n    )\n    \n    # Then, we compute the gradient of the top predicted class for our input image\n    # with respect to the activations of the last conv layer\n    with tf.GradientTape() as tape:\n        last_conv_layer_output, preds = grad_model(img_array)\n        if pred_index is None:\n            pred_index = tf.argmax(preds[0])\n        class_channel = preds[:, pred_index]\n        \n        \n    # This is the gradient of the output neuron (top predicted or chosen)\n    # with regard to the output feature map of the last conv layer\n    grads = tape.gradient(class_channel, last_conv_layer_output)\n\n    # This is a vector where each entry is the mean intensity of the gradient\n    # over a specific feature map channel\n    pooled_grads = tf.reduce_mean(grads, axis=(0, 1))\n    # We multiply each channel in the feature map array\n    # by \"how important this channel is\" with regard to the top predicted class\n    # then sum all the channels to obtain the heatmap class activation\n    last_conv_layer_output = last_conv_layer_output\n    heatmap = last_conv_layer_output @ pooled_grads[..., tf.newaxis]\n    heatmap = tf.squeeze(heatmap)\n    \n    # For visualization purpose, we will also normalize the heatmap between 0 & 1\n    heatmap = tf.maximum(heatmap, 0) / tf.math.reduce_max(heatmap)\n    return heatmap.numpy()","metadata":{"execution":{"iopub.status.busy":"2022-10-23T18:33:18.757357Z","iopub.execute_input":"2022-10-23T18:33:18.757710Z","iopub.status.idle":"2022-10-23T18:33:18.765113Z","shell.execute_reply.started":"2022-10-23T18:33:18.757676Z","shell.execute_reply":"2022-10-23T18:33:18.763956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## superimposed visualization","metadata":{}},{"cell_type":"code","source":"classes.values()","metadata":{"execution":{"iopub.status.busy":"2022-10-23T18:35:12.589676Z","iopub.execute_input":"2022-10-23T18:35:12.590024Z","iopub.status.idle":"2022-10-23T18:35:12.596738Z","shell.execute_reply.started":"2022-10-23T18:35:12.589992Z","shell.execute_reply":"2022-10-23T18:35:12.595599Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def display_gradcam(img, heatmap, cam_path=\"cam.jpg\", alpha=0.4, preds=[0,0,0,0,0], plot=None):\n\n    # Rescale heatmap to a range 0-255\n    heatmap = np.uint8(255 * heatmap)\n\n    # Use jet colormap to colorize heatmap\n    jet = cm.get_cmap(\"jet\")\n\n    # Use RGB values of the colormap\n    jet_colors = jet(np.arange(256))[:, :3]\n    jet_heatmap = jet_colors[heatmap]\n\n    # Create an image with RGB colorized heatmap\n    jet_heatmap = keras.preprocessing.image.array_to_img(jet_heatmap)\n    jet_heatmap = jet_heatmap.resize((img.shape[1], img.shape[0]))\n    jet_heatmap = keras.preprocessing.image.img_to_array(jet_heatmap)\n\n    # Superimpose the heatmap on original image\n    superimposed_img = jet_heatmap * alpha + img\n    superimposed_img = keras.preprocessing.image.array_to_img(superimposed_img)\n\n    # Save the superimposed image\n    #superimposed_img.save(cam_path)\n\n    # Display Grad CAM\n    plot.imshow(superimposed_img)\n    plot.set(title =\n        \" No DR: \\\n        {:.3f}\\nMild: \\\n        {:.3f}\\nModerate: \\\n        {:.3f}\\nSevere: \\\n        {:.3f}\\nProliferative: \\\n        {:.3f}\".format(preds[0], \\\n                    preds[1], \\\n                    preds[2], \\\n                    preds[3],\n                    preds[4])\n    )\n    plot.axis('off')","metadata":{"execution":{"iopub.status.busy":"2022-10-23T18:47:30.266156Z","iopub.execute_input":"2022-10-23T18:47:30.266498Z","iopub.status.idle":"2022-10-23T18:47:30.273843Z","shell.execute_reply.started":"2022-10-23T18:47:30.266464Z","shell.execute_reply":"2022-10-23T18:47:30.272726Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Implement Grad CAM","metadata":{}},{"cell_type":"code","source":"# As in layer_normalization (LayerNorma (None, 1296, 64) ) \n#the last dim is 1296 so 36x36 for heatmap\nnp.sqrt(1296)","metadata":{"execution":{"iopub.status.busy":"2022-10-23T18:42:33.798281Z","iopub.execute_input":"2022-10-23T18:42:33.798633Z","iopub.status.idle":"2022-10-23T18:42:33.805565Z","shell.execute_reply.started":"2022-10-23T18:42:33.798585Z","shell.execute_reply":"2022-10-23T18:42:33.804460Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_image = next(iter(test_loader))[0][5]\n# Prepare image\nimg_array =get_img_array(test_image)\n\nlast_conv_layer_name = 'layer_normalization'\n# Remove last layer's softmax\nvit_classifier.layers[-1].activation = None\n# Print what the top predicted class is\npreds = vit_classifier.predict(img_array)\nprint(\"Predicted:\\n\" + \"No DR: \\\n    {p1}\\nMild: {p2}\\nModerate: \\\n    {p3}\\nSevere: \\\n    {p4}\\nProliferative: {p5}\".format(p1=preds[0][0], \\\n                                            p2=preds[0][1],p3=preds[0][2],p4=preds[0][3],p5=preds[0][4]))\n# Generate class activation heatmap\nheatmap = gradcam_heatmap(img_array, vit_classifier, last_conv_layer_name)\n\nheatmap = np.reshape(heatmap, (36,36))\n# Display heatmap\nplt.matshow(heatmap)\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-10-23T18:47:43.669080Z","iopub.execute_input":"2022-10-23T18:47:43.669406Z","iopub.status.idle":"2022-10-23T18:47:44.261647Z","shell.execute_reply.started":"2022-10-23T18:47:43.669369Z","shell.execute_reply":"2022-10-23T18:47:44.260526Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Heat-Map Visualization over Test-set","metadata":{}},{"cell_type":"code","source":"fig, axis = plt.subplots(3, 2, figsize=(30, 30))\nfor images, ax in zip(next(iter(test_loader))[0][:6], axis.flat):\n    img_array = get_img_array(images)\n    # Remove last layer's softmax\n    vit_classifier.layers[-1].activation = None\n    # Print what the top predicted class is\n    preds = vit_classifier.predict(img_array)\n    heatmap = gradcam_heatmap(img_array, vit_classifier, last_conv_layer_name)\n\n    heatmap = np.reshape(heatmap, (36,36))\n    display_gradcam(images, heatmap, preds=preds[0], plot=ax)","metadata":{"execution":{"iopub.status.busy":"2022-10-23T18:49:15.181455Z","iopub.execute_input":"2022-10-23T18:49:15.181812Z","iopub.status.idle":"2022-10-23T18:49:17.959939Z","shell.execute_reply.started":"2022-10-23T18:49:15.181779Z","shell.execute_reply":"2022-10-23T18:49:17.958197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axis = plt.subplots(3, 2, figsize=(30, 30))\nfor images, ax in zip(next(iter(test_loader))[0][6:12], axis.flat):\n    img_array = get_img_array(images)\n    # Remove last layer's softmax\n    vit_classifier.layers[-1].activation = None\n    # Print what the top predicted class is\n    preds = vit_classifier.predict(img_array)\n    heatmap = gradcam_heatmap(img_array, vit_classifier, last_conv_layer_name)\n\n    heatmap = np.reshape(heatmap, (36,36))\n    display_gradcam(images, heatmap, preds=preds[0], plot=ax)","metadata":{"execution":{"iopub.status.busy":"2022-10-23T18:53:50.967046Z","iopub.execute_input":"2022-10-23T18:53:50.967462Z","iopub.status.idle":"2022-10-23T18:53:53.866726Z","shell.execute_reply.started":"2022-10-23T18:53:50.967417Z","shell.execute_reply":"2022-10-23T18:53:53.862786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}