{"cells":[{"metadata":{},"cell_type":"markdown","source":"# Cassava Leaf - Training on TPU using TFRecords"},{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport os, json, cv2, math, re\nfrom PIL import Image\nfrom sklearn.model_selection import train_test_split, StratifiedKFold\nimport random\n\nimport tensorflow as tf\nfrom keras.preprocessing.image import ImageDataGenerator\nfrom keras.preprocessing import image\nfrom keras import layers, models\nfrom keras.optimizers import Adam\nfrom keras.layers import Conv2D, MaxPooling2D, BatchNormalization, GlobalAveragePooling2D\nfrom keras.layers import Dense, Dropout, Flatten\nfrom keras.callbacks import ModelCheckpoint, EarlyStopping, ReduceLROnPlateau\nimport tensorflow.keras.backend as K\nfrom tensorflow import keras\nfrom kaggle_datasets import KaggleDatasets\n\nprint(\"Tensorflow version \" + tf.__version__)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Notebooks that I found very useful"},{"metadata":{},"cell_type":"markdown","source":"Notebooks that helped me understand training with TPU's\n\n- [dimitreoliveira's notebook](https://www.kaggle.com/dimitreoliveira/flower-classification-with-tpus-eda-and-baseline)\n\n- [Xhlulu's flower competition notebook](https://www.kaggle.com/xhlulu/flowers-tpu-concise-efficientnet-b7)\n\n- [Getting Started with TPU's](https://www.kaggle.com/jessemostipak/getting-started-tpus-cassava-leaf-disease) -- [TPU Docs](https://www.kaggle.com/docs/tpu) -- [TFRecords Basics](https://www.kaggle.com/ryanholbrook/tfrecords-basics)\n\nDATA Aug\n- [TF Data Augmentation Docs](tensorflow.org/tutorials/images/data_augmentation)\n- [TF Image Docs](https://www.tensorflow.org/api_docs/python/tf/image)"},{"metadata":{},"cell_type":"markdown","source":"# Detecting TPU\n\nThe following cell isnt necessary but is nice to double check that we are going to be using the TPU. If we have everything set up correctly the number of replicas should be 8. If we do not have the TPU turned on we will see a value of 1.\n\nADD Google Cloud Software Development Kit (SDK) to Notebook if using a private dataset (add-ons tab). Pretty much always using public data though."},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"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)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Variables\n\n- note: actual file path for data is 'cassava-leaf-disease-tfrecords-center-512x512'"},{"metadata":{"trusted":true},"cell_type":"code","source":"SEED = 64\nAUTOTUNE = tf.data.experimental.AUTOTUNE\n#GCS_PATH = KaggleDatasets().get_gcs_path('cldtfrecords512x512')\nGCS_PATH = KaggleDatasets().get_gcs_path('cassava-leaf-disease-classification')\nBATCH_SIZE = 16 * strategy.num_replicas_in_sync\nIMAGE_SIZE = [512, 512]\nTARGET_SIZE = 512\nCLASSES = ['0', '1', '2', '3', '4']\nEPOCHS = 10\nDROPOUT_RATE = 0.2 #would like to implement this \nXception_NOTOP = '../input/xception/xception_weights_tf_dim_ordering_tf_kernels_notop.h5'\n","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Splitting TFRecords\n\n\nThe number on the end of the tfrecord file corresponds to the number of images in that tfrecord."},{"metadata":{"trusted":true},"cell_type":"code","source":"#this function counts number of images in all TFRecords\ndef count_data_items(filenames):\n    n = [int(re.compile(r'-([0-9]*)\\.').search(filename).group(1)) for filename in filenames]\n    return np.sum(n)\n\n#note: this filepath had to be manually selected\ndatabase_base_path = '../input/cassava-leaf-disease-classification/'\n\n#reading train metadata\ntrain = pd.read_csv(f'{database_base_path}train.csv')\nprint(f'Train samples: {len(train)}')\n\nALL_TRAINING_FILENAMES = tf.io.gfile.glob(GCS_PATH + '/train_tfrecords/*.tfrec')\nNUM_ALL_TRAINING_IMAGES = count_data_items(ALL_TRAINING_FILENAMES)\n\nprint(f'GCS: train images: {NUM_ALL_TRAINING_IMAGES}')\ndisplay(train.head())","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Train vs Test split"},{"metadata":{"trusted":true},"cell_type":"code","source":"TRAINING_FILENAMES, VALIDATION_FILENAMES = train_test_split(\n    ALL_TRAINING_FILENAMES,\n    train_size= 0.80, test_size=0.20,\n    random_state=SEED,\n)\n\nNUM_VALIDATION_IMAGES = count_data_items(VALIDATION_FILENAMES)\nNUM_TRAINING_IMAGES = count_data_items(TRAINING_FILENAMES)\n\nprint(\"Training Images: {}  Validation Image: {}\".format(NUM_TRAINING_IMAGES, NUM_VALIDATION_IMAGES))\nprint(\"Training Percent: {:.2f}  Validation Percent: {:.2f}\".format((NUM_TRAINING_IMAGES/NUM_ALL_TRAINING_IMAGES),\n                                                           (NUM_VALIDATION_IMAGES/NUM_ALL_TRAINING_IMAGES)))\n\nSTEPS_PER_EPOCH =  NUM_TRAINING_IMAGES // BATCH_SIZE\nVALID_STEPS = NUM_VALIDATION_IMAGES // BATCH_SIZE","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Functions"},{"metadata":{"trusted":true},"cell_type":"code","source":"def decode_image(image_data):\n    image = tf.image.decode_jpeg(image_data, channels=3) #decoding jpeg-encoded img to uint8 tensor\n    image = tf.cast(image, tf.float32) / 255.0 #cast int val to float so we can normalize it\n    image = tf.image.resize(image, [*IMAGE_SIZE]) #added this back seeing if it does anything\n    image = tf.reshape(image, [*IMAGE_SIZE, 3]) #resizing to proper shape\n    return image","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def read_tfrecord(example, labeled=True):\n    \"\"\"\n        1. Parse data based on the 'TFREC_FORMAT' map.\n        2. Decode image.\n        3. If 'labeled' returns (image, label) if not (image, name).\n    \"\"\"\n    if labeled:\n        TFREC_FORMAT = {\n            'image': tf.io.FixedLenFeature([], tf.string), \n            'target': tf.io.FixedLenFeature([], tf.int64), \n        }\n    else:\n        TFREC_FORMAT = {\n            'image': tf.io.FixedLenFeature([], tf.string), \n            'image_name': tf.io.FixedLenFeature([], tf.string), \n        }\n    example = tf.io.parse_single_example(example, TFREC_FORMAT)\n    image = decode_image(example['image'])\n    if labeled:\n        label_or_name = tf.cast(example['target'], tf.int32)\n    else:\n        label_or_name =  example['image_name']\n    return image, label_or_name","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def load_dataset(filenames, labeled=True, ordered=False):\n    \"\"\"\n        Create a Tensorflow dataset from TFRecords.\n    \"\"\"\n    ignore_order = tf.data.Options()\n    if not ordered:\n        ignore_order.experimental_deterministic = False\n\n    dataset = tf.data.TFRecordDataset(filenames, num_parallel_reads=AUTOTUNE)\n    dataset = dataset.with_options(ignore_order)\n    dataset = dataset.map(lambda x: read_tfrecord(x, labeled=labeled), num_parallel_calls=AUTOTUNE)\n    return dataset","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Data augmentation\nTrying augmentation technique here called Coarse Dropout which randomly cuts out parts of the image. This is another technique to prevent overfitting."},{"metadata":{"trusted":true},"cell_type":"code","source":"def dropout(image, DIM=512, PROBABILITY = 0.75, CT = 5, SZ = 0.1):\n    \n    # input image - is one image of size [dim,dim,3] not a batch of [b,dim,dim,3]\n    # output - image with CT squares of side size SZ*DIM removed\n    \n    # DO DROPOUT WITH PROBABILITY DEFINED ABOVE\n    P = tf.cast( tf.random.uniform([],0,1)<PROBABILITY, tf.int32)\n    if (P==0)|(CT==0)|(SZ==0): return image\n    \n    for k in range(CT):\n        # CHOOSE RANDOM LOCATION\n        x = tf.cast( tf.random.uniform([],0,DIM),tf.int32)\n        y = tf.cast( tf.random.uniform([],0,DIM),tf.int32)\n        # COMPUTE SQUARE \n        WIDTH = tf.cast( SZ*DIM,tf.int32) * P\n        ya = tf.math.maximum(0,y-WIDTH//2)\n        yb = tf.math.minimum(DIM,y+WIDTH//2)\n        xa = tf.math.maximum(0,x-WIDTH//2)\n        xb = tf.math.minimum(DIM,x+WIDTH//2)\n        # DROPOUT IMAGE\n        one = image[ya:yb,0:xa,:]\n        two = tf.zeros([yb-ya,xb-xa,3]) \n        three = image[ya:yb,xb:DIM,:]\n        middle = tf.concat([one,two,three],axis=1)\n        image = tf.concat([image[0:ya,:,:],middle,image[yb:DIM,:,:]],axis=0)\n            \n    # RESHAPE HACK SO TPU COMPILER KNOWS SHAPE OF OUTPUT TENSOR \n    image = tf.reshape(image,[DIM,DIM,3])\n    return image","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def get_mat(rotation, shear, height_zoom, width_zoom, height_shift, width_shift):\n    # returns 3x3 transformmatrix which transforms indicies\n        \n    # CONVERT DEGREES TO RADIANS\n    rotation = math.pi * rotation / 180.\n    shear = math.pi * shear / 180.\n    \n    # ROTATION MATRIX\n    c1 = tf.math.cos(rotation)\n    s1 = tf.math.sin(rotation)\n    one = tf.constant([1],dtype='float32')\n    zero = tf.constant([0],dtype='float32')\n    rotation_matrix = tf.reshape( tf.concat([c1,s1,zero, -s1,c1,zero, zero,zero,one],axis=0),[3,3] )\n        \n    # SHEAR MATRIX\n    c2 = tf.math.cos(shear)\n    s2 = tf.math.sin(shear)\n    shear_matrix = tf.reshape( tf.concat([one,s2,zero, zero,c2,zero, zero,zero,one],axis=0),[3,3] )    \n    \n    # ZOOM MATRIX\n    zoom_matrix = tf.reshape( tf.concat([one/height_zoom,zero,zero, zero,one/width_zoom,zero, zero,zero,one],axis=0),[3,3] )\n    \n    # SHIFT MATRIX\n    shift_matrix = tf.reshape( tf.concat([one,zero,height_shift, zero,one,width_shift, zero,zero,one],axis=0),[3,3] )\n    \n    return K.dot(K.dot(rotation_matrix, shear_matrix), K.dot(zoom_matrix, shift_matrix))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def transform(image,label):\n    # input image - is one image of size [dim,dim,3] not a batch of [b,dim,dim,3]\n    # output - image randomly rotated, sheared, zoomed, and shifted\n    DIM = IMAGE_SIZE[0]\n    XDIM = DIM%2 #fix for size 331\n    \n    rot = 45. * tf.random.normal([1],dtype='float32')\n    shr = 0. * tf.random.normal([1],dtype='float32') #need to test with and without shear\n    h_zoom = 1.0 + tf.random.normal([1],dtype='float32')/10.\n    w_zoom = 1.0 + tf.random.normal([1],dtype='float32')/10.\n    h_shift = 20. * tf.random.normal([1],dtype='float32') \n    w_shift = 20. * tf.random.normal([1],dtype='float32') \n  \n    # GET TRANSFORMATION MATRIX\n    m = get_mat(rot,shr,h_zoom,w_zoom,h_shift,w_shift) \n\n    # LIST DESTINATION PIXEL INDICES\n    x = tf.repeat( tf.range(DIM//2,-DIM//2,-1), DIM )\n    y = tf.tile( tf.range(-DIM//2,DIM//2),[DIM] )\n    z = tf.ones([DIM*DIM],dtype='int32')\n    idx = tf.stack( [x,y,z] )\n    \n    # ROTATE DESTINATION PIXELS ONTO ORIGIN PIXELS\n    idx2 = K.dot(m,tf.cast(idx,dtype='float32'))\n    idx2 = K.cast(idx2,dtype='int32')\n    idx2 = K.clip(idx2,-DIM//2+XDIM+1,DIM//2)\n    \n    # FIND ORIGIN PIXEL VALUES           \n    idx3 = tf.stack( [DIM//2-idx2[0,], DIM//2-1+idx2[1,]] )\n    d = tf.gather_nd(image,tf.transpose(idx3))\n    \n    #PERFORMING DROPOUT ON RETURN STATEMENT\n        \n    return dropout(tf.reshape(d,[DIM,DIM,3])),label","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Loading Data\n\nMake sure to have the correct augmenter selected in the next function. \n\n- transform - Chris Deottes Augmenter\n- simple_data_augmenter - quicker but more simple augmenter"},{"metadata":{"trusted":true},"cell_type":"code","source":"def decode_jpeg_and_label(image,label):\n    DIM = IMAGE_SIZE[0]\n    image = tf.reshape(image,[DIM,DIM,3])\n    return dropout(image, label)\n\ndef get_training_dataset():\n    dataset = load_dataset(TRAINING_FILENAMES, labeled=True)\n    dataset = dataset.map(transform, num_parallel_calls=AUTOTUNE)\n    dataset = dataset.repeat() # the training dataset must repeat for several epochs\n    dataset = dataset.shuffle(2048) #set higher than input?\n    dataset = dataset.batch(BATCH_SIZE)\n    dataset = dataset.prefetch(AUTOTUNE) # prefetch next batch while training (autotune prefetch buffer size)\n    return dataset\n\ndef get_validation_dataset(ordered=False):\n    dataset = load_dataset(VALIDATION_FILENAMES, labeled=True, ordered=ordered)\n    dataset = dataset.map(transform, num_parallel_calls=AUTOTUNE)\n    dataset = dataset.batch(BATCH_SIZE)\n    dataset = dataset.cache()\n    dataset = dataset.prefetch(AUTOTUNE)\n    return dataset\n\n#not used\ndef get_test_dataset(ordered=False):\n    dataset = load_dataset(TEST_FILENAMES, labeled=False, ordered=ordered)\n    dataset = dataset.batch(BATCH_SIZE)\n    dataset = dataset.prefetch(AUTOTUNE)\n    return dataset","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_dataset = get_training_dataset()\nvalidation_dataset=get_validation_dataset()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Building Model\n\nAlso note that we're using sparse_categorical_crossentropy as our loss function, because we did not one-hot encode our labels."},{"metadata":{},"cell_type":"markdown","source":"### Defining model"},{"metadata":{"trusted":true},"cell_type":"code","source":"def create_model():\n    \n    model = models.Sequential()\n    model.add(keras.applications.Xception(\n        input_shape=(TARGET_SIZE, TARGET_SIZE, 3),\n        weights=Xception_NOTOP,\n        include_top=False)\n    )\n    model.add(layers.GlobalAveragePooling2D())\n    model.add(layers.Dropout(DROPOUT_RATE))\n    model.add(layers.Dense(5, activation = \"softmax\"))# 5 is the dimensionality of the output space \"5 options\"\n    \n    loss = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=False)\n\n    return model","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"In order to ensure that our model is trained on the TPU, we build it using strategy.scope()"},{"metadata":{"trusted":true},"cell_type":"code","source":"with strategy.scope():\n    model = create_model()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Model Compile\nFor this model we will be using Adamax optimizer and since our output labels are categorical we will be using SparseCategoricalAccuracy as a loss function.\n\nUse this crossentropy loss function when there are two or more label classes. We expect labels to be provided as integers.\n\ntf.keras.metrics.SparseCategoricalAccuracy\n"},{"metadata":{"trusted":true},"cell_type":"code","source":"model.compile(\n    optimizer = 'adamax',\n    loss = 'sparse_categorical_crossentropy',\n    metrics = ['accuracy','sparse_categorical_accuracy'])","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Save untrained model"},{"metadata":{"trusted":true},"cell_type":"code","source":"model.save('./Untrained_TPU_model.h5')\nmodel.summary()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Callback\nA callback is an object that can perform actions at various stages of training (e.g. at the start or end of an epoch, before or after a single batch, etc). In this model we will be using 3 callbacks as below:"},{"metadata":{},"cell_type":"markdown","source":"1. ModelCheckpoint : Callback to save the Keras model or model weights at some frequency."},{"metadata":{"trusted":true},"cell_type":"code","source":"model_checkpoint = keras.callbacks.ModelCheckpoint(\n    './best_weights.h5',\n    monitor=\"val_loss\",\n    verbose=1,\n    save_best_only=True,\n    save_weights_only=True,\n    mode=\"min\"\n)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"2.EarlyStopping : Stop training when a monitored metric has stopped improving."},{"metadata":{"trusted":true},"cell_type":"code","source":"early_stopping = keras.callbacks.EarlyStopping(\n    monitor=\"val_loss\",\n    min_delta=0.001,\n    patience=5,\n    verbose=0,\n    mode=\"min\",\n    restore_best_weights=True,\n)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"3.ReduceLROnPlateau : Reduce learning rate when a metric has stopped improving."},{"metadata":{"trusted":true},"cell_type":"code","source":"reduce_lr = keras.callbacks.ReduceLROnPlateau(\n    monitor=\"val_loss\",\n    factor=0.1,\n    patience=2,\n    verbose=1,\n    mode=\"min\",\n    min_delta=0.001,\n)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"4.Stream the epoch loss to a file in JSON format. The file content is not well-formed JSON but rather has a JSON object per line."},{"metadata":{"trusted":true},"cell_type":"code","source":"import json\njson_log = open('/kaggle/working/loss_log.json', mode='w', buffering=1)\njson_logging_callback = keras.callbacks.LambdaCallback(\n    on_train_begin=lambda logs: json_log.write('{ \"train\": [ \\n'),\n    on_epoch_end=lambda epoch,logs: json_log.write(json.dumps({'epoch': epoch, 'loss': logs['loss'], 'acc':logs['accuracy']}) + ',\\n'),\n    on_train_end=lambda logs: json_log.write(']\\n}'),\n)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Start Training\nWe have defined everything we need, it's time to train the model..."},{"metadata":{"trusted":true},"cell_type":"code","source":"history = model.fit(\n    train_dataset,\n    epochs=EPOCHS,\n    steps_per_epoch = STEPS_PER_EPOCH,\n    validation_steps=VALID_STEPS,\n    validation_data=validation_dataset,\n    callbacks = [\n        model_checkpoint, \n        early_stopping, \n        reduce_lr,\n        json_logging_callback],\n    verbose=1)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Submission"},{"metadata":{"trusted":true},"cell_type":"code","source":"submission = pd.DataFrame(columns=['image_id','label'])\nfor image_name in os.listdir(database_base_path + '/test_images'):\n    image_path = os.path.join(database_base_path + '/test_images', image_name)\n    image = tf.keras.preprocessing.image.load_img(image_path)\n    resized_image = image.resize((TARGET_SIZE, TARGET_SIZE))\n    numpied_image = np.expand_dims(resized_image, 0)\n    tensored_image = tf.cast(numpied_image, tf.float32)\n    submission = submission.append(pd.DataFrame({'image_id': image_name,\n                                                 'label': model.predict_classes(tensored_image)}))\n\nsubmission.to_csv('/kaggle/working/submission.csv', index=False)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Performance\nSince the performance metric for this competition is accuracy, let us plot the train and validation accuracy to monitor our model performance."},{"metadata":{"trusted":true},"cell_type":"code","source":"import json\nfile = open('/kaggle/working/loss_log.json', 'r')\ncountriesStr = file.read()\ncountriesStr = countriesStr[::-1].replace(',', '', 1)[::-1]\nfile.close()\n\nwith open('/kaggle/working/loss_log.json', 'w') as file:\n    file.write(countriesStr)\n    file.close()\n\njsonFile = open('/kaggle/working//loss_log.json', 'r')\njson_array = json.load(jsonFile)\n\nfor item in json_array['train']:\n    last_epoch = item['epoch']\n    \nlast_epoch+=1\n\nacc = history.history['accuracy']\nval_acc = history.history['val_accuracy']\nloss = history.history['loss']\nval_loss = history.history['val_loss']\nepochs_range = range(last_epoch)\n\nplt.figure(figsize=(8, 8))\nplt.subplot(1, 2, 1)\nplt.plot(epochs_range, acc, label='Training Accuracy')\nplt.plot(epochs_range, val_acc, label='Validation Accuracy')\nplt.legend(loc='lower right')\nplt.title('Training and Validation Accuracy')\n\nplt.subplot(1, 2, 2)\nplt.plot(epochs_range, loss, label='Training Loss')\nplt.plot(epochs_range, val_loss, label='Validation Loss')\nplt.legend(loc='upper right')\nplt.title('Training and Validation Loss')\nplt.show()\n#json_log.close()","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}