{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpu1vmV38","dataSources":[{"sourceId":21154,"databundleVersionId":1243559,"sourceType":"competition"},{"sourceId":37130068,"sourceType":"kernelVersion"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"\n# Introduction #\n\nThis is notebook for the [**Petals to the Metal**](https://www.kaggle.com/c/tpu-getting-started) competition! In this competition, we are challenged to build a machine learning model to classify 104 types of flowers based on their images.\n\nIn this tutorial notebook, we will build an image classifier in Keras and train it on a [Tensor Processing Unit (TPU)](https://www.kaggle.com/docs/tpu).\n\nThis code is in this notebook is heavily inspired from [**Ryan Holbrok's Starter Notebook**](https://www.kaggle.com/ryanholbrook/create-your-first-submission).","metadata":{}},{"cell_type":"markdown","source":"# Step 1: Imports #\n\nWe begin by importing several Python packages.","metadata":{}},{"cell_type":"code","source":"import math, re, os\nimport numpy as np\nimport tensorflow as tf\n\nprint(\"Tensorflow version \" + tf.__version__)","metadata":{"execution":{"iopub.status.busy":"2023-11-22T05:54:55.368990Z","iopub.execute_input":"2023-11-22T05:54:55.369414Z","iopub.status.idle":"2023-11-22T05:55:09.749465Z","shell.execute_reply.started":"2023-11-22T05:54:55.369385Z","shell.execute_reply":"2023-11-22T05:55:09.748699Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Step 2: Using TPUs #\n\nThe following cell will help us use TPUs to train our model.","metadata":{}},{"cell_type":"code","source":"# Detect TPU, return appropriate distribution strategy\ntry:\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)","metadata":{"execution":{"iopub.status.busy":"2023-11-22T05:55:13.041508Z","iopub.execute_input":"2023-11-22T05:55:13.042091Z","iopub.status.idle":"2023-11-22T05:55:22.513048Z","shell.execute_reply.started":"2023-11-22T05:55:13.042057Z","shell.execute_reply":"2023-11-22T05:55:22.512295Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Step 3: Loading the Competition Data #\n\n## Get GCS Path ##\n\nWhen used with TPUs, datasets need to be stored in a [Google Cloud Storage bucket](https://cloud.google.com/storage/). You can use data from any public GCS bucket by giving its path just like you would data from `'/kaggle/input'`. The following will retrieve the GCS path for this competition's dataset.","metadata":{}},{"cell_type":"code","source":"from kaggle_datasets import KaggleDatasets\n\nGCS_DS_PATH = KaggleDatasets().get_gcs_path('tpu-getting-started')\nprint(GCS_DS_PATH) # what do gcs paths look like?","metadata":{"execution":{"iopub.status.busy":"2023-11-22T05:55:23.805012Z","iopub.execute_input":"2023-11-22T05:55:23.805939Z","iopub.status.idle":"2023-11-22T05:55:23.815144Z","shell.execute_reply.started":"2023-11-22T05:55:23.805892Z","shell.execute_reply":"2023-11-22T05:55:23.814293Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Load Data ##\n\nWhen used with TPUs, datasets are often serialized into [TFRecords](https://www.kaggle.com/ryanholbrook/tfrecords-basics). So, we will make our custom functions accordingly.","metadata":{}},{"cell_type":"code","source":"\nIMAGE_SIZE = [512, 512]\nGCS_PATH = GCS_DS_PATH + '/tfrecords-jpeg-512x512'\nAUTO = tf.data.experimental.AUTOTUNE\n\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') \n\nCLASSES = ['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']                                                                                                                                               # 100 - 102\n\n\ndef decode_image(image_data):\n    image = tf.image.decode_jpeg(image_data, channels=3)\n    image = tf.cast(image, tf.float32) / 255.0  # convert image to floats in [0, 1] range\n    image = tf.reshape(image, [*IMAGE_SIZE, 3]) # explicit size needed for TPU\n    return image\n\ndef read_labeled_tfrecord(example):\n    LABELED_TFREC_FORMAT = {\n        \"image\": tf.io.FixedLenFeature([], tf.string), # tf.string means bytestring\n        \"class\": tf.io.FixedLenFeature([], tf.int64),  # shape [] means single element\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    return image, label # returns a dataset of (image, label) pairs\n\ndef read_unlabeled_tfrecord(example):\n    UNLABELED_TFREC_FORMAT = {\n        \"image\": tf.io.FixedLenFeature([], tf.string), # tf.string means bytestring\n        \"id\": tf.io.FixedLenFeature([], tf.string),  # shape [] means single element\n    }\n    example = tf.io.parse_single_example(example, UNLABELED_TFREC_FORMAT)\n    image = decode_image(example['image'])\n    idnum = example['id']\n    return image, idnum # returns a dataset of image(s)\n\ndef load_dataset(filenames, labeled=True, ordered=False):\n    # Read from TFRecords. For optimal performance, reading from multiple files at once and\n    # disregarding data order. Order does not matter since we will be shuffling the data anyway.\n\n    ignore_order = tf.data.Options()\n    if not ordered:\n        ignore_order.experimental_deterministic = False # disable order, increase speed\n\n    dataset = tf.data.TFRecordDataset(filenames, num_parallel_reads=AUTO) # automatically interleaves reads from multiple files\n    dataset = dataset.with_options(ignore_order) # uses data as soon as it streams in, rather than in its original order\n    dataset = dataset.map(read_labeled_tfrecord if labeled else read_unlabeled_tfrecord, num_parallel_calls=AUTO)\n    # returns a dataset of (image, label) pairs if labeled=True or (image, id) pairs if labeled=False\n    return dataset","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-11-22T05:55:29.670205Z","iopub.execute_input":"2023-11-22T05:55:29.670559Z","iopub.status.idle":"2023-11-22T05:55:29.711193Z","shell.execute_reply.started":"2023-11-22T05:55:29.670529Z","shell.execute_reply":"2023-11-22T05:55:29.710327Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Create Data Pipelines ##\n\nIn this step we'll use the `tf.data` API to define an efficient data pipeline for each of the training, validation, and test splits.","metadata":{}},{"cell_type":"code","source":"\ndef data_augment(image, label):\n    image = tf.image.random_flip_left_right(image)\n    image = tf.image.random_flip_up_down(image)\n    #image = tf.image.random_saturation(image, 0, 2)\n    return image, label   \n\ndef get_training_dataset():\n    dataset = load_dataset(TRAINING_FILENAMES, labeled=True)\n    dataset = dataset.map(data_augment, num_parallel_calls=AUTO)\n    dataset = dataset.repeat() # the training dataset must repeat for several epochs\n    dataset = dataset.shuffle(2048)\n    dataset = dataset.batch(BATCH_SIZE)\n    dataset = dataset.prefetch(AUTO) # 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.batch(BATCH_SIZE)\n    dataset = dataset.cache()\n    dataset = dataset.prefetch(AUTO)\n    return dataset\n\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(AUTO)\n    return dataset\n\ndef count_data_items(filenames):\n    # the number of data items is written in the name of the .tfrec files, i.e. flowers00-230.tfrec = 230 data items\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)\nNUM_VALIDATION_IMAGES = count_data_items(VALIDATION_FILENAMES)\nNUM_TEST_IMAGES = count_data_items(TEST_FILENAMES)\nprint('Dataset: {} training images, {} validation images, {} unlabeled test images'.format(NUM_TRAINING_IMAGES, NUM_VALIDATION_IMAGES, NUM_TEST_IMAGES))\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-11-22T06:18:35.981005Z","iopub.execute_input":"2023-11-22T06:18:35.981476Z","iopub.status.idle":"2023-11-22T06:18:35.993264Z","shell.execute_reply.started":"2023-11-22T06:18:35.981428Z","shell.execute_reply":"2023-11-22T06:18:35.992198Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"This next cell will create the datasets that we'll use with Keras during training and inference. Here, we will scale the size of the batches to the number of TPU cores.","metadata":{}},{"cell_type":"code","source":"# Define the batch size. This will be 16 with TPU off and 128 (=16*8) with TPU on\nBATCH_SIZE = 16 * strategy.num_replicas_in_sync\n\nds_train = get_training_dataset()\nds_valid = get_validation_dataset()\nds_test = get_test_dataset()\n\nprint(\"Training:\", ds_train)\nprint (\"Validation:\", ds_valid)\nprint(\"Test:\", ds_test)","metadata":{"execution":{"iopub.status.busy":"2023-11-22T06:19:34.581497Z","iopub.execute_input":"2023-11-22T06:19:34.581978Z","iopub.status.idle":"2023-11-22T06:19:35.974523Z","shell.execute_reply.started":"2023-11-22T06:19:34.581940Z","shell.execute_reply":"2023-11-22T06:19:35.973425Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Step 4: Explore Data #\n\nLet's take a moment to look at some of the images in the dataset.","metadata":{}},{"cell_type":"code","source":"from matplotlib import pyplot as plt\n\ndef 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: # binary string in this case, these are image ID strings\n        numpy_labels = [None for _ in enumerate(numpy_images)]\n    # If no labels, only image IDs, return None for labels (this is the case for test data)\n    return numpy_images, numpy_labels\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\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    \ndef display_batch_of_images(databatch, predictions=None):\n    \"\"\"This will work with:\n    display_batch_of_images(images)\n    display_batch_of_images(images, predictions)\n    display_batch_of_images((images, labels))\n    display_batch_of_images((images, labels), predictions)\n    \"\"\"\n    # data\n    images, labels = batch_to_numpy_images_and_labels(databatch)\n    if labels is None:\n        labels = [None for _ in enumerate(images)]\n        \n    # auto-squaring: this will drop data that does not fit into square or square-ish rectangle\n    rows = int(math.sqrt(len(images)))\n    cols = len(images)//rows\n        \n    # size and spacing\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\ndef display_training_curves(training, validation, title, subplot):\n    if subplot%10==1: # set up the subplots on the first call\n        plt.subplots(figsize=(10,10), 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(['train', 'valid.'])","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-11-22T05:55:48.202830Z","iopub.execute_input":"2023-11-22T05:55:48.203413Z","iopub.status.idle":"2023-11-22T05:55:49.131120Z","shell.execute_reply.started":"2023-11-22T05:55:48.203377Z","shell.execute_reply":"2023-11-22T05:55:49.130303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"You can display a single batch of images from a dataset with another of our helper functions. The next cell will turn the dataset into an iterator of batches of 20 images.","metadata":{}},{"cell_type":"code","source":"ds_iter = iter(ds_train.unbatch().batch(20))","metadata":{"execution":{"iopub.status.busy":"2023-11-22T05:55:52.809050Z","iopub.execute_input":"2023-11-22T05:55:52.810084Z","iopub.status.idle":"2023-11-22T05:55:52.862141Z","shell.execute_reply.started":"2023-11-22T05:55:52.810024Z","shell.execute_reply":"2023-11-22T05:55:52.861198Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Then we use the Python `next` function to pop out the next batch in the stream and display it with the helper function.","metadata":{}},{"cell_type":"code","source":"one_batch = next(ds_iter)\ndisplay_batch_of_images(one_batch)","metadata":{"execution":{"iopub.status.busy":"2023-11-22T05:55:54.024859Z","iopub.execute_input":"2023-11-22T05:55:54.025194Z","iopub.status.idle":"2023-11-22T05:55:58.108371Z","shell.execute_reply.started":"2023-11-22T05:55:54.025166Z","shell.execute_reply":"2023-11-22T05:55:58.107417Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"By defining `ds_iter` and `one_batch` in separate cells, you only need to rerun the cell above to see a new batch of images.","metadata":{}},{"cell_type":"markdown","source":"# Step 5: Define Model #\n\nFor this problem, we'll use a model called **DenseNet** available in [Keras](https://www.tensorflow.org/api_docs/python/tf/keras/applications).\n\nHere, we import the entire DenseNet except for its top and instead use our own Neural Network in its place to help with the classification.","metadata":{}},{"cell_type":"code","source":"EPOCHS = 15\n\nwith strategy.scope():\n    pretrained_model = tf.keras.applications.DenseNet201(\n        weights='imagenet',\n        include_top=False ,\n        input_shape=[*IMAGE_SIZE, 3]\n    )\n    pretrained_model.trainable = True\n    \n    model = tf.keras.Sequential([\n        pretrained_model,\n        # ... attach a new head to act as a classifier.\n        tf.keras.layers.GlobalAveragePooling2D(),\n#         tf.keras.layers.Dense(480, activation = 'relu'),\n#         tf.keras.layers.Dense(240, activation = 'relu'),\n        tf.keras.layers.Dense(len(CLASSES), activation='softmax')\n    ])","metadata":{"execution":{"iopub.status.busy":"2023-11-22T06:19:43.749342Z","iopub.execute_input":"2023-11-22T06:19:43.749792Z","iopub.status.idle":"2023-11-22T06:20:21.806526Z","shell.execute_reply.started":"2023-11-22T06:19:43.749749Z","shell.execute_reply":"2023-11-22T06:20:21.805477Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The `'sparse_categorical'` versions of the loss and metrics are appropriate for a classification task with more than two labels, like this one.","metadata":{}},{"cell_type":"code","source":"model.compile(\n    optimizer='adam',\n    loss = 'sparse_categorical_crossentropy',\n    metrics=['sparse_categorical_accuracy'],\n)\n\nmodel.summary()","metadata":{"execution":{"iopub.status.busy":"2023-11-22T06:20:21.808009Z","iopub.execute_input":"2023-11-22T06:20:21.808286Z","iopub.status.idle":"2023-11-22T06:20:21.916518Z","shell.execute_reply.started":"2023-11-22T06:20:21.808243Z","shell.execute_reply":"2023-11-22T06:20:21.915708Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Step 6: Training #\n\n## Learning Rate Schedule ##\n\nWe'll train this network with a special learning rate schedule.","metadata":{}},{"cell_type":"code","source":"# Learning Rate Schedule for Fine Tuning #\ndef exponential_lr(epoch,\n                   start_lr = 0.00001, min_lr = 0.00001, max_lr = 0.00005,\n                   rampup_epochs = 5, sustain_epochs = 0,\n                   exp_decay = 0.8):\n\n    def lr(epoch, start_lr, min_lr, max_lr, rampup_epochs, sustain_epochs, exp_decay):\n        # linear increase from start to rampup_epochs\n        if epoch < rampup_epochs:\n            lr = ((max_lr - start_lr) /\n                  rampup_epochs * epoch + start_lr)\n        # constant max_lr during sustain_epochs\n        elif epoch < rampup_epochs + sustain_epochs:\n            lr = max_lr\n        # exponential decay towards min_lr\n        else:\n            lr = ((max_lr - min_lr) *\n                  exp_decay**(epoch - rampup_epochs - sustain_epochs) +\n                  min_lr)\n        return lr\n    return lr(epoch,\n              start_lr,\n              min_lr,\n              max_lr,\n              rampup_epochs,\n              sustain_epochs,\n              exp_decay)\n\nlr_callback = tf.keras.callbacks.LearningRateScheduler(exponential_lr, verbose=True)\n\nrng = [i for i in range(EPOCHS)]\ny = [exponential_lr(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]))","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-11-22T05:57:03.436691Z","iopub.execute_input":"2023-11-22T05:57:03.436976Z","iopub.status.idle":"2023-11-22T05:57:03.674234Z","shell.execute_reply.started":"2023-11-22T05:57:03.436947Z","shell.execute_reply":"2023-11-22T05:57:03.673048Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Fit Model ##\n\nAnd now we're ready to train the model. After defining a few parameters, we're good to go!","metadata":{}},{"cell_type":"code","source":"# Define training epochs\nEPOCHS = 15\nSTEPS_PER_EPOCH = NUM_TRAINING_IMAGES // BATCH_SIZE\n\nhistory = model.fit(\n    ds_train,\n    validation_data=ds_valid,\n    epochs=EPOCHS,\n    steps_per_epoch=STEPS_PER_EPOCH,\n    callbacks=[lr_callback],\n)","metadata":{"execution":{"iopub.status.busy":"2023-11-22T06:20:21.917432Z","iopub.execute_input":"2023-11-22T06:20:21.917670Z","iopub.status.idle":"2023-11-22T06:49:02.768454Z","shell.execute_reply.started":"2023-11-22T06:20:21.917644Z","shell.execute_reply":"2023-11-22T06:49:02.766531Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"This next cell shows how the loss and metrics progressed during training. Thankfully, it converges!","metadata":{}},{"cell_type":"code","source":"display_training_curves(\n    history.history['loss'],\n    history.history['val_loss'],\n    'loss',\n    211,\n)\ndisplay_training_curves(\n    history.history['sparse_categorical_accuracy'],\n    history.history['val_sparse_categorical_accuracy'],\n    'accuracy',\n    212,\n)","metadata":{"execution":{"iopub.status.busy":"2023-11-22T06:49:02.771261Z","iopub.execute_input":"2023-11-22T06:49:02.771627Z","iopub.status.idle":"2023-11-22T06:49:03.417396Z","shell.execute_reply.started":"2023-11-22T06:49:02.771594Z","shell.execute_reply":"2023-11-22T06:49:03.416309Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Step 7: Evaluate Predictions #\n\nBefore making our final predictions on the test set, we'll look at two common ways of validation: plotting the **confusion matrix** and **visual validation**.","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nfrom sklearn.metrics import f1_score, precision_score, recall_score, confusion_matrix\n\ndef display_confusion_matrix(cmat, score, precision, recall):\n    plt.figure(figsize=(15,15))\n    ax = plt.gca()\n    ax.matshow(cmat, cmap='Reds')\n    ax.set_xticks(range(len(CLASSES)))\n    ax.set_xticklabels(CLASSES, fontdict={'fontsize': 7})\n    plt.setp(ax.get_xticklabels(), rotation=45, ha=\"left\", rotation_mode=\"anchor\")\n    ax.set_yticks(range(len(CLASSES)))\n    ax.set_yticklabels(CLASSES, fontdict={'fontsize': 7})\n    plt.setp(ax.get_yticklabels(), rotation=45, ha=\"right\", rotation_mode=\"anchor\")\n    titlestring = \"\"\n    if score is not None:\n        titlestring += 'f1 = {:.3f} '.format(score)\n    if precision is not None:\n        titlestring += '\\nprecision = {:.3f} '.format(precision)\n    if recall is not None:\n        titlestring += '\\nrecall = {:.3f} '.format(recall)\n    if len(titlestring) > 0:\n        ax.text(101, 1, titlestring, fontdict={'fontsize': 18, 'horizontalalignment':'right', 'verticalalignment':'top', 'color':'#804040'})\n    plt.show()\n    \ndef display_training_curves(training, validation, title, subplot):\n    if subplot%10==1: # set up the subplots on the first call\n        plt.subplots(figsize=(10,10), 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(['train', 'valid.'])","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-11-22T06:49:27.431014Z","iopub.execute_input":"2023-11-22T06:49:27.432161Z","iopub.status.idle":"2023-11-22T06:49:27.442342Z","shell.execute_reply.started":"2023-11-22T06:49:27.432117Z","shell.execute_reply":"2023-11-22T06:49:27.441140Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Confusion Matrix ##\n\nThe following cell does some processing on the validation data and then creates the matrix with the `confusion_matrix` function included in [`scikit-learn`](https://scikit-learn.org/stable/index.html).","metadata":{}},{"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()\n\ncm_correct_labels = next(iter(labels_ds.batch(NUM_VALIDATION_IMAGES))).numpy()\ncm_probabilities = model.predict(images_ds)\ncm_predictions = np.argmax(cm_probabilities, axis=-1)\n\nlabels = range(len(CLASSES))\ncmat = confusion_matrix(\n    cm_correct_labels,\n    cm_predictions,\n    labels=labels,\n)\ncmat = (cmat.T / cmat.sum(axis=1)).T # normalize","metadata":{"execution":{"iopub.status.busy":"2023-11-22T06:49:29.941792Z","iopub.execute_input":"2023-11-22T06:49:29.942876Z","iopub.status.idle":"2023-11-22T06:50:14.576110Z","shell.execute_reply.started":"2023-11-22T06:49:29.942827Z","shell.execute_reply":"2023-11-22T06:50:14.574740Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"This cell computes the precision and recall metrics and then displays them with a plot of the confusion matrix.","metadata":{}},{"cell_type":"code","source":"score = f1_score(\n    cm_correct_labels,\n    cm_predictions,\n    labels=labels,\n    average='macro',\n)\nprecision = precision_score(\n    cm_correct_labels,\n    cm_predictions,\n    labels=labels,\n    average='macro',\n)\nrecall = recall_score(\n    cm_correct_labels,\n    cm_predictions,\n    labels=labels,\n    average='macro',\n)\ndisplay_confusion_matrix(cmat, score, precision, recall)","metadata":{"execution":{"iopub.status.busy":"2023-11-22T06:50:14.578061Z","iopub.execute_input":"2023-11-22T06:50:14.578711Z","iopub.status.idle":"2023-11-22T06:50:16.382786Z","shell.execute_reply.started":"2023-11-22T06:50:14.578671Z","shell.execute_reply":"2023-11-22T06:50:16.381733Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visual Validation ##\n\nIt can also be helpful to look at some examples from the validation set and see what class your model predicted. This can help reveal patterns in the kinds of images your model has trouble with.\n\nThis cell will set up the validation set to display 20 images at a time.","metadata":{}},{"cell_type":"code","source":"dataset = get_validation_dataset()\ndataset = dataset.unbatch().batch(20)\nbatch = iter(dataset)","metadata":{"execution":{"iopub.status.busy":"2023-11-22T06:50:16.384092Z","iopub.execute_input":"2023-11-22T06:50:16.384507Z","iopub.status.idle":"2023-11-22T06:50:16.659111Z","shell.execute_reply.started":"2023-11-22T06:50:16.384474Z","shell.execute_reply":"2023-11-22T06:50:16.657761Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"And here is a set of flowers with their predicted species. Run the cell again to see another set.","metadata":{}},{"cell_type":"code","source":"images, labels = next(batch)\nprobabilities = model.predict(images)\npredictions = np.argmax(probabilities, axis=-1)\ndisplay_batch_of_images((images, labels), predictions)","metadata":{"execution":{"iopub.status.busy":"2023-11-22T06:50:16.661925Z","iopub.execute_input":"2023-11-22T06:50:16.662645Z","iopub.status.idle":"2023-11-22T06:50:59.830817Z","shell.execute_reply.started":"2023-11-22T06:50:16.662607Z","shell.execute_reply":"2023-11-22T06:50:59.829506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Step 8: Make Test Predictions #\n\nNow we are ready to make predictions on the test set.","metadata":{}},{"cell_type":"code","source":"test_ds = get_test_dataset(ordered=True)\n\nprint('Computing predictions...')\ntest_images_ds = test_ds.map(lambda image, idnum: image)\nprobabilities = model.predict(test_images_ds)\npredictions = np.argmax(probabilities, axis=-1)\nprint(predictions)","metadata":{"execution":{"iopub.status.busy":"2023-11-22T06:50:59.832190Z","iopub.execute_input":"2023-11-22T06:50:59.832669Z","iopub.status.idle":"2023-11-22T06:51:58.240151Z","shell.execute_reply.started":"2023-11-22T06:50:59.832619Z","shell.execute_reply":"2023-11-22T06:51:58.238798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We'll generate a file `submission.csv`. This file is what we'll submit to get the score on the leaderboard.","metadata":{}},{"cell_type":"code","source":"print('Generating submission.csv file...')\n\n# Get image ids from test set and convert to unicode\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')\n\n# Write the submission file\nnp.savetxt(\n    'submission.csv',\n    np.rec.fromarrays([test_ids, predictions]),\n    fmt=['%s', '%d'],\n    delimiter=',',\n    header='id,label',\n    comments='',\n)\n\n# Look at the first few predictions\n!head submission.csv","metadata":{"execution":{"iopub.status.busy":"2023-11-22T06:51:58.241594Z","iopub.execute_input":"2023-11-22T06:51:58.241959Z","iopub.status.idle":"2023-11-22T06:52:06.099768Z","shell.execute_reply.started":"2023-11-22T06:51:58.241922Z","shell.execute_reply":"2023-11-22T06:52:06.098095Z"},"trusted":true},"execution_count":null,"outputs":[]}]}