{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Transfer Learning on TPU For Flower Classification","metadata":{"papermill":{"duration":0.044088,"end_time":"2022-03-03T14:45:47.040832","exception":false,"start_time":"2022-03-03T14:45:46.996744","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"This notebook demonstrates how to use TPUs with TensorFlow to train a model for classifying flower images.\n\nCopy of [Transfer Learning On TPU For Flower Classification](https://www.kaggle.com/code/defcodeking/transfer-learning-on-tpu-for-flower-classification).","metadata":{"papermill":{"duration":0.042164,"end_time":"2022-03-03T14:45:47.128931","exception":false,"start_time":"2022-03-03T14:45:47.086767","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# Imports","metadata":{"papermill":{"duration":0.042127,"end_time":"2022-03-03T14:45:47.214014","exception":false,"start_time":"2022-03-03T14:45:47.171887","status":"completed"},"tags":[]}},{"cell_type":"code","source":"from __future__ import annotations\n\nimport functools\nimport math\nimport os\nimport warnings\n\nos.environ['TF_CPP_MIN_LOG_LEVEL'] = '4'\n\nfrom kaggle_datasets import KaggleDatasets\nfrom matplotlib import pyplot as plt\nimport numpy as np\nimport optuna\nimport pandas as pd\nimport tensorflow as tf\nfrom tensorflow.keras import layers, optimizers, applications, callbacks, Sequential, Input\n\nwarnings.filterwarnings(\"ignore\")\ntf.random.set_seed(42)","metadata":{"papermill":{"duration":6.891148,"end_time":"2022-03-03T14:45:54.147219","exception":false,"start_time":"2022-03-03T14:45:47.256071","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-06-05T16:15:11.899102Z","iopub.execute_input":"2022-06-05T16:15:11.901287Z","iopub.status.idle":"2022-06-05T16:15:21.007264Z","shell.execute_reply.started":"2022-06-05T16:15:11.899738Z","shell.execute_reply":"2022-06-05T16:15:21.006342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Detect TPU","metadata":{"papermill":{"duration":0.043358,"end_time":"2022-03-03T14:45:54.234722","exception":false,"start_time":"2022-03-03T14:45:54.191364","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"Code borrowed from [Getting started with 100+ flowers on TPU](https://www.kaggle.com/mgornergoogle/getting-started-with-100-flowers-on-tpu) by [Martin Görner](https://www.kaggle.com/mgornergoogle).","metadata":{"papermill":{"duration":0.042394,"end_time":"2022-03-03T14:45:54.320948","exception":false,"start_time":"2022-03-03T14:45:54.278554","status":"completed"},"tags":[]}},{"cell_type":"code","source":"try:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver.connect()\n    strategy = tf.distribute.TPUStrategy(tpu)\nexcept ValueError:  # detect GPUs\n    strategy = tf.distribute.MirroredStrategy()\n\nprint(\"Number of accelerators: \", strategy.num_replicas_in_sync)\nstrategy","metadata":{"execution":{"iopub.execute_input":"2022-03-03T14:45:54.421232Z","iopub.status.busy":"2022-03-03T14:45:54.420375Z","iopub.status.idle":"2022-03-03T14:46:00.115827Z","shell.execute_reply":"2022-03-03T14:46:00.116373Z","shell.execute_reply.started":"2022-03-03T14:38:20.421318Z"},"papermill":{"duration":5.752586,"end_time":"2022-03-03T14:46:00.116570","exception":false,"start_time":"2022-03-03T14:45:54.363984","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Moving Data To Google Cloud Storage (GCS)","metadata":{"papermill":{"duration":0.042417,"end_time":"2022-03-03T14:46:00.202268","exception":false,"start_time":"2022-03-03T14:46:00.159851","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"TPUs require data to be present on GCS. The below utility copies the data to a GCS bucket co-located with the TPU.","metadata":{"papermill":{"duration":0.043442,"end_time":"2022-03-03T14:46:00.288712","exception":false,"start_time":"2022-03-03T14:46:00.245270","status":"completed"},"tags":[]}},{"cell_type":"code","source":"GCS_DS_PATH = KaggleDatasets().get_gcs_path(\"tpu-getting-started\")","metadata":{"execution":{"iopub.execute_input":"2022-03-03T14:46:00.377503Z","iopub.status.busy":"2022-03-03T14:46:00.376858Z","iopub.status.idle":"2022-03-03T14:46:00.881782Z","shell.execute_reply":"2022-03-03T14:46:00.882304Z","shell.execute_reply.started":"2022-03-03T14:38:25.815490Z"},"papermill":{"duration":0.551185,"end_time":"2022-03-03T14:46:00.882495","exception":false,"start_time":"2022-03-03T14:46:00.331310","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Check the URLs for the dataset\n!gsutil ls $GCS_DS_PATH","metadata":{"execution":{"iopub.execute_input":"2022-03-03T14:46:00.973912Z","iopub.status.busy":"2022-03-03T14:46:00.973246Z","iopub.status.idle":"2022-03-03T14:46:04.760757Z","shell.execute_reply":"2022-03-03T14:46:04.760047Z","shell.execute_reply.started":"2022-03-03T14:38:26.348641Z"},"papermill":{"duration":3.835219,"end_time":"2022-03-03T14:46:04.760908","exception":false,"start_time":"2022-03-03T14:46:00.925689","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configuration","metadata":{"papermill":{"duration":0.042605,"end_time":"2022-03-03T14:46:04.847392","exception":false,"start_time":"2022-03-03T14:46:04.804787","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"This section defines some basic configuration that will be used by the rest of the notebook.","metadata":{"papermill":{"duration":0.042675,"end_time":"2022-03-03T14:46:04.933778","exception":false,"start_time":"2022-03-03T14:46:04.891103","status":"completed"},"tags":[]}},{"cell_type":"code","source":"IMG_SIZE = 192\nEPOCHS = 12\nBATCH_SIZE = 16 * strategy.num_replicas_in_sync\nAUTO = tf.data.experimental.AUTOTUNE","metadata":{"execution":{"iopub.execute_input":"2022-03-03T14:46:05.026312Z","iopub.status.busy":"2022-03-03T14:46:05.025248Z","iopub.status.idle":"2022-03-03T14:46:05.027573Z","shell.execute_reply":"2022-03-03T14:46:05.028198Z","shell.execute_reply.started":"2022-03-03T14:38:29.463442Z"},"papermill":{"duration":0.051735,"end_time":"2022-03-03T14:46:05.028371","exception":false,"start_time":"2022-03-03T14:46:04.976636","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BASE_GCS_PATH = GCS_DS_PATH\n\ngcs_fmt = os.path.join(BASE_GCS_PATH, \"tfrecords-jpeg-{}x{}\", \"\")\n\nGCS_PATHS = {\n    192: gcs_fmt.format(192, 192),\n    224: gcs_fmt.format(224, 224),\n    331: gcs_fmt.format(331, 331),\n    512: gcs_fmt.format(512, 512),\n}\nDATA_DIR = GCS_PATHS[IMG_SIZE]","metadata":{"execution":{"iopub.execute_input":"2022-03-03T14:46:05.118163Z","iopub.status.busy":"2022-03-03T14:46:05.117188Z","iopub.status.idle":"2022-03-03T14:46:05.122663Z","shell.execute_reply":"2022-03-03T14:46:05.123244Z","shell.execute_reply.started":"2022-03-03T14:38:29.471080Z"},"papermill":{"duration":0.052209,"end_time":"2022-03-03T14:46:05.123411","exception":false,"start_time":"2022-03-03T14:46:05.071202","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(\"../input/flower-classification-labels/flower_classification_labels.csv\")\nCLASSES = df[\"class\"].tolist()\nCLASSES","metadata":{"_kg_hide-output":true,"papermill":{"duration":0.076699,"end_time":"2022-03-03T14:46:05.243406","exception":false,"start_time":"2022-03-03T14:46:05.166707","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-06-05T16:20:06.900860Z","iopub.execute_input":"2022-06-05T16:20:06.901824Z","iopub.status.idle":"2022-06-05T16:20:06.931119Z","shell.execute_reply.started":"2022-06-05T16:20:06.901783Z","shell.execute_reply":"2022-06-05T16:20:06.930077Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Augmentation","metadata":{"papermill":{"duration":0.043194,"end_time":"2022-03-03T14:46:05.330248","exception":false,"start_time":"2022-03-03T14:46:05.287054","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def data_augment(image, label):\n    image = tf.image.random_flip_left_right(image)\n    image = tf.image.random_brightness(image, 0.2)\n    image = tf.image.random_contrast(image, 0.5, 2.0)\n    return image, label","metadata":{"execution":{"iopub.execute_input":"2022-03-03T14:46:05.422634Z","iopub.status.busy":"2022-03-03T14:46:05.421636Z","iopub.status.idle":"2022-03-03T14:46:05.427767Z","shell.execute_reply":"2022-03-03T14:46:05.427059Z","shell.execute_reply.started":"2022-03-03T14:38:29.514251Z"},"papermill":{"duration":0.052547,"end_time":"2022-03-03T14:46:05.427908","exception":false,"start_time":"2022-03-03T14:46:05.375361","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset Functions","metadata":{"papermill":{"duration":0.042907,"end_time":"2022-03-03T14:46:05.514559","exception":false,"start_time":"2022-03-03T14:46:05.471652","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"This section has code which loads the train, validation and test datasets.","metadata":{"papermill":{"duration":0.043058,"end_time":"2022-03-03T14:46:05.600887","exception":false,"start_time":"2022-03-03T14:46:05.557829","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Decode a JPEG image into a unit8 Tensor \ndef decode_image(image: tf.Tensor, channels: int = 3) -> tf.Tensor:\n    img = tf.image.decode_jpeg(image, channels=channels)\n    img = tf.reshape(img, [IMG_SIZE, IMG_SIZE, 3])\n    return img","metadata":{"execution":{"iopub.execute_input":"2022-03-03T14:46:05.692740Z","iopub.status.busy":"2022-03-03T14:46:05.692071Z","iopub.status.idle":"2022-03-03T14:46:05.693715Z","shell.execute_reply":"2022-03-03T14:46:05.694278Z","shell.execute_reply.started":"2022-03-03T14:38:29.524520Z"},"papermill":{"duration":0.050389,"end_time":"2022-03-03T14:46:05.694437","exception":false,"start_time":"2022-03-03T14:46:05.644048","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Read a TFRecord, extracing the image and either the label or the ID\ndef read_tfrecord(example: tf.Tensor, has_labels: bool = True) -> tuple[tf.Tensor, tf.Tensor]:\n    tfrecord_format = {\"image\": tf.io.FixedLenFeature([], tf.string)}\n\n    if has_labels is True:\n        key = \"class\"\n        tfrecord_format[\"class\"] = tf.io.FixedLenFeature([], tf.int64)\n    else:\n        key = \"id\"\n        tfrecord_format[\"id\"] = tf.io.FixedLenFeature([], tf.string)\n\n    example = tf.io.parse_single_example(example, tfrecord_format)\n    image = decode_image(example[\"image\"])\n    value = example[key]\n    return image, value","metadata":{"execution":{"iopub.execute_input":"2022-03-03T14:46:05.787766Z","iopub.status.busy":"2022-03-03T14:46:05.787056Z","iopub.status.idle":"2022-03-03T14:46:05.793246Z","shell.execute_reply":"2022-03-03T14:46:05.793810Z","shell.execute_reply.started":"2022-03-03T14:38:29.538045Z"},"papermill":{"duration":0.05548,"end_time":"2022-03-03T14:46:05.793967","exception":false,"start_time":"2022-03-03T14:46:05.738487","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Use the list of provided filepaths for TFRecords and build a dataset out of it\ndef get_dataset(filepaths: list[str], has_labels: bool = True, ordered: int = False) -> tf.data.Dataset:\n    options = tf.data.Options()\n    if ordered is False:\n        options.experimental_deterministic = False\n\n    dataset = tf.data.TFRecordDataset(filepaths, num_parallel_reads=AUTO)\n    dataset = dataset.with_options(options)\n\n    reader = functools.partial(read_tfrecord, has_labels=has_labels)\n    dataset = dataset.map(reader, num_parallel_calls=AUTO)\n    return dataset","metadata":{"execution":{"iopub.execute_input":"2022-03-03T14:46:05.884254Z","iopub.status.busy":"2022-03-03T14:46:05.883556Z","iopub.status.idle":"2022-03-03T14:46:05.888829Z","shell.execute_reply":"2022-03-03T14:46:05.889297Z","shell.execute_reply.started":"2022-03-03T14:38:29.551926Z"},"papermill":{"duration":0.052372,"end_time":"2022-03-03T14:46:05.889471","exception":false,"start_time":"2022-03-03T14:46:05.837099","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Make a dataset from the TFRecords stored in the given directory\ndef load_from_dir(\n    directory: str,\n    has_labels: bool = True,\n    ordered: bool = False,\n    repeat: bool = False,\n    cache: bool = False,\n    shuffle: bool = False,\n    augment: bool = False\n) -> tuple[tf.data.Dataset, filenames]:\n    path = os.path.join(DATA_DIR, directory, \"*.tfrec\")\n    filepaths = tf.io.gfile.glob(path)\n    \n    dataset = get_dataset(filepaths, has_labels=has_labels, ordered=ordered)\n    \n    if augment is True:\n        dataset = dataset.map(data_augment, num_parallel_calls=AUTO)\n    \n    if repeat is True:\n        dataset = dataset.repeat()\n    \n    if shuffle is True:\n        dataset = dataset.shuffle(2048)\n        \n    dataset = dataset.batch(BATCH_SIZE)\n    \n    if cache is True:\n        dataset = dataset.cache()\n        \n    dataset = dataset.prefetch(AUTO)\n    return dataset","metadata":{"execution":{"iopub.execute_input":"2022-03-03T14:46:05.981229Z","iopub.status.busy":"2022-03-03T14:46:05.980543Z","iopub.status.idle":"2022-03-03T14:46:05.988315Z","shell.execute_reply":"2022-03-03T14:46:05.988838Z","shell.execute_reply.started":"2022-03-03T14:38:29.568806Z"},"papermill":{"duration":0.055747,"end_time":"2022-03-03T14:46:05.989032","exception":false,"start_time":"2022-03-03T14:46:05.933285","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def n_samples(directory):\n    path = os.path.join(DATA_DIR, directory, \"*.tfrec\")\n    filepaths = tf.io.gfile.glob(path)\n    tot = 0\n    for filepath in filepaths:\n        basename = os.path.basename(filepath)\n        filename, _ = os.path.splitext(basename)\n        tot += int(filename.split(\"-\")[-1])\n    return tot","metadata":{"execution":{"iopub.execute_input":"2022-03-03T14:46:06.080971Z","iopub.status.busy":"2022-03-03T14:46:06.080334Z","iopub.status.idle":"2022-03-03T14:46:06.085372Z","shell.execute_reply":"2022-03-03T14:46:06.085815Z","shell.execute_reply.started":"2022-03-03T14:38:29.583593Z"},"papermill":{"duration":0.052726,"end_time":"2022-03-03T14:46:06.085975","exception":false,"start_time":"2022-03-03T14:46:06.033249","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Plotting functions","metadata":{"_kg_hide-input":false,"papermill":{"duration":0.042681,"end_time":"2022-03-03T14:46:06.171968","exception":false,"start_time":"2022-03-03T14:46:06.129287","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Convert a batch of images to NumPy\ndef batch_to_numpy(\n    batch: tuple[tf.Tensor, tf.Tensor],\n    has_labels=False\n) -> tuple[np.ndarray, np.ndarray]:\n    if has_labels is False:\n        images, _ = batch\n        return images.numpy(), None\n    \n    images, labels = batch\n    return images.numpy(), labels.numpy()","metadata":{"_kg_hide-input":true,"execution":{"iopub.execute_input":"2022-03-03T14:46:06.262788Z","iopub.status.busy":"2022-03-03T14:46:06.261806Z","iopub.status.idle":"2022-03-03T14:46:06.267438Z","shell.execute_reply":"2022-03-03T14:46:06.268032Z","shell.execute_reply.started":"2022-03-03T14:38:29.605298Z"},"papermill":{"duration":0.052324,"end_time":"2022-03-03T14:46:06.268194","exception":false,"start_time":"2022-03-03T14:46:06.215870","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Get the text to display on top of each image\n# With its size and color\ndef get_title(prediction: int, label: int):\n    c = [\"red\", \"black\"]\n    \n    # If test data with no predictions, no text\n    if prediction is None and label is None:\n        return '', c[True]\n    \n    # If test data but with predictions, return predicted label\n    if label is None:\n        return CLASSES[prediction], c[True] \n    \n    actual = CLASSES[label]\n    \n    # If train/validation data with prediction,\n    # Display only the label if correct prediction\n    # Otherwise, display the prediction with the correct label\n    if prediction is not None:\n        correct = prediction == label\n        title = f\"p: {CLASSES[prediction]}\\na: {actual}\"\n        return title, c[correct]\n    \n    # If only label, return as is\n    return f\"{actual}\", c[True]","metadata":{"_kg_hide-input":true,"execution":{"iopub.execute_input":"2022-03-03T14:46:06.364359Z","iopub.status.busy":"2022-03-03T14:46:06.363309Z","iopub.status.idle":"2022-03-03T14:46:06.365471Z","shell.execute_reply":"2022-03-03T14:46:06.365925Z","shell.execute_reply.started":"2022-03-03T14:38:29.619076Z"},"papermill":{"duration":0.054357,"end_time":"2022-03-03T14:46:06.366112","exception":false,"start_time":"2022-03-03T14:46:06.311755","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Make a grid of the given images\ndef plot_grid(\n    images: np.ndarray,\n    labels: np.ndarray | list[None],\n    predictions: tf.Tensor | list[None],\n    spacing: float = 0.1\n) -> None:\n    n_images = len(images)\n    \n    # Make a square grid by taking square root\n    rows = int(math.sqrt(n_images))\n    cols = n_images // rows\n    tot = rows * cols\n\n    # Some parameters borrowed from the Getting Started Notebook.\n    size = 13\n    fontdict = {\"verticalalignment\": \"center\"}\n\n    figsize = (size, size / tot) if rows < cols else (size / tot, size)\n    plt.figure(figsize=figsize)\n\n    # Make a subplot\n    fig, axs = plt.subplots(\n        rows,\n        cols,\n        figsize=(size, size),\n        constrained_layout=True,\n        gridspec_kw={\"wspace\": spacing, \"hspace\": spacing}\n    )\n    plt.axis(\"off\")\n    axs = axs.flatten()\n\n    # Go over each image, label, prediction\n    # And add to subplot\n    zipped = zip(images[:tot], labels[:tot], predictions[:tot], axs)\n    for image, label, prediction, ax in zipped:\n        fontsize = size * spacing / max(rows, cols) * 40 + 3\n        title, color = get_title(prediction, label)\n        ax.imshow(image)\n        ax.set_title(title, fontsize=fontsize, color=color, fontdict=fontdict, pad=fontsize / 1.5)\n        ax.set_axis_off()","metadata":{"_kg_hide-input":true,"execution":{"iopub.execute_input":"2022-03-03T14:46:06.462221Z","iopub.status.busy":"2022-03-03T14:46:06.461511Z","iopub.status.idle":"2022-03-03T14:46:06.466148Z","shell.execute_reply":"2022-03-03T14:46:06.466750Z","shell.execute_reply.started":"2022-03-03T14:38:29.633063Z"},"papermill":{"duration":0.057077,"end_time":"2022-03-03T14:46:06.466918","exception":false,"start_time":"2022-03-03T14:46:06.409841","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plot a batch of images\ndef plot_batch(\n    batch: tuple[tf.Tensor, tf.Tensor],\n    has_labels: bool = True,\n    predictions: tf.Tensor = None\n) -> None:\n    # Convert to Numpy\n    images, labels = batch_to_numpy(batch, has_labels=has_labels)\n\n    n_images = len(images)\n\n    # Fill labels and predictions with None if required\n    labels_ = labels if labels is not None else [None for _ in range(n_images)]\n    predictions_ = (\n        predictions if predictions is not None else [None for _ in range(n_images)]\n    )\n    spacing = 0.1\n    # Plot the images with labels and predictions in a grid\n    plot_grid(images, labels_, predictions_, spacing=spacing)\n    \n    # Handle whitespace\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()","metadata":{"_kg_hide-input":true,"execution":{"iopub.execute_input":"2022-03-03T14:46:06.563124Z","iopub.status.busy":"2022-03-03T14:46:06.562459Z","iopub.status.idle":"2022-03-03T14:46:06.565514Z","shell.execute_reply":"2022-03-03T14:46:06.564886Z","shell.execute_reply.started":"2022-03-03T14:38:29.648918Z"},"papermill":{"duration":0.054852,"end_time":"2022-03-03T14:46:06.565649","exception":false,"start_time":"2022-03-03T14:46:06.510797","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inspect Datasets","metadata":{"papermill":{"duration":0.04293,"end_time":"2022-03-03T14:46:06.651691","exception":false,"start_time":"2022-03-03T14:46:06.608761","status":"completed"},"tags":[]}},{"cell_type":"code","source":"train_data = load_from_dir(\"train\", repeat=True, shuffle=True, augment=True)\nn_train = n_samples(\"train\")\nval_data = load_from_dir(\"val\", cache=True)\nn_val = n_samples(\"val\")","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.execute_input":"2022-03-03T14:46:06.742366Z","iopub.status.busy":"2022-03-03T14:46:06.741659Z","iopub.status.idle":"2022-03-03T14:46:07.407950Z","shell.execute_reply":"2022-03-03T14:46:07.406907Z","shell.execute_reply.started":"2022-03-03T14:38:29.665184Z"},"papermill":{"duration":0.713754,"end_time":"2022-03-03T14:46:07.408108","exception":false,"start_time":"2022-03-03T14:46:06.694354","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"Number of training samples = {n_train}\")\nprint(\"Training data shape and sample labels:\")\nfor image, label in train_data.take(1):\n    print(image.shape, label)\n\nprint(f\"Number of validation samples = {n_val}\")\nprint(\"Validation data shape and sample labels:\")\nfor image, label in val_data.take(1):\n    print(image.shape, label)","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.execute_input":"2022-03-03T14:46:07.503110Z","iopub.status.busy":"2022-03-03T14:46:07.502229Z","iopub.status.idle":"2022-03-03T14:46:11.396196Z","shell.execute_reply":"2022-03-03T14:46:11.395532Z","shell.execute_reply.started":"2022-03-03T14:38:30.463418Z"},"papermill":{"duration":3.94441,"end_time":"2022-03-03T14:46:11.396339","exception":false,"start_time":"2022-03-03T14:46:07.451929","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"itrain = iter(train_data.unbatch().batch(20))","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.execute_input":"2022-03-03T14:46:11.494799Z","iopub.status.busy":"2022-03-03T14:46:11.494135Z","iopub.status.idle":"2022-03-03T14:46:11.507810Z","shell.execute_reply":"2022-03-03T14:46:11.507149Z","shell.execute_reply.started":"2022-03-03T14:38:31.523012Z"},"papermill":{"duration":0.065806,"end_time":"2022-03-03T14:46:11.507957","exception":false,"start_time":"2022-03-03T14:46:11.442151","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_batch(next(itrain))","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.execute_input":"2022-03-03T14:46:11.711476Z","iopub.status.busy":"2022-03-03T14:46:11.710747Z","iopub.status.idle":"2022-03-03T14:46:13.195892Z","shell.execute_reply":"2022-03-03T14:46:13.196426Z","shell.execute_reply.started":"2022-03-03T14:38:31.548052Z"},"papermill":{"duration":1.64396,"end_time":"2022-03-03T14:46:13.196599","exception":false,"start_time":"2022-03-03T14:46:11.552639","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ival = iter(val_data.unbatch().batch(20))","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.execute_input":"2022-03-03T14:46:13.343032Z","iopub.status.busy":"2022-03-03T14:46:13.342416Z","iopub.status.idle":"2022-03-03T14:46:13.355388Z","shell.execute_reply":"2022-03-03T14:46:13.354855Z","shell.execute_reply.started":"2022-03-03T14:38:33.787944Z"},"papermill":{"duration":0.087429,"end_time":"2022-03-03T14:46:13.355563","exception":false,"start_time":"2022-03-03T14:46:13.268134","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_batch(next(ival))","metadata":{"_kg_hide-output":true,"execution":{"iopub.execute_input":"2022-03-03T14:46:13.504024Z","iopub.status.busy":"2022-03-03T14:46:13.503415Z","iopub.status.idle":"2022-03-03T14:46:15.224477Z","shell.execute_reply":"2022-03-03T14:46:15.224995Z","shell.execute_reply.started":"2022-03-03T14:38:33.811803Z"},"papermill":{"duration":1.797571,"end_time":"2022-03-03T14:46:15.225192","exception":false,"start_time":"2022-03-03T14:46:13.427621","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Functions For Building The Model","metadata":{"papermill":{"duration":0.100804,"end_time":"2022-03-03T14:46:15.427340","exception":false,"start_time":"2022-03-03T14:46:15.326536","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"Transfer learning in the form of fine-tuning is used to adapt a pretrained model for this task. The model is initialized with weights for the ImageNet dataset and training is turned on for it. A `GlobalAveragePooling2D` layer is applied to its output. Optionally, additional `Dense` layers can also added. The output is a `Dense` layer with as many units as the number of classes and a softmax activation.\n\nThe pretrained model can either be set to VGG16, VGG19, Xception or ResNet50 by adding the key `core_model` in `params` with an appropriate value. The default is VGG16. The optional `Dense` layers are added by adding the key `dense_out_features` with a list of integers, each integer being the number of units in the layer. As many layers as the length of the list will be added to the model, in addition to the final output layer.","metadata":{"papermill":{"duration":0.099269,"end_time":"2022-03-03T14:46:15.626490","exception":false,"start_time":"2022-03-03T14:46:15.527221","status":"completed"},"tags":[]}},{"cell_type":"code","source":"core_model_map = {\n    \"vgg16\": [\n        applications.vgg16.preprocess_input,\n        applications.VGG16,\n    ],\n    \"xception\": [\n        applications.xception.preprocess_input,\n        applications.Xception,\n    ],\n    \"vgg19\": [\n        applications.vgg19.preprocess_input,\n        applications.VGG19,\n    ],\n    \"resnet50\": [\n        applications.resnet50.preprocess_input,\n        applications.ResNet50,\n    ]\n}","metadata":{"execution":{"iopub.execute_input":"2022-03-03T14:46:15.838480Z","iopub.status.busy":"2022-03-03T14:46:15.837833Z","iopub.status.idle":"2022-03-03T14:46:15.842543Z","shell.execute_reply":"2022-03-03T14:46:15.843036Z","shell.execute_reply.started":"2022-03-03T14:38:35.384784Z"},"papermill":{"duration":0.11626,"end_time":"2022-03-03T14:46:15.843231","exception":false,"start_time":"2022-03-03T14:46:15.726971","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_model_params(params):\n    dense_out_features = []\n    \n    for param, value in params.items():\n        if \"dense\" in param:\n            dense_out_features.append(value)\n    \n    return {\n        \"core_model\": params.get(\"core_model\", \"vgg16\"),\n        \"dense_out_features\": dense_out_features\n    }","metadata":{"execution":{"iopub.execute_input":"2022-03-03T14:46:16.046704Z","iopub.status.busy":"2022-03-03T14:46:16.046005Z","iopub.status.idle":"2022-03-03T14:46:16.050666Z","shell.execute_reply":"2022-03-03T14:46:16.051235Z","shell.execute_reply.started":"2022-03-03T14:38:35.392971Z"},"papermill":{"duration":0.108315,"end_time":"2022-03-03T14:46:16.051430","exception":false,"start_time":"2022-03-03T14:46:15.943115","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def make_model(params):\n    model_params = get_model_params(params)\n    \n    core_model = params.get(\"core_model\", \"vgg16\")\n    dense_out_features = params.get(\"dense_out_features\", [])\n    \n    shape = [IMG_SIZE, IMG_SIZE, 3]\n    \n    with strategy.scope():\n        preproces, core = core_model_map[core_model]\n        \n        ip = layers.Lambda(lambda data: preproces(tf.cast(data, tf.float32)), input_shape=shape)\n        core = core(weights=\"imagenet\", include_top=False)\n        \n        dense = [layers.Dense(features, activation=\"relu\") for features in dense_out_features]\n        \n        model = Sequential(\n            [\n                ip,\n                core,\n                layers.GlobalAveragePooling2D(),\n                *dense,\n                layers.Dense(len(CLASSES), activation=\"softmax\"),\n            ]\n        )\n        \n        optimizer = optimizers.Adam(learning_rate=params.get(\"lr\", 1e-3))\n        loss = \"sparse_categorical_crossentropy\"\n        metric = \"sparse_categorical_accuracy\"\n        model.compile(optimizer=optimizer, loss=loss, metrics=[metric], steps_per_execution=16)\n        \n        return model","metadata":{"execution":{"iopub.execute_input":"2022-03-03T14:46:16.254240Z","iopub.status.busy":"2022-03-03T14:46:16.253547Z","iopub.status.idle":"2022-03-03T14:46:16.262455Z","shell.execute_reply":"2022-03-03T14:46:16.262966Z","shell.execute_reply.started":"2022-03-03T14:38:35.409841Z"},"papermill":{"duration":0.112042,"end_time":"2022-03-03T14:46:16.263151","exception":false,"start_time":"2022-03-03T14:46:16.151109","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training Function","metadata":{"papermill":{"duration":0.100493,"end_time":"2022-03-03T14:46:16.463478","exception":false,"start_time":"2022-03-03T14:46:16.362985","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"The training logic is encapsulated in the function below. It returns the validation loss after the final epoch, the trained model and the training history at the end of training.","metadata":{"papermill":{"duration":0.099696,"end_time":"2022-03-03T14:46:16.664348","exception":false,"start_time":"2022-03-03T14:46:16.564652","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def train(params):\n    train_data = load_from_dir(\"train\", repeat=True, shuffle=True, augment=True)\n    n_train = n_samples(\"train\")\n\n    val_data = load_from_dir(\"val\", cache=True)\n    n_val = n_samples(\"val\")\n    \n    model = make_model(params)    \n\n    steps_per_epoch = n_train // BATCH_SIZE\n    validation_steps = -(-n_val // BATCH_SIZE)\n    \n    early_stopping = callbacks.EarlyStopping(patience=5)\n    pruning_callback = params.get(\"pruning_callback\", [])\n    \n    epochs = params.get(\"epochs\", EPOCHS)\n    \n    history = model.fit(\n        train_data,\n        steps_per_epoch=steps_per_epoch,\n        epochs=epochs,\n        validation_data=val_data,\n        validation_steps=validation_steps,\n        callbacks=[early_stopping, *pruning_callback]\n    )\n    \n    val_loss = history.history[\"val_loss\"]\n\n    return val_loss[-1], model, history\n        ","metadata":{"execution":{"iopub.execute_input":"2022-03-03T14:46:16.870122Z","iopub.status.busy":"2022-03-03T14:46:16.869114Z","iopub.status.idle":"2022-03-03T14:46:16.878078Z","shell.execute_reply":"2022-03-03T14:46:16.877528Z","shell.execute_reply.started":"2022-03-03T14:38:35.430224Z"},"papermill":{"duration":0.11269,"end_time":"2022-03-03T14:46:16.878233","exception":false,"start_time":"2022-03-03T14:46:16.765543","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Optuna Objective\n\nThe hyperparameters are tuned using Optuna. It is used to select the architecture that will be used for transfer learning, the learning rate, the number and sizes of the additional `Dense` layers to be added and the number of epochs. Additionally, Optuna is penalized whenever it chooses parameters that lead to early stopping since early stopping suggests that there is overfitting.","metadata":{"papermill":{"duration":0.100435,"end_time":"2022-03-03T14:46:17.080547","exception":false,"start_time":"2022-03-03T14:46:16.980112","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def objective(trial):\n    params = {\n        \"core_model\": trial.suggest_categorical(\"core_model\", [\"vgg16\", \"xception\", \"vgg19\", \"resnet50\"]),\n        \"lr\": trial.suggest_float(\"lr\", 1e-5, 4e-4),\n        \"epochs\": trial.suggest_int(\"epochs\", 8, 20),\n    }\n    \n    n_dense = trial.suggest_categorical(\"n_dense\", [1, 2, 3, 4, 5])\n    for i in range(n_dense):\n        key = f\"dense{i + 1}_out\"\n        params[f\"dense{i + 1}_out\"] = trial.suggest_int(key, 32, 1024)\n        \n    params[\"pruning_callback\"] = [optuna.integration.TFKerasPruningCallback(trial, \"val_loss\")]\n    \n    score, _, hist = train(params)\n    \n    val_loss = hist.history[\"val_loss\"]\n    \n    # Penalize tuner for being too aggresive\n    if len(val_loss) < params[\"epochs\"]:\n        return np.max(val_loss)\n\n    return score","metadata":{"execution":{"iopub.execute_input":"2022-03-03T14:46:17.289983Z","iopub.status.busy":"2022-03-03T14:46:17.288974Z","iopub.status.idle":"2022-03-03T14:46:17.291345Z","shell.execute_reply":"2022-03-03T14:46:17.291855Z","shell.execute_reply.started":"2022-03-03T14:38:35.444256Z"},"papermill":{"duration":0.111009,"end_time":"2022-03-03T14:46:17.292011","exception":false,"start_time":"2022-03-03T14:46:17.181002","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{"papermill":{"duration":0.099407,"end_time":"2022-03-03T14:46:17.490975","exception":false,"start_time":"2022-03-03T14:46:17.391568","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# Tune Parameters\n\nThe Hyperband pruning technique is a popular and SOTA pruning technique for stopping unpromising trial mid-way.","metadata":{"papermill":{"duration":0.099169,"end_time":"2022-03-03T14:46:17.689856","exception":false,"start_time":"2022-03-03T14:46:17.590687","status":"completed"},"tags":[]}},{"cell_type":"code","source":"pruner = optuna.pruners.HyperbandPruner()\nsampler = optuna.samplers.TPESampler(42, multivariate=True)\nstudy = optuna.create_study(\n    direction=\"minimize\",\n    pruner=pruner,\n    sampler=sampler\n)\nstudy.optimize(objective, n_trials=50, gc_after_trial=True)","metadata":{"_kg_hide-output":true,"execution":{"iopub.execute_input":"2022-03-03T14:46:17.893736Z","iopub.status.busy":"2022-03-03T14:46:17.892777Z","iopub.status.idle":"2022-03-03T16:02:39.486886Z","shell.execute_reply":"2022-03-03T16:02:39.486318Z"},"papermill":{"duration":4581.696764,"end_time":"2022-03-03T16:02:39.487069","exception":false,"start_time":"2022-03-03T14:46:17.790305","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train Final Model Using Best Parameters","metadata":{"papermill":{"duration":0.993889,"end_time":"2022-03-03T16:02:41.490749","exception":false,"start_time":"2022-03-03T16:02:40.496860","status":"completed"},"tags":[]}},{"cell_type":"code","source":"study.best_trial.params","metadata":{"execution":{"iopub.execute_input":"2022-03-03T16:02:43.498816Z","iopub.status.busy":"2022-03-03T16:02:43.498144Z","iopub.status.idle":"2022-03-03T16:02:43.503761Z","shell.execute_reply":"2022-03-03T16:02:43.504248Z","shell.execute_reply.started":"2022-03-03T14:08:29.714019Z"},"papermill":{"duration":1.011546,"end_time":"2022-03-03T16:02:43.504410","exception":false,"start_time":"2022-03-03T16:02:42.492864","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"best_trial = study.best_trial\n_, model, history = train(best_trial.params)","metadata":{"execution":{"iopub.execute_input":"2022-03-03T16:02:45.569425Z","iopub.status.busy":"2022-03-03T16:02:45.568731Z","iopub.status.idle":"2022-03-03T16:05:23.468047Z","shell.execute_reply":"2022-03-03T16:05:23.466863Z","shell.execute_reply.started":"2022-03-03T14:08:32.064177Z"},"papermill":{"duration":158.96073,"end_time":"2022-03-03T16:05:23.468258","exception":false,"start_time":"2022-03-03T16:02:44.507528","status":"completed"},"tags":[],"_kg_hide-output":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss Curve","metadata":{"papermill":{"duration":1.053957,"end_time":"2022-03-03T16:05:25.587083","exception":false,"start_time":"2022-03-03T16:05:24.533126","status":"completed"},"tags":[]}},{"cell_type":"code","source":"train_loss = history.history[\"loss\"]\nval_loss = history.history[\"val_loss\"]\nplt.plot(train_loss, label=\"Training\")\nplt.plot(val_loss, label=\"Validation\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.legend()\nplt.show()","metadata":{"execution":{"iopub.execute_input":"2022-03-03T16:05:27.768129Z","iopub.status.busy":"2022-03-03T16:05:27.767359Z","iopub.status.idle":"2022-03-03T16:05:28.019397Z","shell.execute_reply":"2022-03-03T16:05:28.018386Z","shell.execute_reply.started":"2022-03-01T11:02:20.245584Z"},"papermill":{"duration":1.375456,"end_time":"2022-03-03T16:05:28.019551","exception":false,"start_time":"2022-03-03T16:05:26.644095","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Accuracy Curve","metadata":{"papermill":{"duration":1.052326,"end_time":"2022-03-03T16:05:30.139030","exception":false,"start_time":"2022-03-03T16:05:29.086704","status":"completed"},"tags":[]}},{"cell_type":"code","source":"train_loss = history.history[\"sparse_categorical_accuracy\"]\nval_loss = history.history[\"val_sparse_categorical_accuracy\"]\nplt.plot(train_loss, label=\"Training\")\nplt.plot(val_loss, label=\"Validation\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Accuracy\")\nplt.legend()\nplt.show()","metadata":{"execution":{"iopub.execute_input":"2022-03-03T16:05:32.273283Z","iopub.status.busy":"2022-03-03T16:05:32.272234Z","iopub.status.idle":"2022-03-03T16:05:32.500053Z","shell.execute_reply":"2022-03-03T16:05:32.499417Z","shell.execute_reply.started":"2022-03-01T11:02:23.281151Z"},"papermill":{"duration":1.2991,"end_time":"2022-03-03T16:05:32.500211","exception":false,"start_time":"2022-03-03T16:05:31.201111","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predictions","metadata":{"papermill":{"duration":1.05385,"end_time":"2022-03-03T16:05:34.646885","exception":false,"start_time":"2022-03-03T16:05:33.593035","status":"completed"},"tags":[]}},{"cell_type":"code","source":"test_data = load_from_dir(\"test\", has_labels=False, ordered=True)\nn_test = n_samples(\"test\")\ntest_steps = -(-n_test // BATCH_SIZE)","metadata":{"execution":{"iopub.execute_input":"2022-03-03T16:05:36.765688Z","iopub.status.busy":"2022-03-03T16:05:36.764957Z","iopub.status.idle":"2022-03-03T16:05:36.941218Z","shell.execute_reply":"2022-03-03T16:05:36.940583Z","shell.execute_reply.started":"2022-03-01T11:02:28.845683Z"},"papermill":{"duration":1.239388,"end_time":"2022-03-03T16:05:36.941378","exception":false,"start_time":"2022-03-03T16:05:35.701990","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Computing predictions...')\ntest_images = test_data.map(lambda image, idnum: image)\nprobabilities = model.predict(test_images, steps=test_steps)\npredictions = np.argmax(probabilities, axis=-1)\nprint(predictions)\n\nprint('Generating submission.csv file...')\ntest_ids = test_data.map(lambda image, idnum: idnum).unbatch()\ntest_ids = next(iter(test_ids.batch(n_test))).numpy().astype('U') # all in one batch\nnp.savetxt('submission.csv', np.rec.fromarrays([test_ids, predictions]), fmt=['%s', '%d'], delimiter=',', header='id,label', comments='')","metadata":{"execution":{"iopub.execute_input":"2022-03-03T16:05:39.102619Z","iopub.status.busy":"2022-03-03T16:05:39.101881Z","iopub.status.idle":"2022-03-03T16:06:12.184410Z","shell.execute_reply":"2022-03-03T16:06:12.184960Z","shell.execute_reply.started":"2022-03-01T11:03:10.117014Z"},"papermill":{"duration":34.144251,"end_time":"2022-03-03T16:06:12.185133","exception":false,"start_time":"2022-03-03T16:05:38.040882","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visual Validation","metadata":{"papermill":{"duration":1.092342,"end_time":"2022-03-03T16:06:14.334219","exception":false,"start_time":"2022-03-03T16:06:13.241877","status":"completed"},"tags":[]}},{"cell_type":"code","source":"val_data = load_from_dir(\"val\")\nbatches = val_data.unbatch().batch(20)\nibatches = iter(batches)","metadata":{"execution":{"iopub.execute_input":"2022-03-03T16:06:16.451559Z","iopub.status.busy":"2022-03-03T16:06:16.450879Z","iopub.status.idle":"2022-03-03T16:06:16.560482Z","shell.execute_reply":"2022-03-03T16:06:16.559914Z","shell.execute_reply.started":"2022-03-03T14:35:02.145303Z"},"papermill":{"duration":1.171521,"end_time":"2022-03-03T16:06:16.560625","exception":false,"start_time":"2022-03-03T16:06:15.389104","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images, labels = next(ibatches)\nprobabilities = model.predict(tf.cast(images, tf.float32))\npredictions = np.argmax(probabilities, axis=-1)\nplot_batch((images, labels), predictions=predictions)","metadata":{"execution":{"iopub.execute_input":"2022-03-03T16:06:18.702205Z","iopub.status.busy":"2022-03-03T16:06:18.701554Z","iopub.status.idle":"2022-03-03T16:06:37.833861Z","shell.execute_reply":"2022-03-03T16:06:37.834388Z","shell.execute_reply.started":"2022-03-03T14:35:08.114101Z"},"papermill":{"duration":20.185685,"end_time":"2022-03-03T16:06:37.834574","exception":false,"start_time":"2022-03-03T16:06:17.648889","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]}]}