{"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":31328,"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, progressively improving a model across multiple versions, from a simple baseline to more sophisticated architectures. Each version builds on the lessons learned from the previous one.\n\n![](https://github.com/ThomasPRZilliox/ML-explained/raw/ebb86f09feb83d4ea98ea6040579ebb71c4f0daa/fc-effnet.png)\n\n\n| Version    | Description | Public performances - F1 Score| Notebook run time\n| -------- | ------- | ------- | ------- |\n| 16 - EffNet + FT + ED + DA | Transfer learning with a pretrained EfficientNetV2S backbone (ImageNet weights), frozen then fine-tuned + Extended Data + Data augmentation (Random Flip) + More epochs | 0.93238 | 1 hour and 33 minutes|\n| 15 - EffNet + FT + ED + DA | Transfer learning with a pretrained EfficientNetV2S backbone (ImageNet weights), frozen then fine-tuned + Extended Data + Data augmentation (Random Flip + Random Erasing) + More epochs | 0.93211 | 1 hour and 36 minutes|\n| 14 - EffNet + FT + ED + DA | Transfer learning with a pretrained EfficientNetV2S backbone (ImageNet weights), frozen then fine-tuned + Extended Data + Data augmentation (Random Flip + Random Erasing) | 0.92312 | 1 hour and 6 minutes|\n| 13 - EffNet + FT + ED + DA | Transfer learning with a pretrained EfficientNetV2S backbone (ImageNet weights), frozen then fine-tuned + Extended Data + Data augmentation (Random Flip) | 0.92815 | 1 hour and 4 minutes|\n| 12 - EffNet + FT + ED | Transfer learning with a pretrained EfficientNetV2S backbone (ImageNet weights), frozen then fine-tuned + Extended Data | 0.92719 | 59 minutes and 35 seconds|\n| 11 - EffNet + FT + DA | Transfer learning with a pretrained EfficientNetV2S backbone (ImageNet weights), frozen then fine-tuned + Data augmentation (Random Flip + Random Erasing) | 0.88278 | 18 minutes and 14 seconds|\n| 10 - EffNet + FT + DA | Transfer learning with a pretrained EfficientNetV2S backbone (ImageNet weights), frozen then fine-tuned + Data augmentation (Random Flip) | 0.88041 | 16 minutes and 53 seconds |\n| 9 - EffNet + FT + DA | Transfer learning with a pretrained EfficientNetV2S backbone (ImageNet weights), frozen then fine-tuned + Data augmentation (Random Crop Resize) | 0.86860 | 17 minutes and 54 seconds |\n| 8 - EffNet + FT + DA | Transfer learning with a pretrained EfficientNetV2S backbone (ImageNet weights), frozen then fine-tuned + Data augmentation (Random Erasing) | 0.87651 | 17 minutes and 45 seconds |\n| 7 - EffNet + FT | Transfer learning with a pretrained EfficientNetV2S backbone (ImageNet weights), frozen then fine-tuned | 0.87275 | 17 minutes and 19 seconds |\n| 6 - EffNet + DA  | Transfer learning with a pretrained EfficientNetV2S backbone (ImageNet weights) + Data augmentation (Random Erasing) | 0.85141 | 9 minutes and 50 seconds|\n| 5 - EffNet + DA  | Transfer learning with a pretrained EfficientNetV2S backbone (ImageNet weights) + Data augmentation (Random Crop and Resize) | 0.85060 |9 minutes and 22 seconds |\n| 4 - EffNet + DA  | Transfer learning with a pretrained EfficientNetV2S backbone (ImageNet weights) + Data augmentation (Random Flip) | 0.84821 | 9 minutes and 52 seconds |\n| 3 - EffNet   | Transfer learning with a pretrained EfficientNetV2S backbone (ImageNet weights) | 0.85614 | 10 minutes and 24 seconds |\n| 2 - VGG16    | Transfer learning with a pretrained VGG16 backbone (ImageNet weights)  | 0.56577 | 9 minutes and 47 seconds |\n| 1 - Baseline    | Small custom CNN trained from scratch   | 0.04336 | 5 minutes and 32 seconds |","metadata":{}},{"cell_type":"markdown","source":"# Setting up the notebook","metadata":{}},{"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__)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-04-14T17:00:02.494369Z","iopub.execute_input":"2026-04-14T17:00:02.494851Z","iopub.status.idle":"2026-04-14T17:00:30.014008Z","shell.execute_reply.started":"2026-04-14T17:00:02.494823Z","shell.execute_reply":"2026-04-14T17:00:30.013350Z"}},"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-04-14T17:00:30.015617Z","iopub.execute_input":"2026-04-14T17:00:30.016152Z","iopub.status.idle":"2026-04-14T17:00:31.301747Z","shell.execute_reply.started":"2026-04-14T17:00:30.016123Z","shell.execute_reply":"2026-04-14T17:00:31.300686Z"}},"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-04-14T17:00:31.303313Z","iopub.execute_input":"2026-04-14T17:00:31.304153Z","iopub.status.idle":"2026-04-14T17:00:31.410606Z","shell.execute_reply.started":"2026-04-14T17:00:31.304121Z","shell.execute_reply":"2026-04-14T17:00:31.409916Z"}},"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-04-14T17:00:31.411533Z","iopub.execute_input":"2026-04-14T17:00:31.412372Z","iopub.status.idle":"2026-04-14T17:00:31.418604Z","shell.execute_reply.started":"2026-04-14T17:00:31.412346Z","shell.execute_reply":"2026-04-14T17:00:31.417744Z"}},"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-04-14T17:00:31.419570Z","iopub.execute_input":"2026-04-14T17:00:31.419874Z","iopub.status.idle":"2026-04-14T17:00:31.500633Z","shell.execute_reply.started":"2026-04-14T17:00:31.419852Z","shell.execute_reply":"2026-04-14T17:00:31.499997Z"}},"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-04-14T17:00:31.501548Z","iopub.execute_input":"2026-04-14T17:00:31.501853Z","iopub.status.idle":"2026-04-14T17:00:31.517288Z","shell.execute_reply.started":"2026-04-14T17:00:31.501830Z","shell.execute_reply":"2026-04-14T17:00:31.516453Z"}},"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-04-14T17:00:31.519690Z","iopub.execute_input":"2026-04-14T17:00:31.520014Z","iopub.status.idle":"2026-04-14T17:00:31.523593Z","shell.execute_reply.started":"2026-04-14T17:00:31.519990Z","shell.execute_reply":"2026-04-14T17:00:31.522985Z"}},"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-04-14T17:00:31.538398Z","iopub.execute_input":"2026-04-14T17:00:31.538707Z","iopub.status.idle":"2026-04-14T17:00:31.548938Z","shell.execute_reply.started":"2026-04-14T17:00:31.538666Z","shell.execute_reply":"2026-04-14T17:00:31.548131Z"}},"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    image = tf.cast(image, tf.float32) # For EfficientNetV2S  expect [0,255]\n    # image = tf.cast(image, tf.float32) / 255.0  # convert image to floats in [0, 1] range ->  VGG\n    image = tf.reshape(image, [*IMAGE_SIZE, 3]) # explicit size needed for TPU\n    return image\n\ndef read_labeled_tfrecord(example):\n    LABELED_TFREC_FORMAT = {\n        \"image\": tf.io.FixedLenFeature([], tf.string), # tf.string means bytestring\n        \"class\": tf.io.FixedLenFeature([], tf.int64),  # shape [] means single element\n    }\n    example = tf.io.parse_single_example(example, LABELED_TFREC_FORMAT)\n    image = decode_image(example['image'])\n    label = tf.cast(example['class'], tf.int32)\n    return image, label # returns a dataset of (image, label) pairs\n\ndef read_unlabeled_tfrecord(example):\n    UNLABELED_TFREC_FORMAT = {\n        \"image\": tf.io.FixedLenFeature([], tf.string), # tf.string means bytestring\n        \"id\": tf.io.FixedLenFeature([], tf.string),  # shape [] means single element\n        # 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-04-14T17:00:31.549932Z","iopub.execute_input":"2026-04-14T17:00:31.550320Z","iopub.status.idle":"2026-04-14T17:00:31.568615Z","shell.execute_reply.started":"2026-04-14T17:00:31.550297Z","shell.execute_reply":"2026-04-14T17:00:31.567745Z"}},"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-04-14T17:00:31.570581Z","iopub.execute_input":"2026-04-14T17:00:31.570927Z","iopub.status.idle":"2026-04-14T17:00:49.662237Z","shell.execute_reply.started":"2026-04-14T17:00:31.570895Z","shell.execute_reply":"2026-04-14T17:00:49.661398Z"}},"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.3,\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-04-14T17:00:49.663375Z","iopub.execute_input":"2026-04-14T17:00:49.663713Z","iopub.status.idle":"2026-04-14T17:00:49.671807Z","shell.execute_reply.started":"2026-04-14T17:00:49.663690Z","shell.execute_reply":"2026-04-14T17:00:49.670996Z"}},"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-04-14T17:00:49.672765Z","iopub.execute_input":"2026-04-14T17:00:49.672967Z","iopub.status.idle":"2026-04-14T17:00:49.823106Z","shell.execute_reply.started":"2026-04-14T17:00:49.672947Z","shell.execute_reply":"2026-04-14T17:00:49.822440Z"}},"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-04-14T17:00:49.824207Z","iopub.execute_input":"2026-04-14T17:00:49.824643Z","iopub.status.idle":"2026-04-14T17:00:51.118927Z","shell.execute_reply.started":"2026-04-14T17:00:49.824614Z","shell.execute_reply":"2026-04-14T17:00:51.118286Z"}},"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-04-14T17:00:51.119802Z","iopub.execute_input":"2026-04-14T17:00:51.120417Z","iopub.status.idle":"2026-04-14T17:00:51.224660Z","shell.execute_reply.started":"2026-04-14T17:00:51.120381Z","shell.execute_reply":"2026-04-14T17:00:51.223928Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train the model","metadata":{}},{"cell_type":"code","source":"EPOCHS = 1     # Version to update the display \nFT_EPOCHS = 1\n# EPOCHS = 15\n# FT_EPOCHS = 15","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T17:00:51.225608Z","iopub.execute_input":"2026-04-14T17:00:51.225954Z","iopub.status.idle":"2026-04-14T17:00:51.229793Z","shell.execute_reply.started":"2026-04-14T17:00:51.225913Z","shell.execute_reply":"2026-04-14T17:00:51.229243Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model configuration","metadata":{}},{"cell_type":"code","source":"with strategy.scope():\n    pretrained_model =tf.keras.applications.EfficientNetV2S(\n        weights='imagenet',\n        include_top=False,\n        input_shape=[*IMAGE_SIZE, 3]\n    )\n\n    pretrained_model.trainable = False\n    \n    model = tf.keras.Sequential([\n        # To a base pretrained on ImageNet to extract features from images...\n        pretrained_model,\n        # ... attach a new head to act as a classifier.\n        tf.keras.layers.GlobalAveragePooling2D(),\n        tf.keras.layers.Dense(len(CLASSES), activation='softmax')\n    ])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T17:00:51.230779Z","iopub.execute_input":"2026-04-14T17:00:51.231034Z","iopub.status.idle":"2026-04-14T17:01:00.643819Z","shell.execute_reply.started":"2026-04-14T17:00:51.231013Z","shell.execute_reply":"2026-04-14T17:01:00.642886Z"}},"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-04-14T17:01:00.644881Z","iopub.execute_input":"2026-04-14T17:01:00.645285Z","iopub.status.idle":"2026-04-14T17:01:00.680011Z","shell.execute_reply.started":"2026-04-14T17:01:00.645240Z","shell.execute_reply":"2026-04-14T17:01:00.679439Z"}},"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)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T17:01:00.680971Z","iopub.execute_input":"2026-04-14T17:01:00.681403Z","iopub.status.idle":"2026-04-14T17:03:35.332709Z","shell.execute_reply.started":"2026-04-14T17:01:00.681379Z","shell.execute_reply":"2026-04-14T17:03:35.330611Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tf_viz.plot_history(history)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T17:03:35.333344Z","iopub.status.idle":"2026-04-14T17:03:35.333601Z","shell.execute_reply.started":"2026-04-14T17:03:35.333481Z","shell.execute_reply":"2026-04-14T17:03:35.333495Z"}},"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":"# Unfreeze the whole backbone\npretrained_model.trainable = True\n\n# Freeze the early layers as they contain generic features (edges, textures)\n# If you want to only unfreeze some layers, see https://www.tensorflow.org/tutorials/images/transfer_learning#un-freeze_the_top_layers_of_the_model\nfor layer in pretrained_model.layers[:-30]:\n    layer.trainable = False\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(learning_rate=1e-5),\n    loss='sparse_categorical_crossentropy',\n    metrics=['accuracy']\n)\n\n\nmodel.summary()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T17:03:35.334872Z","iopub.status.idle":"2026-04-14T17:03:35.335309Z","shell.execute_reply.started":"2026-04-14T17:03:35.335106Z","shell.execute_reply":"2026-04-14T17:03:35.335131Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"history_ft = model.fit(\n    ds_train,\n    validation_data=ds_valid,\n    epochs=FT_EPOCHS,\n    steps_per_epoch=STEPS_PER_EPOCH\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T17:03:35.336870Z","iopub.status.idle":"2026-04-14T17:03:35.337256Z","shell.execute_reply.started":"2026-04-14T17:03:35.337051Z","shell.execute_reply":"2026-04-14T17:03:35.337074Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tf_viz.plot_history(history_ft)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T17:03:35.338392Z","iopub.status.idle":"2026-04-14T17:03:35.338777Z","shell.execute_reply.started":"2026-04-14T17:03:35.338587Z","shell.execute_reply":"2026-04-14T17:03:35.338610Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Validation","metadata":{}},{"cell_type":"code","source":"tf_viz.plot_history_with_fine_tuning(history_ft,history_ft)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T17:03:35.340940Z","iopub.status.idle":"2026-04-14T17:03:35.341352Z","shell.execute_reply.started":"2026-04-14T17:03:35.341143Z","shell.execute_reply":"2026-04-14T17:03:35.341163Z"}},"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,"execution":{"iopub.status.busy":"2026-04-14T17:03:35.342373Z","iopub.status.idle":"2026-04-14T17:03:35.342607Z","shell.execute_reply.started":"2026-04-14T17:03:35.342495Z","shell.execute_reply":"2026-04-14T17:03:35.342508Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tf_viz.plot_confusion_matrix(all_labels, all_preds,labels,CLASSES)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T17:03:35.344068Z","iopub.status.idle":"2026-04-14T17:03:35.344481Z","shell.execute_reply.started":"2026-04-14T17:03:35.344250Z","shell.execute_reply":"2026-04-14T17:03:35.344308Z"}},"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,"execution":{"iopub.status.busy":"2026-04-14T17:03:35.345224Z","iopub.status.idle":"2026-04-14T17:03:35.345621Z","shell.execute_reply.started":"2026-04-14T17:03:35.345503Z","shell.execute_reply":"2026-04-14T17:03:35.345519Z"}},"outputs":[],"execution_count":null}]}