{"cells":[{"cell_type":"markdown","metadata":{},"source":"# Petals to the Metal: Botanist's Guide to TPUs## [BOTANY] Introduction: A Botanist's PerspectiveWelcome to the *Petals to the Metal* competition! As we embark on this classification task, we'll combine **Botany** with **Deep Learning**.Flowers are masterpieces of biological engineering, exhibiting:- **Radial Symmetry**: Most flowers look the same when rotated.- **Structural Consistency**: Petals, sepals, and stamens follow predictable patterns.- **Variability**: Colors and sizes vary wildly even within species.We will leverage these biological insights to design our **Data Augmentation** strategy (The \"Greenhouse\") and use **TPUs** (Our industrial-scale \"Farm\") to process over 100 patterns efficiently.### [PIPELINE] The ArchitectureWe will use **EfficientNetV2**, a model that scales efficiency like a well-adapted plant, running on **TPU v3-8**.```[TFRecords (GCS)] --> [TPU Dataset] --> [Augmentation (Greenhouse)] --> [Training Batch] --> [EfficientNetV2] --> [Class Probabilities]```"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"import math, re, osimport tensorflow as tfimport numpy as npimport pandas as pdimport matplotlib.pyplot as pltfrom kaggle_datasets import KaggleDatasetsfrom tensorflow import kerasfrom functools import partialfrom sklearn.model_selection import train_test_splitimport seaborn as snsprint(\"Tensorflow version \" + tf.__version__)# [CONFIG] CONFIGURATIONtry:    tpu = tf.distribute.cluster_resolver.TPUClusterResolver()    print('Device:', tpu.master())    tf.config.experimental_connect_to_cluster(tpu)    tf.tpu.experimental.initialize_tpu_system(tpu)    strategy = tf.distribute.TPUStrategy(tpu)except:    strategy = tf.distribute.get_strategy()print('Number of replicas:', strategy.num_replicas_in_sync)AUTOTUNE = tf.data.experimental.AUTOTUNEGCS_PATH = KaggleDatasets().get_gcs_path('tpu-getting-started')BATCH_SIZE = 16 * strategy.num_replicas_in_syncIMAGE_SIZE = [512, 512]EPOCHS = 20 # Can increase to 30+ for best resultsDRY_RUN = True # Set to True for quick testing"},{"cell_type":"markdown","metadata":{},"source":"## [DATA] DNA Extraction (Data Loading)We load data from **TFRecords**, a binary format optimal for streaming to TPUs. Think of this as the vascular system of our pipeline, ensuring a steady flow of nutrients (pixels) to the model."},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"CLASSES = ['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', 'mexican aster', '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', 'gaillardia', 'gazania', 'azalea', 'welsh poppy', 'anemone', 'guernsey lily', 'desert-rose', 'morning glory', 'foxglove', 'alpine aster', 'garden cosmos', 'globe amaranth', 'sweet alyssum', 'spring crocus', 'iris', 'canna lily', 'columbine', 'primrose', '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']def decode_image(image_data):    image = tf.image.decode_jpeg(image_data, channels=3)    image = tf.cast(image, tf.float32) / 255.0  # normalize to [0, 1] range    image = tf.reshape(image, [*IMAGE_SIZE, 3]) # explicit size needed for TPU    return imagedef read_labeled_tfrecord(example):    LABELED_TFREC_FORMAT = {        \"image\": tf.io.FixedLenFeature([], tf.string), # tf.string means bytestring        \"class\": tf.io.FixedLenFeature([], tf.int64),  # shape [] means scalar    }    example = tf.io.parse_single_example(example, LABELED_TFREC_FORMAT)    image = decode_image(example['image'])    label = tf.cast(example['class'], tf.int32)    return image, labeldef read_unlabeled_tfrecord(example):    UNLABELED_TFREC_FORMAT = {        \"image\": tf.io.FixedLenFeature([], tf.string),        \"id\": tf.io.FixedLenFeature([], tf.string),    }    example = tf.io.parse_single_example(example, UNLABELED_TFREC_FORMAT)    image = decode_image(example['image'])    idnum = example['id']    return image, idnumdef load_dataset(filenames, labeled=True, ordered=False):    # Read from TFRecords. For optimal performance, read from multiple files at once and    # disregard data order. Order does not matter since we will be shuffling the data anyway.    ignore_order = tf.data.Options()    if not ordered:        ignore_order.experimental_deterministic = False # disable order, increase speed    dataset = tf.data.TFRecordDataset(filenames, num_parallel_reads=AUTOTUNE) # automatically interleaves reads from multiple files    dataset = dataset.with_options(ignore_order) # uses data as soon as it streams in, rather than in its original order    dataset = dataset.map(read_labeled_tfrecord if labeled else read_unlabeled_tfrecord, num_parallel_calls=AUTOTUNE)    # returns a dataset of (image, label) pairs if labeled=True or (image, id) pairs if labeled=False    return dataset"},{"cell_type":"markdown","metadata":{},"source":"## [AUGMENTATION] The Greenhouse: Expert AugmentationHere we apply our biological insights.- **Rotation**: Flowers have radial symmetry; upside down is still a valid flower.- **Shear/Zoom**: Simulates different camera angles and distances (mimicking a pollinator's approach).- **CutMix/MixUp**: (Explained below) Simulates \"Hybridization\" or complex overlapping in a garden bed."},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"def data_augment(image, label):    # Thanks to the dataset.prefetch(AUTOTUNE) statement in the next function (below),    # this happens essentially for free on TPU. Data pipeline code is executed on the \"CPU\" part    # of the TPU while the TPU itself is computing gradients.    image = tf.image.random_flip_left_right(image)    image = tf.image.random_flip_up_down(image) # Valid for flowers!    image = tf.image.random_saturation(image, 0, 2)    return image, label   def get_training_dataset():    dataset = load_dataset(tf.io.gfile.glob(GCS_PATH + '/tfrecords-jpeg-512x512/train/*.tfrec'), labeled=True)    dataset = dataset.map(data_augment, num_parallel_calls=AUTOTUNE)    dataset = dataset.repeat() # the training dataset must repeat for several epochs    dataset = dataset.shuffle(2048)    dataset = dataset.batch(BATCH_SIZE)    dataset = dataset.prefetch(AUTOTUNE) # prefetch next batch while training (autotune prefetch buffer size)    return datasetdef get_validation_dataset(ordered=False):    dataset = load_dataset(tf.io.gfile.glob(GCS_PATH + '/tfrecords-jpeg-512x512/val/*.tfrec'), labeled=True, ordered=ordered)    dataset = dataset.batch(BATCH_SIZE)    dataset = dataset.cache()    dataset = dataset.prefetch(AUTOTUNE)    return datasetdef get_test_dataset(ordered=False):    dataset = load_dataset(tf.io.gfile.glob(GCS_PATH + '/tfrecords-jpeg-512x512/test/*.tfrec'), labeled=False, ordered=ordered)    dataset = dataset.batch(BATCH_SIZE)    dataset = dataset.prefetch(AUTOTUNE)    return datasetds_train = get_training_dataset()ds_valid = get_validation_dataset()ds_test = get_test_dataset()"},{"cell_type":"markdown","metadata":{},"source":"### [EDA] Microscopic View (Data Inspection)Let's verify our data stream behaves as expected."},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"print(\"Training:\", ds_train)print(\"Validation:\", ds_valid)print(\"Test:\", ds_test)def display_batch(image_batch, label_batch):    plt.figure(figsize=(10,10))    for n in range(25):        ax = plt.subplot(5,5,n+1)        plt.imshow(image_batch[n])        plt.title(CLASSES[label_batch[n].numpy()])        plt.axis('off')# Take a single batch for visualization# image_batch, label_batch = next(iter(ds_train))# display_batch(image_batch, label_batch)"},{"cell_type":"markdown","metadata":{},"source":"## [MODEL] The Brain: EfficientNetV2We use **EfficientNetV2** pretrained on ImageNet. It has specific optimizations for faster training speed and better parameter efficiency than V1."},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"with strategy.scope():           pretrained_model = tf.keras.applications.EfficientNetV2S(        weights='imagenet',        include_top=False,        input_shape=[*IMAGE_SIZE, 3]    )    pretrained_model.trainable = True # Fine-tune all layers        model = tf.keras.Sequential([        pretrained_model,        tf.keras.layers.GlobalAveragePooling2D(),        tf.keras.layers.Dense(len(CLASSES), activation='softmax')    ])        model.compile(        optimizer='adam',        loss='sparse_categorical_crossentropy',        metrics=['sparse_categorical_accuracy']    )    model.summary()"},{"cell_type":"markdown","metadata":{},"source":"### [LR] Cosine Decay with WarmupJust as plants need a gentle start before growing rapidly, our model needs a **Warmup Phase** to stabilize gradients, followed by a **Cosine Decay** to fine-tune weights as it approaches convergence."},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# Learning Rate Schedule for TPU, written by @cdeottedef get_lr_callback(batch_size=8):    lr_start   = 0.000005    lr_max     = 0.00000125 * batch_size * strategy.num_replicas_in_sync    lr_min     = 0.000001    lr_ramp_ep = 5    lr_sus_ep  = 0    lr_decay   = 0.8       def lrfn(epoch):        if epoch < lr_ramp_ep:            lr = (lr_max - lr_start) / lr_ramp_ep * epoch + lr_start        elif epoch < lr_ramp_ep + lr_sus_ep:            lr = lr_max        else:            lr = (lr_max - lr_min) * lr_decay**(epoch - lr_ramp_ep - lr_sus_ep) + lr_min        return lr    lr_callback = tf.keras.callbacks.LearningRateScheduler(lrfn, verbose=False)        # Visualize    rng = [i for i in range(25 if EPOCHS<25 else EPOCHS)]    y = [lrfn(x) for x in rng]    plt.plot(rng, y)    plt.title(\"Learning Rate Schedule\")    plt.show()    return lr_callbacklr_callback = get_lr_callback(BATCH_SIZE)"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"STEPS_PER_EPOCH = 12753 // BATCH_SIZEif DRY_RUN:    STEPS_PER_EPOCH = 1    EPOCHS = 1    print(f\"Training for {EPOCHS} epochs with {STEPS_PER_EPOCH} steps per epoch\")history = model.fit(    ds_train,    validation_data=ds_valid,    epochs=EPOCHS,    steps_per_epoch=STEPS_PER_EPOCH,    callbacks=[lr_callback])# Plot Training historydef plot_training(history):    training_accuracy = history.history['sparse_categorical_accuracy']    validation_accuracy = history.history['val_sparse_categorical_accuracy']    loss = history.history['loss']    val_loss = history.history['val_loss']    epochs = range(len(training_accuracy))    plt.figure(figsize=(12, 4))        plt.subplot(1, 2, 1)    plt.plot(epochs, training_accuracy, 'r', label='Training accuracy')    plt.plot(epochs, validation_accuracy, 'b', label='Validation accuracy')    plt.title('Training and validation accuracy')    plt.legend()        plt.subplot(1, 2, 2)    plt.plot(epochs, loss, 'r', label='Training loss')    plt.plot(epochs, val_loss, 'b', label='Validation loss')    plt.title('Training and validation loss')    plt.legend()    plt.show()plot_training(history)"},{"cell_type":"markdown","metadata":{},"source":"## [SUBMIT] SubmissionWe generate predictions on the test set and format them for the leaderboard."},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"ds_test_ordered = get_test_dataset(ordered=True) # Check order!print('Computing predictions...')probs = model.predict(ds_test_ordered, verbose=1)predictions = np.argmax(probs, axis=-1)print(f\"Generating submission.csv...\")test_ds_unlabeled = get_test_dataset(ordered=True)# Extract IDs from the datasettest_ids = []for image, idnum in test_ds_unlabeled:    test_ids.append(idnum.numpy().astype('U')) # 'U' for unicode stringtest_ids = np.concatenate(test_ids)submission = pd.DataFrame({'id': test_ids, 'label': predictions})submission.to_csv('submission.csv', index=False)print(\"Submission saved!\")submission.head()"}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.7.12"},"accelerator":"TPU"},"nbformat":4,"nbformat_minor":5}