{"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.8.17"},"papermill":{"default_parameters":{},"duration":5389.22296,"end_time":"2023-10-19T15:56:10.510034","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2023-10-19T14:26:21.287074","version":"2.4.0"},"kaggle":{"accelerator":"tpu1vmV38","dataSources":[{"sourceId":21154,"databundleVersionId":1243559,"sourceType":"competition"},{"sourceId":1138814,"sourceType":"datasetVersion","datasetId":601927}],"dockerImageVersionId":30589,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Kaggle Inputs - visualising paths","metadata":{"papermill":{"duration":0.008516,"end_time":"2023-10-19T14:26:23.506189","exception":false,"start_time":"2023-10-19T14:26:23.497673","status":"completed"},"tags":[]}},{"cell_type":"code","source":"!pip install -qU wandb","metadata":{"execution":{"iopub.execute_input":"2023-10-19T14:26:23.521002Z","iopub.status.busy":"2023-10-19T14:26:23.520495Z","iopub.status.idle":"2023-10-19T14:26:32.117609Z","shell.execute_reply":"2023-10-19T14:26:32.116504Z"},"papermill":{"duration":8.607081,"end_time":"2023-10-19T14:26:32.120159","exception":false,"start_time":"2023-10-19T14:26:23.513078","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!git clone https://github.com/rishigami/Swin-Transformer-TF\n    \nimport sys\nsys.path.append('/kaggle/working/Swin-Transformer-TF')","metadata":{"execution":{"iopub.execute_input":"2023-10-19T14:26:32.135711Z","iopub.status.busy":"2023-10-19T14:26:32.135437Z","iopub.status.idle":"2023-10-19T14:26:33.918648Z","shell.execute_reply":"2023-10-19T14:26:33.917778Z"},"papermill":{"duration":1.793677,"end_time":"2023-10-19T14:26:33.921038","exception":false,"start_time":"2023-10-19T14:26:32.127361","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 0 Imports","metadata":{"papermill":{"duration":0.008282,"end_time":"2023-10-19T14:26:33.937270","exception":false,"start_time":"2023-10-19T14:26:33.928988","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import math, re, os, random\nimport numpy as np\nimport pandas as pd\nimport wandb\nfrom wandb.integration.keras import WandbCallback\nfrom matplotlib import pyplot as plt\nfrom sklearn.metrics import f1_score, precision_score, \\\n                            recall_score, confusion_matrix\n\nimport tensorflow as tf\nfrom tensorflow_addons.metrics import F1Score\nfrom tensorflow.keras import layers as L\nfrom tensorflow.keras import callbacks\nfrom swintransformer import SwinTransformer\n\nfrom kaggle_datasets import KaggleDatasets\n\nprint(\"TF version \" + tf.__version__)","metadata":{"execution":{"iopub.execute_input":"2023-10-19T14:26:33.955041Z","iopub.status.busy":"2023-10-19T14:26:33.954706Z","iopub.status.idle":"2023-10-19T14:27:17.884013Z","shell.execute_reply":"2023-10-19T14:27:17.883149Z"},"papermill":{"duration":43.941308,"end_time":"2023-10-19T14:27:17.886537","exception":false,"start_time":"2023-10-19T14:26:33.945229","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"try:\n    from kaggle_secrets import UserSecretsClient\n    user_secrets = UserSecretsClient()\n    api_key = user_secrets.get_secret('wandb_key')\n    wandb.login(key=api_key)\n    anonymous = None\nexcept:\n    wandb.login(anonymous='must')\n    print('To use your W&B account,\\nGo to Add-ons -> Secrets and provide your \\\n           W&B access token. Use the Label name as WANDB. \\nGet your W&B access \\\n           token from here: https://wandb.ai/authorize')","metadata":{"execution":{"iopub.execute_input":"2023-10-19T14:27:17.902748Z","iopub.status.busy":"2023-10-19T14:27:17.902099Z","iopub.status.idle":"2023-10-19T14:27:19.471567Z","shell.execute_reply":"2023-10-19T14:27:19.470306Z"},"papermill":{"duration":1.58011,"end_time":"2023-10-19T14:27:19.473737","exception":false,"start_time":"2023-10-19T14:27:17.893627","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 1 TPU - Distribution Strategy","metadata":{"papermill":{"duration":0.007093,"end_time":"2023-10-19T14:27:19.488223","exception":false,"start_time":"2023-10-19T14:27:19.481130","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"> A TPU has eight different cores and each of these cores acts as its own accelerator. (A TPU is sort of like having eight GPUs in one machine.) \n> We tell TensorFlow how to make use of all these cores at once through a distribution strategy. \n> Run the following cell to create the distribution strategy that we'll later apply to our model.\n\n> We'll use the distribution strategy when we create our neural network model. Then, TensorFlow will distribute the training among the eight TPU cores by creating eight different replicas of the model, one for each core.\n","metadata":{"papermill":{"duration":0.00677,"end_time":"2023-10-19T14:27:19.502760","exception":false,"start_time":"2023-10-19T14:27:19.495990","status":"completed"},"tags":[]}},{"cell_type":"code","source":"AUTO = tf.data.experimental.AUTOTUNE\n\n# 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":{"execution":{"iopub.execute_input":"2023-10-19T14:27:19.519085Z","iopub.status.busy":"2023-10-19T14:27:19.518775Z","iopub.status.idle":"2023-10-19T14:27:28.644355Z","shell.execute_reply":"2023-10-19T14:27:28.643147Z"},"papermill":{"duration":9.14108,"end_time":"2023-10-19T14:27:28.651097","exception":false,"start_time":"2023-10-19T14:27:19.510017","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 2 Configuration","metadata":{"papermill":{"duration":0.008874,"end_time":"2023-10-19T14:27:28.669463","exception":false,"start_time":"2023-10-19T14:27:28.660589","status":"completed"},"tags":[]}},{"cell_type":"code","source":"SEED = 42\n\nIMAGE_SIZE = [224, 224]\nEPOCHS = 22\nBATCH_SIZE = 32 * strategy.num_replicas_in_sync\n\nSWIN_TYPE = 'large'","metadata":{"execution":{"iopub.execute_input":"2023-10-19T14:27:28.688577Z","iopub.status.busy":"2023-10-19T14:27:28.688303Z","iopub.status.idle":"2023-10-19T14:27:28.692812Z","shell.execute_reply":"2023-10-19T14:27:28.691934Z"},"papermill":{"duration":0.016081,"end_time":"2023-10-19T14:27:28.694432","exception":false,"start_time":"2023-10-19T14:27:28.678351","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed):\n    np.random.seed(seed)\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    tf.random.set_seed(seed)\n    \nset_seed(SEED)","metadata":{"execution":{"iopub.execute_input":"2023-10-19T14:27:28.713303Z","iopub.status.busy":"2023-10-19T14:27:28.713046Z","iopub.status.idle":"2023-10-19T14:27:28.718157Z","shell.execute_reply":"2023-10-19T14:27:28.717301Z"},"papermill":{"duration":0.016539,"end_time":"2023-10-19T14:27:28.719824","exception":false,"start_time":"2023-10-19T14:27:28.703285","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 2 Loading the Competition Data","metadata":{"execution":{"iopub.execute_input":"2023-10-17T13:25:00.929228Z","iopub.status.busy":"2023-10-17T13:25:00.928911Z","iopub.status.idle":"2023-10-17T13:25:00.933320Z","shell.execute_reply":"2023-10-17T13:25:00.932509Z","shell.execute_reply.started":"2023-10-17T13:25:00.929203Z"},"papermill":{"duration":0.008701,"end_time":"2023-10-19T14:27:28.737491","exception":false,"start_time":"2023-10-19T14:27:28.728790","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"> When used with TPUs, datasets need to be stored in a Google Cloud Storage bucket. You can use data from any public GCS bucket by giving its path just like you would data from '/kaggle/input'. The following will retrieve the GCS path for this competition's dataset.\n\n> When used with TPUs, datasets are often serialized into TFRecords. This is a format convenient for distributing data to each of the TPUs cores. We've hidden the cell that reads the TFRecords for our dataset since the process is a bit long. You could come back to it later for some guidance on using your own datasets with TPUs.","metadata":{"papermill":{"duration":0.009173,"end_time":"2023-10-19T14:27:28.755220","exception":false,"start_time":"2023-10-19T14:27:28.746047","status":"completed"},"tags":[]}},{"cell_type":"code","source":"GCS_DS_PATH = KaggleDatasets().get_gcs_path('tpu-getting-started')\nGCS_DS_PATH_EXT = KaggleDatasets().get_gcs_path('tf-flower-photo-tfrec')\nGCS_DS_PATH_ERASED = KaggleDatasets().get_gcs_path('tpu-getting-started-erased-512')\nGCS_DS_PATH_EXT_ERASED = KaggleDatasets().get_gcs_path('tf-flower-photo-tfrec-erased-512')\nprint(GCS_DS_PATH) # what do gcs paths look like?\nprint(GCS_DS_PATH_EXT)\nprint(GCS_DS_PATH_ERASED)\nprint(GCS_DS_PATH_EXT_ERASED)","metadata":{"execution":{"iopub.execute_input":"2023-10-19T14:27:28.773972Z","iopub.status.busy":"2023-10-19T14:27:28.773732Z","iopub.status.idle":"2023-10-19T14:27:28.779169Z","shell.execute_reply":"2023-10-19T14:27:28.778365Z"},"papermill":{"duration":0.017154,"end_time":"2023-10-19T14:27:28.781110","exception":false,"start_time":"2023-10-19T14:27:28.763956","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DS_PATH = '/kaggle/input/tpu-getting-started'\nDS_PATH_EXT = '/kaggle/input/tf-flower-photo-tfrec'\n\nDATA_PATH_SELECT = { # available image sizes\n    192: DS_PATH + '/tfrecords-jpeg-192x192',\n    224: DS_PATH + '/tfrecords-jpeg-224x224',\n    331: DS_PATH + '/tfrecords-jpeg-331x331',\n    512: DS_PATH + '/tfrecords-jpeg-512x512'\n}\nDATA_PATH = DATA_PATH_SELECT[IMAGE_SIZE[0]]\n\n# External data\nDATA_PATH_SELECT_EXT = {\n    192: '/tfrecords-jpeg-192x192',\n    224: '/tfrecords-jpeg-224x224',\n    331: '/tfrecords-jpeg-331x331',\n    512: '/tfrecords-jpeg-512x512'\n}\nDATA_PATH_EXT = DATA_PATH_SELECT_EXT[IMAGE_SIZE[0]]\n\nIMAGENET_FILES = tf.io.gfile.glob(DS_PATH_EXT + '/imagenet' + DATA_PATH_EXT + '/*.tfrec')\nINATURELIST_FILES = tf.io.gfile.glob(DS_PATH_EXT + '/inaturalist' + DATA_PATH_EXT + '/*.tfrec')\nOPENIMAGE_FILES = tf.io.gfile.glob(DS_PATH_EXT + '/openimage' + DATA_PATH_EXT + '/*.tfrec')\nOXFORD_FILES = tf.io.gfile.glob(DS_PATH_EXT + '/oxford_102' + DATA_PATH_EXT + '/*.tfrec')\nTENSORFLOW_FILES = tf.io.gfile.glob(DS_PATH_EXT + '/tf_flowers' + DATA_PATH_EXT + '/*.tfrec')\n\nADDITIONAL_TRAINING_FILENAMES = IMAGENET_FILES + INATURELIST_FILES + OPENIMAGE_FILES + OXFORD_FILES + TENSORFLOW_FILES  \n\n\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\nTRAINING_FILENAMES = tf.io.gfile.glob(DATA_PATH + '/train/*.tfrec')\nTRAINING_FILENAMES = TRAINING_FILENAMES + ADDITIONAL_TRAINING_FILENAMES\nVALIDATION_FILENAMES = tf.io.gfile.glob(DATA_PATH + '/val/*.tfrec')\nTEST_FILENAMES = tf.io.gfile.glob(DATA_PATH + '/test/*.tfrec') # predictions on this dataset should be submitted for the competition ","metadata":{"execution":{"iopub.execute_input":"2023-10-19T14:27:28.800331Z","iopub.status.busy":"2023-10-19T14:27:28.800096Z","iopub.status.idle":"2023-10-19T14:27:28.891215Z","shell.execute_reply":"2023-10-19T14:27:28.890243Z"},"papermill":{"duration":0.103142,"end_time":"2023-10-19T14:27:28.893197","exception":false,"start_time":"2023-10-19T14:27:28.790055","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 3 Visualization functions","metadata":{"papermill":{"duration":0.009653,"end_time":"2023-10-19T14:27:28.912944","exception":false,"start_time":"2023-10-19T14:27:28.903291","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# numpy and matplotlib defaults\nnp.set_printoptions(threshold=15, linewidth=80)\n\ndef batch_to_numpy_images_and_labels(data):\n    images, labels = data\n    numpy_images = images.numpy()\n    numpy_labels = labels.numpy()\n    if numpy_labels.dtype == object: # binary string in this case, these are image ID strings\n        numpy_labels = [None for _ in enumerate(numpy_images)]\n    # If no labels, only image IDs, return None for labels (this is the case for test data)\n    return numpy_images, numpy_labels\n\ndef title_from_label_and_target(label, correct_label):\n    if correct_label is None:\n        return CLASSES[label], True\n    correct = (label == correct_label)\n    return \"{} [{}{}{}]\".format(CLASSES[label], 'OK' if correct else 'NO', u\"\\u2192\" if not correct else '',\n                                CLASSES[correct_label] if not correct else ''), correct\n\ndef display_one_flower(image, title, subplot, red=False, titlesize=16):\n    plt.subplot(*subplot)\n    plt.axis('off')\n    plt.imshow(image)\n    if len(title) > 0:\n        plt.title(title, fontsize=int(titlesize) if not red else int(titlesize/1.2), color='red' if red else 'black', fontdict={'verticalalignment':'center'}, pad=int(titlesize/1.5))\n    return (subplot[0], subplot[1], subplot[2]+1)\n\ndef display_batch_of_images(databatch, predictions=None):\n    \"\"\"This will work with:\n    display_batch_of_images(images)\n    display_batch_of_images(images, predictions)\n    display_batch_of_images((images, labels))\n    display_batch_of_images((images, labels), predictions)\n    \"\"\"\n    # data\n    images, labels = batch_to_numpy_images_and_labels(databatch)\n    if labels is None:\n        labels = [None for _ in enumerate(images)]\n        \n    # auto-squaring: this will drop data that does not fit into square or square-ish rectangle\n    rows = int(math.sqrt(len(images)))\n    cols = len(images)//rows\n        \n    # size and spacing\n    FIGSIZE = 13.0\n    SPACING = 0.1\n    subplot=(rows,cols,1)\n    if rows < cols:\n        plt.figure(figsize=(FIGSIZE,FIGSIZE/cols*rows))\n    else:\n        plt.figure(figsize=(FIGSIZE/rows*cols,FIGSIZE))\n    \n    # display\n    for i, (image, label) in enumerate(zip(images[:rows*cols], labels[:rows*cols])):\n        title = '' if label is None else CLASSES[label]\n        correct = True\n        if predictions is not None:\n            title, correct = title_from_label_and_target(predictions[i], label)\n        dynamic_titlesize = FIGSIZE*SPACING/max(rows,cols)*40+3 # magic formula tested to work from 1x1 to 10x10 images\n        subplot = display_one_flower(image, title, subplot, not correct, titlesize=dynamic_titlesize)\n    \n    #layout\n    plt.tight_layout()\n    if label is None and predictions is None:\n        plt.subplots_adjust(wspace=0, hspace=0)\n    else:\n        plt.subplots_adjust(wspace=SPACING, hspace=SPACING)\n    plt.show()\n    \ndef display_confusion_matrix(cmat, score, precision, recall):\n    plt.figure(figsize=(15,15))\n    ax = plt.gca()\n    ax.matshow(cmat, cmap='Reds')\n    ax.set_xticks(range(len(CLASSES)))\n    ax.set_xticklabels(CLASSES, fontdict={'fontsize': 7})\n    plt.setp(ax.get_xticklabels(), rotation=45, ha=\"left\", rotation_mode=\"anchor\")\n    ax.set_yticks(range(len(CLASSES)))\n    ax.set_yticklabels(CLASSES, fontdict={'fontsize': 7})\n    plt.setp(ax.get_yticklabels(), rotation=45, ha=\"right\", rotation_mode=\"anchor\")\n    titlestring = \"\"\n    if score is not None:\n        titlestring += 'f1 = {:.3f} '.format(score)\n    if precision is not None:\n        titlestring += '\\nprecision = {:.3f} '.format(precision)\n    if recall is not None:\n        titlestring += '\\nrecall = {:.3f} '.format(recall)\n    if len(titlestring) > 0:\n        ax.text(101, 1, titlestring, fontdict={'fontsize': 18, 'horizontalalignment':'right', 'verticalalignment':'top', 'color':'#804040'})\n    plt.show()\n    \ndef display_training_curves(training, validation, title, subplot):\n    if subplot%10==1: # set up the subplots on the first call\n        plt.subplots(figsize=(10,10), facecolor='#F0F0F0')\n        plt.tight_layout()\n    ax = plt.subplot(subplot)\n    ax.set_facecolor('#F8F8F8')\n    ax.plot(training)\n    ax.plot(validation)\n    ax.set_title('model '+ title)\n    ax.set_ylabel(title)\n    #ax.set_ylim(0.28,1.05)\n    ax.set_xlabel('epoch')\n    ax.legend(['train', 'valid.'])","metadata":{"execution":{"iopub.execute_input":"2023-10-19T14:27:28.932686Z","iopub.status.busy":"2023-10-19T14:27:28.932397Z","iopub.status.idle":"2023-10-19T14:27:28.952172Z","shell.execute_reply":"2023-10-19T14:27:28.951363Z"},"papermill":{"duration":0.031731,"end_time":"2023-10-19T14:27:28.953647","exception":false,"start_time":"2023-10-19T14:27:28.921916","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 4 Random erasing (blockout) augmentation","metadata":{"papermill":{"duration":0.008857,"end_time":"2023-10-19T14:27:28.971980","exception":false,"start_time":"2023-10-19T14:27:28.963123","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# https://www.kaggle.com/tusharkendre/tpu-flowers\ndef random_erasing(img, sl=0.1, sh=0.2, rl=0.4, p=0.3):\n    h = tf.shape(img)[0]\n    w = tf.shape(img)[1]\n    c = tf.shape(img)[2]\n    origin_area = tf.cast(h*w, tf.float32)\n\n    e_size_l = tf.cast(tf.round(tf.sqrt(origin_area * sl * rl)), tf.int32)\n    e_size_h = tf.cast(tf.round(tf.sqrt(origin_area * sh / rl)), tf.int32)\n\n    e_height_h = tf.minimum(e_size_h, h)\n    e_width_h = tf.minimum(e_size_h, w)\n\n    erase_height = tf.random.uniform(shape=[], minval=e_size_l, maxval=e_height_h, dtype=tf.int32)\n    erase_width = tf.random.uniform(shape=[], minval=e_size_l, maxval=e_width_h, dtype=tf.int32)\n\n    erase_area = tf.zeros(shape=[erase_height, erase_width, c])\n    erase_area = tf.cast(erase_area, tf.uint8)\n\n    pad_h = h - erase_height\n    pad_top = tf.random.uniform(shape=[], minval=0, maxval=pad_h, dtype=tf.int32)\n    pad_bottom = pad_h - pad_top\n\n    pad_w = w - erase_width\n    pad_left = tf.random.uniform(shape=[], minval=0, maxval=pad_w, dtype=tf.int32)\n    pad_right = pad_w - pad_left\n\n    erase_mask = tf.pad([erase_area], [[0,0],[pad_top, pad_bottom], [pad_left, pad_right], [0,0]], constant_values=1)\n    erase_mask = tf.squeeze(erase_mask, axis=0)\n    erased_img = tf.multiply(tf.cast(img,tf.float32), tf.cast(erase_mask, tf.float32))\n\n    return tf.cond(tf.random.uniform([], 0, 1) > p, lambda: tf.cast(img, img.dtype), lambda:  tf.cast(erased_img, img.dtype))","metadata":{"execution":{"iopub.execute_input":"2023-10-19T14:27:28.991205Z","iopub.status.busy":"2023-10-19T14:27:28.990959Z","iopub.status.idle":"2023-10-19T14:27:29.001284Z","shell.execute_reply":"2023-10-19T14:27:29.000470Z"},"papermill":{"duration":0.022015,"end_time":"2023-10-19T14:27:29.002890","exception":false,"start_time":"2023-10-19T14:27:28.980875","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The `random_erasing` function appears to be implementing the \"Random Erasing\" data augmentation technique for images. It generates an augmented version of an input image by randomly erasing a rectangular or square region within the image with a certain probability. Here's a breakdown of what this function does:\n\n1. Calculate the height (`h`), width (`w`), and number of channels (`c`) of the input image `img`.\n\n2. Determine the area of the original image `origin_area` (product of `h` and `w`).\n\n3. Calculate the lower and upper limits for the size of the erasing region:\n   - `e_size_l`: The lower limit is determined by scaling a small value (`sl`) with respect to the original image area and a \"regularization factor\" (`rl`). It's rounded to the nearest integer.\n   - `e_size_h`: The upper limit is calculated similarly using `sh` and `rl`.\n\n4. Generate random values for the height and width of the erasing region:\n   - `erase_height`: Randomly selected height for erasing, bounded by the limits `e_size_l` and `e_height_h`.\n   - `erase_width`: Randomly selected width for erasing, bounded by the same limits.\n\n5. Create an erase area with the selected height and width. This area is initially filled with zeros.\n\n6. Calculate the amount of padding required around the erasing region to fit it back into the original image:\n   - `pad_h`: The remaining height after erasing.\n   - `pad_top` and `pad_bottom`: Randomly determined padding values for the top and bottom of the erasing region.\n\n7. Calculate the padding required for the width:\n   - `pad_w`: The remaining width after erasing.\n   - `pad_left` and `pad_right`: Randomly determined padding values for the left and right of the erasing region.\n\n8. Create an `erase_mask` that combines the erase area with padding. This mask is used to erase the selected region by setting the corresponding pixel values to zero.\n\n9. Multiply the original image with the `erase_mask` to create an erased version of the image, where the selected region is zeroed out.\n\n10. Finally, with a certain probability (`p`), return the erased image; otherwise, return the original image.\n\nThis function can be applied to each image in a dataset as part of a data augmentation pipeline to create more diverse training examples for neural networks. It helps the model become more robust and better generalize to different types of image occlusions and variations. The random parameters ensure that each augmented image is unique.","metadata":{"papermill":{"duration":0.008844,"end_time":"2023-10-19T14:27:29.020638","exception":false,"start_time":"2023-10-19T14:27:29.011794","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# 5 Datasets functions","metadata":{"papermill":{"duration":0.008754,"end_time":"2023-10-19T14:27:29.038272","exception":false,"start_time":"2023-10-19T14:27:29.029518","status":"completed"},"tags":[]}},{"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) / 255.0  # convert image to floats in [0, 1] range\n    image = tf.reshape(image, [*IMAGE_SIZE, 3]) # explicit size needed for TPU\n    return image\n\ndef onehot(image,label):\n    return image,tf.one_hot(label, len(CLASSES))\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\n\ndef data_augment(image, label):\n    image = tf.image.random_flip_left_right(image)\n    image = random_erasing(image)\n    return image, label\n\ndef data_hflip(image, idnum):\n    image = tf.image.flip_left_right(image)\n    return image, idnum\n\ndef get_training_dataset(do_onehot=False):\n    dataset = load_dataset(TRAINING_FILENAMES, labeled=True)\n    dataset = dataset.map(data_augment, num_parallel_calls=AUTO)\n    if do_onehot:\n        dataset = dataset.map(onehot, num_parallel_calls=AUTO)\n    dataset = dataset.repeat() # the training dataset must repeat for several epochs\n    dataset = dataset.shuffle(2048, reshuffle_each_iteration=True)\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, do_onehot=False):\n    dataset = load_dataset(VALIDATION_FILENAMES, labeled=True, ordered=ordered)\n    if do_onehot:\n        dataset = dataset.map(onehot, num_parallel_calls=AUTO)\n    dataset = dataset.batch(BATCH_SIZE)\n    dataset = dataset.cache()\n    dataset = dataset.prefetch(AUTO) # prefetch next batch while training (autotune prefetch buffer size)\n    return dataset\n\ndef get_test_dataset(ordered=False, augmented=False):\n    dataset = load_dataset(TEST_FILENAMES, labeled=False, ordered=ordered)\n    dataset = dataset.map(data_hflip, num_parallel_calls=AUTO)\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 count_data_items(filenames):\n    # the number of data items is written in the name of the .tfrec 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)\nSTEPS_PER_EPOCH = NUM_TRAINING_IMAGES // BATCH_SIZE\nprint(f'Dataset: {NUM_TRAINING_IMAGES} training images, {NUM_VALIDATION_IMAGES} validation images, {NUM_TEST_IMAGES} unlabeled test images')","metadata":{"execution":{"iopub.execute_input":"2023-10-19T14:27:29.057107Z","iopub.status.busy":"2023-10-19T14:27:29.056832Z","iopub.status.idle":"2023-10-19T14:27:29.074352Z","shell.execute_reply":"2023-10-19T14:27:29.073537Z"},"papermill":{"duration":0.029041,"end_time":"2023-10-19T14:27:29.076047","exception":false,"start_time":"2023-10-19T14:27:29.047006","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Datasets visualization","metadata":{"papermill":{"duration":0.008921,"end_time":"2023-10-19T14:27:29.093859","exception":false,"start_time":"2023-10-19T14:27:29.084938","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# data dump\nprint(\"Training data shapes:\")\nfor image, label in get_training_dataset().take(3):\n    print(image.numpy().shape, label.numpy().shape)\nprint(\"Training data label examples:\", label.numpy())\nprint(\"Validation data shapes:\")\nfor image, label in get_validation_dataset().take(3):\n    print(image.numpy().shape, label.numpy().shape)\nprint(\"Validation data label examples:\", label.numpy())\nprint(\"Test data shapes:\")\nfor image, idnum in get_test_dataset().take(3):\n    print(image.numpy().shape, idnum.numpy().shape)\nprint(\"Test data IDs:\", idnum.numpy().astype('U')) # U=unicode string","metadata":{"execution":{"iopub.execute_input":"2023-10-19T14:27:29.113120Z","iopub.status.busy":"2023-10-19T14:27:29.112835Z","iopub.status.idle":"2023-10-19T14:27:31.765205Z","shell.execute_reply":"2023-10-19T14:27:31.763920Z"},"papermill":{"duration":2.664722,"end_time":"2023-10-19T14:27:31.767507","exception":false,"start_time":"2023-10-19T14:27:29.102785","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Peek at training data\ntraining_dataset = get_training_dataset()\ntraining_dataset = training_dataset.unbatch().batch(20)\ntrain_batch = iter(training_dataset)","metadata":{"execution":{"iopub.execute_input":"2023-10-19T14:27:31.790118Z","iopub.status.busy":"2023-10-19T14:27:31.789806Z","iopub.status.idle":"2023-10-19T14:27:31.915278Z","shell.execute_reply":"2023-10-19T14:27:31.914163Z"},"papermill":{"duration":0.138911,"end_time":"2023-10-19T14:27:31.917837","exception":false,"start_time":"2023-10-19T14:27:31.778926","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# run this cell again for next set of images\ndisplay_batch_of_images(next(train_batch))","metadata":{"execution":{"iopub.execute_input":"2023-10-19T14:27:31.939357Z","iopub.status.busy":"2023-10-19T14:27:31.939075Z","iopub.status.idle":"2023-10-19T14:27:34.705230Z","shell.execute_reply":"2023-10-19T14:27:34.704071Z"},"papermill":{"duration":2.806256,"end_time":"2023-10-19T14:27:34.734155","exception":false,"start_time":"2023-10-19T14:27:31.927899","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# peer at test data\ntest_dataset = get_test_dataset()\ntest_dataset = test_dataset.unbatch().batch(20)\ntest_batch = iter(test_dataset)","metadata":{"execution":{"iopub.execute_input":"2023-10-19T14:27:34.793169Z","iopub.status.busy":"2023-10-19T14:27:34.792803Z","iopub.status.idle":"2023-10-19T14:27:34.868892Z","shell.execute_reply":"2023-10-19T14:27:34.867610Z"},"papermill":{"duration":0.110529,"end_time":"2023-10-19T14:27:34.871047","exception":false,"start_time":"2023-10-19T14:27:34.760518","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# run this cell again for next set of images\ndisplay_batch_of_images(next(test_batch))","metadata":{"execution":{"iopub.execute_input":"2023-10-19T14:27:34.937479Z","iopub.status.busy":"2023-10-19T14:27:34.937059Z","iopub.status.idle":"2023-10-19T14:27:36.604528Z","shell.execute_reply":"2023-10-19T14:27:36.603091Z"},"papermill":{"duration":1.738022,"end_time":"2023-10-19T14:27:36.640046","exception":false,"start_time":"2023-10-19T14:27:34.902024","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define and train model","metadata":{"papermill":{"duration":0.061966,"end_time":"2023-10-19T14:27:36.761181","exception":false,"start_time":"2023-10-19T14:27:36.699215","status":"completed"},"tags":[]}},{"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):\n    LR_START = 0.00001\n    LR_MAX = 0.00005 * strategy.num_replicas_in_sync\n    LR_MIN = 0.00001\n    LR_RAMPUP_EPOCHS = 5\n    LR_SUSTAIN_EPOCHS = 0\n    LR_EXP_DECAY = .8\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 callbacks.LearningRateScheduler(lrfn, verbose=0)","metadata":{"execution":{"iopub.execute_input":"2023-10-19T14:27:36.887401Z","iopub.status.busy":"2023-10-19T14:27:36.886213Z","iopub.status.idle":"2023-10-19T14:27:36.894178Z","shell.execute_reply":"2023-10-19T14:27:36.893216Z"},"papermill":{"duration":0.073337,"end_time":"2023-10-19T14:27:36.895969","exception":false,"start_time":"2023-10-19T14:27:36.822632","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_and_fit_model(print_summary=False):\n    with strategy.scope():\n        model = tf.keras.Sequential([\n            SwinTransformer(f'swin_{SWIN_TYPE}_{IMAGE_SIZE[0]}', \n                                         include_top=False, \n                                         pretrained=True),\n            L.Dense(len(CLASSES), activation='softmax')\n        ])\n        \n        model.compile(\n            optimizer='adam',\n            loss = 'categorical_crossentropy',\n            metrics=[F1Score(len(CLASSES), average='macro')]\n        )\n\n    if print_summary:\n        model.summary()\n        \n    os.makedirs('checkpoints', exist_ok=True)\n    \n    lr_callback = get_lr_callback()\n    chk_callback = callbacks.ModelCheckpoint(f'checkpoints/swin_{SWIN_TYPE}_best',\n                     save_weights_only=True, monitor='val_f1_score',\n                     mode='max', save_best_only=True, verbose=1)\n    \n    wandb.init(project='flower-classification-tpu-public', \n               name='swin_large_v6',\n               job_type='train', \n               reinit=True)\n    log_callback = WandbCallback(save_model=False)\n\n    _ = model.fit(get_training_dataset(do_onehot=True), \n                  steps_per_epoch=STEPS_PER_EPOCH, \n                  epochs=EPOCHS, \n                  validation_data=get_validation_dataset(do_onehot=True),\n                  callbacks=[lr_callback, chk_callback, log_callback],\n                  verbose=2)\n    model.load_weights(f'checkpoints/swin_{SWIN_TYPE}_best')\n    wandb.finish()\n    return model","metadata":{"execution":{"iopub.execute_input":"2023-10-19T14:27:37.009105Z","iopub.status.busy":"2023-10-19T14:27:37.008754Z","iopub.status.idle":"2023-10-19T14:27:37.018304Z","shell.execute_reply":"2023-10-19T14:27:37.017337Z"},"papermill":{"duration":0.068901,"end_time":"2023-10-19T14:27:37.020300","exception":false,"start_time":"2023-10-19T14:27:36.951399","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = load_and_fit_model()","metadata":{"execution":{"iopub.execute_input":"2023-10-19T14:27:37.133775Z","iopub.status.busy":"2023-10-19T14:27:37.133302Z","iopub.status.idle":"2023-10-19T15:54:27.770304Z","shell.execute_reply":"2023-10-19T15:54:27.768997Z"},"papermill":{"duration":5210.697589,"end_time":"2023-10-19T15:54:27.773194","exception":false,"start_time":"2023-10-19T14:27:37.075605","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Plot confusion matrix and predict on test dataset","metadata":{"papermill":{"duration":0.082319,"end_time":"2023-10-19T15:54:27.942480","exception":false,"start_time":"2023-10-19T15:54:27.860161","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def predict(dataset, model):\n    print('Calculating predictions...')\n    images_ds = dataset.map(lambda image, idnum: image)\n    preds = model.predict(images_ds,verbose=0)\n    preds = np.argmax(preds, axis=1)\n    return preds","metadata":{"execution":{"iopub.execute_input":"2023-10-19T15:54:28.112949Z","iopub.status.busy":"2023-10-19T15:54:28.112524Z","iopub.status.idle":"2023-10-19T15:54:28.118419Z","shell.execute_reply":"2023-10-19T15:54:28.117464Z"},"papermill":{"duration":0.093363,"end_time":"2023-10-19T15:54:28.120165","exception":false,"start_time":"2023-10-19T15:54:28.026802","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_ds = get_validation_dataset(ordered=True)\ncm_predictions = predict(valid_ds, model)\n\nlabels_ds = valid_ds.map(lambda image, label: label).unbatch()\ncm_correct_labels = next(iter(labels_ds.batch(NUM_VALIDATION_IMAGES))).numpy() # get everything as one batch\n\ncmat = confusion_matrix(cm_correct_labels, cm_predictions, labels=range(len(CLASSES)))\nscore = f1_score(cm_correct_labels, cm_predictions, labels=range(len(CLASSES)), average='macro')\nprecision = precision_score(cm_correct_labels, cm_predictions, labels=range(len(CLASSES)), average='macro')\nrecall = recall_score(cm_correct_labels, cm_predictions, labels=range(len(CLASSES)), average='macro')\n#cmat = (cmat.T / cmat.sum(axis=1)).T # normalized\ndisplay_confusion_matrix(cmat, score, precision, recall)","metadata":{"execution":{"iopub.execute_input":"2023-10-19T15:54:28.291338Z","iopub.status.busy":"2023-10-19T15:54:28.290964Z","iopub.status.idle":"2023-10-19T15:55:29.114124Z","shell.execute_reply":"2023-10-19T15:55:29.112706Z"},"papermill":{"duration":61.005356,"end_time":"2023-10-19T15:55:29.209989","exception":false,"start_time":"2023-10-19T15:54:28.204633","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_ds = get_test_dataset(ordered=True) # since we are splitting the dataset and iterating separately on images and ids, order matters.\n\npredictions = predict(test_ds, model)\n\nprint('Generating submission file...')\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') # all in one batch\n                 \nsub_df = pd.DataFrame({'id': test_ids, 'label': predictions})\nsub_df.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.execute_input":"2023-10-19T15:55:29.387896Z","iopub.status.busy":"2023-10-19T15:55:29.387604Z","iopub.status.idle":"2023-10-19T15:56:00.065512Z","shell.execute_reply":"2023-10-19T15:56:00.064121Z"},"papermill":{"duration":30.770703,"end_time":"2023-10-19T15:56:00.068253","exception":false,"start_time":"2023-10-19T15:55:29.297550","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":0.093115,"end_time":"2023-10-19T15:56:00.264131","exception":false,"start_time":"2023-10-19T15:56:00.171016","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]}]}