{"cells":[{"metadata":{"papermill":{"duration":0.050907,"end_time":"2020-12-04T20:29:39.441792","exception":false,"start_time":"2020-12-04T20:29:39.390885","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# Machine Learning Final Project Notebook with EfficientNet B7\n## *Petals to the Metal: Flower Classification on TPU*\n### *By Xuanzhi Huang, Rahul Paul*"},{"metadata":{"papermill":{"duration":0.046053,"end_time":"2020-12-04T20:29:39.534500","exception":false,"start_time":"2020-12-04T20:29:39.488447","status":"completed"},"tags":[]},"cell_type":"markdown","source":"## Step 1: Some pre-setting"},{"metadata":{"papermill":{"duration":0.046725,"end_time":"2020-12-04T20:29:39.627606","exception":false,"start_time":"2020-12-04T20:29:39.580881","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Package preliminary"},{"metadata":{"papermill":{"duration":0.04647,"end_time":"2020-12-04T20:29:39.720925","exception":false,"start_time":"2020-12-04T20:29:39.674455","status":"completed"},"tags":[]},"cell_type":"markdown","source":"#### Install package \"efficientnet\" so that we can build a Efficient Net model."},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:29:39.819815Z","iopub.status.busy":"2020-12-04T20:29:39.818948Z","iopub.status.idle":"2020-12-04T20:29:50.832709Z","shell.execute_reply":"2020-12-04T20:29:50.831378Z"},"papermill":{"duration":11.065524,"end_time":"2020-12-04T20:29:50.832880","exception":false,"start_time":"2020-12-04T20:29:39.767356","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"!pip install efficientnet","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.056449,"end_time":"2020-12-04T20:29:50.949062","exception":false,"start_time":"2020-12-04T20:29:50.892613","status":"completed"},"tags":[]},"cell_type":"markdown","source":"#### Then import all the packages we need."},{"metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.execute_input":"2020-12-04T20:29:51.076170Z","iopub.status.busy":"2020-12-04T20:29:51.073680Z","iopub.status.idle":"2020-12-04T20:29:58.088966Z","shell.execute_reply":"2020-12-04T20:29:58.088324Z"},"papermill":{"duration":7.087275,"end_time":"2020-12-04T20:29:58.089113","exception":false,"start_time":"2020-12-04T20:29:51.001838","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"# These are some basic packages\nimport random, re, math, os\nimport numpy as np\nimport pandas as pd\nimport tensorflow as tf\nimport matplotlib.pyplot as plt\n\n\n# These are for data processing\nimport tensorflow_addons as tfa\nfrom kaggle_datasets import KaggleDatasets\n\n\n# These are for model training\nfrom tensorflow.keras.mixed_precision import experimental as mixed_precision\nimport tensorflow.keras.backend as K\nimport efficientnet.tfkeras as efn\n\n\n# These are performance metrics\nfrom sklearn.metrics import f1_score, precision_score, recall_score, confusion_matrix\n\n\n# These are for class weights\nimport datetime\nimport tqdm\nimport json\nfrom collections import Counter\nimport gc","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.052865,"end_time":"2020-12-04T20:29:58.195556","exception":false,"start_time":"2020-12-04T20:29:58.142691","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Detect the hardware and tell the appropriate distribution strategy"},{"metadata":{"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","execution":{"iopub.execute_input":"2020-12-04T20:29:58.315600Z","iopub.status.busy":"2020-12-04T20:29:58.308875Z","iopub.status.idle":"2020-12-04T20:30:02.643881Z","shell.execute_reply":"2020-12-04T20:30:02.643242Z"},"papermill":{"duration":4.395969,"end_time":"2020-12-04T20:30:02.644001","exception":false,"start_time":"2020-12-04T20:29:58.248032","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"try:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver()\n    print('Running on TPU ', tpu.master())\nexcept ValueError:\n    tpu = None\n\nif tpu:\n    tf.config.experimental_connect_to_cluster(tpu)\n    tf.tpu.experimental.initialize_tpu_system(tpu)\n    strategy = tf.distribute.experimental.TPUStrategy(tpu)\nelse:\n    strategy = tf.distribute.get_strategy()\n\nprint(\"REPLICAS: \", strategy.num_replicas_in_sync)\n\n# Make the system tune the number of threads for us\nAUTO = tf.data.experimental.AUTOTUNE","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.052818,"end_time":"2020-12-04T20:30:02.750693","exception":false,"start_time":"2020-12-04T20:30:02.697875","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Configuration for image size, training epoch, batch size, and random seed"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:30:02.863868Z","iopub.status.busy":"2020-12-04T20:30:02.862839Z","iopub.status.idle":"2020-12-04T20:30:02.866306Z","shell.execute_reply":"2020-12-04T20:30:02.865577Z"},"papermill":{"duration":0.062478,"end_time":"2020-12-04T20:30:02.866428","exception":false,"start_time":"2020-12-04T20:30:02.803950","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"IMAGE_SIZE = [512, 512]\nEPOCHS = 16\nBATCH_SIZE = 16 * strategy.num_replicas_in_sync\nSEED = 100","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.053272,"end_time":"2020-12-04T20:30:02.973365","exception":false,"start_time":"2020-12-04T20:30:02.920093","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Set the data access"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:30:03.091975Z","iopub.status.busy":"2020-12-04T20:30:03.091061Z","iopub.status.idle":"2020-12-04T20:30:04.037382Z","shell.execute_reply":"2020-12-04T20:30:04.036694Z"},"papermill":{"duration":1.010299,"end_time":"2020-12-04T20:30:04.037522","exception":false,"start_time":"2020-12-04T20:30:03.027223","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"GCS_DS_PATH = KaggleDatasets().get_gcs_path('tpu-getting-started')\n# These are available image sizes in the data set\nGCS_PATH_SELECT = { \n    192: GCS_DS_PATH + '/tfrecords-jpeg-192x192',\n    224: GCS_DS_PATH + '/tfrecords-jpeg-224x224',\n    331: GCS_DS_PATH + '/tfrecords-jpeg-331x331',\n    512: GCS_DS_PATH + '/tfrecords-jpeg-512x512'\n}\nGCS_PATH = GCS_PATH_SELECT[IMAGE_SIZE[0]]\nTRAINING_FILENAMES = tf.io.gfile.glob(GCS_PATH + '/train/*.tfrec')\nVALIDATION_FILENAMES = tf.io.gfile.glob(GCS_PATH + '/val/*.tfrec')\nTEST_FILENAMES = tf.io.gfile.glob(GCS_PATH + '/test/*.tfrec')","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.052972,"end_time":"2020-12-04T20:30:04.144739","exception":false,"start_time":"2020-12-04T20:30:04.091767","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Add more mixed precision and/or XLA (refer to Chris Deotte's notebook)\n[Rotation Augmentation GPU/TPU - [0.96+]](https://www.kaggle.com/cdeotte/rotation-augmentation-gpu-tpu-0-96)"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:30:04.262165Z","iopub.status.busy":"2020-12-04T20:30:04.261184Z","iopub.status.idle":"2020-12-04T20:30:04.264801Z","shell.execute_reply":"2020-12-04T20:30:04.264043Z"},"papermill":{"duration":0.06617,"end_time":"2020-12-04T20:30:04.264925","exception":false,"start_time":"2020-12-04T20:30:04.198755","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"# Add more mixed precision and/or XLA to allow the TPU memory to handle larger batch sizes \n# and can speed up the training process\nMIXED_PRECISION = False\nXLA_ACCELERATE = False\n\nif MIXED_PRECISION:\n    if tpu: policy = tf.keras.mixed_precision.experimental.Policy('mixed_bfloat16')\n    else: policy = tf.keras.mixed_precision.experimental.Policy('mixed_float16')\n    mixed_precision.set_policy(policy)\n    print('Mixed precision enabled')\n\nif XLA_ACCELERATE:\n    tf.config.optimizer.set_jit(True)\n    print('Accelerated Linear Algebra enabled')","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.054462,"end_time":"2020-12-04T20:30:04.373166","exception":false,"start_time":"2020-12-04T20:30:04.318704","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Show all the classes we have"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:30:04.494338Z","iopub.status.busy":"2020-12-04T20:30:04.493295Z","iopub.status.idle":"2020-12-04T20:30:04.496811Z","shell.execute_reply":"2020-12-04T20:30:04.496191Z"},"papermill":{"duration":0.070112,"end_time":"2020-12-04T20:30:04.496935","exception":false,"start_time":"2020-12-04T20:30:04.426823","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"CLASSES = ['pink primrose',    'hard-leaved pocket orchid', 'canterbury bells', 'sweet pea',     'wild geranium',     'tiger lily',           'moon orchid',              'bird of paradise', 'monkshood',        'globe thistle',         # 00 - 09\n           'snapdragon',       \"colt's foot\",               'king protea',      'spear thistle', 'yellow iris',       'globe-flower',         'purple coneflower',        'peruvian lily',    'balloon flower',   'giant white arum lily', # 10 - 19\n           'fire lily',        'pincushion flower',         'fritillary',       'red ginger',    'grape hyacinth',    'corn poppy',           'prince of wales feathers', 'stemless gentian', 'artichoke',        'sweet william',         # 20 - 29\n           'carnation',        'garden phlox',              'love in the mist', 'cosmos',        'alpine sea holly',  'ruby-lipped cattleya', 'cape flower',              'great masterwort', 'siam tulip',       'lenten rose',           # 30 - 39\n           'barberton daisy',  'daffodil',                  'sword lily',       'poinsettia',    'bolero deep blue',  'wallflower',           'marigold',                 'buttercup',        'daisy',            'common dandelion',      # 40 - 49\n           'petunia',          'wild pansy',                'primula',          'sunflower',     'lilac hibiscus',    'bishop of llandaff',   'gaura',                    'geranium',         'orange dahlia',    'pink-yellow dahlia',    # 50 - 59\n           'cautleya spicata', 'japanese anemone',          'black-eyed susan', 'silverbush',    'californian poppy', 'osteospermum',         'spring crocus',            'iris',             'windflower',       'tree poppy',            # 60 - 69\n           'gazania',          'azalea',                    'water lily',       'rose',          'thorn apple',       'morning glory',        'passion flower',           'lotus',            'toad lily',        'anthurium',             # 70 - 79\n           'frangipani',       'clematis',                  'hibiscus',         'columbine',     'desert-rose',       'tree mallow',          'magnolia',                 'cyclamen ',        'watercress',       'canna lily',            # 80 - 89\n           'hippeastrum ',     'bee balm',                  'pink quill',       'foxglove',      'bougainvillea',     'camellia',             'mallow',                   'mexican petunia',  'bromelia',         'blanket flower',        # 90 - 99\n           'trumpet creeper',  'blackberry lily',           'common tulip',     'wild rose'] ","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.053515,"end_time":"2020-12-04T20:30:04.604558","exception":false,"start_time":"2020-12-04T20:30:04.551043","status":"completed"},"tags":[]},"cell_type":"markdown","source":"## Step 2: Set some visualization functions"},{"metadata":{"papermill":{"duration":0.054064,"end_time":"2020-12-04T20:30:04.712659","exception":false,"start_time":"2020-12-04T20:30:04.658595","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Set training and validation curve functions to show the changes in loss and accuracy"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:30:04.830626Z","iopub.status.busy":"2020-12-04T20:30:04.829663Z","iopub.status.idle":"2020-12-04T20:30:04.832905Z","shell.execute_reply":"2020-12-04T20:30:04.832190Z"},"papermill":{"duration":0.066304,"end_time":"2020-12-04T20:30:04.833055","exception":false,"start_time":"2020-12-04T20:30:04.766751","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"def plot_train_valid_curves(training, validation, title, subplot):\n    \n    if subplot % 10 == 1:\n        plt.subplots(figsize = (15,15), facecolor = '#F0F0F0')\n        plt.tight_layout()\n    ax = plt.subplot(subplot)\n    ax.set_facecolor('#F8F8F8')\n    ax.plot(training)\n    ax.plot(validation)\n    ax.set_title('model '+ title)\n    ax.set_ylabel(title)\n    ax.set_xlabel('epoch')\n    ax.legend(['training', 'validation'])","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.063288,"end_time":"2020-12-04T20:30:04.950817","exception":false,"start_time":"2020-12-04T20:30:04.887529","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Set a function to plot confusion matrix"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:30:05.102370Z","iopub.status.busy":"2020-12-04T20:30:05.101540Z","iopub.status.idle":"2020-12-04T20:30:05.104992Z","shell.execute_reply":"2020-12-04T20:30:05.104425Z"},"papermill":{"duration":0.088005,"end_time":"2020-12-04T20:30:05.105134","exception":false,"start_time":"2020-12-04T20:30:05.017129","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"def display_confusion_matrix(cmat, score, precision, recall):\n    \n    plt.figure(figsize = (20,20))  # Specify the size of confusion matrix\n    ax = plt.gca()\n    ax.matshow(cmat, cmap = 'Reds')  # Draw a matrix\n    ax.set_xticks(range(len(CLASSES)))  # Set the range of X coordinate according to #classes\n    ax.set_xticklabels(CLASSES, fontdict={'fontsize': 7})  # Set the font size of X coordinate\n    # Rotate labels on X coordinate to make them look better\n    plt.setp(ax.get_xticklabels(), rotation = 45, ha = \"left\", rotation_mode = \"anchor\")\n    ax.set_yticks(range(len(CLASSES)))  # Set the range of Y coordinate according to #classes\n    ax.set_yticklabels(CLASSES, fontdict={'fontsize': 7})  # Set the font size of Y coordinate\n    # Rotate labels on Y coordinate to make them look better\n    plt.setp(ax.get_yticklabels(), rotation = 45, ha = \"right\", rotation_mode = \"anchor\")\n    # Round F1 score, precision, and recall to the nearest fourth decimal place\n    titlestring = \"\"\n    if score is not None:\n        titlestring += 'f1 = {:.4f} '.format(score)\n    if precision is not None:\n        titlestring += '\\nprecision = {:.4f} '.format(precision)\n    if recall is not None:\n        titlestring += '\\nrecall = {:.4f} '.format(recall)\n    # Add some comments about F1 score, precision, and recall on the plot\n    if len(titlestring) > 0:\n        ax.text(101, 1, titlestring, fontdict = {'fontsize': 18, 'horizontalalignment': 'right', 'verticalalignment': 'top', 'color': 'Blue'})\n    plt.show()","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.054403,"end_time":"2020-12-04T20:30:05.214096","exception":false,"start_time":"2020-12-04T20:30:05.159693","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Show the beautiful flowers (refer to Dimitre Oliveira)\n[Flower with TPUs - Advanced augmentations](https://www.kaggle.com/dimitreoliveira/flower-with-tpus-advanced-augmentations)"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:30:05.368024Z","iopub.status.busy":"2020-12-04T20:30:05.354459Z","iopub.status.idle":"2020-12-04T20:30:05.371438Z","shell.execute_reply":"2020-12-04T20:30:05.370723Z"},"papermill":{"duration":0.102619,"end_time":"2020-12-04T20:30:05.371577","exception":false,"start_time":"2020-12-04T20:30:05.268958","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"def batch_to_numpy_images_and_labels(data):\n    images, labels = data\n    numpy_images = images.numpy()\n    numpy_labels = labels.numpy()\n    if numpy_labels.dtype == object:\n        numpy_labels = [None for _ in enumerate(numpy_images)]\n    # If no labels, only image IDs, return None for labels (for test data)\n    return numpy_images, numpy_labels\n\n\ndef title_from_label_and_target_(label, correct_label):\n    if correct_label is None:\n        return CLASSES[label], True\n    correct = (label == correct_label)\n    return \"{} [{}{}{}]\".format(CLASSES[label], 'OK' if correct else 'NO', u\"\\u2192\" if not correct else '',\n                                CLASSES[correct_label] if not correct else ''), correct\n\n\ndef display_one_flower(image, title, subplot, red = False, titlesize = 16):\n    plt.subplot(*subplot)\n    plt.axis('off')\n    plt.imshow(image)\n    if len(title) > 0:\n        plt.title(title, fontsize = int(titlesize) if not red else int(titlesize / 1.2), color = 'red' if red else 'black', fontdict={'verticalalignment':'center'}, pad=int(titlesize/1.5))\n    return (subplot[0], subplot[1], subplot[2] + 1)\n\n\ndef display_batch_of_images(databatch, predictions = None):\n    images, labels = batch_to_numpy_images_and_labels(databatch)\n    if labels is None:\n        labels = [None for _ in enumerate(images)]\n    rows = int(math.sqrt(len(images)))\n    cols = len(images) // rows\n    FIGSIZE = 13.0\n    SPACING = 0.1\n    subplot = (rows,cols,1)\n    if rows < cols:\n        plt.figure(figsize = (FIGSIZE, FIGSIZE / cols*rows))\n    else:\n        plt.figure(figsize = (FIGSIZE / rows * cols,FIGSIZE))\n    \n    # Display\n    for i, (image, label) in enumerate(zip(images[:rows*cols], labels[:rows*cols])):\n        title = '' if label is None else CLASSES[label]\n        correct = True\n        if predictions is not None:\n            title, correct = title_from_label_and_target_(predictions[i], label)\n        dynamic_titlesize = FIGSIZE * SPACING / max(rows,cols) * 40 + 3\n        subplot = display_one_flower(image, title, subplot, not correct, titlesize = dynamic_titlesize)\n    \n    # Layout\n    plt.tight_layout()\n    if label is None and predictions is None:\n        plt.subplots_adjust(wspace = 0, hspace = 0)\n    else:\n        plt.subplots_adjust(wspace = SPACING, hspace = SPACING)\n    plt.show()\n    \n\n# Visualize model predictions (on training and validation sets)\n# Images of flowers with labels telling whether prediction is true will be shown\ndef dataset_to_numpy_util(dataset, N):\n    dataset = dataset.unbatch().batch(N)\n    for images, labels in dataset:\n        numpy_images = images.numpy()\n        numpy_labels = labels.numpy()\n        break;  \n    return numpy_images, numpy_labels\n\ndef title_from_label_and_target(label, correct_label):\n    label = np.argmax(label, axis = -1)\n    correct = (label == correct_label)\n    return \"{} [{}{}{}]\".format(CLASSES[label], str(correct), ', should be ' if not correct else '',\n                                CLASSES[correct_label] if not correct else ''), correct\n\ndef display_one_flower_eval(image, title, subplot, red = False):\n    plt.subplot(subplot)\n    plt.axis('off')\n    plt.imshow(image)\n    plt.title(title, fontsize = 14, color = 'red' if red else 'black')\n    return subplot + 1\n\ndef display_9_images_with_predictions(images, predictions, labels):\n    subplot = 331\n    plt.figure(figsize = (13,13))\n    for i, image in enumerate(images):\n        title, correct = title_from_label_and_target(predictions[i], labels[i])\n        subplot = display_one_flower_eval(image, title, subplot, not correct)\n        if i >= 8:\n            break;\n    plt.tight_layout()\n    plt.subplots_adjust(wspace = 0.1, hspace = 0.1)\n    plt.show()","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.054284,"end_time":"2020-12-04T20:30:05.480670","exception":false,"start_time":"2020-12-04T20:30:05.426386","status":"completed"},"tags":[]},"cell_type":"markdown","source":"## Step 3: Set functions to gain training set, validation set, and test set"},{"metadata":{"papermill":{"duration":0.053853,"end_time":"2020-12-04T20:30:05.588933","exception":false,"start_time":"2020-12-04T20:30:05.535080","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Decode images and convert pixels to floats between 0 and 1"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:30:05.705094Z","iopub.status.busy":"2020-12-04T20:30:05.704327Z","iopub.status.idle":"2020-12-04T20:30:05.707538Z","shell.execute_reply":"2020-12-04T20:30:05.706910Z"},"papermill":{"duration":0.064253,"end_time":"2020-12-04T20:30:05.707680","exception":false,"start_time":"2020-12-04T20:30:05.643427","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"def decode_image(image_data):\n    \n    image = tf.image.decode_jpeg(image_data, channels = 3)\n    image = tf.cast(image, tf.float32) / 255.0\n    # Reshape the images to fit the size required by TPU\n    image = tf.reshape(image, [*IMAGE_SIZE, 3])\n    \n    return image","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.054176,"end_time":"2020-12-04T20:30:05.816598","exception":false,"start_time":"2020-12-04T20:30:05.762422","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Set a function to read labeled tfrec files (i.e. training & validation set)"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:30:05.939877Z","iopub.status.busy":"2020-12-04T20:30:05.938762Z","iopub.status.idle":"2020-12-04T20:30:05.942389Z","shell.execute_reply":"2020-12-04T20:30:05.941631Z"},"papermill":{"duration":0.071164,"end_time":"2020-12-04T20:30:05.942515","exception":false,"start_time":"2020-12-04T20:30:05.871351","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"def read_labeled_tfrecord(example):\n    \n    LABELED_TFREC_FORMAT = {\n        \"image\": tf.io.FixedLenFeature([], tf.string),\n        \"class\": tf.io.FixedLenFeature([], tf.int64),\n    }\n    example = tf.io.parse_single_example(example, LABELED_TFREC_FORMAT)\n    image = decode_image(example['image'])\n    label = tf.cast(example['class'], tf.int32)\n    \n    return image, label\n\n\n# This is for data visualization\ndef read_labeled_id_tfrecord(example):\n    \n    LABELED_ID_TFREC_FORMAT = {\n        \"image\": tf.io.FixedLenFeature([], tf.string),\n        \"class\": tf.io.FixedLenFeature([], tf.int64),\n        \"id\": tf.io.FixedLenFeature([], tf.string),\n    }\n    example = tf.io.parse_single_example(example, LABELED_ID_TFREC_FORMAT)\n    image = decode_image(example['image'])\n    label = tf.cast(example['class'], tf.int32)\n    idnum =  example['id']\n    \n    return image, label, idnum","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.05488,"end_time":"2020-12-04T20:30:06.052133","exception":false,"start_time":"2020-12-04T20:30:05.997253","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Set a function to read unlabeled tfrec files (i.e. test set)"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:30:06.170628Z","iopub.status.busy":"2020-12-04T20:30:06.169658Z","iopub.status.idle":"2020-12-04T20:30:06.172762Z","shell.execute_reply":"2020-12-04T20:30:06.172125Z"},"papermill":{"duration":0.065374,"end_time":"2020-12-04T20:30:06.172882","exception":false,"start_time":"2020-12-04T20:30:06.107508","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"def read_unlabeled_tfrecord(example):\n    \n    UNLABELED_TFREC_FORMAT = {\n        \"image\": tf.io.FixedLenFeature([], tf.string),\n        \"id\": tf.io.FixedLenFeature([], tf.string),\n    }\n    example = tf.io.parse_single_example(example, UNLABELED_TFREC_FORMAT)\n    image = decode_image(example['image'])\n    idnum = example['id']\n    \n    return image, idnum","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.05439,"end_time":"2020-12-04T20:30:06.283408","exception":false,"start_time":"2020-12-04T20:30:06.229018","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Load image data"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:30:06.401545Z","iopub.status.busy":"2020-12-04T20:30:06.400513Z","iopub.status.idle":"2020-12-04T20:30:06.403886Z","shell.execute_reply":"2020-12-04T20:30:06.403144Z"},"papermill":{"duration":0.065751,"end_time":"2020-12-04T20:30:06.404012","exception":false,"start_time":"2020-12-04T20:30:06.338261","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"# For best performance, read from multiple tfrec files at once\n# Disregard data's order, since data will be shuffled\ndef load_dataset(filenames, labeled = True, ordered = False):\n    \n    ignore_order = tf.data.Options()\n    if not ordered:\n        ignore_order.experimental_deterministic = False  # Disable order to increase running speed\n    # Automatically interleaves reading\n    dataset = tf.data.TFRecordDataset(filenames, num_parallel_reads = AUTO)\n    # Use data in the shuffled order\n    dataset = dataset.with_options(ignore_order)\n    # Returns a dataset of (image, label) pairs if labeled = True (i.e. training & validation set)\n    # or (image, id) pair if labeld = False (i.e. test set)\n    dataset = dataset.map(read_labeled_id_tfrecord if labeled else read_unlabeled_tfrecord, num_parallel_calls=AUTO)\n    \n    return dataset","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.054317,"end_time":"2020-12-04T20:30:06.513633","exception":false,"start_time":"2020-12-04T20:30:06.459316","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Data augmentation"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:30:06.634737Z","iopub.status.busy":"2020-12-04T20:30:06.631721Z","iopub.status.idle":"2020-12-04T20:30:06.637907Z","shell.execute_reply":"2020-12-04T20:30:06.637339Z"},"papermill":{"duration":0.069603,"end_time":"2020-12-04T20:30:06.638047","exception":false,"start_time":"2020-12-04T20:30:06.568444","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"# Randomly make some changes to the images and return the new images and labels\ndef data_augment(image, label):\n        \n    # Set seed for data augmentation\n    seed = 100\n    \n    # Randomly resize and then crop images\n    image = tf.image.resize(image, [720, 720])\n    image = tf.image.random_crop(image, [512, 512, 3], seed = seed)\n\n    # Randomly reset brightness of images\n    image = tf.image.random_brightness(image, 0.6, seed = seed)\n    \n    # Randomly reset saturation of images\n    image = tf.image.random_saturation(image, 3, 5, seed = seed)\n        \n    # Randomly reset contrast of images\n    image = tf.image.random_contrast(image, 0.3, 0.5, seed = seed)\n\n    # Randomly reset hue of images, but this will make the colors really weird, which we think will not happen\n    # in common photography\n    # image = tf.image.random_hue(image, 0.5, seed = seed)\n    \n    # Blur images\n    image = tfa.image.mean_filter2d(image, filter_shape = 10)\n    \n    # Randomly flip images\n    image = tf.image.random_flip_left_right(image, seed = seed)\n    image = tf.image.random_flip_up_down(image, seed = seed)\n    \n    # Fail to rotate and transform images due to some bug in TensorFlow\n    # angle = random.randint(0, 180)\n    # image = tfa.image.rotate(image, tf.constant(np.pi * angle / 180))\n    # image = tfa.image.transform(image, [1.0, 1.0, -250, 0.0, 1.0, 0.0, 0.0, 0.0])\n    \n    return image, label","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.054534,"end_time":"2020-12-04T20:30:06.747699","exception":false,"start_time":"2020-12-04T20:30:06.693165","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Gain training set"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:30:06.868289Z","iopub.status.busy":"2020-12-04T20:30:06.867135Z","iopub.status.idle":"2020-12-04T20:30:06.870567Z","shell.execute_reply":"2020-12-04T20:30:06.869820Z"},"papermill":{"duration":0.068229,"end_time":"2020-12-04T20:30:06.870692","exception":false,"start_time":"2020-12-04T20:30:06.802463","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"def get_training_dataset():\n   \n    train = load_dataset(TRAINING_FILENAMES, labeled = True)\n    train = train.map(lambda image, label, idnum: [image, label])\n    train = train.repeat()\n    train = train.shuffle(2048)\n    train = train.batch(BATCH_SIZE)\n    train = train.prefetch(AUTO)\n    \n    return train\n\n\n# This function is for data visualization\ndef get_training_dataset_preview(ordered = True):\n    \n    train = load_dataset(TRAINING_FILENAMES, labeled = True, ordered = ordered)\n    train = train.batch(BATCH_SIZE)\n    train = train.cache()\n    train = train.prefetch(AUTO)\n    \n    return train","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.055128,"end_time":"2020-12-04T20:30:07.028536","exception":false,"start_time":"2020-12-04T20:30:06.973408","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Gain validation set"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:30:07.148614Z","iopub.status.busy":"2020-12-04T20:30:07.147681Z","iopub.status.idle":"2020-12-04T20:30:07.150566Z","shell.execute_reply":"2020-12-04T20:30:07.149846Z"},"papermill":{"duration":0.065351,"end_time":"2020-12-04T20:30:07.150685","exception":false,"start_time":"2020-12-04T20:30:07.085334","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"def get_validation_dataset(ordered = False):\n\n    validation = load_dataset(VALIDATION_FILENAMES, labeled = True, ordered = ordered)\n    validation = validation.map(lambda image, label, idnum: [image, label])\n    validation = validation.batch(BATCH_SIZE)\n    validation = validation.cache()\n    # Prefetch next batch while training (autotune prefetch buffer size)\n    validation = validation.prefetch(AUTO)\n    \n    return validation","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.062416,"end_time":"2020-12-04T20:30:07.268999","exception":false,"start_time":"2020-12-04T20:30:07.206583","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Gain test set"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:30:07.389084Z","iopub.status.busy":"2020-12-04T20:30:07.388352Z","iopub.status.idle":"2020-12-04T20:30:07.392173Z","shell.execute_reply":"2020-12-04T20:30:07.391448Z"},"papermill":{"duration":0.06697,"end_time":"2020-12-04T20:30:07.392315","exception":false,"start_time":"2020-12-04T20:30:07.325345","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"def get_test_dataset(ordered = False):\n    \n    test = load_dataset(TEST_FILENAMES, labeled = False, ordered = ordered)\n    test = test.batch(BATCH_SIZE)\n    test = test.prefetch(AUTO)\n    \n    return test","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.055251,"end_time":"2020-12-04T20:30:07.503333","exception":false,"start_time":"2020-12-04T20:30:07.448082","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Count the number of images"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:30:07.625473Z","iopub.status.busy":"2020-12-04T20:30:07.624615Z","iopub.status.idle":"2020-12-04T20:30:07.629967Z","shell.execute_reply":"2020-12-04T20:30:07.629356Z"},"papermill":{"duration":0.071079,"end_time":"2020-12-04T20:30:07.630088","exception":false,"start_time":"2020-12-04T20:30:07.559009","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"def 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\nNUM_TRAINING_IMAGES = count_data_items(TRAINING_FILENAMES)  # Number of images in training set\nNUM_VALIDATION_IMAGES = count_data_items(VALIDATION_FILENAMES)  # Number of images in validation set\nNUM_TEST_IMAGES = count_data_items(TEST_FILENAMES)  # Number of images in test set\nSTEPS_PER_EPOCH = NUM_TRAINING_IMAGES // BATCH_SIZE  # Steps of each epoch\nprint('Dataset: {} training images, {} validation images, {} unlabeled test images'.format(NUM_TRAINING_IMAGES, NUM_VALIDATION_IMAGES, NUM_TEST_IMAGES))","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.056244,"end_time":"2020-12-04T20:30:07.743071","exception":false,"start_time":"2020-12-04T20:30:07.686827","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Show the beautiful flowers in training set before data augmentation"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:30:07.866256Z","iopub.status.busy":"2020-12-04T20:30:07.865176Z","iopub.status.idle":"2020-12-04T20:30:21.620443Z","shell.execute_reply":"2020-12-04T20:30:21.621051Z"},"papermill":{"duration":13.821997,"end_time":"2020-12-04T20:30:21.621224","exception":false,"start_time":"2020-12-04T20:30:07.799227","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"train_dataset_aug = get_training_dataset()\ndisplay_batch_of_images(next(iter(train_dataset_aug.unbatch().batch(20))))\ndisplay_batch_of_images(next(iter(train_dataset_aug.unbatch().batch(20))))\ndisplay_batch_of_images(next(iter(train_dataset_aug.unbatch().batch(20))))","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.164442,"end_time":"2020-12-04T20:30:21.956918","exception":false,"start_time":"2020-12-04T20:30:21.792476","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Show the beautiful flowers in validation set before data augmentation"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:30:22.288286Z","iopub.status.busy":"2020-12-04T20:30:22.287531Z","iopub.status.idle":"2020-12-04T20:30:32.940511Z","shell.execute_reply":"2020-12-04T20:30:32.941109Z"},"papermill":{"duration":10.823431,"end_time":"2020-12-04T20:30:32.941299","exception":false,"start_time":"2020-12-04T20:30:22.117868","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"validation_dataset_aug = get_validation_dataset()\ndisplay_batch_of_images(next(iter(validation_dataset_aug.unbatch().batch(20))))\ndisplay_batch_of_images(next(iter(validation_dataset_aug.unbatch().batch(20))))\ndisplay_batch_of_images(next(iter(validation_dataset_aug.unbatch().batch(20))))","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.265212,"end_time":"2020-12-04T20:30:33.474668","exception":false,"start_time":"2020-12-04T20:30:33.209456","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Show the beautiful flowers in test set before data augmentation"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:30:33.994027Z","iopub.status.busy":"2020-12-04T20:30:33.992912Z","iopub.status.idle":"2020-12-04T20:30:45.071795Z","shell.execute_reply":"2020-12-04T20:30:45.072418Z"},"papermill":{"duration":11.344111,"end_time":"2020-12-04T20:30:45.072599","exception":false,"start_time":"2020-12-04T20:30:33.728488","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"test_dataset_aug = get_test_dataset()\ndisplay_batch_of_images(next(iter(test_dataset_aug.unbatch().batch(20))))\ndisplay_batch_of_images(next(iter(test_dataset_aug.unbatch().batch(20))))\ndisplay_batch_of_images(next(iter(test_dataset_aug.unbatch().batch(20))))","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.384763,"end_time":"2020-12-04T20:30:45.852388","exception":false,"start_time":"2020-12-04T20:30:45.467625","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Show example augmentation"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:30:46.622279Z","iopub.status.busy":"2020-12-04T20:30:46.621453Z","iopub.status.idle":"2020-12-04T20:30:51.501419Z","shell.execute_reply":"2020-12-04T20:30:51.501986Z"},"papermill":{"duration":5.271472,"end_time":"2020-12-04T20:30:51.502143","exception":false,"start_time":"2020-12-04T20:30:46.230671","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"row = 3\ncol = 4\nall_elements = get_training_dataset().unbatch()\none_element = tf.data.Dataset.from_tensors(next(iter(all_elements)))\n# Map the images to the data augmentation function for image processing\naugmented_element = one_element.repeat().map(data_augment).batch(row * col)\n\nfor (img, label) in augmented_element:\n    plt.figure(figsize = (15, int(15 * row / col)))\n    for j in range(row * col):\n        plt.subplot(row, col, j + 1)\n        plt.axis('off')\n        plt.imshow(img[j, ])\n    plt.show()\n    break","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.386027,"end_time":"2020-12-04T20:30:52.302666","exception":false,"start_time":"2020-12-04T20:30:51.916639","status":"completed"},"tags":[]},"cell_type":"markdown","source":"## Step 4: Build the model and make prediction"},{"metadata":{"papermill":{"duration":0.384104,"end_time":"2020-12-04T20:30:53.072136","exception":false,"start_time":"2020-12-04T20:30:52.688032","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Customize learning rate scheduler and visualize it (refer to Chris Deotte)"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:30:53.858886Z","iopub.status.busy":"2020-12-04T20:30:53.856882Z","iopub.status.idle":"2020-12-04T20:30:54.052967Z","shell.execute_reply":"2020-12-04T20:30:54.052318Z"},"papermill":{"duration":0.59278,"end_time":"2020-12-04T20:30:54.053094","exception":false,"start_time":"2020-12-04T20:30:53.460314","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"def lrfn(epoch):\n    \n    LR_START = 0.00001\n    LR_MAX = 0.00005 * strategy.num_replicas_in_sync\n    LR_MIN = 0.00001\n    LR_RAMPUP_EPOCHS = 5\n    LR_SUSTAIN_EPOCHS = 0\n    LR_EXP_DECAY = .8\n    \n    if epoch < LR_RAMPUP_EPOCHS:\n        lr = (LR_MAX - LR_START) / LR_RAMPUP_EPOCHS * epoch + LR_START\n    elif epoch < LR_RAMPUP_EPOCHS + LR_SUSTAIN_EPOCHS:\n        lr = LR_MAX\n    else:\n        lr = (LR_MAX - LR_MIN) * LR_EXP_DECAY ** (epoch - LR_RAMPUP_EPOCHS - LR_SUSTAIN_EPOCHS) + LR_MIN\n    return lr\n\nlr_callback = tf.keras.callbacks.LearningRateScheduler(lrfn, verbose = True)\n\n# Visualization changes in learning rate\nrng = [i for i in range(25 if EPOCHS<25 else EPOCHS)]\ny = [lrfn(x) for x in rng]\nplt.plot(rng, y)\nprint(\"Learning rate schedule: {:.3g} to {:.3g} to {:.3g}\".format(y[0], max(y), y[-1]))","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.38404,"end_time":"2020-12-04T20:30:54.820889","exception":false,"start_time":"2020-12-04T20:30:54.436849","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Build the model and load it into TPU"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:30:55.642370Z","iopub.status.busy":"2020-12-04T20:30:55.641342Z","iopub.status.idle":"2020-12-04T20:32:28.361279Z","shell.execute_reply":"2020-12-04T20:32:28.360370Z"},"papermill":{"duration":93.137474,"end_time":"2020-12-04T20:32:28.361415","exception":false,"start_time":"2020-12-04T20:30:55.223941","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"with strategy.scope():\n    # Create EfficientNetB7 model\n    enet = efn.EfficientNetB7(\n        input_shape = (512, 512, 3),\n        weights = 'imagenet',  # Use the preset parameters of ImageNet\n        include_top = False  # Drop the fully connected network on the top\n    )\n\n    model = tf.keras.Sequential([\n        enet,\n        tf.keras.layers.GlobalAveragePooling2D(),\n        tf.keras.layers.Dense(len(CLASSES), activation = 'softmax')\n    ])\n\n    model.compile(\n        optimizer=tf.keras.optimizers.Adam(),  # Use Adam Algorithm for optimization\n        # For multiclassification, we can use cross entropy or sparse cross entropy as our loss function \n        # These two cross entropy are the same in essence, but they are applied in different scenarios\n        # If our target is one-hot encoded, it is better to use cross entropy\n        # If our target is an integer, sparse cross entropy is a better choice, and this is our case\n        loss = 'sparse_categorical_crossentropy', \n        metrics = ['sparse_categorical_accuracy']\n    )\n\n    model.summary()\n    # Save the model\n    model.save('ML_finalproject.h5')","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.446508,"end_time":"2020-12-04T20:32:29.269985","exception":false,"start_time":"2020-12-04T20:32:28.823477","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Calculate weight for each class (refer to Roman Weilguny)\n[TPU flowers - First Love](https://www.kaggle.com/romanweilguny/tpu-flowers-first-love)"},{"metadata":{"papermill":{"duration":0.464065,"end_time":"2020-12-04T20:32:30.168966","exception":false,"start_time":"2020-12-04T20:32:29.704901","status":"completed"},"tags":[]},"cell_type":"markdown","source":"#### As the classes may not be uniformly distributed, add weights to classes"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:32:31.079498Z","iopub.status.busy":"2020-12-04T20:32:31.078687Z","iopub.status.idle":"2020-12-04T20:32:53.656030Z","shell.execute_reply":"2020-12-04T20:32:53.655380Z"},"papermill":{"duration":23.056655,"end_time":"2020-12-04T20:32:53.656178","exception":false,"start_time":"2020-12-04T20:32:30.599523","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"gc.enable()\n\ndef get_training_dataset_raw():\n    dataset = load_dataset(TRAINING_FILENAMES, labeled = True, ordered = False)\n    return dataset\n\nraw_training_dataset = get_training_dataset_raw()\n\nlabel_counter = Counter()\nfor images, labels, id in raw_training_dataset:\n    label_counter.update([labels.numpy()])\n\ndel raw_training_dataset    \n\nTARGET_NUM_PER_CLASS = 122\n\ndef get_weight_for_class(class_id):\n    counting = label_counter[class_id]\n    weight = TARGET_NUM_PER_CLASS / counting\n    return weight\n\nweight_per_class = {class_id: get_weight_for_class(class_id) for class_id in range(104)}","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.445565,"end_time":"2020-12-04T20:32:54.533359","exception":false,"start_time":"2020-12-04T20:32:54.087794","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Train the model"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T20:32:55.445856Z","iopub.status.busy":"2020-12-04T20:32:55.444624Z","iopub.status.idle":"2020-12-04T21:04:42.867686Z","shell.execute_reply":"2020-12-04T21:04:42.868600Z"},"papermill":{"duration":1907.866883,"end_time":"2020-12-04T21:04:42.868902","exception":false,"start_time":"2020-12-04T20:32:55.002019","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"history = model.fit(\n    get_training_dataset(),\n    steps_per_epoch = STEPS_PER_EPOCH,\n    epochs = EPOCHS,\n    callbacks = [lr_callback],\n    validation_data = get_validation_dataset(),\n    class_weight = weight_per_class\n)","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":1.126388,"end_time":"2020-12-04T21:04:45.080241","exception":false,"start_time":"2020-12-04T21:04:43.953853","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Show how loss and accuracy changes on training set"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T21:04:47.302226Z","iopub.status.busy":"2020-12-04T21:04:47.301356Z","iopub.status.idle":"2020-12-04T21:04:47.844544Z","shell.execute_reply":"2020-12-04T21:04:47.845118Z"},"papermill":{"duration":1.668528,"end_time":"2020-12-04T21:04:47.845321","exception":false,"start_time":"2020-12-04T21:04:46.176793","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"plot_train_valid_curves(history.history['loss'], history.history['val_loss'], 'loss', 211)  # Loss curve\nplot_train_valid_curves(history.history['sparse_categorical_accuracy'], \n                        history.history['val_sparse_categorical_accuracy'], 'accuracy', 212)  # Accuracy curve","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":1.074597,"end_time":"2020-12-04T21:04:50.090095","exception":false,"start_time":"2020-12-04T21:04:49.015498","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Check model's performance on validation set"},{"metadata":{"papermill":{"duration":1.081977,"end_time":"2020-12-04T21:04:52.248145","exception":false,"start_time":"2020-12-04T21:04:51.166168","status":"completed"},"tags":[]},"cell_type":"markdown","source":"#### Get the correct labels and predicted labels"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T21:04:54.476297Z","iopub.status.busy":"2020-12-04T21:04:54.475174Z","iopub.status.idle":"2020-12-04T21:05:45.702068Z","shell.execute_reply":"2020-12-04T21:05:45.702826Z"},"papermill":{"duration":52.378771,"end_time":"2020-12-04T21:05:45.702999","exception":false,"start_time":"2020-12-04T21:04:53.324228","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"cmdataset = get_validation_dataset(ordered = True)\nimages_ds = cmdataset.map(lambda image, label: image)\nlabels_ds = cmdataset.map(lambda image, label: label).unbatch() \ncm_correct_labels = next(iter(labels_ds.batch(NUM_VALIDATION_IMAGES))).numpy()  # Get everything as one batch\ncm_probabilities = model.predict(images_ds)  # The probability that each image is of each class\ncm_predictions = np.argmax(cm_probabilities, axis = -1)  # The class of the largest probability is what we need\nprint(\"Correct labels: \", cm_correct_labels.shape, cm_correct_labels)\nprint(\"Predicted labels: \", cm_predictions.shape, cm_predictions)","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":1.073499,"end_time":"2020-12-04T21:05:47.853550","exception":false,"start_time":"2020-12-04T21:05:46.780051","status":"completed"},"tags":[]},"cell_type":"markdown","source":"#### Draw the confusion matrix, compute F1 score, precision, and recall"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T21:05:50.041721Z","iopub.status.busy":"2020-12-04T21:05:50.040534Z","iopub.status.idle":"2020-12-04T21:05:52.630617Z","shell.execute_reply":"2020-12-04T21:05:52.629931Z"},"papermill":{"duration":3.702326,"end_time":"2020-12-04T21:05:52.630799","exception":false,"start_time":"2020-12-04T21:05:48.928473","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"cmat = confusion_matrix(cm_correct_labels, cm_predictions, labels = range(len(CLASSES)))\nscore = f1_score(cm_correct_labels, cm_predictions, labels = range(len(CLASSES)), average = 'macro')\nprecision = precision_score(cm_correct_labels, cm_predictions, labels = range(len(CLASSES)), average = 'macro')\nrecall = recall_score(cm_correct_labels, cm_predictions, labels = range(len(CLASSES)), average = 'macro')\ncmat = (cmat.T / cmat.sum(axis = 1)).T\ndisplay_confusion_matrix(cmat, score, precision, recall)\nprint('f1 score: {:.3f}, precision: {:.3f}, recall: {:.3f}'.format(score, precision, recall))","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":1.104288,"end_time":"2020-12-04T21:05:54.865786","exception":false,"start_time":"2020-12-04T21:05:53.761498","status":"completed"},"tags":[]},"cell_type":"markdown","source":"* ### Make prediction"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-04T21:05:57.046750Z","iopub.status.busy":"2020-12-04T21:05:57.045863Z","iopub.status.idle":"2020-12-04T21:06:36.017733Z","shell.execute_reply":"2020-12-04T21:06:36.016687Z"},"papermill":{"duration":40.055905,"end_time":"2020-12-04T21:06:36.017961","exception":false,"start_time":"2020-12-04T21:05:55.962056","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"test_ds = get_test_dataset(ordered = True)\n\ntest_images_ds = test_ds.map(lambda image, idnum: image)\nprobabilities = model.predict(test_images_ds)  # Compute the probability that each image is of each class\npredictions = np.argmax(probabilities, axis = -1)  # Use the one with largest probability as the predicted class\nprint(predictions)\n\n# Generate submission file, remember to name it by \"submission.csv\"\ntest_ids_ds = test_ds.map(lambda image, idnum: idnum).unbatch()\ntest_ids = next(iter(test_ids_ds.batch(NUM_TEST_IMAGES))).numpy().astype('U')\ntest = pd.DataFrame({\"id\": test_ids, \"label\": predictions})\nprint(test.head)\ntest.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}