{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd \nimport numpy as np \nimport tensorflow as tf\nimport tensorflow_addons as tfa\nimport glob, warnings\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix, classification_report\nimport seaborn as sns\nimport os\nimport gc\nfrom sklearn.model_selection import train_test_split\nwarnings.filterwarnings('ignore')\nprint('TensorFlow Version ' + tf.__version__)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-06-19T10:45:42.851628Z","iopub.execute_input":"2023-06-19T10:45:42.852229Z","iopub.status.idle":"2023-06-19T10:45:50.464443Z","shell.execute_reply.started":"2023-06-19T10:45:42.852070Z","shell.execute_reply":"2023-06-19T10:45:50.463377Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_save_name = 'classifier.pt'\npath = F\"/kaggle/working/{model_save_name}\"","metadata":{"execution":{"iopub.status.busy":"2023-06-22T11:28:38.997724Z","iopub.execute_input":"2023-06-22T11:28:38.998241Z","iopub.status.idle":"2023-06-22T11:28:39.004720Z","shell.execute_reply.started":"2023-06-22T11:28:38.998201Z","shell.execute_reply":"2023-06-22T11:28:39.003266Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split","metadata":{"execution":{"iopub.status.busy":"2023-06-19T11:15:33.392451Z","iopub.execute_input":"2023-06-19T11:15:33.392947Z","iopub.status.idle":"2023-06-19T11:15:33.399187Z","shell.execute_reply.started":"2023-06-19T11:15:33.392908Z","shell.execute_reply":"2023-06-19T11:15:33.398123Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load the training data CSV file\nDF_TRAIN = pd.read_csv('/kaggle/input/diabetic-retinopathy-224x224-gaussian-filtered/train.csv', dtype=str)\nDF_TRAIN['image_path'] = \"./train_imgs_reshaped/\" + DF_TRAIN[\"id_code\"] + \".png\"","metadata":{"execution":{"iopub.status.busy":"2023-06-19T11:16:28.430543Z","iopub.execute_input":"2023-06-19T11:16:28.431016Z","iopub.status.idle":"2023-06-19T11:16:28.450327Z","shell.execute_reply.started":"2023-06-19T11:16:28.430981Z","shell.execute_reply":"2023-06-19T11:16:28.448494Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IMAGE_SIZE = 224\nBATCH_SIZE = 16\nEPOCHS = 7\nTRAIN_PATH = './train_imgs_reshaped'\n\nDF_TRAIN['image_path'] = \"./train_imgs_reshaped/\" + DF_TRAIN[\"id_code\"] + \".png\"\n\nDF_TEST = pd.DataFrame(glob.glob(os.path.join(TRAIN_PATH, \"*.png\")), columns=['image_path'])\n\nDF_TRAIN['image_path'] = DF_TRAIN[\"id_code\"] \nTEST_IMAGES = glob.glob(TEST_PATH )\nDF_TEST = pd.DataFrame(TEST_IMAGES, columns = ['image_path'])\n\nclasses = {0 : \"No DR\",\n           1 : \"Mild\",\n           2 : \"Moderate\",\n           3 : \"Severe\",\n           4 : \"Proliferative\"}\n","metadata":{"execution":{"iopub.status.busy":"2023-06-19T11:03:54.407516Z","iopub.execute_input":"2023-06-19T11:03:54.407964Z","iopub.status.idle":"2023-06-19T11:03:54.446848Z","shell.execute_reply.started":"2023-06-19T11:03:54.407929Z","shell.execute_reply":"2023-06-19T11:03:54.445795Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DF_TRAIN.head()","metadata":{"execution":{"iopub.status.busy":"2023-06-19T11:16:33.443929Z","iopub.execute_input":"2023-06-19T11:16:33.444371Z","iopub.status.idle":"2023-06-19T11:16:33.457526Z","shell.execute_reply.started":"2023-06-19T11:16:33.444334Z","shell.execute_reply":"2023-06-19T11:16:33.455894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define constants and parameters\nTRAIN_PATH = './train_imgs_reshaped'\nIMAGE_SIZE = 224\nBATCH_SIZE = 16\nEPOCHS = 7","metadata":{"execution":{"iopub.status.busy":"2023-06-19T11:16:41.932766Z","iopub.execute_input":"2023-06-19T11:16:41.933274Z","iopub.status.idle":"2023-06-19T11:16:41.939204Z","shell.execute_reply.started":"2023-06-19T11:16:41.933188Z","shell.execute_reply":"2023-06-19T11:16:41.938059Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def data_augment(image):\n    p_spatial = tf.random.uniform([], 0, 1.0, dtype = tf.float32)\n    #p_rotate = tf.random.uniform([], 0, 1.0, dtype = tf.float32)\n    p_pixel_1 = tf.random.uniform([], 0, 1.0, dtype = tf.float32)\n    p_pixel_2 = tf.random.uniform([], 0, 1.0, dtype = tf.float32)\n    p_pixel_3 = tf.random.uniform([], 0, 1.0, dtype = tf.float32)\n    \n    # Flips\n    #image =tf.image.random_flip_left_right(image)\n    #image =tf.image.random_flip_up_down(image)\n    \n    if p_spatial > .75:\n        image = tf.image.transpose(image)\n        \n  #Rotates\n  #if p_rotate > .75:\n    #image = tf.image.rot90(image, k = 3) # rotate 270°\n  #if p_rotate > .5:\n    #image = tf.image.rot90(image, k = 2) # rotate 180°\n #if p_rotate > .25:\n    #image = tf.image.rot90(image, k = 1) # rotate 90°\n    \n  #pixel_level transforms\n    if p_pixel_1 >= .4:\n        image = tf.image.random_saturation(image, lower = .7, upper = 1.3)\n    if p_pixel_2 >= .4:\n        image = tf.image.random_contrast(image, lower = .8, upper = 1.2)\n    if p_pixel_3 >= .4:\n        image = tf.image.random_brightness(image, max_delta = .1)\n\n    return image\n    ","metadata":{"execution":{"iopub.status.busy":"2023-06-19T11:04:00.033719Z","iopub.execute_input":"2023-06-19T11:04:00.034137Z","iopub.status.idle":"2023-06-19T11:04:00.046943Z","shell.execute_reply.started":"2023-06-19T11:04:00.034102Z","shell.execute_reply":"2023-06-19T11:04:00.045502Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"datagen = tf.keras.preprocessing.image.ImageDataGenerator(rescale = 1./255,\n                                                         samplewise_center = True,\n                                                         samplewise_std_normalization = True,\n                                                         validation_split = 0.2,\n                                                         preprocessing_function = data_augment)\n\ntrain_gen = datagen.flow_from_dataframe(dataframe = DF_TRAIN,\n                                       directory = TRAIN_PATH,\n                                       x_col = 'image_path',\n                                       y_col = 'diagnosis',\n                                       subset = 'training',\n                                       batch_size = BATCH_SIZE,\n                                       seed = 1,\n                                       color_mode = 'rgb',\n                                       shuffle = True,\n                                       class_mode = 'categorical',\n                                       target_size = (IMAGE_SIZE, IMAGE_SIZE))\nvalid_gen = datagen.flow_from_dataframe(dataframe = DF_TRAIN,\n                                       directory = TRAIN_PATH,\n                                       x_col = 'image_path',\n                                       y_col = 'diagnosis',\n                                       subset = 'validation',\n                                       batch_size = BATCH_SIZE,\n                                       seed = 1,\n                                       color_mode = 'rgb',\n                                       shuffle = False,\n                                       class_mode = 'categorical',\n                                       target_size = (IMAGE_SIZE, IMAGE_SIZE))\n\ntest_gen = datagen.flow_from_dataframe(dataframe = DF_TEST,\n                                       #directory = TRAIN_PATH\n                                       x_col = 'image_path',\n                                       y_col = None,\n                                       #subset = 'validation',\n                                       batch_size = BATCH_SIZE,\n                                       seed = 1,\n                                       color_mode = 'rgb',\n                                       shuffle = False,\n                                       class_mode = None,\n                                       target_size = (IMAGE_SIZE, IMAGE_SIZE))\n","metadata":{"execution":{"iopub.status.busy":"2023-06-19T11:18:44.494990Z","iopub.execute_input":"2023-06-19T11:18:44.495441Z","iopub.status.idle":"2023-06-19T11:18:44.641201Z","shell.execute_reply.started":"2023-06-19T11:18:44.495404Z","shell.execute_reply":"2023-06-19T11:18:44.639486Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"images = [train_gen[0][0][i] for i in range(16)]\nfig, axes = plt.subplots(3, 5, figsize = (10, 10))\n\naxes = axes .flatten()\n\nfor img, ax in zip(images, axes):\n    ax.imshow(img.reshape(IMAGE_SIZE, IMAGE_SIZE, 3))\n    ax.axis('off')\n    \nplt.tight_layout()\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2023-06-17T12:20:50.545648Z","iopub.execute_input":"2023-06-17T12:20:50.546144Z","iopub.status.idle":"2023-06-17T12:20:57.331780Z","shell.execute_reply.started":"2023-06-17T12:20:50.546107Z","shell.execute_reply":"2023-06-17T12:20:57.330517Z"}}},{"cell_type":"code","source":"!pip install --quiet vit-keras\n  \nfrom vit_keras import vit ","metadata":{"execution":{"iopub.status.busy":"2023-06-17T12:21:02.721577Z","iopub.execute_input":"2023-06-17T12:21:02.722051Z","iopub.status.idle":"2023-06-17T12:21:21.723625Z","shell.execute_reply.started":"2023-06-17T12:21:02.722013Z","shell.execute_reply":"2023-06-17T12:21:21.721980Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"vit_model = vit.vit_b32(\n        image_size = IMAGE_SIZE,\n        activation = 'softmax',\n        pretrained = True,\n        include_top = False,\n        pretrained_top = False,\n        classes = 5)","metadata":{"execution":{"iopub.status.busy":"2023-06-17T12:21:24.457866Z","iopub.execute_input":"2023-06-17T12:21:24.458421Z","iopub.status.idle":"2023-06-17T12:21:31.568210Z","shell.execute_reply.started":"2023-06-17T12:21:24.458368Z","shell.execute_reply":"2023-06-17T12:21:31.566836Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from vit_keras import visualize\n\nx = test_gen.next()\nimage = x[0]\n\nattention_map = visualize.attention_map(model = vit_model, image = image)\n\n#Plot results \nfig, (ax1, ax2) = plt.subplots(ncols = 2)\nax1.axis('off')\nax2.axis('off')\nax1.set_title('Original')\nax2.set_title('Attention.Map')\n_ = ax1.imshow(image)\n_ = ax2.imshow(attention_map)\n\n\n","metadata":{"execution":{"iopub.status.busy":"2023-06-17T12:21:34.257699Z","iopub.execute_input":"2023-06-17T12:21:34.258172Z","iopub.status.idle":"2023-06-17T12:21:34.318686Z","shell.execute_reply.started":"2023-06-17T12:21:34.258126Z","shell.execute_reply":"2023-06-17T12:21:34.316790Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = tf.keras.Sequential([\n        vit_model,\n        tf.keras.layers.Flatten(),\n        tf.keras.layers.BatchNormalization(),\n        tf.keras.layers.Dense(11, activation = tfa.activations.gelu),\n        tf.keras.layers.BatchNormalization(),\n        tf.keras.layers.Dense(5, 'softmax')\n    ],\n    name = 'vision_transformer')\n\nmodel.summary()","metadata":{"execution":{"iopub.status.busy":"2023-05-22T18:07:44.432260Z","iopub.execute_input":"2023-05-22T18:07:44.432709Z","iopub.status.idle":"2023-05-22T18:07:46.180042Z","shell.execute_reply.started":"2023-05-22T18:07:44.432669Z","shell.execute_reply":"2023-05-22T18:07:46.178783Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_b = tf.keras.Sequential([\n        vit_model,\n        tf.keras.layers.LayerNormalization(),\n        tf.keras.layers.Flatten(),\n        tf.keras.layers.BatchNormalization(),\n        tf.keras.layers.Dense(11, activation = tfa.activations.gelu),\n        tf.keras.layers.BatchNormalization(),\n        tf.keras.layers.Dense(5, 'softmax')\n    ],\n    name = 'vision_transformer')\n\nmodel_b.summary()","metadata":{"execution":{"iopub.status.busy":"2023-05-22T18:07:49.613110Z","iopub.execute_input":"2023-05-22T18:07:49.613557Z","iopub.status.idle":"2023-05-22T18:07:51.358994Z","shell.execute_reply.started":"2023-05-22T18:07:49.613519Z","shell.execute_reply":"2023-05-22T18:07:51.357803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"learning_rate = 1e-4\n\noptimizer = tfa.optimizers.RectifiedAdam(learning_rate = learning_rate)\n\nmodel.compile(optimizer = optimizer,\n             loss = tf.keras.losses.CategoricalCrossentropy(label_smoothing = 0.2),\n             metrics = ['accuracy'])\nSTEP_SIZE_TRAIN = train_gen.n // train_gen.batch_size\nSTEP_SIZE_VALID = valid_gen.n // valid_gen.batch_size\n\nreduce_lr = tf.keras.callbacks.ReduceLROnPlateau(monitor = 'val_accuracy',\n                                                factor = 0.2,\n                                                patience = 2,\n                                                verbose = 1,\n                                                min_delta = 1e-4,\n                                                min_lr = 1e-6,\n                                                mode = 'max')\n\nearlystopping = tf.keras.callbacks.EarlyStopping(monitor = 'val_accuracy',\n                                                min_delta = 1e-4,\n                                                patience = 5,\n                                                mode = 'max',\n                                                restore_best_weights = True,\n                                                verbose = 1)\n\ncheckpointer = tf.keras.callbacks.ModelCheckpoint(filepath = './model.hdf5',\n                                                 monitor = 'val_accuracy',\n                                                 verbose = 1,\n                                                 save_best_only = True,\n                                                 save_weights_only = True,\n                                                 mode = 'max')\ncallbacks = [earlystopping, reduce_lr, checkpointer]\n\nmodel.fit(x = train_gen,\n         steps_per_epoch = STEP_SIZE_TRAIN,\n         validation_data = valid_gen,\n         validation_steps = STEP_SIZE_VALID,\n         epochs = EPOCHS,\n         callbacks = callbacks)","metadata":{"execution":{"iopub.status.busy":"2023-05-23T08:11:18.925830Z","iopub.execute_input":"2023-05-23T08:11:18.926406Z","iopub.status.idle":"2023-05-23T08:11:19.060964Z","shell.execute_reply.started":"2023-05-23T08:11:18.926207Z","shell.execute_reply":"2023-05-23T08:11:19.058908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.load_weights(\"./model.hdf5\")","metadata":{"execution":{"iopub.status.busy":"2023-02-06T17:02:31.030836Z","iopub.execute_input":"2023-02-06T17:02:31.031257Z","iopub.status.idle":"2023-02-06T17:02:31.112467Z","shell.execute_reply.started":"2023-02-06T17:02:31.031225Z","shell.execute_reply":"2023-02-06T17:02:31.110932Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.summary()","metadata":{"execution":{"iopub.status.busy":"2023-02-06T11:30:33.844375Z","iopub.execute_input":"2023-02-06T11:30:33.844876Z","iopub.status.idle":"2023-02-06T11:30:33.870671Z","shell.execute_reply.started":"2023-02-06T11:30:33.844832Z","shell.execute_reply":"2023-02-06T11:30:33.868451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.get_layer(\"vit-b32\").summary()","metadata":{"execution":{"iopub.status.busy":"2023-02-06T11:30:39.827853Z","iopub.execute_input":"2023-02-06T11:30:39.828338Z","iopub.status.idle":"2023-02-06T11:30:39.850363Z","shell.execute_reply.started":"2023-02-06T11:30:39.828302Z","shell.execute_reply":"2023-02-06T11:30:39.849033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predicted_classes = np.argmax(model.predict(valid_gen, steps = valid_gen.n // valid_gen.batch_size + 1), axis = 1)\ntrue_classes = valid_gen.classes\nclass_labels = list(valid_gen.class_indices.keys())  \n\nconfusionmatrix = confusion_matrix(true_classes, predicted_classes)\nplt.figure(figsize = (16, 16))\nsns.heatmap(confusionmatrix, cmap = 'Blues', annot = True, cbar = True)\n\nprint(classification_report(true_classes, predicted_classes))","metadata":{"execution":{"iopub.status.busy":"2023-02-06T11:30:48.646444Z","iopub.execute_input":"2023-02-06T11:30:48.646922Z","iopub.status.idle":"2023-02-06T11:33:27.866263Z","shell.execute_reply.started":"2023-02-06T11:30:48.646886Z","shell.execute_reply":"2023-02-06T11:33:27.864455Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Creates a confusion matrix\ncm = confusion_matrix(true_classes, predicted_classes)\n#Transform to df for easier plotting\ncm_df = pd.DataFrame(cm,\n                     index = ['NO DR', 'Mild','Moderate','Severe', 'Poliferative DR'],\n                     columns = ['NO DR', 'Mild','Moderate','Severe', 'Poliferative DR'])\nplt.figure(figsize=(6,4))\nsns.heatmap(cm_df, square=True, annot=True, cmap='Greens', fmt='d' , cbar=False)\nplt.title('Fine tuned ViT Model Confusion Matrix')\nplt.ylabel('True label')\nplt.xlabel('Predicted label')\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2023-02-06T11:49:13.681546Z","iopub.execute_input":"2023-02-06T11:49:13.682091Z","iopub.status.idle":"2023-02-06T11:49:13.886159Z","shell.execute_reply.started":"2023-02-06T11:49:13.682042Z","shell.execute_reply":"2023-02-06T11:49:13.884953Z"},"trusted":true},"execution_count":null,"outputs":[]}],"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"}}