{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpu1vmV38","dataSources":[{"sourceId":21154,"databundleVersionId":1243559,"sourceType":"competition"}],"dockerImageVersionId":30734,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Imports and distribution strategy","metadata":{}},{"cell_type":"code","source":"# copied from https://www.kaggle.com/code/stefanrothermel/tensors-in-bloom-tpu-awakening\n\nimport os\nimport subprocess\n\nenv = os.environ.copy()\nenv[\"PATH\"] = f\"{os.path.expanduser('~')}/.local/bin:{env['PATH']}\"\n\nsubprocess.check_call(\n        [\"uv\", \"pip\", \"install\", \"--system\",\n         \"tensorflow-tpu==2.18.0\",\n         \"--find-links\", \"https://storage.googleapis.com/libtpu-tf-releases/index.html\"],\n        env=env\n)\nsubprocess.check_call(\n        [\"uv\", \"pip\", \"install\", \"--system\", \"ml_dtypes>=0.5.1\"],\n        env=env\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T14:53:17.060388Z","iopub.execute_input":"2026-06-10T14:53:17.060595Z","iopub.status.idle":"2026-06-10T14:53:17.355438Z","shell.execute_reply.started":"2026-06-10T14:53:17.060577Z","shell.execute_reply":"2026-06-10T14:53:17.354689Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport re\nos.environ[\"KERAS_BACKEND\"] = \"tensorflow\"\nimport tensorflow as tf\nimport keras\nimport numpy as np\nimport matplotlib.pyplot as plt\ntry :\n    resolver = tf.distribute.cluster_resolver.TPUClusterResolver(tpu='local')\n    tf.config.experimental_connect_to_cluster(resolver)\n    tf.tpu.experimental.initialize_tpu_system(resolver)\n    strategy = tf.distribute.TPUStrategy(resolver)\n    TPU = True\nexcept ValueError:\n    TPU = False\n    strategy = tf.distribute.MirroredStrategy()\nfrom IPython.display import clear_output\nclear_output()\nprint(\"tensorflow version :\", tf.__version__)\nprint(\"keras version :\", keras.__version__)\nprint(\"Number of accelerators =\", strategy.num_replicas_in_sync)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2026-06-10T14:53:17.356204Z","iopub.execute_input":"2026-06-10T14:53:17.356395Z","iopub.status.idle":"2026-06-10T14:53:44.452535Z","shell.execute_reply.started":"2026-06-10T14:53:17.356379Z","shell.execute_reply":"2026-06-10T14:53:44.451422Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Preparing datasets (load, parse , augment)","metadata":{}},{"cell_type":"code","source":"IMAGE_SIZE = (512, 512)\nBATCH_SIZE = 32 if TPU else 8\nBASE_DIR = f\"/kaggle/input/tpu-getting-started/tfrecords-jpeg-{IMAGE_SIZE[0]}x{IMAGE_SIZE[1]}\"\nDATASETS = {}\ncardinalities = {}\nCLASSES = ['pink primrose',    'hard-leaved pocket orchid', 'canterbury bells', 'sweet pea',     'wild geranium',     'tiger lily',           'moon orchid',              'bird of paradise', 'monkshood',        'globe thistle',         # 00 - 09\n           'snapdragon',       \"colt's foot\",               'king protea',      'spear thistle', 'yellow iris',       'globe-flower',         'purple coneflower',        'peruvian lily',    'balloon flower',   'giant white arum lily', # 10 - 19\n           'fire lily',        'pincushion flower',         'fritillary',       'red ginger',    'grape hyacinth',    'corn poppy',           'prince of wales feathers', 'stemless gentian', 'artichoke',        'sweet william',         # 20 - 29\n           'carnation',        'garden phlox',              'love in the mist', 'cosmos',        'alpine sea holly',  'ruby-lipped cattleya', 'cape flower',              'great masterwort', 'siam tulip',       'lenten rose',           # 30 - 39\n           'barberton daisy',  'daffodil',                  'sword lily',       'poinsettia',    'bolero deep blue',  'wallflower',           'marigold',                 'buttercup',        'daisy',            'common dandelion',      # 40 - 49\n           'petunia',          'wild pansy',                'primula',          'sunflower',     'lilac hibiscus',    'bishop of llandaff',   'gaura',                    'geranium',         'orange dahlia',    'pink-yellow dahlia',    # 50 - 59\n           'cautleya spicata', 'japanese anemone',          'black-eyed susan', 'silverbush',    'californian poppy', 'osteospermum',         'spring crocus',            'iris',             'windflower',       'tree poppy',            # 60 - 69\n           'gazania',          'azalea',                    'water lily',       'rose',          'thorn apple',       'morning glory',        'passion flower',           'lotus',            'toad lily',        'anthurium',             # 70 - 79\n           'frangipani',       'clematis',                  'hibiscus',         'columbine',     'desert-rose',       'tree mallow',          'magnolia',                 'cyclamen ',        'watercress',       'canna lily',            # 80 - 89\n           'hippeastrum ',     'bee balm',                  'pink quill',       'foxglove',      'bougainvillea',     'camellia',             'mallow',                   'mexican petunia',  'bromelia',         'blanket flower',        # 90 - 99\n           'trumpet creeper',  'blackberry lily',           'common tulip',     'wild rose']                                                                                                                                               # 100 - 102\n\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\ndef get_parse_function(split):\n    feature_description = {\n        'class': tf.io.FixedLenFeature([], tf.int64, default_value=0),\n        'image': tf.io.FixedLenFeature([], tf.string, default_value=''),\n        'id': tf.io.FixedLenFeature([], tf.string, default_value=''),\n    }\n    def _parse_function(example_proto, split=split):\n        example = tf.io.parse_single_example(example_proto, feature_description)\n        image = example['image']\n        image = tf.io.decode_jpeg(image, channels=3)\n        image = tf.reshape(image , [*IMAGE_SIZE, 3])\n        _id = example['id']\n        label = example[\"class\"]\n        if split == 'test':\n            return image, _id\n        else:\n            label = tf.one_hot(label, len(CLASSES))\n            return image, label\n    return _parse_function\n\n# def augment(images, labels):\n#     augmenter = keras.layers.RandAugment()\n#     inputs = {\"images\": images, \"labels\": labels}\n#     outputs = augmenter(inputs, training=True)\n#     return outputs['images'], outputs['labels']\n\nfor split in ['train', 'val', 'test']:\n    filelists = tf.io.gfile.glob(BASE_DIR + f'/{split}/*.tfrec')\n    cardinalities[split] = count_data_items(filelists)\n    DATASETS[split] = tf.data.TFRecordDataset(filelists)\\\n                             .map(get_parse_function(split), num_parallel_calls=tf.data.AUTOTUNE)\\\n                             .batch(BATCH_SIZE)\n    # if split == 'train':\n    #     DATASETS[split] = DATASETS[split].map(augment, num_parallel_calls=tf.data.AUTOTUNE)\\\n    #                                      .repeat(-1).prefetch(tf.data.AUTOTUNE)\n    # else:\n    #     DATASETS[split] = DATASETS[split].repeat(1).prefetch(tf.data.AUTOTUNE)\n\n    DATASETS[split] = DATASETS[split].repeat(1).prefetch(tf.data.AUTOTUNE)\n    \nfor split, dataset in DATASETS.items():\n    for image, id_or_label in dataset.take(1):\n        print(\"---\",split,\"---\")\n        print('cardinality from tf Dataset:', dataset.cardinality().numpy())\n        print('cardinality :', cardinalities[split])\n        print(\"images shape :\",image.shape)\n        print(\"label/id shape :\",id_or_label.shape)\nDATASETS['train']","metadata":{"execution":{"iopub.status.busy":"2026-06-10T14:53:44.453178Z","iopub.execute_input":"2026-06-10T14:53:44.453362Z","iopub.status.idle":"2026-06-10T14:53:44.710785Z","shell.execute_reply.started":"2026-06-10T14:53:44.453347Z","shell.execute_reply":"2026-06-10T14:53:44.709829Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Visualize training images","metadata":{}},{"cell_type":"code","source":"from tensorflow.errors import InvalidArgumentError\n\nbatch = next(iter(DATASETS['train']))\nimage_batch = batch[0]\nlabel_batch = batch[1]\n\nncol = 4\nnrow = 4\n\nrandom_start = 0 if BATCH_SIZE < (nrow*ncol) else np.random.randint(0,BATCH_SIZE-nrow*ncol)\nplt.figure(figsize=(4.5*ncol, 4.5*nrow))\nfor i in range(nrow*ncol):\n    index = random_start+i\n    try:\n        image = image_batch[index].numpy().astype(\"uint8\")\n    except InvalidArgumentError:\n        break\n    ax = plt.subplot(nrow, ncol, i + 1)\n    plt.imshow(image)\n    label = label_batch[index]\n    plt.title(f\"({index})\\nclass : {CLASSES[np.argmax(label)]}\")\n    plt.axis(\"off\")","metadata":{"execution":{"iopub.status.busy":"2026-06-10T14:53:44.711329Z","iopub.execute_input":"2026-06-10T14:53:44.711505Z","iopub.status.idle":"2026-06-10T14:53:46.463329Z","shell.execute_reply.started":"2026-06-10T14:53:44.711490Z","shell.execute_reply":"2026-06-10T14:53:46.462254Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train the classification model","metadata":{}},{"cell_type":"code","source":"# Learning rate schedule for TPU, GPU and CPU.\n# Using an LR ramp up because fine-tuning a pre-trained model.\n# Starting with a high LR would break the pre-trained weights.\ndef get_lr_callback(plot_schedule=False, EPOCHS=30):\n    LR_START = 0.00001\n    LR_MAX = 0.00005 * strategy.num_replicas_in_sync\n    LR_MIN = 0.000001\n    LR_RAMPUP_EPOCHS = 5\n    LR_SUSTAIN_EPOCHS = 0\n    LR_EXP_DECAY = .86\n\n    def lrfn(epoch):\n        if epoch < LR_RAMPUP_EPOCHS:\n            lr = (LR_MAX - LR_START) / LR_RAMPUP_EPOCHS * epoch + LR_START\n        elif epoch < LR_RAMPUP_EPOCHS + LR_SUSTAIN_EPOCHS:\n            lr = LR_MAX\n        else:\n            lr = (LR_MAX - LR_MIN) * LR_EXP_DECAY**(epoch - LR_RAMPUP_EPOCHS - LR_SUSTAIN_EPOCHS) + LR_MIN\n        return lr\n    \n    if plot_schedule:\n        rng = [i for i in range(25 if EPOCHS < 25 else EPOCHS)]\n        y = [lrfn(x) for x in rng]\n        plt.plot(rng, y)\n\n    return keras.callbacks.LearningRateScheduler(lrfn, verbose=0)\nget_lr_callback(plot_schedule=True)","metadata":{"execution":{"iopub.status.busy":"2026-06-10T14:53:46.463821Z","iopub.execute_input":"2026-06-10T14:53:46.463982Z","iopub.status.idle":"2026-06-10T14:53:46.578362Z","shell.execute_reply.started":"2026-06-10T14:53:46.463966Z","shell.execute_reply":"2026-06-10T14:53:46.577467Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CustomGlobalPooling2D(tf.keras.layers.Layer):\n    '''\n    modified from : https://github.com/csvance/keras-global-weighted-pooling/blob/master/gwp.py#L51\n    reference : https://arxiv.org/abs/1809.08264\n    '''\n    def __init__(self, reduce_func=\"max\", **kwargs):\n        super().__init__(**kwargs)\n        self.w = None\n        if reduce_func ==\"mean\":\n            self.reduce_func = tf.reduce_mean\n        else:\n            self.reduce_func = tf.reduce_max\n\n    def build(self, input_shape):\n        self.w = self.add_weight(name='w',\n                                  shape=(input_shape[1], input_shape[2], 1),\n                                  initializer='ones',\n                                  trainable=True)\n        super().build(input_shape)\n\n    def compute_output_shape(self, input_shape):\n        return input_shape[0], input_shape[3],\n\n    def call(self, x):\n        outputs = self.reduce_func(x*self.w, axis=(1, 2))\n        return outputs\n    \nepochs = 30\nverbose = 2\ntrain_steps = cardinalities['train'] // BATCH_SIZE + 1\n\nwith strategy.scope():\n    inputs = keras.layers.Input(shape=[*IMAGE_SIZE,3], name='input_layer')\n    backbone = keras.applications.ConvNeXtSmall(\n                    name=\"convnext_small\",\n                    include_top=False,\n                    weights=\"imagenet\",\n                    input_shape=[*IMAGE_SIZE,3],\n                )\n    dropout = keras.layers.Dropout(0.5, name='dropout')\n    pooling_layer = CustomGlobalPooling2D(name='pooling_layer')\n    classifier_layer = keras.layers.Dense(1024, name='classifier_layer')\n    output_layer = keras.layers.Dense(len(CLASSES), activation='softmax', name='output_layer')\n    \n    features = dropout(backbone(inputs))\n    features = pooling_layer(features)\n    features = classifier_layer(features)\n    outputs = output_layer(features)\n    model = keras.Model(inputs, outputs, name='flower_classifier')\n\n    model.compile(\n        loss='categorical_crossentropy',\n        optimizer=keras.optimizers.Adam(),\n        metrics=[\n            keras.metrics.CategoricalAccuracy(name='accuracy'),\n            keras.metrics.F1Score(average='macro',name='macro_F1')\n        ]\n    )\n    model.summary(expand_nested=False)\n    h = model.fit(\n            DATASETS['train'],\n            validation_data=DATASETS['val'],\n            epochs=epochs,\n            steps_per_epoch=train_steps,\n            verbose=verbose,\n            callbacks=[\n                get_lr_callback(EPOCHS=epochs),\n                keras.callbacks.ModelCheckpoint(\n                                    'checkpoint.weights.h5',\n                                    monitor=\"val_macro_F1\",\n                                    verbose=2,\n                                    save_best_only=True,\n                                    save_weights_only=True,\n                                    mode=\"max\",\n                                ),\n            ]\n        )\n    model.load_weights('/kaggle/working/checkpoint.weights.h5')","metadata":{"execution":{"iopub.status.busy":"2026-06-10T14:53:46.578816Z","iopub.execute_input":"2026-06-10T14:53:46.578968Z","iopub.status.idle":"2026-06-10T15:09:01.140610Z","shell.execute_reply.started":"2026-06-10T14:53:46.578954Z","shell.execute_reply":"2026-06-10T15:09:01.139149Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Visualize the training history","metadata":{}},{"cell_type":"code","source":"def plot_history_metrics(history):\n    \n    loss = history.history.pop('loss')\n    val_loss = history.history.pop('val_loss')\n\n    epochs = range(len(loss))\n    plt.plot(epochs, loss, 'r', label='Training Loss')\n    plt.plot(epochs, val_loss, 'b', label='Validation Loss')\n    plt.legend()\n    plt.title('Training and validation loss')\n    \n    plt.figure()\n    for key, values in history.history.items():\n        plt.plot(epochs, values, label=key)\n    plt.title('Training and validation metrics')\n    plt.legend()\n    plt.show()\n    \nplot_history_metrics(h)","metadata":{"execution":{"iopub.status.busy":"2026-06-10T15:09:01.141276Z","iopub.execute_input":"2026-06-10T15:09:01.141457Z","iopub.status.idle":"2026-06-10T15:09:01.393867Z","shell.execute_reply.started":"2026-06-10T15:09:01.141440Z","shell.execute_reply":"2026-06-10T15:09:01.392462Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"code","source":"import pandas as pd\n\nsubmission = pd.read_csv('/kaggle/input/tpu-getting-started/sample_submission.csv')\nsubmission","metadata":{"execution":{"iopub.status.busy":"2026-06-10T15:09:01.394479Z","iopub.execute_input":"2026-06-10T15:09:01.394656Z","iopub.status.idle":"2026-06-10T15:09:01.452113Z","shell.execute_reply.started":"2026-06-10T15:09:01.394639Z","shell.execute_reply":"2026-06-10T15:09:01.450906Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"prediction = pd.DataFrame()\ntest_images = np.array(list(DATASETS['test'].unbatch().map(lambda x, y: x).as_numpy_iterator()))\nprediction['id'] = np.array(list(DATASETS['test'].unbatch().map(lambda x, y: y).as_numpy_iterator())).astype(str)\nprint(test_images.shape)\nwith strategy.scope():\n    preds = model.predict(test_images, verbose=verbose)\nprediction['label'] = np.array([np.argmax(p) for p in preds])\nprediction","metadata":{"execution":{"iopub.status.busy":"2026-06-10T15:09:01.452650Z","iopub.execute_input":"2026-06-10T15:09:01.452825Z","iopub.status.idle":"2026-06-10T15:09:58.525808Z","shell.execute_reply.started":"2026-06-10T15:09:01.452810Z","shell.execute_reply":"2026-06-10T15:09:58.524358Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"my_submission = pd.merge(submission, prediction, on='id')\nmy_submission = my_submission.drop('label_x', axis=1).rename(columns={'label_y':'label'})\nmy_submission","metadata":{"execution":{"iopub.status.busy":"2026-06-10T15:09:58.526820Z","iopub.execute_input":"2026-06-10T15:09:58.527013Z","iopub.status.idle":"2026-06-10T15:09:58.542128Z","shell.execute_reply.started":"2026-06-10T15:09:58.526997Z","shell.execute_reply":"2026-06-10T15:09:58.541130Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"my_submission.to_csv('submission.csv',index=False)","metadata":{"execution":{"iopub.status.busy":"2026-06-10T15:09:58.542614Z","iopub.execute_input":"2026-06-10T15:09:58.542768Z","iopub.status.idle":"2026-06-10T15:09:58.554347Z","shell.execute_reply.started":"2026-06-10T15:09:58.542755Z","shell.execute_reply":"2026-06-10T15:09:58.553445Z"},"trusted":true},"outputs":[],"execution_count":null}]}