{"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":"# Inference Notebook is found here\n[Inference Notebook](https://www.kaggle.com/shubham219/experiment-with-models-in-keras-inference)","metadata":{}},{"cell_type":"markdown","source":"# Importing All The Required Liraries","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport re\nimport pandas as pd\nimport numpy as np\n\nimport tensorflow as tf\nimport tensorflow.keras.layers as tfl\n\nimport shutil\nfrom functools import partial\n\nfrom matplotlib import pyplot as plt\n\nfrom sklearn.model_selection import train_test_split \nfrom sklearn.utils import class_weight\nfrom sklearn.model_selection import KFold\nfrom sklearn.utils import shuffle\nfrom sklearn.utils import class_weight\n\n\nprint(\"Tensorflow version -\",tf.__version__)\nprint(\"Python version\")\n!python --version","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-07-08T17:15:51.782651Z","iopub.execute_input":"2023-07-08T17:15:51.783094Z","iopub.status.idle":"2023-07-08T17:15:52.766182Z","shell.execute_reply.started":"2023-07-08T17:15:51.783059Z","shell.execute_reply":"2023-07-08T17:15:52.765031Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# TPU Config","metadata":{}},{"cell_type":"code","source":"try:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver()\n    print(\"Device:\", tpu.master())\n    tf.config.experimental_connect_to_cluster(tpu)\n    tf.tpu.experimental.initialize_tpu_system(tpu)\n    strategy = tf.distribute.experimental.TPUStrategy(tpu)\nexcept:\n    strategy = tf.distribute.get_strategy()\nprint(\"Number of replicas:\", strategy.num_replicas_in_sync)","metadata":{"execution":{"iopub.status.busy":"2023-07-08T17:15:52.769257Z","iopub.execute_input":"2023-07-08T17:15:52.769659Z","iopub.status.idle":"2023-07-08T17:15:52.777289Z","shell.execute_reply.started":"2023-07-08T17:15:52.769622Z","shell.execute_reply":"2023-07-08T17:15:52.776163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Variables","metadata":{}},{"cell_type":"code","source":"EPOCHS = 100\nIMAGE_SIZE = (512, 512)\nAUTOTUNE = tf.data.experimental.AUTOTUNE\nSEED = 123\nBATCH_SIZE = 16*strategy.num_replicas_in_sync","metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","execution":{"iopub.status.busy":"2023-07-08T17:15:52.778962Z","iopub.execute_input":"2023-07-08T17:15:52.779666Z","iopub.status.idle":"2023-07-08T17:15:52.787177Z","shell.execute_reply.started":"2023-07-08T17:15:52.779635Z","shell.execute_reply":"2023-07-08T17:15:52.786044Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Reading Metadata","metadata":{}},{"cell_type":"code","source":"train_img_dir = '/kaggle/input/cassava-leaf-disease-classification/train_images/'\n\ndata = pd.read_csv('../input/cassava-leaf-disease-classification/train.csv')\n\ndata['image_id'] = data['image_id'].apply(lambda x : train_img_dir+x)\n\n\nwith open('../input/cassava-leaf-disease-classification/label_num_to_disease_map.json') as file:\n    text = file.read()\nprint(text)\ndata.head()","metadata":{"execution":{"iopub.status.busy":"2023-07-08T17:15:52.790347Z","iopub.execute_input":"2023-07-08T17:15:52.790734Z","iopub.status.idle":"2023-07-08T17:15:52.831267Z","shell.execute_reply.started":"2023-07-08T17:15:52.790681Z","shell.execute_reply":"2023-07-08T17:15:52.830367Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Distribution Of Classes\nDataset is very imbalanced","metadata":{}},{"cell_type":"code","source":"figure = plt.figure(figsize=(8,4))\n(data['label'].value_counts()/len(data)*100).plot(kind='bar')\nplt.title(\"Distribution of Classes\")\nplt.ylabel('% count of classes')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-07-08T17:15:52.832761Z","iopub.execute_input":"2023-07-08T17:15:52.833105Z","iopub.status.idle":"2023-07-08T17:15:53.088535Z","shell.execute_reply.started":"2023-07-08T17:15:52.833076Z","shell.execute_reply":"2023-07-08T17:15:53.087653Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Class Weights","metadata":{}},{"cell_type":"code","source":"tot_len = len(data)\nclass_weights = data['label'].value_counts()\nclass_weights = tot_len/class_weights\nclass_weights = class_weights.to_dict()\nclass_weights = {k:v for k,v in sorted(class_weights.items()) }\nclass_weights","metadata":{"execution":{"iopub.status.busy":"2023-07-08T17:15:53.089814Z","iopub.execute_input":"2023-07-08T17:15:53.090768Z","iopub.status.idle":"2023-07-08T17:15:53.102296Z","shell.execute_reply.started":"2023-07-08T17:15:53.090732Z","shell.execute_reply":"2023-07-08T17:15:53.101019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Splitting the data into training and validation ","metadata":{}},{"cell_type":"code","source":"data = shuffle(data, random_state=SEED)\n\nx_train, x_valid, y_train, y_valid = train_test_split(data['image_id'], data['label'],\n                                                      stratify=data['label'],\n                                                      test_size=0.1\n                                                     )\nprint(\"X_Train Shape: \", x_train.shape)\nprint(\"X_Validation Shape: \", x_valid.shape)\n\nprint(\"Disctribution of labels\")\nprint(y_train.value_counts()/len(y_train))\nprint(y_valid.value_counts()/len(y_valid))","metadata":{"execution":{"iopub.status.busy":"2023-07-08T17:15:53.104083Z","iopub.execute_input":"2023-07-08T17:15:53.104574Z","iopub.status.idle":"2023-07-08T17:15:53.131531Z","shell.execute_reply.started":"2023-07-08T17:15:53.104534Z","shell.execute_reply":"2023-07-08T17:15:53.130654Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def augmentation(image, label):\n    aug = tf.keras.models.Sequential()\n    aug.add(tfl.RandomRotation(0.4))\n    aug.add(tfl.RandomFlip('horizontal'))\n    aug.add(tfl.RandomFlip('vertical'))\n    aug.add(tfl.RandomZoom(0.2))\n    image = aug(image)\n    return image, label\n\ndef resnet_preprocess(image, label):\n    preprocess = tf.keras.applications.resnet50.preprocess_input\n    image = preprocess(image)\n    return image, label\n    \n    \ndef load_image_and_label_from_path(image_path, label):\n    img = tf.io.read_file(image_path)\n    img = tf.image.decode_jpeg(img, channels=3)\n    return img, label\n","metadata":{"execution":{"iopub.status.busy":"2023-07-08T17:15:53.132592Z","iopub.execute_input":"2023-07-08T17:15:53.132841Z","iopub.status.idle":"2023-07-08T17:15:53.140712Z","shell.execute_reply.started":"2023-07-08T17:15:53.132820Z","shell.execute_reply":"2023-07-08T17:15:53.139800Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds = tf.data.Dataset.from_tensor_slices((x_train.values, \n                                              y_train.values))\n\nvalidation_ds = tf.data.Dataset.from_tensor_slices((x_valid.values,\n                                                    y_valid.values))\n\n\ntrain_ds = train_ds.map(load_image_and_label_from_path,\n                        num_parallel_calls=AUTOTUNE)\n\ntrain_ds = train_ds.map(augmentation,\n                        num_parallel_calls=AUTOTUNE)\n\ntrain_ds = train_ds.map(resnet_preprocess,\n                        num_parallel_calls=AUTOTUNE)\n\ntrain_ds = train_ds.repeat()\n\ntrain_ds = train_ds.batch(BATCH_SIZE)\n\ntrain_ds = train_ds.prefetch(buffer_size=AUTOTUNE)\n\n\n# Validation DataSet \nvalidation_ds = validation_ds.map(load_image_and_label_from_path, \n                                    num_parallel_calls=AUTOTUNE)\n\nvalidation_ds = validation_ds.map(augmentation,\n                        num_parallel_calls=AUTOTUNE)\n\nvalidation_ds = validation_ds.map(resnet_preprocess,\n                        num_parallel_calls=AUTOTUNE)\n\nvalidation_ds = validation_ds.repeat()\n\n\nvalidation_ds = validation_ds.batch(BATCH_SIZE)\n\nvalidation_ds = validation_ds.prefetch(buffer_size=AUTOTUNE)\n\n\nprint(train_ds)","metadata":{"execution":{"iopub.status.busy":"2023-07-08T17:15:53.142048Z","iopub.execute_input":"2023-07-08T17:15:53.142799Z","iopub.status.idle":"2023-07-08T17:15:53.774110Z","shell.execute_reply.started":"2023-07-08T17:15:53.142769Z","shell.execute_reply":"2023-07-08T17:15:53.773156Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for img, label in train_ds.take(1):\n    print(\"Image Size: \",img.shape)\n    print(\"Labels :\", label.shape)\n    print(label)","metadata":{"execution":{"iopub.status.busy":"2023-07-08T17:15:53.777821Z","iopub.execute_input":"2023-07-08T17:15:53.778107Z","iopub.status.idle":"2023-07-08T17:15:55.458767Z","shell.execute_reply.started":"2023-07-08T17:15:53.778083Z","shell.execute_reply":"2023-07-08T17:15:55.457664Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-07-08T17:15:55.460277Z","iopub.execute_input":"2023-07-08T17:15:55.460622Z","iopub.status.idle":"2023-07-08T17:15:55.772833Z","shell.execute_reply.started":"2023-07-08T17:15:55.460589Z","shell.execute_reply":"2023-07-08T17:15:55.771655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualing Image Samples","metadata":{}},{"cell_type":"code","source":"disease_labels = {\n      0: \"Cassava Bacterial Blight (CBB)\",\n      1: \"Cassava Brown Streak Disease (CBSD)\",\n      2: \"Cassava Green Mottle (CGM)\",\n      3: \"Cassava Mosaic Disease (CMD)\",\n      4: \"Healthy\"\n    }\n\nplt.figure(figsize=(12, 12))\n\nfor i, (img, lab) in enumerate(train_ds.take(9)):\n    ax = plt.subplot(3,3,i+1)\n    plt.imshow(np.array(img[0]).astype(np.int8))\n    plt.title(disease_labels[np.asarray(lab)[0]] + \" - \"+str(np.asarray(lab)[0]))\n    plt.axis(\"off\")\n    ","metadata":{"execution":{"iopub.status.busy":"2023-07-08T17:15:55.777222Z","iopub.execute_input":"2023-07-08T17:15:55.777514Z","iopub.status.idle":"2023-07-08T17:16:09.554870Z","shell.execute_reply.started":"2023-07-08T17:15:55.777483Z","shell.execute_reply":"2023-07-08T17:16:09.553763Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualizing Images After Augmentation","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(12, 12))\n\nfor i, (img, lab) in enumerate(train_ds.take(1)):\n    print(\"Shape of image: \", img.shape)\n\n    for i in range(9): # Plot first image\n        if i==0:\n            ax = plt.subplot(3,3,i+1)\n            plt.imshow(np.array(img[0]).astype(int))\n        else: # plot augmentated images\n            plt.axis('off')\n            ax = plt.subplot(3,3,i+1)\n            image, label = augmentation(img, lab)\n            plt.imshow(np.array(image[0]).astype(int))\n            plt.axis('off')\n","metadata":{"execution":{"iopub.status.busy":"2023-07-08T17:16:09.556546Z","iopub.execute_input":"2023-07-08T17:16:09.556881Z","iopub.status.idle":"2023-07-08T17:16:14.403550Z","shell.execute_reply.started":"2023-07-08T17:16:09.556852Z","shell.execute_reply":"2023-07-08T17:16:14.402725Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model Preparation - With Pretraied Resnet50","metadata":{}},{"cell_type":"code","source":"def make_model(input_shape):\n\n    base_model = tf.keras.applications.ResNet50(input_shape=IMAGE_SHAPE,\n                                                include_top=False,\n                                                weights='imagenet'\n                                               )\n\n    print(\"Total Layers: \", len(base_model.layers))\n    freeze_layers_at = 140\n\n    for layer in base_model.layers[:freeze_layers_at]:\n        layer.trainable=False\n\n    inputs = tf.keras.Input(shape=input_shape)\n    x = base_model(inputs, training=False)\n    x = tfl.GlobalAveragePooling2D()(x)\n    x = tfl.Flatten()(x)\n    x = tfl.Dense(256, activation='relu', kernel_regularizer='l2',\n                  bias_regularizer=tf.keras.regularizers.L1L2(l1=0.01, l2=0.001))(x)\n    x = tfl.Dropout(0.5)(x)\n    output = tfl.Dense(5, activation='softmax')(x)\n\n    model = tf.keras.Model(inputs, output)\n\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-07-08T17:16:14.404773Z","iopub.execute_input":"2023-07-08T17:16:14.406221Z","iopub.status.idle":"2023-07-08T17:16:14.416125Z","shell.execute_reply.started":"2023-07-08T17:16:14.406187Z","shell.execute_reply":"2023-07-08T17:16:14.414771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with strategy.scope():\n    \n    IMAGE_SHAPE = IMAGE_SIZE+(3,)\n    model = make_model(IMAGE_SHAPE)\n    \n    model.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=0.001),\n                  loss = tf.keras.losses.SparseCategoricalCrossentropy(),\n                  metrics=['sparse_categorical_accuracy']\n                 )\n\nmodel.summary()","metadata":{"execution":{"iopub.status.busy":"2023-07-08T17:16:14.417513Z","iopub.execute_input":"2023-07-08T17:16:14.418166Z","iopub.status.idle":"2023-07-08T17:16:20.228934Z","shell.execute_reply.started":"2023-07-08T17:16:14.418133Z","shell.execute_reply":"2023-07-08T17:16:20.228030Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Network Architecture","metadata":{}},{"cell_type":"code","source":"display(tf.keras.utils.plot_model(model))","metadata":{"execution":{"iopub.status.busy":"2023-07-08T17:16:20.230409Z","iopub.execute_input":"2023-07-08T17:16:20.231442Z","iopub.status.idle":"2023-07-08T17:16:20.439439Z","shell.execute_reply.started":"2023-07-08T17:16:20.231409Z","shell.execute_reply":"2023-07-08T17:16:20.438511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define training epochs\ncheckpoint_cb = tf.keras.callbacks.ModelCheckpoint(\"cassava_base.h5\", \n                                                   save_best_only=True)\n\nearly_stopping_cb = tf.keras.callbacks.EarlyStopping(patience=20,\n                                                     monitor='val_loss',\n                                                     mode='min',\n                                                     restore_best_weights=True\n                                                    )\n\nreduce_lr = tf.keras.callbacks.ReduceLROnPlateau(monitor = 'val_loss',\n                               factor = 0.2,\n                               patience = 2,\n                               min_lr = 1e-6,\n                               mode = 'min',\n                               verbose = 1\n                              )\n    \n    \nEPOCHS = 50\ntrain_images_cnt = len(x_train)\n\nSTEPS_PER_EPOCH = train_images_cnt // BATCH_SIZE\n\nvalid_steps=len(x_valid)//BATCH_SIZE\n\nhistory = model.fit(train_ds,\n                    validation_data=validation_ds,\n                    validation_steps = valid_steps,\n                    epochs=EPOCHS,\n                    steps_per_epoch=STEPS_PER_EPOCH,\n#                     class_weight=class_weights,\n                    callbacks=[checkpoint_cb, early_stopping_cb],\n)","metadata":{"execution":{"iopub.status.busy":"2023-07-08T17:18:33.686059Z","iopub.execute_input":"2023-07-08T17:18:33.686489Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# CHECKING THE METRICS\n\nprint('Train_Cat-Acc: ', max(history.history['sparse_categorical_accuracy']))\nprint('Val_Cat-Acc: ', max(history.history['val_sparse_categorical_accuracy']))","metadata":{"execution":{"iopub.status.busy":"2023-07-08T17:16:20.516412Z","iopub.status.idle":"2023-07-08T17:16:20.516849Z","shell.execute_reply.started":"2023-07-08T17:16:20.516636Z","shell.execute_reply":"2023-07-08T17:16:20.516656Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.title(\"Loss Over Epochs\")\nplt.plot(history.history['loss'])\nplt.plot(history.history['val_loss'])\nplt.xlabel(\"Epochs\")\nplt.ylabel(\"Loss\")\n\n# PLOTTING RESULTS (Train vs Validation FOLDER 1)\n\ndef model_plots(acc,val_acc,loss,val_loss):\n    \n    fig, (ax1, ax2) = plt.subplots(1,2, figsize= (15,10))\n    fig.suptitle(\" MODEL'S METRICS VISUALIZATION \", fontsize=20)\n\n    ax1.plot(range(1, len(acc) + 1), acc)\n    ax1.plot(range(1, len(val_acc) + 1), val_acc)\n    ax1.set_title('History of Accuracy', fontsize=15)\n    ax1.set_xlabel('Epochs', fontsize=15)\n    ax1.set_ylabel('Accuracy', fontsize=15)\n    ax1.legend(['training', 'validation'])\n\n\n    ax2.plot(range(1, len(loss) + 1), loss)\n    ax2.plot(range(1, len(val_loss) + 1), val_loss)\n    ax2.set_title('History of Loss', fontsize=15)\n    ax2.set_xlabel('Epochs', fontsize=15)\n    ax2.set_ylabel('Loss', fontsize=15)\n    ax2.legend(['training', 'validation'])\n    plt.show()\n    \n\nmodel_plots(history.history['sparse_categorical_accuracy'],\n            history.history['val_sparse_categorical_accuracy'],\n            history.history['loss'],\n            history.history['val_loss'])","metadata":{"execution":{"iopub.status.busy":"2023-07-08T17:16:20.521812Z","iopub.status.idle":"2023-07-08T17:16:20.522262Z","shell.execute_reply.started":"2023-07-08T17:16:20.522037Z","shell.execute_reply":"2023-07-08T17:16:20.522057Z"},"trusted":true},"execution_count":null,"outputs":[]}]}