{"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":"markdown","source":"# 2. Setting up Training Pipeline\nThis section discusses how to train efficient net model using TFRecords prepared before!","metadata":{}},{"cell_type":"markdown","source":"### Import libraries","metadata":{}},{"cell_type":"code","source":"%%capture\n!pip install -q efficientnet\n!pip install tensorflow_addons\nimport re\nimport os\nimport numpy as np\nimport pandas as pd\nimport random\nimport math\nimport tensorflow as tf\nimport efficientnet.tfkeras as efn\nfrom sklearn import metrics\nfrom sklearn.model_selection import KFold, train_test_split\nfrom tensorflow.keras import backend as K\nimport tensorflow_addons as tfa\nfrom tqdm.notebook import tqdm\nfrom kaggle_datasets import KaggleDatasets\nimport matplotlib.pyplot as plt","metadata":{"execution":{"iopub.status.busy":"2022-10-11T20:56:22.819360Z","iopub.execute_input":"2022-10-11T20:56:22.819729Z","iopub.status.idle":"2022-10-11T20:56:41.395944Z","shell.execute_reply.started":"2022-10-11T20:56:22.819639Z","shell.execute_reply":"2022-10-11T20:56:41.394990Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Load the data","metadata":{}},{"cell_type":"markdown","source":"These are the support function to load our data. In general, the pipeline is:\n- Create a tf.data.TFRecordDataset\n- Perform reading each TFRecord and add them to our dataset","metadata":{}},{"cell_type":"code","source":"# Function to decode our images\ndef decode_image(image_data):\n    image = tf.image.decode_jpeg(image_data, channels = 3)\n    image = tf.image.resize(image, IMAGE_SIZE)\n    image = tf.cast(image, tf.float32) / 255.0\n    return image\n\n# This function parse our images and also get the target variable\ndef read_labeled_tfrecord(example):\n    LABELED_TFREC_FORMAT = {\n        \"image_name\": tf.io.FixedLenFeature([], tf.string),\n        \"image\": tf.io.FixedLenFeature([], tf.string),\n        \"target\": tf.io.FixedLenFeature([], tf.int64),\n#         \"matches\": tf.io.FixedLenFeature([], tf.string)\n    }\n\n    example = tf.io.parse_single_example(example, LABELED_TFREC_FORMAT)\n    posting_id = example['image_name']\n    image = decode_image(example['image'])\n#     label_group = tf.one_hot(tf.cast(example['label_group'], tf.int32), depth = N_CLASSES)\n    label_group = tf.cast(example['target'], tf.int32)\n#     matches = example['matches']\n    matches = 1\n    return posting_id, image, label_group, matches\n\n# This function loads TF Records and parse them into tensors\ndef load_dataset(filenames, ordered = False, cache=False):\n    \n    ignore_order = tf.data.Options()\n    if not ordered:\n        ignore_order.experimental_deterministic = False \n        \n    dataset = tf.data.TFRecordDataset(filenames, num_parallel_reads = AUTO)\n    if cache:\n        dataset = dataset.cache()\n    dataset = dataset.with_options(ignore_order)\n    dataset = dataset.map(read_labeled_tfrecord, num_parallel_calls = AUTO) \n    return dataset\n\n# This function is to get our training tensors\ndef get_training_dataset(filenames, ordered = False):\n    dataset = load_dataset(filenames, ordered = ordered)\n#     dataset = dataset.map(data_augment, num_parallel_calls = AUTO)\n    dataset = dataset.map(lambda posting_id, image, label_group, matches: (image, label_group))\n    dataset = dataset.repeat()\n    dataset = dataset.shuffle(2048)\n    dataset = dataset.batch(BATCH_SIZE)\n    dataset = dataset.prefetch(AUTO)\n    return dataset\n\n# Function to count how many photos we have in\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","metadata":{"execution":{"iopub.status.busy":"2022-10-11T20:56:41.397978Z","iopub.execute_input":"2022-10-11T20:56:41.398677Z","iopub.status.idle":"2022-10-11T20:56:41.411485Z","shell.execute_reply.started":"2022-10-11T20:56:41.398617Z","shell.execute_reply":"2022-10-11T20:56:41.410503Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's plot some of our images","metadata":{}},{"cell_type":"code","source":"# For tf.dataset\nAUTO = tf.data.experimental.AUTOTUNE\nIMAGE_SIZE = [256, 256]\nBATCH_SIZE = 16\n\n### Note that: These GCS paths are updated every week. To get updated GCS path, follow the instructions mentioned above.\n# Data access\nGCS_PATHS = {\n    0: 'gs://kds-0fda8ca7f7331dee1d9147bf884973011e12f1e944eca58bda194bae',\n    1: 'gs://kds-083aaa0624c4997f64c5f05c5fe2fe598fcfd64442a80efa3bbabef5',\n    2: 'gs://kds-f44bf20389d4c6de78cc6db2d4d5aa6b4ae7181a901b471ae8aac20e',\n    3: 'gs://kds-c149899075770ba3606ea9fd6275350e99b9b203bec95d87e771455d',\n    4: 'gs://kds-948a45c11f937432bc63408731ac3bde0797a307e5e9273949130dad'\n}\n# Training filenames directory\nTRAINING_FILENAMES = []\nfor fold in GCS_PATHS:\n    TRAINING_FILENAMES += tf.io.gfile.glob(GCS_PATHS[fold] + '/*.tfrec')\n\nNUM_TRAINING_IMAGES = count_data_items(TRAINING_FILENAMES)\nprint(f'Dataset: {NUM_TRAINING_IMAGES} training images')\n\n\nrow = 10; col = 2;\nrow = min(row,BATCH_SIZE//col)\nN_TRAIN = count_data_items(TRAINING_FILENAMES)\nprint(N_TRAIN)\nds = get_training_dataset(TRAINING_FILENAMES, ordered = False)\n\nfor (sample,label) in ds:\n    img = sample\n    plt.figure(figsize=(25,int(25*row/col)))\n    for j in range(row*col):\n        plt.subplot(row,col,j+1)\n        plt.title(label[j].numpy())\n        plt.axis('off')\n        plt.imshow(img[j,])\n    plt.show()\n    break\nprint(img.shape)","metadata":{"execution":{"iopub.status.busy":"2022-10-11T20:56:41.412895Z","iopub.execute_input":"2022-10-11T20:56:41.413203Z","iopub.status.idle":"2022-10-11T20:56:55.623135Z","shell.execute_reply.started":"2022-10-11T20:56:41.413173Z","shell.execute_reply":"2022-10-11T20:56:55.622379Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Setting up our training pipeline\nNice, we have finished plotting our data. Now let's config our TPU and set some hyperparameters for our model.","metadata":{}},{"cell_type":"markdown","source":"### Setting up the TPU","metadata":{}},{"cell_type":"code","source":"try:\n    # TPU detection. No parameters necessary if TPU_NAME environment variable is\n    # set: this is always the case on Kaggle.\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.TPUStrategy(tpu)\nelse:\n    # Default distribution strategy in Tensorflow. Works on CPU and single GPU.\n    strategy = tf.distribute.get_strategy()\n\nprint(\"REPLICAS: \", strategy.num_replicas_in_sync)","metadata":{"execution":{"iopub.status.busy":"2022-10-11T21:01:38.437667Z","iopub.execute_input":"2022-10-11T21:01:38.438960Z","iopub.status.idle":"2022-10-11T21:01:44.296224Z","shell.execute_reply.started":"2022-10-11T21:01:38.438896Z","shell.execute_reply":"2022-10-11T21:01:44.294614Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Hyperparams","metadata":{}},{"cell_type":"code","source":"# Configuration\nEPOCHS = 4\nBATCH_SIZE = 16 * strategy.num_replicas_in_sync\n# Seed\nSEED = 42\nFOLD_TO_RUN = [0,1,2,3,4]\n# Learning rate\nLR = 0.001\n# Verbosity\nVERBOSE = 2\n# Number of classes\nN_CLASSES = 81313\n# Number of folds\nFOLDS = 5\n# EfficientNet\nEFF_NET = 0\n# Freeze Batch Norm\nFREEZE_BATCH_NORM = False\nSNAPSHOT_THRESHOLD = 0","metadata":{"execution":{"iopub.status.busy":"2022-10-11T21:05:43.406256Z","iopub.execute_input":"2022-10-11T21:05:43.406914Z","iopub.status.idle":"2022-10-11T21:05:43.414394Z","shell.execute_reply.started":"2022-10-11T21:05:43.406855Z","shell.execute_reply":"2022-10-11T21:05:43.412803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Our evaluation function","metadata":{}},{"cell_type":"code","source":"# Function to get our f1 score\ndef f1_score(y_true, y_pred):\n    y_true = y_true.apply(lambda x: set(x.split()))\n    y_pred = y_pred.apply(lambda x: set(x.split()))\n    intersection = np.array([len(x[0] & x[1]) for x in zip(y_true, y_pred)])\n    len_y_pred = y_pred.apply(lambda x: len(x)).values\n    len_y_true = y_true.apply(lambda x: len(x)).values\n    f1 = 2 * intersection / (len_y_pred + len_y_true)\n    return f1","metadata":{"execution":{"iopub.status.busy":"2022-10-11T21:05:43.827370Z","iopub.execute_input":"2022-10-11T21:05:43.827690Z","iopub.status.idle":"2022-10-11T21:05:43.836617Z","shell.execute_reply.started":"2022-10-11T21:05:43.827660Z","shell.execute_reply":"2022-10-11T21:05:43.834985Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"To keep our code reproducable, we use seed_everything()","metadata":{}},{"cell_type":"code","source":"# Function to seed everything\ndef seed_everything(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    tf.random.set_seed(seed)","metadata":{"execution":{"iopub.status.busy":"2022-10-11T21:05:45.064874Z","iopub.execute_input":"2022-10-11T21:05:45.065215Z","iopub.status.idle":"2022-10-11T21:05:45.071117Z","shell.execute_reply.started":"2022-10-11T21:05:45.065182Z","shell.execute_reply":"2022-10-11T21:05:45.070071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Now, we will perform Augmentation","metadata":{}},{"cell_type":"code","source":"def arcface_format(posting_id, image, label_group, matches):\n    return posting_id, {'inp1': image, 'inp2': label_group}, label_group, matches\n\n# Data augmentation function\ndef data_augment(posting_id, image, label_group, matches):\n    image = tf.image.random_flip_left_right(image)\n    image = tf.image.random_hue(image, 0.01)\n    image = tf.image.random_saturation(image, 0.70, 1.30)\n    image = tf.image.random_contrast(image, 0.80, 1.20)\n    image = tf.image.random_brightness(image, 0.10)\n    return posting_id, image, label_group, matches\n\n# This function is to get our training tensors\ndef get_training_dataset(filenames, ordered = False):\n    dataset = load_dataset(filenames, ordered = ordered)\n    dataset = dataset.map(data_augment, num_parallel_calls = AUTO)\n    dataset = dataset.map(arcface_format, num_parallel_calls = AUTO)\n    dataset = dataset.map(lambda posting_id, image, label_group, matches: (image, label_group))\n    dataset = dataset.repeat()\n    dataset = dataset.shuffle(2048)\n    dataset = dataset.batch(BATCH_SIZE)\n    dataset = dataset.prefetch(AUTO)\n    return dataset\n\n# This function is to get our validation tensors\ndef get_validation_dataset(filenames, ordered = True):\n    dataset = load_dataset(filenames, ordered = ordered)\n    dataset = dataset.map(arcface_format, num_parallel_calls = AUTO)\n    dataset = dataset.batch(BATCH_SIZE)\n    dataset = dataset.prefetch(AUTO) \n    return dataset","metadata":{"execution":{"iopub.status.busy":"2022-10-11T21:29:42.850191Z","iopub.execute_input":"2022-10-11T21:29:42.850558Z","iopub.status.idle":"2022-10-11T21:29:42.861953Z","shell.execute_reply.started":"2022-10-11T21:29:42.850523Z","shell.execute_reply":"2022-10-11T21:29:42.861097Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### We use simple constant learning rate for now","metadata":{}},{"cell_type":"code","source":"def get_lr_callback():\n    lr_constant     = 0.000005\n   \n    def lrfn(epoch):\n        return lr_constant\n\n    lr_callback = tf.keras.callbacks.LearningRateScheduler(lrfn, verbose=False)\n    return lr_callback","metadata":{"execution":{"iopub.status.busy":"2022-10-11T21:29:43.614871Z","iopub.execute_input":"2022-10-11T21:29:43.615350Z","iopub.status.idle":"2022-10-11T21:29:43.619590Z","shell.execute_reply.started":"2022-10-11T21:29:43.615317Z","shell.execute_reply":"2022-10-11T21:29:43.618866Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Arcmarginproduct class keras layer\nclass ArcMarginProduct(tf.keras.layers.Layer):\n    '''\n    Implements large margin arc distance.\n\n    Reference:\n        https://arxiv.org/pdf/1801.07698.pdf\n        https://github.com/lyakaap/Landmark2019-1st-and-3rd-Place-Solution/\n            blob/master/src/modeling/metric_learning.py\n    '''\n    def __init__(self, n_classes, s=30, m=0.50, easy_margin=False,\n                 ls_eps=0.0, **kwargs):\n\n        super(ArcMarginProduct, self).__init__(**kwargs)\n\n        self.n_classes = n_classes\n        self.s = s\n        self.m = m\n        self.ls_eps = ls_eps\n        self.easy_margin = easy_margin\n        self.cos_m = tf.math.cos(m)\n        self.sin_m = tf.math.sin(m)\n        self.th = tf.math.cos(math.pi - m)\n        self.mm = tf.math.sin(math.pi - m) * m\n\n    def get_config(self):\n\n        config = super().get_config().copy()\n        config.update({\n            'n_classes': self.n_classes,\n            's': self.s,\n            'm': self.m,\n            'ls_eps': self.ls_eps,\n            'easy_margin': self.easy_margin,\n        })\n        return config\n\n    def build(self, input_shape):\n        super(ArcMarginProduct, self).build(input_shape[0])\n\n        self.W = self.add_weight(\n            name='W',\n            shape=(int(input_shape[0][-1]), self.n_classes),\n            initializer='glorot_uniform',\n            dtype='float32',\n            trainable=True,\n            regularizer=None)\n\n    def call(self, inputs):\n        X, y = inputs\n        y = tf.cast(y, dtype=tf.int32)\n        cosine = tf.matmul(\n            tf.math.l2_normalize(X, axis=1),\n            tf.math.l2_normalize(self.W, axis=0)\n        )\n        sine = tf.math.sqrt(1.0 - tf.math.pow(cosine, 2))\n        phi = cosine * self.cos_m - sine * self.sin_m\n        if self.easy_margin:\n            phi = tf.where(cosine > 0, phi, cosine)\n        else:\n            phi = tf.where(cosine > self.th, phi, cosine - self.mm)\n        one_hot = tf.cast(\n            tf.one_hot(y, depth=self.n_classes),\n            dtype=cosine.dtype\n        )\n        if self.ls_eps > 0:\n            one_hot = (1 - self.ls_eps) * one_hot + self.ls_eps / self.n_classes\n\n        output = (one_hot * phi) + ((1.0 - one_hot) * cosine)\n        output *= self.s\n        return output\n\nEFNS = [efn.EfficientNetB0, efn.EfficientNetB1, efn.EfficientNetB2, efn.EfficientNetB3, \n        efn.EfficientNetB4, efn.EfficientNetB5, efn.EfficientNetB6, efn.EfficientNetB7]\n\ndef freeze_BN(model):\n    # Unfreeze layers while leaving BatchNorm layers frozen\n    for layer in model.layers:\n        if not isinstance(layer, tf.keras.layers.BatchNormalization):\n            layer.trainable = True\n        else:\n            layer.trainable = False\n\n# Function to create our EfficientNetB3 model\ndef get_model():\n\n    with strategy.scope():\n\n        margin = ArcMarginProduct(\n            n_classes = N_CLASSES, \n            s = 30, \n            m = 0.5, \n            name='head/arc_margin', \n            dtype='float32'\n            )\n\n        inp = tf.keras.layers.Input(shape = (*IMAGE_SIZE, 3), name = 'inp1')\n        label = tf.keras.layers.Input(shape = (), name = 'inp2')\n        x = EFNS[EFF_NET](weights = 'imagenet', include_top = False)(inp)\n        x = tf.keras.layers.GlobalAveragePooling2D()(x)\n        x = margin([x, label])\n        \n        output = tf.keras.layers.Softmax(dtype='float32')(x)\n\n        model = tf.keras.models.Model(inputs = [inp, label], outputs = [output])\n\n        opt = tf.keras.optimizers.Adam(learning_rate = LR)\n        if FREEZE_BATCH_NORM:\n            freeze_BN(model)\n\n        model.compile(\n            optimizer = opt,\n            loss = [tf.keras.losses.SparseCategoricalCrossentropy()],\n            metrics = [tf.keras.metrics.SparseCategoricalAccuracy()]\n            ) \n        \n        return model\nget_model().summary()","metadata":{"execution":{"iopub.status.busy":"2022-10-11T21:31:43.415915Z","iopub.execute_input":"2022-10-11T21:31:43.416994Z","iopub.status.idle":"2022-10-11T21:32:00.874683Z","shell.execute_reply.started":"2022-10-11T21:31:43.416945Z","shell.execute_reply":"2022-10-11T21:32:00.873477Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def is_interactive():\n    return 'runtime'    in get_ipython().config.IPKernelApp.connection_file\nIS_INTERACTIVE = is_interactive()\nprint(IS_INTERACTIVE)","metadata":{"execution":{"iopub.status.busy":"2022-10-11T21:32:19.305714Z","iopub.execute_input":"2022-10-11T21:32:19.307060Z","iopub.status.idle":"2022-10-11T21:32:19.314355Z","shell.execute_reply.started":"2022-10-11T21:32:19.307014Z","shell.execute_reply":"2022-10-11T21:32:19.313198Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Snapshot(tf.keras.callbacks.Callback):\n    \n    def __init__(self,snapshot_min_epoch,fold):\n        super(Snapshot, self).__init__()\n        self.snapshot_min_epoch = snapshot_min_epoch\n        self.fold = fold\n        \n        \n    def on_epoch_end(self, epoch, logs=None):\n        # logs is a dictionary\n#         print(f\"epoch: {epoch}, train_acc: {logs['acc']}, valid_acc: {logs['val_acc']}\")\n        if epoch >=self.snapshot_min_epoch: # your custom condition         \n            self.model.save_weights(f\"EF{EFF_NET}_fold{self.fold}_epoch{epoch}.h5\")","metadata":{"execution":{"iopub.status.busy":"2022-10-11T21:32:19.532428Z","iopub.execute_input":"2022-10-11T21:32:19.533148Z","iopub.status.idle":"2022-10-11T21:32:19.540642Z","shell.execute_reply.started":"2022-10-11T21:32:19.533105Z","shell.execute_reply":"2022-10-11T21:32:19.539530Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"VERBOSE = 1\nseed_everything(SEED)\ntrain = np.array(TRAINING_FILENAMES)\nSTEPS_PER_EPOCH = count_data_items(train) // BATCH_SIZE\ntrain_dataset = get_training_dataset(train, ordered = False)\nmodel = get_model()\nsnap = Snapshot(snapshot_min_epoch=SNAPSHOT_THRESHOLD,fold=0)\nhistory = model.fit(train_dataset,\n                        steps_per_epoch = STEPS_PER_EPOCH,\n                        epochs = EPOCHS,\n                        callbacks = [snap,get_lr_callback()], \n                        verbose = VERBOSE)","metadata":{"execution":{"iopub.status.busy":"2022-10-11T21:32:20.100610Z","iopub.execute_input":"2022-10-11T21:32:20.100980Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}