{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import math, re, os\nimport tensorflow as tf\nimport numpy as np\nimport pandas as pd\nfrom matplotlib import pyplot as plt\nfrom kaggle_datasets import KaggleDatasets\nfrom sklearn.metrics import f1_score, precision_score, recall_score, confusion_matrix\nfrom tensorflow.keras.utils import to_categorical\nfrom sklearn.utils import class_weight\nprint(\"Tensorflow version \" + tf.__version__)\nAUTO = tf.data.experimental.AUTOTUNE","metadata":{"execution":{"iopub.status.busy":"2021-12-20T05:11:17.368885Z","iopub.execute_input":"2021-12-20T05:11:17.369146Z","iopub.status.idle":"2021-12-20T05:11:17.375596Z","shell.execute_reply.started":"2021-12-20T05:11:17.369119Z","shell.execute_reply":"2021-12-20T05:11:17.374816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append('../input/swintransformertf')\nfrom swintransformer import SwinTransformer","metadata":{"execution":{"iopub.status.busy":"2021-12-20T05:11:17.541669Z","iopub.execute_input":"2021-12-20T05:11:17.542456Z","iopub.status.idle":"2021-12-20T05:11:17.547181Z","shell.execute_reply.started":"2021-12-20T05:11:17.54241Z","shell.execute_reply":"2021-12-20T05:11:17.546461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir learner_models","metadata":{"execution":{"iopub.status.busy":"2021-12-20T05:11:17.730343Z","iopub.execute_input":"2021-12-20T05:11:17.730789Z","iopub.status.idle":"2021-12-20T05:11:18.515972Z","shell.execute_reply.started":"2021-12-20T05:11:17.730757Z","shell.execute_reply":"2021-12-20T05:11:18.514519Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# pretrained_model.summary()","metadata":{"execution":{"iopub.status.busy":"2021-12-20T05:11:18.518352Z","iopub.execute_input":"2021-12-20T05:11:18.518718Z","iopub.status.idle":"2021-12-20T05:11:18.523979Z","shell.execute_reply.started":"2021-12-20T05:11:18.518663Z","shell.execute_reply":"2021-12-20T05:11:18.522879Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# TPU or GPU detection","metadata":{}},{"cell_type":"code","source":"# NEW on TPU in TensorFlow 24: shorter cross-compatible TPU/GPU/multi-GPU/cluster-GPU detection code\n\ntry: # detect TPUs\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver.connect() # TPU detection\n    strategy = tf.distribute.TPUStrategy(tpu)\nexcept ValueError: # detect GPUs\n    strategy = tf.distribute.MirroredStrategy() # for GPU or multi-GPU machines\n    #strategy = tf.distribute.get_strategy() # default strategy that works on CPU and single GPU\n    #strategy = tf.distribute.experimental.MultiWorkerMirroredStrategy() # for clusters of multi-GPU machines\n\nprint(\"Number of accelerators: \", strategy.num_replicas_in_sync)","metadata":{"execution":{"iopub.status.busy":"2021-12-20T05:11:19.412656Z","iopub.execute_input":"2021-12-20T05:11:19.413146Z","iopub.status.idle":"2021-12-20T05:11:25.18337Z","shell.execute_reply.started":"2021-12-20T05:11:19.413063Z","shell.execute_reply":"2021-12-20T05:11:25.182367Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Competition data access\nTPUs read data directly from Google Cloud Storage (GCS). This Kaggle utility will copy the dataset to a GCS bucket co-located with the TPU. If you have multiple datasets attached to the notebook, you can pass the name of a specific dataset to the get_gcs_path function. The name of the dataset is the name of the directory it is mounted in. Use `!ls /kaggle/input/` to list attached datasets.","metadata":{}},{"cell_type":"code","source":"from kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nuser_credential = user_secrets.get_gcloud_credential()\nuser_secrets.set_tensorflow_credential(user_credential)","metadata":{"execution":{"iopub.status.busy":"2021-12-20T05:11:25.184991Z","iopub.execute_input":"2021-12-20T05:11:25.185212Z","iopub.status.idle":"2021-12-20T05:11:25.400356Z","shell.execute_reply.started":"2021-12-20T05:11:25.185186Z","shell.execute_reply":"2021-12-20T05:11:25.399373Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# GCS_DS_PATH = KaggleDatasets().get_gcs_path(\"flower-classification\") # you can list the bucket with \"!gsutil ls $GCS_DS_PATH\"\n# GCS_DS_PATH = KaggleDatasets().get_gcs_path(\"photos-tfrecord-901\")\nGCS_DS_PATH = KaggleDatasets().get_gcs_path(\"photos-tfrecord-901\")","metadata":{"execution":{"iopub.status.busy":"2021-12-20T05:11:26.110828Z","iopub.execute_input":"2021-12-20T05:11:26.11112Z","iopub.status.idle":"2021-12-20T05:11:26.949628Z","shell.execute_reply.started":"2021-12-20T05:11:26.11109Z","shell.execute_reply":"2021-12-20T05:11:26.94875Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configuration","metadata":{}},{"cell_type":"code","source":"IMAGE_SIZE = [384, 384] # At this size, a GPU will run out of memory. Use the TPU.\n                        # For GPU training, please select 224 x 224 px image size.\nEPOCHS = 12\nBATCH_SIZE = 8 * strategy.num_replicas_in_sync\n\nGCS_PATH = GCS_DS_PATH + '/records_901'\n\nALL_FILENAMES = tf.io.gfile.glob(GCS_PATH + '/*.tfrec')\n\n# GCS_PATH_SELECT = { # available image sizes\n#     192: GCS_DS_PATH + '/tfrecords-jpeg-192x192',\n#     224: GCS_DS_PATH + '/tfrecords-jpeg-224x224',\n#     331: GCS_DS_PATH + '/tfrecords-jpeg-331x331',\n#     512: GCS_DS_PATH + '/tfrecords-jpeg-512x512'\n# }\n# GCS_PATH = GCS_PATH_SELECT[IMAGE_SIZE[0]]\n\n# TRAINING_FILENAMES = [f'train_0{i}-901.tfrec' for i in range(10)]\n# VALIDATION_FILENAMES = [f'train_0{i}-901.tfrec' for i in range(10, 11)]\n\nTRAINING_FILENAMES = ALL_FILENAMES[:-1]\nVALIDATION_FILENAMES = ALL_FILENAMES[-1:]\nMETA_FEATURE_NAMES = ['Subject Focus', 'Eyes', 'Face', 'Near', 'Action', 'Accessory', 'Group', 'Collage', 'Human','Occlusion', 'Info', 'Blur']\n\ndata_path = '../input/petfinder-pawpularity-score'\ntrain_data_path = os.path.join(data_path, 'train')\ntrain_data_files = os.listdir(train_data_path)\n\ntest_data_path = os.path.join(data_path, 'test')\ntest_data_files = os.listdir(test_data_path)\n\ndf_true = pd.read_csv(os.path.join(data_path, 'train.csv'))\ndf_test = pd.read_csv(os.path.join(data_path, 'test.csv'))\nsample_submission = pd.read_csv(os.path.join(data_path, 'sample_submission.csv'))\n\n\n# class_weights = class_weight.compute_class_weight('balanced',\n#                                                  np.unique(df_true['Pawpularity']),\n#                                                  df_true['Pawpularity'])\n\n# class_weights = {i - 1: class_weights[i - 1] for i in np.unique(df_true['Pawpularity'])}\n# class_weights\n# class_weights = class_weight.compute_class_weight('balanced',\n#                                                  np.unique(df_true['Pawpularity']),\n#                                                  df_true['Pawpularity'])\n\n# class_weights = {i - 1: class_weights[i - 1] for i in np.unique(df_true['Pawpularity'])}\n# class_weights\n\n# TRAINING_FILENAMES = tf.io.gfile.glob(GCS_PATH + '/train/*.tfrec')\n# VALIDATION_FILENAMES = tf.io.gfile.glob(GCS_PATH + '/val/*.tfrec')\n# TEST_FILENAMES = tf.io.gfile.glob(GCS_PATH + '/test/*.tfrec') # predictions on this dataset should be submitted for the competition\n\n# 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":{"execution":{"iopub.status.busy":"2021-12-20T05:11:37.340824Z","iopub.execute_input":"2021-12-20T05:11:37.341392Z","iopub.status.idle":"2021-12-20T05:11:37.518393Z","shell.execute_reply.started":"2021-12-20T05:11:37.341355Z","shell.execute_reply":"2021-12-20T05:11:37.517384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualization utilities\ndata -> pixels, nothing of much interest for the machine learning practitioner in this section.","metadata":{}},{"cell_type":"code","source":"TRAINING_FILENAMES","metadata":{"execution":{"iopub.status.busy":"2021-12-20T05:11:40.908783Z","iopub.execute_input":"2021-12-20T05:11:40.90909Z","iopub.status.idle":"2021-12-20T05:11:40.91405Z","shell.execute_reply.started":"2021-12-20T05:11:40.909057Z","shell.execute_reply":"2021-12-20T05:11:40.913453Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Datasets","metadata":{}},{"cell_type":"code","source":"\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def decode_image(image_data):\n    image = tf.image.decode_jpeg(image_data, channels=3)  # image format uint8 [0,255]\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        'image_name': tf.io.FixedLenFeature([], tf.string),  # shape [] means single element\n        'Subject Focus': tf.io.FixedLenFeature([],tf.int64),\n        'Eyes': tf.io.FixedLenFeature([],tf.int64),\n        'Face': tf.io.FixedLenFeature([],tf.int64),\n        'Near': tf.io.FixedLenFeature([],tf.int64),\n        'Action': tf.io.FixedLenFeature([],tf.int64),\n        'Accessory': tf.io.FixedLenFeature([],tf.int64),\n        'Group': tf.io.FixedLenFeature([],tf.int64),\n        'Collage': tf.io.FixedLenFeature([],tf.int64),\n        'Human': tf.io.FixedLenFeature([],tf.int64),\n        'Occlusion': tf.io.FixedLenFeature([],tf.int64),\n        'Info': tf.io.FixedLenFeature([],tf.int64),\n        'Blur': tf.io.FixedLenFeature([],tf.int64),\n        'Pawpularity': tf.io.FixedLenFeature([],tf.int64)\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    \n    meta_features =  tf.reshape([example[feature] for feature in META_FEATURE_NAMES], [len(META_FEATURE_NAMES)])\n    meta_features = tf.cast(meta_features, tf.uint8)\n    label = example[\"Pawpularity\"]\n#     label = label - 1\n#     label = tf.one_hot(label, 100)\n\n    label = tf.cast(label, tf.float64)\n    label = label / 100.0\n    return {'inputs_swin': image, 'inputs_meta': meta_features}, label # returns a dataset of (image, label) pairs\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, 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    # data augmentation. Thanks to the dataset.prefetch(AUTO) statement in the next function (below),\n    # this happens essentially for free on TPU. Data pipeline code is executed on the \"CPU\" part\n    # of the TPU while the TPU itself is computing gradients.\n    image['inputs_swin'] = tf.image.resize(image['inputs_swin'], (384, 384))\n    image['inputs_swin'] = tf.image.random_flip_left_right(image['inputs_swin'])\n    image['inputs_swin'] = tf.image.random_saturation(image['inputs_swin'], 0, 2)\n    return image, label   \n\ndef get_training_dataset(RECORD_FILENAME):\n    dataset = load_dataset(RECORD_FILENAME, 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, drop_remainder=True)\n    dataset = dataset.prefetch(AUTO) # prefetch next batch while training (autotune prefetch buffer size)\n    return dataset\n\ndef get_validation_dataset(RECORD_FILENAME, ordered=False):\n    dataset = load_dataset(RECORD_FILENAME, labeled=True, ordered=ordered)\n#     dataset = dataset.map(data_augment, num_parallel_calls=AUTO)\n    dataset = dataset.batch(BATCH_SIZE, drop_remainder=True)\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):\n    dataset = load_dataset(TEST_FILENAMES, labeled=False, ordered=ordered)\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[0:1])\nNUM_VALIDATION_IMAGES = count_data_items(VALIDATION_FILENAMES[0:])\n# NUM_TEST_IMAGES = count_data_items(TEST_FILENAMES)\nSTEPS_PER_EPOCH = NUM_TRAINING_IMAGES // BATCH_SIZE\nVALIDATION_STEPS = NUM_VALIDATION_IMAGES // BATCH_SIZE # The \"-(-//)\" trick rounds up instead of down :-)\n# TEST_STEPS = -(-NUM_TEST_IMAGES // BATCH_SIZE)             # The \"-(-//)\" trick rounds up instead of down :-)\nprint('Dataset: {} training images, {} validation images, {} unlabeled test images'.format(NUM_TRAINING_IMAGES, NUM_VALIDATION_IMAGES, 0))","metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","execution":{"iopub.status.busy":"2021-12-20T05:15:58.424681Z","iopub.execute_input":"2021-12-20T05:15:58.425205Z","iopub.status.idle":"2021-12-20T05:15:58.453129Z","shell.execute_reply.started":"2021-12-20T05:15:58.425149Z","shell.execute_reply":"2021-12-20T05:15:58.452166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"f\"Steps per epoch {STEPS_PER_EPOCH} validation steps {VALIDATION_STEPS}\"","metadata":{"execution":{"iopub.status.busy":"2021-12-20T05:16:03.881534Z","iopub.execute_input":"2021-12-20T05:16:03.882065Z","iopub.status.idle":"2021-12-20T05:16:03.888578Z","shell.execute_reply.started":"2021-12-20T05:16:03.882019Z","shell.execute_reply":"2021-12-20T05:16:03.88759Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset visualizations","metadata":{}},{"cell_type":"code","source":"# data dump\nprint(\"Training data shapes:\")\nfor image, label in get_training_dataset(TRAINING_FILENAMES[0]).take(3):\n    print(image[\"inputs_swin\"].numpy().shape, image[\"inputs_meta\"].numpy().shape, label.numpy().shape)\nprint(\"Training data label examples for this record:\", label.numpy())\nprint(\"Validation data shapes:\")\nfor image, label in get_validation_dataset(VALIDATION_FILENAMES[0]).take(3):\n    print(image[\"inputs_swin\"].numpy().shape, image[\"inputs_meta\"].numpy().shape, label.numpy().shape)\nprint(\"Validation data label examples for this record:\", label.numpy())\n# print(\"Test data shapes:\")\n# for image, idnum in get_test_dataset().take(3):\n#     print(image.numpy().shape, idnum.numpy().shape)\n# print(\"Test data IDs:\", idnum.numpy().astype('U')) # U=unicode string","metadata":{"execution":{"iopub.status.busy":"2021-12-20T05:11:54.181885Z","iopub.execute_input":"2021-12-20T05:11:54.182713Z","iopub.status.idle":"2021-12-20T05:11:58.738905Z","shell.execute_reply.started":"2021-12-20T05:11:54.182663Z","shell.execute_reply":"2021-12-20T05:11:58.737859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Peek at training data\n# training_dataset = get_training_dataset()\n# training_dataset = training_dataset.unbatch().batch(20)\n# train_batch = iter(training_dataset)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_dataset","metadata":{"trusted":true},"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":{"trusted":true},"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":{"trusted":true},"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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model\nYou can select these models:  \n`swin_tiny_224`    \n`swin_small_224`  \n`swin_base_224`  \n`swin_base_384`  \n`swin_large_224`  \n`swin_large_384`  ","metadata":{}},{"cell_type":"code","source":"from tensorflow.keras import Model, Input\nfrom tensorflow.keras.models import save_model\nfrom tensorflow.keras.layers import Dense, Flatten, Dropout, Concatenate, BatchNormalization, Lambda\nfrom tensorflow.keras.layers.experimental.preprocessing import RandomFlip, RandomRotation, RandomContrast, Resizing, Rescaling, RandomWidth, RandomHeight, RandomZoom, RandomContrast, RandomCrop\n\n\ndef create_model(LR):\n    with strategy.scope():\n        img_adjust_layer = tf.keras.layers.Lambda(lambda data: tf.keras.applications.imagenet_utils.preprocess_input(tf.cast(data, tf.float32),\n                                                                                                                     mode=\"torch\"),\n                                                  input_shape=[384, 384, 3])\n        pretrained_model = SwinTransformer('swin_base_384', num_classes=1, include_top=False, pretrained=True, use_tpu=True)\n        pretrained_model.trainable = False\n\n#         efficient_net = tf.keras.applications.EfficientNetB0(\n#             include_top=False,\n#             weights=\"imagenet\",\n#             input_shape=(384, 384, 3),\n#             pooling=\"max\",\n#         )\n#         efficient_net.trainable = False\n\n        inputs_swin = Input((384, 384, 3), name=\"inputs_swin\")\n        inputs_meta = Input((12, ), name=\"inputs_meta\")\n\n    #     aug_layers = RandomFlip()(inputs_swin)\n        aug_layers = img_adjust_layer(inputs_swin)\n\n        model_left = pretrained_model(aug_layers)\n\n        concat = tf.keras.layers.concatenate([model_left, inputs_meta])\n\n        concat = Dense(4096, activation=\"relu\") (concat)\n        concat = Dropout(0.2) (concat)\n        concat = Dense(256, activation=\"relu\") (concat)\n        concat = Dropout(0.2) (concat)\n        concat = Dense(1, activation=\"sigmoid\") (concat)\n#         concat = Dense(100, activation=\"softmax\") (concat)\n\n        model_concat = Model(inputs=[inputs_swin, inputs_meta], outputs=concat)\n        BCE = tf.keras.losses.BinaryCrossentropy(label_smoothing=0.001)\n#         SCCE = tf.keras.losses.SparseCategoricalCrossentropy()\n#         CCE = tf.keras.losses.CategoricalCrossentropy(label_smoothing=0.001)\n        RMSE = tf.keras.metrics.RootMeanSquaredError()\n        \n    model_concat.compile(\n        optimizer=tf.keras.optimizers.Adam(learning_rate=LR, epsilon=1e-8),\n        loss = BCE,\n        metrics=[RMSE]\n    )\n    \n    return model_concat, pretrained_model, BCE, RMSE\n    \n\n    \n#     model = tf.keras.Sequential([\n#         img_adjust_layer,\n#         pretrained_model,\n#         tf.keras.layers.Dense(1, activation='sigmoid')\n#     ])\n    \n","metadata":{"execution":{"iopub.status.busy":"2021-12-20T05:12:14.795097Z","iopub.execute_input":"2021-12-20T05:12:14.795441Z","iopub.status.idle":"2021-12-20T05:12:14.809585Z","shell.execute_reply.started":"2021-12-20T05:12:14.795404Z","shell.execute_reply":"2021-12-20T05:12:14.808888Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_callbacks(MIN_LR):\n    earlystopping = tf.keras.callbacks.EarlyStopping(monitor = \"val_root_mean_squared_error\",\n                                                     min_delta = 0.0005,\n                                                     patience = 6,\n                                                     mode = 'min',\n                                                     restore_best_weights = True,\n                                                     verbose = 1)\n\n    reduce_lr = tf.keras.callbacks.ReduceLROnPlateau(monitor = \"val_root_mean_squared_error\", factor=0.2,\n                                  patience=3, min_lr=MIN_LR)\n    \n    return [earlystopping, reduce_lr]\n\n","metadata":{"execution":{"iopub.status.busy":"2021-12-20T05:12:21.479154Z","iopub.execute_input":"2021-12-20T05:12:21.479854Z","iopub.status.idle":"2021-12-20T05:12:21.488539Z","shell.execute_reply.started":"2021-12-20T05:12:21.479807Z","shell.execute_reply":"2021-12-20T05:12:21.487312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Stratified training ","metadata":{}},{"cell_type":"code","source":"VALIDATION_FILENAMES","metadata":{"execution":{"iopub.status.busy":"2021-12-20T05:12:24.156103Z","iopub.execute_input":"2021-12-20T05:12:24.156436Z","iopub.status.idle":"2021-12-20T05:12:24.162481Z","shell.execute_reply.started":"2021-12-20T05:12:24.156401Z","shell.execute_reply":"2021-12-20T05:12:24.161783Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"EPOCHS = 30\nNUM_FOLDS = 10\nimport gc\n\nfor i in range(NUM_FOLDS):\n    model, pretrained_model, LOSS, RMSE = create_model(0.0001)\n    \n    callbacks = get_callbacks(0.000001)\n    \n    history_fit = model.fit(get_training_dataset(TRAINING_FILENAMES[i:i+1]), \n                           steps_per_epoch=STEPS_PER_EPOCH, \n                           epochs=EPOCHS,\n                           validation_data=get_validation_dataset(VALIDATION_FILENAMES), \n                           validation_steps=VALIDATION_STEPS, \n                           callbacks=callbacks)\n#                            class_weight=class_weights)\n    \n    pretrained_model.trainable = True\n    pretrained_model.layers[3].trainable = False\n    \n    model.compile(\n        optimizer=tf.keras.optimizers.Adam(learning_rate=1e-5, epsilon=1e-8),\n        loss = LOSS,\n        metrics=[RMSE]\n    )\n    \n    callbacks_finetune = get_callbacks(1e-7)\n\n    history_finetune = model.fit(get_training_dataset(TRAINING_FILENAMES[i:i+1]), \n                               steps_per_epoch=STEPS_PER_EPOCH, \n                               epochs=EPOCHS,\n                               validation_data=get_validation_dataset(VALIDATION_FILENAMES), \n                               validation_steps=VALIDATION_STEPS, \n                               callbacks=callbacks)\n#                                class_weight=class_weights)\n    \n    model.save_weights(f\"learner_models/learner_model_{i}.h5\")\n    \n    del model\n    del pretrained_model\n    gc.collect()\n    \n    \n    ","metadata":{"execution":{"iopub.status.busy":"2021-12-20T05:21:31.206178Z","iopub.execute_input":"2021-12-20T05:21:31.206508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.evaluate(get_validation_dataset(VALIDATION_FILENAMES))","metadata":{"execution":{"iopub.status.busy":"2021-12-20T05:11:12.152166Z","iopub.status.idle":"2021-12-20T05:11:12.152671Z","shell.execute_reply.started":"2021-12-20T05:11:12.15242Z","shell.execute_reply":"2021-12-20T05:11:12.152445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## With metadata, add 4096 neurons","metadata":{}},{"cell_type":"code","source":"\nhistory = model_concat.fit(get_training_dataset(), steps_per_epoch=STEPS_PER_EPOCH, epochs=EPOCHS,\n                    validation_data=get_validation_dataset(), validation_steps=VALIDATION_STEPS, callbacks=callbacks)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss = plt.figure()\n\nplt.plot(history.history['loss'], label=\"loss\")\nplt.plot(history.history['val_loss'], label=\"val_loss\")\nplt.legend()\nplt.title(\"loss\")\n\nerror = plt.figure()\nplt.plot(history.history['root_mean_squared_error'], label=\"root_mean_squared_error\")\nplt.plot(history.history['val_root_mean_squared_error'], label=\"val_root_mean_squared_error\")\nplt.legend()\nplt.title(\"RMSE\")\n\nLR = plt.figure()\nplt.plot(history.history['lr'])\nplt.title(\"Learning Rate\")\n\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Without metadata ","metadata":{}},{"cell_type":"code","source":"history = model_concat.fit(get_training_dataset(), steps_per_epoch=STEPS_PER_EPOCH, epochs=EPOCHS,\n                    validation_data=get_validation_dataset(), validation_steps=VALIDATION_STEPS, callbacks=callbacks)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss = plt.figure()\n\nplt.plot(history.history['loss'], label=\"loss\")\nplt.plot(history.history['val_loss'], label=\"val_loss\")\nplt.legend()\nplt.title(\"loss\")\n\nerror = plt.figure()\nplt.plot(history.history['root_mean_squared_error'], label=\"root_mean_squared_error\")\nplt.plot(history.history['val_root_mean_squared_error'], label=\"val_root_mean_squared_error\")\nplt.legend()\nplt.title(\"RMSE\")\n\nLR = plt.figure()\nplt.plot(history.history['lr'])\nplt.title(\"Learning Rate\")\n\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Using 2048 + 256 neurons","metadata":{}},{"cell_type":"code","source":"history = model_concat.fit(get_training_dataset(), steps_per_epoch=STEPS_PER_EPOCH, epochs=EPOCHS,\n                    validation_data=get_validation_dataset(), validation_steps=VALIDATION_STEPS, callbacks=callbacks)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = model_concat.fit(get_training_dataset(), steps_per_epoch=STEPS_PER_EPOCH, epochs=EPOCHS,\n                    validation_data=get_validation_dataset(), validation_steps=VALIDATION_STEPS, callbacks=callbacks)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_concat.evaluate(get_validation_dataset())","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Using 1024 neurons","metadata":{}},{"cell_type":"code","source":"loss = plt.figure()\n\nplt.plot(history.history['loss'], label=\"loss\")\nplt.plot(history.history['val_loss'], label=\"val_loss\")\nplt.legend()\nplt.title(\"loss\")\n\nerror = plt.figure()\nplt.plot(history.history['root_mean_squared_error'], label=\"root_mean_squared_error\")\nplt.plot(history.history['val_root_mean_squared_error'], label=\"val_root_mean_squared_error\")\nplt.legend()\nplt.title(\"RMSE\")\n\nLR = plt.figure()\nplt.plot(history.history['lr'])\nplt.title(\"Learning Rate\")\n\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Using 2048 + 256 neurons","metadata":{}},{"cell_type":"code","source":"loss = plt.figure()\n\nplt.plot(history.history['loss'], label=\"loss\")\nplt.plot(history.history['val_loss'], label=\"val_loss\")\nplt.legend()\nplt.title(\"loss\")\n\nerror = plt.figure()\nplt.plot(history.history['root_mean_squared_error'], label=\"root_mean_squared_error\")\nplt.plot(history.history['val_root_mean_squared_error'], label=\"val_root_mean_squared_error\")\nplt.legend()\nplt.title(\"RMSE\")\n\nLR = plt.figure()\nplt.plot(history.history['lr'])\nplt.title(\"Learning Rate\")\n\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# After top layer has converged, unfreeze the previous layers and fine-tune","metadata":{}},{"cell_type":"code","source":"pretrained_model.trainable = True\npretrained_model.layers[3].trainable = False","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_concat.compile(\n    optimizer=tf.keras.optimizers.Adam(learning_rate=1e-5, epsilon=1e-8),\n    loss = BCE,\n    metrics=[RMSE]\n)\n\nearlystopping = tf.keras.callbacks.EarlyStopping(monitor = \"val_root_mean_squared_error\",\n                                                 min_delta = 0.0005,\n                                                 patience = 6,\n                                                 mode = 'min',\n                                                 restore_best_weights = True,\n                                                 verbose = 1)\n\nreduce_lr = tf.keras.callbacks.ReduceLROnPlateau(monitor = \"val_root_mean_squared_error\", factor=0.2,\n                              patience=3, min_lr=1e-7)\n\nEPOCHS=30\ncallbacks = [earlystopping, reduce_lr]\n\n\nhistory = model_concat.fit(get_training_dataset(), steps_per_epoch=STEPS_PER_EPOCH, epochs=EPOCHS,\n                    validation_data=get_validation_dataset(), validation_steps=VALIDATION_STEPS, callbacks=callbacks)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss = plt.figure()\n\nplt.plot(history.history['loss'], label=\"loss\")\nplt.plot(history.history['val_loss'], label=\"val_loss\")\nplt.legend()\nplt.title(\"loss\")\n\nerror = plt.figure()\nplt.plot(history.history['root_mean_squared_error'], label=\"root_mean_squared_error\")\nplt.plot(history.history['val_root_mean_squared_error'], label=\"val_root_mean_squared_error\")\nplt.legend()\nplt.title(\"RMSE\")\n\nLR = plt.figure()\nplt.plot(history.history['lr'])\nplt.title(\"Learning Rate\")\n\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_concat.save_weights(\"model_concat_weights_3.h5\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_concat.evaluate(get_validation_dataset())","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Without meta inputs","metadata":{}},{"cell_type":"code","source":"model_concat.compile(\n    optimizer=tf.keras.optimizers.Adam(learning_rate=1e-5, epsilon=1e-8),\n    loss = BCE,\n    metrics=[RMSE]\n)\n\nearlystopping = tf.keras.callbacks.EarlyStopping(monitor = \"val_root_mean_squared_error\",\n                                                 min_delta = 0.0005,\n                                                 patience = 6,\n                                                 mode = 'min',\n                                                 restore_best_weights = True,\n                                                 verbose = 1)\n\nreduce_lr = tf.keras.callbacks.ReduceLROnPlateau(monitor = \"val_root_mean_squared_error\", factor=0.2,\n                              patience=3, min_lr=1e-7)\n\nEPOCHS=30\ncallbacks = [earlystopping, reduce_lr]\n\n\nhistory = model_concat.fit(get_training_dataset(), steps_per_epoch=STEPS_PER_EPOCH, epochs=EPOCHS,\n                    validation_data=get_validation_dataset(), validation_steps=VALIDATION_STEPS, callbacks=callbacks)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_concat.evaluate(get_validation_dataset())","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss = plt.figure()\n\nplt.plot(history.history['loss'], label=\"loss\")\nplt.plot(history.history['val_loss'], label=\"val_loss\")\nplt.legend()\nplt.title(\"loss\")\n\nerror = plt.figure()\nplt.plot(history.history['root_mean_squared_error'], label=\"root_mean_squared_error\")\nplt.plot(history.history['val_root_mean_squared_error'], label=\"val_root_mean_squared_error\")\nplt.legend()\nplt.title(\"RMSE\")\n\nLR = plt.figure()\nplt.plot(history.history['lr'])\nplt.title(\"Learning Rate\")\n\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_concat.save_weights(\"model_concat_weights_2.h5\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Using 2048 + 256 neurons","metadata":{}},{"cell_type":"code","source":"model_concat.compile(\n    optimizer=tf.keras.optimizers.Adam(learning_rate=1e-5, epsilon=1e-8),\n    loss = BCE,\n    metrics=[RMSE]\n)\n\nearlystopping = tf.keras.callbacks.EarlyStopping(monitor = \"val_root_mean_squared_error\",\n                                                 min_delta = 0.0005,\n                                                 patience = 6,\n                                                 mode = 'min',\n                                                 restore_best_weights = True,\n                                                 verbose = 1)\n\nreduce_lr = tf.keras.callbacks.ReduceLROnPlateau(monitor = \"val_root_mean_squared_error\", factor=0.2,\n                              patience=3, min_lr=1e-6)\n\nEPOCHS=30\ncallbacks = [earlystopping, reduce_lr]\n\n\nhistory = model_concat.fit(get_training_dataset(), steps_per_epoch=STEPS_PER_EPOCH, epochs=EPOCHS,\n                    validation_data=get_validation_dataset(), validation_steps=VALIDATION_STEPS, callbacks=callbacks)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss = plt.figure()\n\nplt.plot(history.history['loss'], label=\"loss\")\nplt.plot(history.history['val_loss'], label=\"val_loss\")\nplt.legend()\nplt.title(\"loss\")\n\nerror = plt.figure()\nplt.plot(history.history['root_mean_squared_error'], label=\"root_mean_squared_error\")\nplt.plot(history.history['val_root_mean_squared_error'], label=\"val_root_mean_squared_error\")\nplt.legend()\nplt.title(\"RMSE\")\n\nLR = plt.figure()\nplt.plot(history.history['lr'])\nplt.title(\"Learning Rate\")\n\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_concat.evaluate(get_validation_dataset())","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Using 1024 neurons","metadata":{}},{"cell_type":"code","source":"model_concat.compile(\n    optimizer=tf.keras.optimizers.Adam(learning_rate=1e-5, epsilon=1e-8),\n    loss = BCE,\n    metrics=[RMSE]\n)\n\nearlystopping = tf.keras.callbacks.EarlyStopping(monitor = \"val_root_mean_squared_error\",\n                                                 min_delta = 0.0005,\n                                                 patience = 5,\n                                                 mode = 'min',\n                                                 restore_best_weights = True,\n                                                 verbose = 1)\n\nreduce_lr = tf.keras.callbacks.ReduceLROnPlateau(monitor = \"val_root_mean_squared_error\", factor=0.2,\n                              patience=3, min_lr=1e-6)\n\nEPOCHS=30\ncallbacks = [earlystopping, reduce_lr]\n\n\nhistory = model_concat.fit(get_training_dataset(), steps_per_epoch=STEPS_PER_EPOCH, epochs=EPOCHS,\n                    validation_data=get_validation_dataset(), validation_steps=VALIDATION_STEPS, callbacks=callbacks)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss = plt.figure()\n\nplt.plot(history.history['loss'], label=\"loss\")\nplt.plot(history.history['val_loss'], label=\"val_loss\")\nplt.legend()\nplt.title(\"loss\")\n\nerror = plt.figure()\nplt.plot(history.history['root_mean_squared_error'], label=\"root_mean_squared_error\")\nplt.plot(history.history['val_root_mean_squared_error'], label=\"val_root_mean_squared_error\")\nplt.legend()\nplt.title(\"RMSE\")\n\nLR = plt.figure()\nplt.plot(history.history['lr'])\nplt.title(\"Learning Rate\")\n\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Using no neurons","metadata":{}},{"cell_type":"code","source":"model_concat.compile(\n    optimizer=tf.keras.optimizers.Adam(learning_rate=1e-5, epsilon=1e-8),\n    loss = BCE,\n    metrics=[RMSE]\n)\n\nhistory = model_concat.fit(get_training_dataset(), steps_per_epoch=STEPS_PER_EPOCH, epochs=EPOCHS,\n                    validation_data=get_validation_dataset(), validation_steps=VALIDATION_STEPS, callbacks=callbacks)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss = plt.figure()\n\nplt.plot(history.history['loss'], label=\"loss\")\nplt.plot(history.history['val_loss'], label=\"val_loss\")\nplt.legend()\nplt.title(\"loss\")\n\nerror = plt.figure()\nplt.plot(history.history['val_root_mean_squared_error'], label=\"val_root_mean_squared_error\")\nplt.plot(history.history['root_mean_squared_error'], label=\"root_mean_squared_error\")\nplt.legend()\nplt.title(\"RMSE\")\n\nLR = plt.figure()\nplt.plot(history.history['lr'])\nplt.title(\"Learning Rate\")\n\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_concat.optimizer.lr.numpy()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = model_concat.fit(get_validation_dataset(), steps_per_epoch=VALIDATION_STEPS, epochs=1,\n                    validation_data=get_validation_dataset(), validation_steps=VALIDATION_STEPS, callbacks=callbacks)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_concat.save_weights(\"model_concat_weights_val.h5\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tf.saved_model.save(model_concat, 'model_concat.h5')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# save_locally = tf.saved_model.SaveOptions(experimental_io_device='/job:localhost')\n# model.save('./model', options=save_locally) # saving in Tensorflow's \"SavedModel\" format\n\nmodel_concat.save(\"model\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = model_concat.fit(get_training_dataset(), steps_per_epoch=STEPS_PER_EPOCH, epochs=EPOCHS,\n                    validation_data=get_validation_dataset(), validation_steps=VALIDATION_STEPS, callbacks=callbacks)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = model_concat.fit(get_training_dataset(), steps_per_epoch=STEPS_PER_EPOCH, epochs=EPOCHS,\n                    validation_data=get_validation_dataset(), validation_steps=VALIDATION_STEPS, callbacks=callbacks)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display_training_curves(history.history['loss'], history.history['val_loss'], 'loss', 211)\ndisplay_training_curves(history.history['root_mean_squared_error'], history.history['val_root_mean_squared_error'], 'accuracy', 212)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Confusion matrix","metadata":{}},{"cell_type":"code","source":"cmdataset = get_validation_dataset(ordered=True) # since we are splitting the dataset and iterating separately on images and labels, order matters.\nimages_ds = cmdataset.map(lambda image, label: image)\nlabels_ds = cmdataset.map(lambda image, label: label).unbatch()\ncm_correct_labels = next(iter(labels_ds.batch(NUM_VALIDATION_IMAGES))).numpy() # get everything as one batch\ncm_probabilities = model.predict(images_ds, steps=VALIDATION_STEPS)\ncm_predictions = np.argmax(cm_probabilities, axis=-1)\nprint(\"Correct   labels: \", cm_correct_labels.shape, cm_correct_labels)\nprint(\"Predicted labels: \", cm_predictions.shape, cm_predictions)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cmat = 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')\ncmat = (cmat.T / cmat.sum(axis=1)).T # normalized\ndisplay_confusion_matrix(cmat, score, precision, recall)\nprint('f1 score: {:.3f}, precision: {:.3f}, recall: {:.3f}'.format(score, precision, recall))","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}