{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":21154,"databundleVersionId":1243559},{"sourceType":"datasetVersion","sourceId":1138814,"datasetId":601927,"databundleVersionId":1169433}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Flower Classification\n\nThis notebook explores flower classification on the [Petals to the Metal](https://www.kaggle.com/c/tpu-getting-started) Kaggle competition, using a Vision Transformer (ViT).\n\nHere is a summary of the lessons learned:\n* Ensemble heads don't help; a full-model ensemble might, but is likely too resource-intensive.\n* Sequentially unfreezing layers during fine-tuning improved performance.\n* A cosine decay learning rate schedule with warm-up yielded better fine-tuning results.\n* Data augmentation helped on the original dataset but appeared to confuse the model on extended data.\n* Transformers 5.x dropped TensorFlow support — pin to transformers==4.44.0.\n* Keras doesn't summarize layers correctly in this setup; a workaround is needed.\n\n\nMore details on the Medium article...","metadata":{}},{"cell_type":"markdown","source":"# Setting up the notebook","metadata":{}},{"cell_type":"markdown","source":"Required as transformers 5.x has dropped TF support entirely","metadata":{}},{"cell_type":"code","source":"!pip install transformers==4.44.0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:18:22.312462Z","iopub.execute_input":"2026-05-04T13:18:22.312742Z","iopub.status.idle":"2026-05-04T13:18:25.80715Z","shell.execute_reply.started":"2026-05-04T13:18:22.312717Z","shell.execute_reply":"2026-05-04T13:18:25.806302Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport os\n\nimport matplotlib.pyplot as plt\nfrom collections import Counter\n\nimport re\n\nimport tensorflow as tf\nprint(\"Tensorflow version \" + tf.__version__)\n\nfrom transformers import TFAutoModel","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:18:25.809057Z","iopub.execute_input":"2026-05-04T13:18:25.809424Z","iopub.status.idle":"2026-05-04T13:18:25.814974Z","shell.execute_reply.started":"2026-05-04T13:18:25.809396Z","shell.execute_reply":"2026-05-04T13:18:25.81408Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## A note on ml_utils 🛠️\n\nThroughout this notebook, I use [ml_utils](https://github.com/ThomasPRZilliox/ml_utils): a small personal toolkit \nI've been building to keep ML notebooks clean and avoid rewriting the same boilerplate over and over. \nIt currently covers TensorFlow utilities for visualisation and computer vision, and it's a work in progress!\n\nIf you find it useful, a ⭐ on GitHub is always appreciated. And if something doesn't work as expected \nor you have ideas for improvements, feel free to [open an issue](https://github.com/ThomasPRZilliox/ml_utils/issues), feedback is very welcome! 🙏","metadata":{}},{"cell_type":"code","source":"!git clone https://github.com/ThomasPRZilliox/ml_utils.git","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:18:25.81599Z","iopub.execute_input":"2026-05-04T13:18:25.816369Z","iopub.status.idle":"2026-05-04T13:18:26.045897Z","shell.execute_reply.started":"2026-05-04T13:18:25.816346Z","shell.execute_reply":"2026-05-04T13:18:26.04511Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import ml_utils.tf.visuals as tf_viz\nimport ml_utils.tf.computer_vision as ml_cv","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:18:26.048101Z","iopub.execute_input":"2026-05-04T13:18:26.048354Z","iopub.status.idle":"2026-05-04T13:18:26.053891Z","shell.execute_reply.started":"2026-05-04T13:18:26.048328Z","shell.execute_reply":"2026-05-04T13:18:26.052932Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## TPU Detection & Distribution Strategy\nFor more information [https://www.tensorflow.org/api_docs/python/tf/distribute/Strategy](https://www.tensorflow.org/api_docs/python/tf/distribute/Strategy)\n\nThe following blocks implement the following: Connects to the TPU cluster → initializes the TPU system → creates a TPUStrategy for distributed training across TPU cores","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":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:18:26.054918Z","iopub.execute_input":"2026-05-04T13:18:26.05528Z","iopub.status.idle":"2026-05-04T13:18:26.067666Z","shell.execute_reply.started":"2026-05-04T13:18:26.055257Z","shell.execute_reply":"2026-05-04T13:18:26.066834Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Prepare the dataset","metadata":{}},{"cell_type":"markdown","source":"## Get the dataset from Kaggle","metadata":{}},{"cell_type":"code","source":"IMAGE_SIZE = [224, 224]\nGCS_PATH = '/kaggle/input/competitions/tpu-getting-started/' + '/tfrecords-jpeg-224x224'\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') ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:18:26.068876Z","iopub.execute_input":"2026-05-04T13:18:26.069169Z","iopub.status.idle":"2026-05-04T13:18:26.139991Z","shell.execute_reply.started":"2026-05-04T13:18:26.069148Z","shell.execute_reply":"2026-05-04T13:18:26.139203Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Extra training data \n\nTo boost the training set, we will also use the external dataset \n[tf-flower-photo-tfrec](https://www.kaggle.com/datasets/kirillblinov/tf-flower-photo-tfrec) \nalongside the competition data. It provides additional flower images in TFRecord format \n(ready to plug straight into the pipeline), with most duplicates removed. Importantly, \nthe `_no_test` folders are used to make sure none of the competition's test images \nleak into training.","metadata":{}},{"cell_type":"code","source":"GCS_PATH_imagenet = '/kaggle/input/datasets/kirillblinov/tf-flower-photo-tfrec/imagenet/tfrecords-jpeg-224x224'\nIMAGENET_FILENAMES = tf.io.gfile.glob(GCS_PATH_imagenet + '/*.tfrec')\n\nGCS_PATH_inaturalist = '/kaggle/input/datasets/kirillblinov/tf-flower-photo-tfrec/inaturalist/tfrecords-jpeg-224x224'\nINATURALIST_FILENAMES = tf.io.gfile.glob(GCS_PATH_inaturalist + '/*.tfrec')\n\nGCS_PATH_openimage = '/kaggle/input/datasets/kirillblinov/tf-flower-photo-tfrec/openimage/tfrecords-jpeg-224x224'\nOPENIMAGE_FILENAMES = tf.io.gfile.glob(GCS_PATH_openimage + '/*.tfrec')\n\nGCS_PATH_oxford_102 = '/kaggle/input/datasets/kirillblinov/tf-flower-photo-tfrec/oxford_102/tfrecords-jpeg-224x224'\nOXFORD_102_FILENAMES = tf.io.gfile.glob(GCS_PATH_oxford_102 + '/*.tfrec')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:18:26.141001Z","iopub.execute_input":"2026-05-04T13:18:26.141308Z","iopub.status.idle":"2026-05-04T13:18:26.234726Z","shell.execute_reply.started":"2026-05-04T13:18:26.141272Z","shell.execute_reply":"2026-05-04T13:18:26.234114Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TRAINING_FILENAMES = TRAINING_FILENAMES + IMAGENET_FILENAMES + INATURALIST_FILENAMES + OPENIMAGE_FILENAMES + OXFORD_102_FILENAMES","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:18:26.235628Z","iopub.execute_input":"2026-05-04T13:18:26.235835Z","iopub.status.idle":"2026-05-04T13:18:26.239592Z","shell.execute_reply.started":"2026-05-04T13:18:26.235814Z","shell.execute_reply":"2026-05-04T13:18:26.239008Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Define the classes","metadata":{}},{"cell_type":"code","source":"CLASSES = ['pink primrose',    'hard-leaved pocket orchid', 'canterbury bells', 'sweet pea',     'wild geranium',     'tiger lily',           'moon orchid',              'bird of paradise', 'monkshood',        'globe thistle',         # 00 - 09\n           'snapdragon',       \"colt's foot\",               'king protea',      'spear thistle', 'yellow iris',       'globe-flower',         'purple coneflower',        'peruvian lily',    'balloon flower',   'giant white arum lily', # 10 - 19\n           'fire lily',        'pincushion flower',         'fritillary',       'red ginger',    'grape hyacinth',    'corn poppy',           'prince of wales feathers', 'stemless gentian', 'artichoke',        'sweet william',         # 20 - 29\n           'carnation',        'garden phlox',              'love in the mist', 'cosmos',        'alpine sea holly',  'ruby-lipped cattleya', 'cape flower',              'great masterwort', 'siam tulip',       'lenten rose',           # 30 - 39\n           'barberton daisy',  'daffodil',                  'sword lily',       'poinsettia',    'bolero deep blue',  'wallflower',           'marigold',                 'buttercup',        'daisy',            'common dandelion',      # 40 - 49\n           'petunia',          'wild pansy',                'primula',          'sunflower',     'lilac hibiscus',    'bishop of llandaff',   'gaura',                    'geranium',         'orange dahlia',    'pink-yellow dahlia',    # 50 - 59\n           'cautleya spicata', 'japanese anemone',          'black-eyed susan', 'silverbush',    'californian poppy', 'osteospermum',         'spring crocus',            'iris',             'windflower',       'tree poppy',            # 60 - 69\n           'gazania',          'azalea',                    'water lily',       'rose',          'thorn apple',       'morning glory',        'passion flower',           'lotus',            'toad lily',        'anthurium',             # 70 - 79\n           'frangipani',       'clematis',                  'hibiscus',         'columbine',     'desert-rose',       'tree mallow',          'magnolia',                 'cyclamen ',        'watercress',       'canna lily',            # 80 - 89\n           'hippeastrum ',     'bee balm',                  'pink quill',       'foxglove',      'bougainvillea',     'camellia',             'mallow',                   'mexican petunia',  'bromelia',         'blanket flower',        # 90 - 99\n           'trumpet creeper',  'blackberry lily',           'common tulip',     'wild rose']                                                                                                                                               # 100 - 102","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:18:26.240649Z","iopub.execute_input":"2026-05-04T13:18:26.240941Z","iopub.status.idle":"2026-05-04T13:18:26.254941Z","shell.execute_reply.started":"2026-05-04T13:18:26.240905Z","shell.execute_reply":"2026-05-04T13:18:26.254116Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Definition of dataset helper function\n\nThe definition are from the \"Getting started\" notebook : [A Simple Petals TF 2.2 notebook](https://www.kaggle.com/code/philculliton/a-simple-petals-tf-2-2-notebook)","metadata":{}},{"cell_type":"code","source":"def decode_image(image_data):\n    image = tf.image.decode_jpeg(image_data, channels=3)\n    \n    # scale to [0, 1] first\n    image = tf.cast(image, tf.float32) / 255.0 \n    \n    # Normalize with ImageNet mean and std (what ViT was pretrained with)\n    mean = tf.constant([0.485, 0.456, 0.406])\n    std  = tf.constant([0.229, 0.224, 0.225])\n    image = (image - mean) / std\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        # class is missing, this competitions's challenge is to predict flower classes for the test dataset\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":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:18:26.257669Z","iopub.execute_input":"2026-05-04T13:18:26.257937Z","iopub.status.idle":"2026-05-04T13:18:26.273488Z","shell.execute_reply.started":"2026-05-04T13:18:26.257907Z","shell.execute_reply":"2026-05-04T13:18:26.272676Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Quick inspection of the training dataset","metadata":{}},{"cell_type":"code","source":"# Extract all labels from the training dataset\ntrain_dataset = load_dataset(TRAINING_FILENAMES, labeled=True, ordered=True)\n\n# Collect just the labels\nall_labels = []\nfor _, label in train_dataset:\n    all_labels.append(label.numpy())\n\n# Count per class\ncounter = Counter(all_labels)\n\n\n# Visual\nplt.figure(figsize=(20, 4))\nplt.bar(range(len(CLASSES)), [counter.get(i, 0) for i in range(len(CLASSES))])\nplt.xticks(range(len(CLASSES)), CLASSES, rotation=90, fontsize=7)\nplt.title('Class Distribution')\nplt.ylabel('Sample count')\nplt.tight_layout()\nplt.show()\n\n# Quick imbalance ratio\ncounts = [counter.get(i, 0) for i in range(len(CLASSES))]\nprint(f\"\\nMax/Min ratio: {max(counts)/min(counts):.1f}x\")\nprint(f\"Most common: {CLASSES[np.argmax(counts)]} ({max(counts)})\")\nprint(f\"Least common: {CLASSES[np.argmin(counts)]} ({min(counts)})\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:18:26.274504Z","iopub.execute_input":"2026-05-04T13:18:26.27483Z","iopub.status.idle":"2026-05-04T13:18:38.510888Z","shell.execute_reply.started":"2026-05-04T13:18:26.274795Z","shell.execute_reply":"2026-05-04T13:18:38.510169Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data pipelines","metadata":{}},{"cell_type":"code","source":"def data_augment(image, label):\n    image = tf.image.random_flip_left_right(image)\n    # image = ml_cv.random_crop_resize(image,cropped_factor=0.9)\n    # image = ml_cv.random_erasing(image, prob=0.5, min_area=0.02, max_area=0.05,\n    #                min_ratio=0.3, max_ratio=3.0, mode=\"black\")\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\n    # 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))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:18:38.511835Z","iopub.execute_input":"2026-05-04T13:18:38.512166Z","iopub.status.idle":"2026-05-04T13:18:38.520723Z","shell.execute_reply.started":"2026-05-04T13:18:38.512141Z","shell.execute_reply":"2026-05-04T13:18:38.519959Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Load the data","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":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:18:38.521682Z","iopub.execute_input":"2026-05-04T13:18:38.521974Z","iopub.status.idle":"2026-05-04T13:18:38.638221Z","shell.execute_reply.started":"2026-05-04T13:18:38.521938Z","shell.execute_reply":"2026-05-04T13:18:38.637605Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Training data shapes:\")\nfor image, label in ds_train.take(3):\n    print(image.numpy().shape, label.numpy().shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:18:38.639102Z","iopub.execute_input":"2026-05-04T13:18:38.639456Z","iopub.status.idle":"2026-05-04T13:18:40.823311Z","shell.execute_reply.started":"2026-05-04T13:18:38.639433Z","shell.execute_reply":"2026-05-04T13:18:40.822598Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Test data shapes:\")\nfor image, idnum in ds_test.take(3):\n    print(image.numpy().shape, idnum.numpy().shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:18:40.824333Z","iopub.execute_input":"2026-05-04T13:18:40.824703Z","iopub.status.idle":"2026-05-04T13:18:40.898643Z","shell.execute_reply.started":"2026-05-04T13:18:40.824677Z","shell.execute_reply":"2026-05-04T13:18:40.897895Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train the model","metadata":{}},{"cell_type":"code","source":"EPOCHS = 10     \nFT_EPOCHS = 10","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:18:40.899503Z","iopub.execute_input":"2026-05-04T13:18:40.89981Z","iopub.status.idle":"2026-05-04T13:18:40.9035Z","shell.execute_reply.started":"2026-05-04T13:18:40.899785Z","shell.execute_reply":"2026-05-04T13:18:40.902623Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model configuration","metadata":{}},{"cell_type":"code","source":"callbacks = [\n    tf.keras.callbacks.EarlyStopping(\n        monitor='val_loss',\n        patience=3,\n        restore_best_weights=True\n    ),\n    tf.keras.callbacks.ReduceLROnPlateau(\n        monitor='val_loss',\n        factor=0.5,\n        patience=2\n    )\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:18:40.90446Z","iopub.execute_input":"2026-05-04T13:18:40.904787Z","iopub.status.idle":"2026-05-04T13:18:40.915644Z","shell.execute_reply.started":"2026-05-04T13:18:40.90475Z","shell.execute_reply":"2026-05-04T13:18:40.914967Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ViTCLSExtractor(tf.keras.layers.Layer):\n    def __init__(self, vit_backbone, **kwargs):\n        super().__init__(**kwargs)\n        self.vit_backbone = vit_backbone\n\n    def call(self, pixel_values, training=False):\n        outputs = self.vit_backbone(pixel_values=pixel_values, training=training)\n        return outputs.last_hidden_state[:, 0, :]\n\n    def count_params(self):\n        return sum(tf.size(w).numpy() for w in self.vit_backbone.weights)\n\n\n    # Override these so Keras walks into vit_backbone properly\n    @property\n    def trainable_weights(self):\n        return self.vit_backbone.trainable_weights\n\n    @property\n    def non_trainable_weights(self):\n        return self.vit_backbone.non_trainable_weights\n\n\nwith strategy.scope():\n    vit_backbone = TFAutoModel.from_pretrained(\"google/vit-base-patch16-224\")\n    vit_extractor = ViTCLSExtractor(vit_backbone)\n\n    vit_backbone.vit.embeddings.trainable = False\n    vit_backbone.vit.layernorm.trainable = False\n    vit_backbone.vit.pooler.trainable = False\n\n    all_encoder_layers = vit_backbone.vit.encoder.layer\n\n    for layer in all_encoder_layers:\n        layer.trainable = False\n\n    inputs = tf.keras.Input(shape=[*IMAGE_SIZE, 3])\n    x = tf.keras.layers.Permute((3, 1, 2))(inputs)\n    vit_output = vit_extractor(x)\n    outputs = tf.keras.layers.Dropout(0.3)(vit_output)\n    outputs = tf.keras.layers.Dense(256, activation='gelu',kernel_regularizer=tf.keras.regularizers.L2(1e-4))(outputs)\n    outputs = tf.keras.layers.Dropout(0.2)(outputs)\n    outputs = tf.keras.layers.Dense(\n        len(CLASSES), activation='softmax',\n        kernel_regularizer=tf.keras.regularizers.L2(1e-4)\n    )(outputs)\n    model = tf.keras.Model(inputs=inputs, outputs=outputs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:18:40.916551Z","iopub.execute_input":"2026-05-04T13:18:40.917339Z","iopub.status.idle":"2026-05-04T13:18:43.661195Z","shell.execute_reply.started":"2026-05-04T13:18:40.917307Z","shell.execute_reply":"2026-05-04T13:18:43.66028Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.compile(\n    optimizer='adam',\n    loss='sparse_categorical_crossentropy',\n    metrics=['accuracy']\n)\n\n\nmodel.summary()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:18:43.662405Z","iopub.execute_input":"2026-05-04T13:18:43.662775Z","iopub.status.idle":"2026-05-04T13:18:43.759169Z","shell.execute_reply.started":"2026-05-04T13:18:43.662728Z","shell.execute_reply":"2026-05-04T13:18:43.758546Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"STEPS_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=callbacks\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:18:43.759971Z","iopub.execute_input":"2026-05-04T13:18:43.760264Z","iopub.status.idle":"2026-05-04T13:34:20.850462Z","shell.execute_reply.started":"2026-05-04T13:18:43.760241Z","shell.execute_reply":"2026-05-04T13:34:20.849802Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tf_viz.plot_history(history)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:34:20.851554Z","iopub.execute_input":"2026-05-04T13:34:20.851825Z","iopub.status.idle":"2026-05-04T13:34:21.167751Z","shell.execute_reply.started":"2026-05-04T13:34:20.851798Z","shell.execute_reply":"2026-05-04T13:34:21.167105Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Fine Tuning\n\nTo allow the model to learn more feature we will unfreeze all layers then freeze the earlier one, which contain generic feature such as edge or texture detection. If you want to only unfreeze some top layers, see the [TF documentation](https://www.tensorflow.org/tutorials/images/transfer_learning#un-freeze_the_top_layers_of_the_model).\n","metadata":{}},{"cell_type":"code","source":"# Use a cosine decay schedule with warm-up\nlr_schedule = tf.keras.optimizers.schedules.CosineDecay(\n    initial_learning_rate=tf.constant(0.0, dtype=tf.float32),\n    warmup_target=tf.constant(1e-5, dtype=tf.float32),\n    decay_steps=int(FT_EPOCHS * STEPS_PER_EPOCH* 0.5),\n    warmup_steps=int(FT_EPOCHS * STEPS_PER_EPOCH* 0.5),\n    alpha=tf.constant(1e-5, dtype=tf.float32),\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:34:21.168908Z","iopub.execute_input":"2026-05-04T13:34:21.169258Z","iopub.status.idle":"2026-05-04T13:34:21.173922Z","shell.execute_reply.started":"2026-05-04T13:34:21.169214Z","shell.execute_reply":"2026-05-04T13:34:21.173152Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ft_callbacks = [tf.keras.callbacks.EarlyStopping(\n        monitor='val_loss',\n        patience=3,\n        restore_best_weights=True\n    )]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:34:21.175409Z","iopub.execute_input":"2026-05-04T13:34:21.17591Z","iopub.status.idle":"2026-05-04T13:34:21.185369Z","shell.execute_reply.started":"2026-05-04T13:34:21.175874Z","shell.execute_reply":"2026-05-04T13:34:21.184614Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### Unfreeze the last 2 blocks","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# The last N_UNFREEZE blocks are never touched, so they stay trainable.\nvit_backbone.trainable = False\n\nN_UNFREEZE = 2\n\nall_encoder_layers = vit_backbone.vit.encoder.layer\nprint(all_encoder_layers)\nfor layer in all_encoder_layers[-N_UNFREEZE:]:\n    layer.trainable = True\n\n\n\n# Recompile with a much lower learning rate, too high and you'll destroy the pretrained weights\nmodel.compile(\n    optimizer=tf.keras.optimizers.Adam(lr_schedule),\n    loss='sparse_categorical_crossentropy',\n    metrics=['accuracy']\n)\n\n# # Check what's actually frozen\n# for i, layer in enumerate(vit_backbone.vit.encoder.layer):\n#     n_trainable = len(layer.trainable_weights)\n#     n_frozen = len(layer.non_trainable_weights)\n#     print(f\"Block {i:2d}: trainable={n_trainable:3d}  frozen={n_frozen:3d}\")\n\nhistory_ft = model.fit(\n    ds_train,\n    validation_data=ds_valid,\n    epochs=FT_EPOCHS,\n    steps_per_epoch=STEPS_PER_EPOCH,\n    callbacks=ft_callbacks\n)\n\ntf_viz.plot_history(history_ft)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T13:34:21.206888Z","iopub.execute_input":"2026-05-04T13:34:21.207175Z","iopub.status.idle":"2026-05-04T13:34:21.327665Z","shell.execute_reply.started":"2026-05-04T13:34:21.207143Z","shell.execute_reply":"2026-05-04T13:34:21.327097Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Unfreeze the last 4 blocks","metadata":{}},{"cell_type":"code","source":"vit_backbone.trainable = False\n\nN_UNFREEZE = 4\n\nall_encoder_layers = vit_backbone.vit.encoder.layer\nprint(all_encoder_layers)\nfor layer in all_encoder_layers[-N_UNFREEZE:]:\n    layer.trainable = True\n\n\n\n# Recompile with a much lower learning rate, too high and you'll destroy the pretrained weights\nmodel.compile(\n    optimizer=tf.keras.optimizers.Adam(lr_schedule),\n    loss='sparse_categorical_crossentropy',\n    metrics=['accuracy']\n)\n\nhistory_ft = model.fit(\n    ds_train,\n    validation_data=ds_valid,\n    epochs=FT_EPOCHS,\n    steps_per_epoch=STEPS_PER_EPOCH,\n    callbacks=ft_callbacks\n)\n\ntf_viz.plot_history(history_ft)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Unfreeze the last 6 blocks","metadata":{}},{"cell_type":"code","source":"vit_backbone.trainable = False\n\nN_UNFREEZE = 6\n\nall_encoder_layers = vit_backbone.vit.encoder.layer\nprint(all_encoder_layers)\nfor layer in all_encoder_layers[-N_UNFREEZE:]:\n    layer.trainable = True\n\n\n\n# Recompile with a much lower learning rate, too high and you'll destroy the pretrained weights\nmodel.compile(\n    optimizer=tf.keras.optimizers.Adam(lr_schedule),\n    loss='sparse_categorical_crossentropy',\n    metrics=['accuracy']\n)\n\nhistory_ft = model.fit(\n    ds_train,\n    validation_data=ds_valid,\n    epochs=FT_EPOCHS,\n    steps_per_epoch=STEPS_PER_EPOCH,\n    callbacks=ft_callbacks\n)\n\ntf_viz.plot_history(history_ft)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Validation","metadata":{}},{"cell_type":"code","source":"tf_viz.plot_history_with_fine_tuning(history,history_ft)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Get predictions and true labels from validation set\nall_preds  = []\nall_labels = []\n\nfor images, labels in ds_valid:\n    preds = model.predict(images, verbose=0)\n    all_preds.extend(np.argmax(preds, axis=-1))\n    all_labels.extend(labels.numpy())\n\nall_preds  = np.array(all_preds)\nall_labels = np.array(all_labels)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tf_viz.plot_confusion_matrix(all_labels, all_preds,labels,CLASSES)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Generate submission","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)\n\nprint('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":{"trusted":true},"outputs":[],"execution_count":null}]}