{"cells":[{"metadata":{},"cell_type":"markdown","source":"## Simple and easy to understand notebook for leveraging TPUs for Image Classification"},{"metadata":{},"cell_type":"markdown","source":"## Importing libraries"},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"from kaggle_datasets import KaggleDatasets\n\nimport cv2\nimport json\nimport datetime\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\nfrom sklearn.metrics import accuracy_score\nfrom sklearn.model_selection import train_test_split\n\nimport tensorflow as tf\nimport tensorflow.keras.layers as L\nfrom tensorflow.keras import models\nfrom tensorflow.keras.optimizers import Adam\nfrom tensorflow.keras.preprocessing import image\nfrom tensorflow.keras.preprocessing.image import ImageDataGenerator\nfrom tensorflow.keras.callbacks import ModelCheckpoint, EarlyStopping, ReduceLROnPlateau","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Initializing the validating the allocation of TPU (And the no of replicas)"},{"metadata":{"trusted":true},"cell_type":"code","source":"try:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver()\n    tf.config.experimental_connect_to_cluster(tpu)\n    tf.tpu.experimental.initialize_tpu_system(tpu)\n    strategy = tf.distribute.experimental.TPUStrategy(tpu)\n    print(\"Running on TPU:\", tpu.master())\n    \nexcept ValueError:\n    strategy = tf.distribute.get_strategy()\n    \nprint(f\"Running on {strategy.num_replicas_in_sync} replicas\")","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Functions for minor Data Augmentations and Image Resizing"},{"metadata":{"trusted":true},"cell_type":"code","source":"def decode_image(path, label = None, target_size = (512, 512)):\n    \n    img = tf.image.decode_jpeg(tf.io.read_file(path), channels = 3)\n    img = tf.cast(img, tf.float32) / 255.0\n    img = tf.image.resize(img, target_size)\n    \n    return img if label is None else img, label\n\ndef data_augment(img, label = None):\n    \n    img = tf.image.random_flip_left_right(img)\n    #img = tf.image.random_saturation(img, 5, 10)\n    #img = tf.image.random_brightness(img, 0.2)\n    #img = tf.image.random_contrast(img, 0.2, 0.5)\n    #img = tf.image.random_jpeg_quality(img, 75, 95)\n    img = tf.image.rot90(img, k=1) \n    img = tf.image.random_brightness(img, max_delta =.2)\n\n    return img if label is None else img, label","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## TPU-specific Hyperparameters"},{"metadata":{"trusted":true},"cell_type":"code","source":"AUTO = tf.data.experimental.AUTOTUNE\nBATCH_SIZE = strategy.num_replicas_in_sync * 16\nGCS_DS_PATH = KaggleDatasets().get_gcs_path('cassava-leaf-disease-classification')","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Loading the data.."},{"metadata":{"trusted":true},"cell_type":"code","source":"load_dir = \"/kaggle/input/cassava-leaf-disease-classification/\"\ndf = pd.read_csv(load_dir + 'train.csv')\ndf['paths'] = GCS_DS_PATH + \"/train_images/\" + df.image_id\nsub_df = pd.read_csv(load_dir + 'sample_submission.csv')","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Segregation into Train & Test splits"},{"metadata":{"trusted":true},"cell_type":"code","source":"train_df, valid_df = train_test_split(df, test_size = 0.25, random_state = 42)\n\ntrain_dataset = (tf.data.Dataset\n                 .from_tensor_slices((train_df.paths, train_df.label))\n                 .map(decode_image, num_parallel_calls = AUTO).cache()\n                 .map(data_augment, num_parallel_calls = AUTO).repeat()\n                 .shuffle(1024).batch(BATCH_SIZE).prefetch(AUTO))\n\nvalid_dataset = (tf.data.Dataset\n                 .from_tensor_slices((valid_df.paths, valid_df.label))\n                 .map(decode_image, num_parallel_calls=AUTO).cache()\n                 .batch(BATCH_SIZE).prefetch(AUTO))","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Defining the model under strategy.scope() - TPU"},{"metadata":{"trusted":true},"cell_type":"code","source":"with strategy.scope():\n    model = tf.keras.Sequential([ tf.keras.applications.Xception(include_top = False, weights = 'imagenet', input_shape = (512, 512, 3)), L.GlobalAveragePooling2D(), L.Dense(512, activation = 'relu'), L.Dropout(0.5), L.Dense(5, activation='softmax') ])    \n    model.compile(optimizer = 'adam', loss='sparse_categorical_crossentropy', metrics = ['sparse_categorical_accuracy'])\n    model.summary()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Defining the Callbacks"},{"metadata":{"trusted":true},"cell_type":"code","source":"checkpoint = ModelCheckpoint('TPU-Xception-Im512.h5', save_best_only = True, monitor = 'val_loss', mode = 'min', verbose = 1)\nearly_stop = EarlyStopping(monitor = 'val_loss', min_delta = 0.001, patience = 10, mode = 'min', verbose = 1, restore_best_weights = True)\nreduce_lr = ReduceLROnPlateau(monitor = 'val_loss', factor = 0.3, patience = 3, min_delta = 0.001, mode = 'min', verbose = 1)\n\nsteps_per_epoch = train_df.shape[0] // BATCH_SIZE","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Training the model - 20 Epochs"},{"metadata":{"trusted":true},"cell_type":"code","source":"history = model.fit(train_dataset, epochs = 30, validation_data = valid_dataset, callbacks = [checkpoint, early_stop, reduce_lr], steps_per_epoch = steps_per_epoch)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Plotting the metrics\n\n### Accuracy Graph"},{"metadata":{"trusted":true},"cell_type":"code","source":"plt.plot(history.history['sparse_categorical_accuracy'])\nplt.plot(history.history['val_sparse_categorical_accuracy'])\nplt.title('Model accuracy')\nplt.ylabel('Accuracy')\nplt.xlabel('Epoch')\nplt.legend(['Train', 'Test'], loc='upper left')\nplt.show()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Loss Graph"},{"metadata":{"trusted":true},"cell_type":"code","source":"plt.plot(history.history['loss'])\nplt.plot(history.history['val_loss'])\nplt.title('Model loss')\nplt.ylabel('Accuracy')\nplt.xlabel('Epoch')\nplt.legend(['Train', 'Test'], loc='upper left')\nplt.show()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Creating the Submission File"},{"metadata":{"trusted":true},"cell_type":"code","source":"import os\nfrom tensorflow.keras.preprocessing.image import load_img\nfrom tensorflow.keras.preprocessing.image import img_to_array\n\nBASE_DIR = \"/kaggle/input/cassava-leaf-disease-classification/\"\npreds = []\n\nfor image_id in sub_df.image_id:\n    \n    image = load_img(os.path.join(BASE_DIR, \"test_images\", image_id), target_size = (512, 512))\n    \n    #Image Processes & Reshaping\n    image = img_to_array(image)    \n    image = image.reshape(1, 512, 512, 3)\n\n    #Center Pixel Data & Normalize into float\n    image = image.astype('float32')\n    preds.append(np.argmax(model.predict(image)))\n\nsub_df['label'] = preds\n\nsub_df.to_csv('submission.csv', index = False)","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}