{"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":"# Flower classification\nNOTE: This notebook is WIP - still actively working on finishing it locally. This is a preview.\n\nThe task here is to train a network to classify images of flowers into 104 species based on photos of their petals. This task is a part of a Kaggle Getting Started competition - more information can be found here: https://www.kaggle.com/competitions/tpu-getting-started/data.\nI will train a Vision Transformer to perform this task. Large parts of the data importing and pre-processing code has been adapted from the code by Novia Putris, https://www.kaggle.com/code/noviaps/petal-to-metal, while the Vision Transformer is based on the example in https://keras.io/examples/vision/image_classification_with_vision_transformer/. Afterwards, I will compare the performance (and training time) of the standard Vision Transfer to a Convolutional Vision Transformer.\n\n***\n## Importing libraries","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\nimport re\nimport pandas as pd\n\n# These imports will be used to pre-process the data\nfrom functools import partial\nfrom tensorflow import data, convert_to_tensor\nfrom tensorflow.image import random_brightness, random_flip_left_right, random_flip_up_down, decode_jpeg, extract_patches, resize\nfrom tensorflow.io import gfile, parse_single_example, FixedLenFeature\nfrom tensorflow import string as tf_string\nfrom tensorflow import int32 as tf_int32\nfrom tensorflow import int64 as tf_int64\nfrom tensorflow import float32 as tf_float\nfrom tensorflow import reshape as tf_reshape\nfrom tensorflow import cast as tf_cast\n\n# These imports will be used to build, train, and use the network\nfrom tensorflow import shape as tf_shape\nfrom tensorflow import range as tf_range\nfrom tensorflow.keras.layers import Layer, Dense, Embedding, Input, LayerNormalization, MultiHeadAttention, Add, Dropout, GlobalAveragePooling1D, GlobalMaxPool1D\nfrom tensorflow.keras.regularizers import L1L2, L1, L2\nfrom tensorflow.keras import Model\nfrom tensorflow.keras.losses import SparseCategoricalCrossentropy\nfrom tensorflow.keras.optimizers import Adam\nfrom tensorflow.keras.metrics import SparseCategoricalAccuracy\nfrom tensorflow.keras.callbacks import ModelCheckpoint, EarlyStopping, Callback\nfrom tensorflow.keras.utils import plot_model\nimport keras_tuner\nfrom os.path import normpath, isfile, abspath, dirname\nfrom os import listdir\nimport dill\nfrom time import time\nfrom cvt_tensorflow import CvT\nfrom sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay, classification_report","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Defining some network and pre-processing constants","metadata":{}},{"cell_type":"code","source":"BATCH_SIZE = 1\nIMAGE_SIZE = [512, 512]\nCHANNELS = 3\nIMAGE_DIMS = [512, 512, 3]\nCLASSES = ['pink primrose', 'hard-leaved pocket orchid', 'canterbury bells', 'sweet pea', 'wild geranium', 'tiger lily', 'moon orchid', 'bird of paradise', 'monkshood', 'globe thistle', 'snapdragon', \"colt's foot\", 'king protea', 'spear thistle', 'yellow iris', 'globe-flower', 'purple coneflower', 'peruvian lily', 'balloon flower', 'giant white arum lily', 'fire lily', 'pincushion flower', 'fritillary', 'red ginger', 'grape hyacinth', 'corn poppy', 'prince of wales feathers', 'stemless gentian', 'artichoke', 'sweet william', 'carnation', 'garden phlox', 'love in the mist', 'cosmos', 'alpine sea holly', 'ruby-lipped cattleya', 'cape flower', 'great masterwort', 'siam tulip', 'lenten rose', 'barberton daisy', 'daffodil', 'sword lily', 'poinsettia', 'bolero deep blue', 'wallflower', 'marigold', 'buttercup', 'daisy', 'common dandelion', 'petunia', 'wild pansy', 'primula', 'sunflower', 'lilac hibiscus', 'bishop of llandaff', 'gaura', 'geranium', 'orange dahlia', 'pink-yellow dahlia', 'cautleya spicata', 'japanese anemone', 'black-eyed susan', 'silverbush', 'californian poppy', 'osteospermum', 'spring crocus', 'iris', 'windflower', 'tree poppy', 'gazania', 'azalea', 'water lily', 'rose', 'thorn apple', 'morning glory', 'passion flower', 'lotus', 'toad lily', 'anthurium', 'frangipani', 'clematis', 'hibiscus', 'columbine', 'desert-rose', 'tree mallow', 'magnolia', 'cyclamen ', 'watercress', 'canna lily', 'hippeastrum ', 'bee balm', 'pink quill', 'foxglove', 'bougainvillea', 'camellia', 'mallow', 'mexican petunia', 'bromelia', 'blanket flower', 'trumpet creeper', 'blackberry lily', 'common tulip', 'wild rose']\nNUM_CLASSES = len(CLASSES)\n\n# Enable float16 mixed precision, to lower the computational complexity (heavily hardware-dependent)\n#mixed_precision.set_global_policy('mixed_float16')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"***\n## Preparing the data\nThe images are saved in a storage-efficient TFRec format. I will first unpack them into tf.data.Dataset objects. Note that I will only be using a few basic image augmentation procedures to artificially increase the size of the dataset - I do not want the network to learn incorrect representations of what each of the flowers look like.","metadata":{}},{"cell_type":"code","source":"data_dir = normpath(\"R:\\\\data\")\nsub_data = gfile.glob(str(normpath(f'{data_dir}\\\\test\\\\*.tfrec')))\ntrain_data = gfile.glob(str(normpath(f'{data_dir}\\\\train\\\\*.tfrec')))\nval_data = gfile.glob(str(normpath(f'{data_dir}\\\\val\\\\*.tfrec')))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def import_dataset(files, labeled: bool, apply_augmentation: bool=False):\n    \"\"\"Extract images from the TFRec files and add them into a tf.data Dataset, optionally applying augmentation to them.\n    Parameters:\n        files (tf.io.gfile.glob object): The list of filenames to be added into the dataset.\n        labeled (bool): Whether the data has a target label or not.\n        apply_augmentation (bool): Whether to apply augmentation (random horizontal and vertical flips, random brightness change within a delta of 20%) to copies of the images. Default: False\n    Returns:\n        dataset (tf.data.Datset object): The files as a tf.data Dataset.\"\"\"\n\n    def read_tfrecord(example, labeled: bool):\n        \"\"\"Extracts the images from each entry of the TFRec entry.\"\"\"\n        tfrecord_format = {\n            'image': FixedLenFeature([], tf_string),\n            'class': FixedLenFeature([], tf_int64)\n        } if labeled else {\n            'image': FixedLenFeature([], tf_string),\n            'id': FixedLenFeature([], tf_string)\n        }\n\n        example = parse_single_example(example, tfrecord_format)\n        image = decode_jpeg(example['image'], channels=CHANNELS)\n        image = tf_cast(image, tf_int32)\n        image = tf_reshape(image, [*IMAGE_SIZE, CHANNELS])\n        return image, example['class'] if labeled else example['id']\n\n\n    ignore_order = data.Options()\n    ignore_order.experimental_deterministic = False\n    dataset = data.TFRecordDataset(files, num_parallel_reads=-1)\n    dataset = dataset.with_options(ignore_order)\n    read_tfrecord = partial(read_tfrecord, labeled=labeled)\n    dataset = dataset.map(read_tfrecord, num_parallel_calls=-1)\n\n    if apply_augmentation:\n        # Random horizontal flips\n        dataset = dataset.map(lambda image, label: (random_flip_left_right(image), label), num_parallel_calls=-1).shuffle(2048)\n        # Random vertical flips\n        dataset = dataset.map(lambda image, label: (random_flip_up_down(image), label), num_parallel_calls=-1).shuffle(2048)\n        # Random brightness\n        #dataset = dataset.map(lambda image, label: (random_brightness(image, 0.2), label), num_parallel_calls=-1).shuffle(2048)\n\n    # Normalize\n    dataset = dataset.map(lambda image, label: (tf_cast(image, tf_float) / 255.0, label), num_parallel_calls=-1)\n\n    return dataset","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds_len = np.sum([int(re.compile(r'-([0-9]*)\\.').search(filename).group(1)) for filename in train_data])\ntrain_ds = import_dataset(train_data, apply_augmentation=True, labeled=True)\n\nval_ds_len = np.sum([int(re.compile(r'-([0-9]*)\\.').search(filename).group(1)) for filename in val_data])\nval_ds = import_dataset(val_data, apply_augmentation=False, labeled=True)\n\ntest_ds = val_ds.take(int(0.1 * val_ds_len)).batch(1).cache().prefetch(-1)  # Test set is 10% of validation set\nval_ds = val_ds.skip(int(0.1 * val_ds_len))\n\nsub_ds = import_dataset(sub_data, apply_augmentation=False, labeled=False).batch(1).cache().prefetch(-1)\n\n# Unzipping the test dataset\ntest_ds_labels = test_ds.map(lambda x, y: y, num_parallel_calls=-1)\ntest_ds = test_ds.map(lambda x, y: x, num_parallel_calls=-1)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def hp_opt_dataset(source_ds, size: int):\n    \"\"\"Extracts a subset of a given size from a source dataset for hyperparameter optimization. To be used by the Keras Tuner to produce a different selection of samples from the source dataset for each trial.\n    Recommended size: For training data, 20 to 25% of the source dataset; For validation data, 15 to 20% of the source dataset.\"\"\"\n    return source_ds.shuffle(2048).take(size).batch(BATCH_SIZE).cache().prefetch(-1)","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's look at the structure of the datasets.","metadata":{}},{"cell_type":"code","source":"np.set_printoptions(threshold=15, linewidth=80)\n\nprint(\"Train data shapes:\")\nfor image, idnum in train_ds.take(3):\n    print(image.numpy().shape, idnum.numpy().shape)\n    print(\"Train data IDs:\", idnum.numpy().astype('U')) # U=unicode string","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Sub data shapes:\")\nfor image, idnum in sub_ds.take(3):\n    print(image.numpy().shape, idnum.numpy().shape)\n    print(\"Sub data IDs:\", idnum.numpy().astype('U')) # U=unicode string","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's look at some examples from the training data.","metadata":{}},{"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: # binary string in this case,\n        # 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\n    # 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\n    # or square-ish rectangle\n    rows = int(np.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 # magic formula tested to work from 1x1 to 10x10 images\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()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_iter = iter(train_ds.batch(20))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"one_batch = next(ds_iter)\ndisplay_batch_of_images(one_batch)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"***\n## Building the model\n### The vision transformer\n#### Image patches\nFirst, I will need to define a Patch class, which will be used to separate each image into N x N smaller images, which will serve as the tokens fed into the network, and between which self-attention will be computed.","metadata":{}},{"cell_type":"code","source":"class Patches(Layer):\n    def __init__(self, patch_size):\n        super().__init__()\n        self.patch_size = patch_size\n\n    def call(self, images):\n        batch_size = tf_shape(images)[0]\n        patches = extract_patches(\n            images=images,\n            sizes=[1, self.patch_size, self.patch_size, 1],\n            strides=[1, self.patch_size, self.patch_size, 1],\n            rates=[1, 1, 1, 1],\n            padding=\"VALID\",\n        )\n        patch_dims = patches.shape[-1]\n        patches = tf_reshape(patches, [batch_size, -1, patch_dims])\n        return patches\n\n    def get_config(self):\n        config = super().get_config()\n        config.update({\n            'patch_size': self.patch_size\n        })\n        return config","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's look at an example of an image split into 32x32 patches.","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(8, 8))\nimage = next(iter(train_ds))[0].numpy() * 255.0\nplt.imshow(image.astype(\"uint8\"))\nplt.axis(\"off\")\n\nresized_image = resize(\n    convert_to_tensor([image]), size=IMAGE_SIZE)\n\npatches = Patches(32)(resized_image)\nprint(f\"Image size: {IMAGE_SIZE}\")\nprint(\"Patch size: 32 x 32\")\nprint(f\"Patches per image: {patches.shape[1]}\")\nprint(f\"Elements per patch: {patches.shape[-1]}\")\n\nn = int(np.sqrt(patches.shape[1]))\nplt.figure(figsize=(8, 8))\nfor i, patch in enumerate(patches[0]):\n    ax = plt.subplot(n, n, i + 1)\n    patch_img = tf_reshape(patch, (32, 32, 3))\n    plt.imshow(patch_img.numpy().astype(\"uint8\"))\n    plt.axis(\"off\")","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Patch encoder\nThis layer will be used to embed the patches into vectors, where patches related to each other will be geometrically closer in the embedding space, and it will also include the positions of where in an image each patch is located.","metadata":{}},{"cell_type":"code","source":"class PatchEncoder(Layer):\n    def __init__(self, num_patches, projection_dim):\n        super().__init__()\n        self.num_patches = num_patches\n        self.projection = Dense(units=projection_dim)\n        self.position_embedding = Embedding(\n            input_dim=num_patches, output_dim=projection_dim\n        )\n\n    def call(self, patch):\n        positions = tf_range(start=0, limit=self.num_patches, delta=1)\n        encoded = self.projection(patch) + self.position_embedding(positions)\n        return encoded\n\n    def get_config(self):\n        config = super().get_config()\n        config.update({\n            'num_patches': self.num_patches\n        })\n        return config","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Multi-layer perceptron\nDefining a function to generate a basic MLP network, which will be used within the vision transformer.","metadata":{}},{"cell_type":"code","source":"def mlp(x, hidden_units, dropout_rate: float, activation: str):\n    for units in hidden_units:\n        x = Dense(units, activation=activation)(x)\n        x = Dropout(dropout_rate)(x)\n    return x","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Generating the network\nI will now define a function that will generate the transformer network, where the hyperparameters will be fine-tuned using Keras Tuner.","metadata":{}},{"cell_type":"code","source":"def generate_vit(hp):\n    \"\"\"Builds the network model - to be used by the keras_tuner hyperparameter tuner.\n    Parameters:\n        hp (dict-like): A dictionary of tunable hyperparameters and their values for the current iteration - to be fed in by keras_tuner\n    Returns:\n        model (KerasTensor): A compiled network\"\"\"\n\n    patch_size = hp.Fixed('patch_size', 32)\n    num_patches = hp.Fixed('num_patches', (IMAGE_DIMS[0] // patch_size) ** 2)\n    projection_dim = hp.Choice('projection_dim', [512, 256, 128, 64, 32])\n    normalization_epsilon = hp.Choice('norm_epsilon', [1e-3, 1e-6])\n\n    inputs = Input(shape=IMAGE_DIMS)\n    patches = Patches(patch_size)(inputs)\n    encoded_patches = PatchEncoder(num_patches, projection_dim)(patches)\n\n    for i in range(hp.Int('num_layers', min_value=6, max_value=16, step=2)):\n        x1 = LayerNormalization(epsilon=normalization_epsilon)(encoded_patches)\n\n        attention_output = MultiHeadAttention(num_heads=hp.Int(f'num_heads{i}', min_value=3, max_value=6),\n                                              key_dim=projection_dim,\n                                              dropout=hp.Float(f'attention_dropout{i}', min_value=0.1, max_value=0.7, step=0.2),\n                                              kernel_regularizer=L1L2(l1=hp.Choice(f'l1_reg{i}',\n                                                                                   np.geomspace(1e-6, 1e-1, 6).tolist() + [0.0]),\n                                                                      l2=hp.Choice(f'l2_reg{i}',\n                                                                                   np.geomspace(1e-6, 1e-1, 6).tolist() + [0.0])))(x1, x1)\n\n        x2 = Add()([attention_output, encoded_patches])\n        x3 = LayerNormalization(epsilon=normalization_epsilon)(x2)\n        x3 = mlp(x3,\n                 hidden_units=[projection_dim * 2, projection_dim],\n                 dropout_rate=hp.Float(f'mlp_dropout{i}', min_value=0.1, max_value=0.7, step=0.2),\n                 activation=hp.Choice(f'mlp_act{i}', ['gelu', 'relu', 'leaky_relu']))\n        encoded_patches = Add()([x3, x2])\n\n    representation = LayerNormalization(epsilon=normalization_epsilon)(encoded_patches)\n    if hp.Boolean('pooling_type'):\n        representation = GlobalAveragePooling1D()(representation)\n    else:\n        representation = GlobalMaxPool1D()(representation)\n    representation = Dropout(hp.Float('representation_dropout', min_value=0.1, max_value=0.7, step=0.2))(representation)\n\n    mlp_head_units = hp.Choice('mlp_head_units', [4096, 2048, 1024, 512, 256])\n    features = mlp(representation,\n                   hidden_units=[mlp_head_units, mlp_head_units / 2],\n                   dropout_rate=hp.Float('mlp_head_dropout', min_value=0.1, max_value=0.7, step=0.2),\n                   activation=hp.Choice('mlp_head_act', ['gelu', 'relu', 'leaky_relu']))\n    logits = Dense(NUM_CLASSES)(features)\n    model = Model(inputs=inputs, outputs=logits)\n\n\n    model.compile(loss=SparseCategoricalCrossentropy(from_logits=True),\n                  metrics=[SparseCategoricalAccuracy(name='accuracy')],\n                  optimizer=Adam(learning_rate=hp.Float('lr', min_value=1e-6, max_value=1e-3, sampling='log'),\n                                 epsilon=hp.Choice('opt_epsilon', [1.0, 0.1, 1e-7]),\n                                 decay=hp.Float('wd', min_value=0.0001, max_value=0.1, sampling='log')))\n    return model","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tuner = keras_tuner.Hyperband(hypermodel=generate_vit, objective='val_accuracy', max_epochs=100, directory=normpath('E:\\\\flower_classification_tuner'))\ntuner.search_space_summary()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tuner.search(hp_opt_dataset(train_ds, int(0.2 * train_ds_len)), validation_data=hp_opt_dataset(val_ds, int(0.2 * int(0.9 * val_ds_len))), epochs=100, callbacks=[EarlyStopping(monitor='val_accuracy', patience=5, min_delta=0.001)])\ntuner.results_summary(1)","metadata":{"pycharm":{"is_executing":true}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"I'll save the best-performing model into the model variable and train it for more epochs.","metadata":{}},{"cell_type":"code","source":"model = generate_vit(tuner.get_best_hyperparameters()[0])\nmodel.summary()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"First, I'll plot the network structure.","metadata":{}},{"cell_type":"code","source":"plot_model(model, show_shapes=True, show_dtype=True)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"I will now train the network. I will be using the Keras EarlyStopping callback and stop the training if validation loss has stopped improving for 20 epochs, instead of necessarily having to reach a fixed 1000 epochs. I'll also define and use a custom callback that will keep a (approximate) track of the training time, and one that will log the training metrics history.","metadata":{}},{"cell_type":"code","source":"class TimeTracker(Callback):\n    def __init__(self, filepath, logs=None):\n        self.filepath = filepath\n        if isfile(filepath):\n            with open(filepath, 'rb') as file:\n                self.times = dill.load(file)\n        else:\n            self.times = []\n\n    def on_epoch_begin(self, epoch, logs=None):\n        self.epoch_time_start = time()\n\n    def on_epoch_end(self, epoch, logs=None):\n        self.times.append(time() - self.epoch_time_start)\n        with open(self.filepath, 'wb') as file:\n            dill.dump(self.times, file)","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HistoryLogger(Callback):\n    def __init__(self, filepath, logs=None):\n        self.filepath = filepath\n        if isfile(filepath):\n            with open(filepath, 'rb') as file:\n                self.history = dill.load(file)\n        else:\n            self.history = {}\n\n    def on_epoch_end(self, epoch, logs=None):\n        if epoch == 0:\n            for key in logs.keys():\n                self.history[key] = [logs[key][0]]\n        else:\n            for key in logs.keys():\n                self.history[key].append(logs[key][0])\n        with open(self.filepath, 'wb') as file:\n            dill.dump(self.history, file)","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds = train_ds.batch(BATCH_SIZE).cache().prefetch(-1)\nval_ds = val_ds.batch(BATCH_SIZE).cache().prefetch(-1)","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"callbacks = [\n    ModelCheckpoint(filepath=normpath('E:\\\\vit_checkpoints\\\\checkpoint_{epoch:03d}'), save_weights_only=True),\n    EarlyStopping(monitor='val_accuracy', patience=20, min_delta=0.0001),\n    TimeTracker(normpath(f'{abspath(dirname(__file__))}\\\\training_time_history.pkl')),\n    HistoryLogger(normpath(f'{abspath(dirname(__file__))}\\\\training_history.pkl'))]\n\nresults = model.fit(x=train_ds, validation_data=val_ds, callbacks=callbacks, epochs=1000)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Continue training from checkpoint if not finished\ndef get_init_epoch(path: str) -> int:\n    file_list = listdir(path)\n    file_list = [int(file.split('_')[-1].split('.')[0]) for file in file_list if file.startswith('checkpoint_')]\n    return max(file_list)\n\ninit_epoch = get_init_epoch(normpath(\"E:\\\\vit_checkpoints\"))\nmodel.load_weights(normpath(f'E:\\\\vit_checkpoints\\\\checkpoint_{init_epoch:03d}'))\nresults = model.fit(x=train_ds, validation_data=val_ds, callbacks=callbacks, epochs=1000, initial_epoch=init_epoch)","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"I will now plot the training results, and print the total training time.","metadata":{}},{"cell_type":"code","source":"with open(normpath('E:\\\\vit_checkpoints\\\\training_history.pkl'), 'rb') as file:\n    results = dill.load(file)","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"epochs = range(len(results['loss']))\nf, axs = plt.subplots(1, 2, figsize=(15,5))\n\nfor i, metric in enumerate(['loss', 'accuracy']):\n    axs[i].plot(epochs, results[metric], \"b\", label=f\"Training {metric}\")\n    axs[i].plot(epochs, results[f'val_{metric}'], \"g\", label=f\"Validation {metric}\")\n    axs[i].set_xlabel('Epochs')\n    axs[i].set_ylabel(metric)\n    axs[i].legend()\n    axs[i].grid()\n\nplt.title('Training and validation loss and accuracy')\nplt.tight_layout()\nplt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open(normpath(f'{abspath(dirname(__file__))}\\\\training_time_history.pkl'), 'rb') as file:\n    times = dill.load(file)\n    print(f'Total training time: {sum(times)}')","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"interpret\n\nI will now evaluate the model on the test set. First, I'll reload the weights of the best-performing epoch of the model.","metadata":{}},{"cell_type":"code","source":"model.load_weights(normpath(f'E:\\\\vit_checkpoints\\\\checkpoint_{np.argmin(results[\"val_accuracy\"]):03d}'))","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_pred = model.predict(test_ds)\ntest_pred = np.argmax(test_pred, axis=1)\ntest_pred","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Classification report of predicting the labels on the test set vs. true labels\npd.DataFrame(classification_report([x for x in test_ds_labels.as_numpy_iterator()], test_pred, target_names=CLASSES, zero_division=0, output_dict=True)).transpose()","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ConfusionMatrixDisplay(confusion_matrix([x for x in test_ds_labels.as_numpy_iterator()], test_pred, labels=range(len(CLASSES))), display_labels=CLASSES).plot()","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"interpret\n\n***\n## Using a pre-trained CvT model\nFor the sake of experiment, I will try using a Convolutional ViT model, which should primarily significantly decrease training time, but it would also be interesting to see how it affects the accuracy, especially when I'm not using a pre-trained model.","metadata":{}},{"cell_type":"code","source":"def generate_cvt(hp):\n    num_stages = hp.Int('num_stages', min_value=6, max_value=16, step=2)\n    starting_projection_dim_options = [32, 64, 128, 256, 512]\n    starting_projection_dim = hp.Int('starting_projection_dim', min_value=0, max_value=len(starting_projection_dim_options) - 1)\n    projection_dims = starting_projection_dim_options[starting_projection_dim:]\n    if len(projection_dims) < num_stages:\n        projection_dims = projection_dims + [projection_dims[-1]] * (num_stages - len(projection_dims))\n    elif len(projection_dims) > num_stages:\n        projection_dims = projection_dims[:num_stages]\n\n    model = CvT(in_chans=hp.Fixed('channels', CHANNELS),\n                num_classes=hp.Fixed('num_classes', NUM_CLASSES),\n                act_layer=hp.Choice('act_layer', ['gelu', 'relu', 'leaky_relu']),\n                classifier_activation=hp.Fixed('classifier_activation', None),\n                spec={\n                    'INIT': hp.Choice('init', ['trunc_norm', 'xavier']),\n                    'NUM_STAGES': num_stages,\n                    'PATCH_SIZE': [16] + ([8] * (num_stages - 1)),\n                    'PATCH_STRIDE': [8] + ([4] * (num_stages - 1)),\n                    'PATCH_PADDING': [2] + ([1] * (num_stages - 1)),\n                    'DIM_EMBED': projection_dims,\n                    'NUM_HEADS': [hp.Int(f'num_heads{i}', min_value=3, max_value=6) for i in range(num_stages)],\n                    'DEPTH': [hp.Int(f'depth{i}', min_value=1, max_value=13, step=3) for i in range(num_stages)],\n                    'MLP_RATIO': hp.Choice('mlp_ratio', [2.0, 4.0]),\n                    'ATTN_DROP_RATE': [hp.Float(f'attention_dropout{i}', min_value=0.1, max_value=0.7, step=0.2) for i in range(num_stages)],\n                    'DROP_RATE': [hp.Float(f'mlp_dropout{i}', min_value=0.1, max_value=0.7, step=0.2) for i in range(num_stages)],\n                    'DROP_PATH_RATE': [hp.Float(f'path_dropout{i}', min_value=0.1, max_value=0.7, step=0.2) for i in range(num_stages)],\n                    'QKV_BIAS': [hp.Choice(f'qkv_bias{i}', [True, False]) for i in range(num_stages)],\n                    'CLS_TOKEN': [hp.Choice(f'cls_token{i}', [True, False]) for i in range(num_stages)],\n                    'QKV_PROJ_METHOD': [hp.Choice(f'qkv_proj_method{i}', ['dw_bn', 'avg']) for i in range(num_stages)],\n                    'KERNEL_QKV': [3] * num_stages,\n                    'PADDING_KV': [1] * num_stages,\n                    'STRIDE_KV': [2] * num_stages,\n                    'PADDING_Q': [1] * num_stages,\n                    'STRIDE_Q': [1] * num_stages\n                })\n\n    model.compile(loss=SparseCategoricalCrossentropy(from_logits=True),\n                  metrics=[SparseCategoricalAccuracy(name='accuracy')],\n                  optimizer=Adam(learning_rate=hp.Float('lr', min_value=1e-6, max_value=1e-3, sampling='log'),\n                                 epsilon=hp.Choice('opt_epsilon', [1.0, 0.1, 1e-7]),\n                                 weight_decay=hp.Float('wd', min_value=0.0001, max_value=0.1, sampling='log')))\n    return model","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tuner2 = keras_tuner.Hyperband(hypermodel=generate_cvt, objective='val_accuracy', max_epochs=100, directory=normpath('E:\\\\flower_classification_cvt_tuner'))\ntuner2.search_space_summary()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tuner2.search(hp_opt_train_ds, validation_data=val_ds, epochs=100, callbacks=[EarlyStopping(monitor='val_accuracy', patience=5, min_delta=0.001)])\ntuner2.results_summary(1)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model2 = generate_cvt(tuner2.get_best_hyperparameters()[0])\nmodel2.summary()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"callbacks = [\n    ModelCheckpoint(filepath=normpath('E:\\\\vit_checkpoints_cvt\\\\checkpoint_{epoch:03d}'), save_weights_only=True),\n    EarlyStopping(monitor='val_accuracy', patience=20, min_delta=0.0001),\n    TimeTracker(normpath(f'{abspath(dirname(__file__))}\\\\training_time_history_cvt.pkl')),\n    HistoryLogger(normpath(f'{abspath(dirname(__file__))}\\\\training_history_cvt.pkl'))]\n\nresults2 = model2.fit(x=train_ds, validation_data=val_ds, callbacks=callbacks, epochs=1000, batch_size=BATCH_SIZE)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Continue training from checkpoint if not finished\ninit_epoch = get_init_epoch(normpath(\"E:\\\\vit_checkpoints_cvt\"))\nmodel2.load_weights(normpath(f'E:\\\\vit_checkpoints_cvt\\\\checkpoint_{init_epoch:03d}'))\nresults2 = model2.fit(x=train_ds, validation_data=val_ds, callbacks=callbacks, epochs=1000, initial_epoch=init_epoch)","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"I will now plot the training results, and print the total training time.","metadata":{}},{"cell_type":"code","source":"with open(normpath('E:\\\\vit_checkpoints_cvt\\\\training_history.pkl'), 'rb') as file:\n    results2 = dill.load(file)","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"epochs = range(len(results2['loss']))\nf, axs = plt.subplots(1, 2, figsize=(15,5))\n\nfor i, metric in enumerate(['loss', 'accuracy']):\n    axs[i].plot(epochs, results2[metric], \"b\", label=f\"Training {metric}\")\n    axs[i].plot(epochs, results2[f'val_{metric}'], \"g\", label=f\"Validation {metric}\")\n    axs[i].set_xlabel('Epochs')\n    axs[i].set_ylabel(metric)\n    axs[i].legend()\n    axs[i].grid()\n\nplt.title('Training and validation loss and accuracy')\nplt.tight_layout()\nplt.show()","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open(normpath(f'{abspath(dirname(__file__))}\\\\training_time_history_cvt.pkl'), 'rb') as file:\n    times = dill.load(file)\n    print(f'Total training time: {sum(times)}')","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"interpret\n\nI will now evaluate the model on the test set. First, I'll reload the weights of the best-performing epoch of the model.","metadata":{}},{"cell_type":"code","source":"model2.load_weights(normpath(f'E:\\\\vit_checkpoints_cvt\\\\checkpoint_{np.argmin(results[\"val_accuracy\"]):03d}'))","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_pred = model2.predict(test_ds)\ntest_pred = np.argmax(test_pred, axis=1)\ntest_pred","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Classification report of predicting the labels on the test set vs. true labels\npd.DataFrame(classification_report([x for x in test_ds_labels.as_numpy_iterator()], test_pred, target_names=CLASSES, zero_division=0, output_dict=True)).transpose()","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ConfusionMatrixDisplay(confusion_matrix([x for x in test_ds_labels.as_numpy_iterator()], test_pred, labels=range(len(CLASSES))), display_labels=CLASSES).plot()","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"interpret and compare","metadata":{}},{"cell_type":"markdown","source":"***\n### Making predictions\nI will now load the best-performing epoch of the best-performing model to make class predictions on the test data, and then write the results to the submission file.","metadata":{}},{"cell_type":"code","source":"argmax_1 = np.argmax(results['val_accuracy'])\nargmax_2 = np.argmax(results2['val_accuracy'])\nif results['val_accuracy'][argmax_1] >= results2['val_accuracy'][argmax_2]:\n    model = model.load_weights(normpath(f'E:\\\\vit_checkpoints\\\\checkpoint_{argmax_1:03d}'))\nelse:\n    model = model2\n    model.load_weights(normpath(f'E:\\\\vit_checkpoints_cvt\\\\checkpoint_{argmax_2:03d}'))","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_images_ds = sub_ds.map(lambda image, idnum: image)\nprobabilities = model.predict(test_images_ds)\npredictions = np.argmax(probabilities, axis=-1)","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_sub_images = np.sum([int(re.compile(r'-([0-9]*)\\.').search(filename).group(1)) for filename in sub_data])\n\nsub_ids_ds = sub_ds.map(lambda image, idnum: idnum).unbatch()\nsub_ids = next(iter(sub_ids_ds.batch(num_sub_images))).numpy().astype('U')\n\nnp.savetxt(\n    'submission.csv',\n    np.rec.fromarrays([sub_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":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]}]}