{"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":"# Imported packages","metadata":{}},{"cell_type":"code","source":"import io\nimport glob\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nimport seaborn as sns; sns.set()\nfrom sklearn.metrics import f1_score\n\nimport tensorflow as tf\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.models as models\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nfrom torchvision.transforms import Compose, Lambda, ToTensor, Normalize, Resize, RandomCrop, TenCrop, RandomHorizontalFlip","metadata":{"execution":{"iopub.status.busy":"2023-05-12T05:57:56.293773Z","iopub.execute_input":"2023-05-12T05:57:56.294070Z","iopub.status.idle":"2023-05-12T05:58:10.963819Z","shell.execute_reply.started":"2023-05-12T05:57:56.294041Z","shell.execute_reply":"2023-05-12T05:58:10.962680Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Set configuration","metadata":{}},{"cell_type":"code","source":"cfg = {\n     \"train_files\": '/kaggle/input/tpu-getting-started/tfrecords-jpeg-224x224/train/*.tfrec',\n     \"valid_files\": '/kaggle/input/tpu-getting-started/tfrecords-jpeg-224x224/val/*.tfrec',\n     \"test_files\": '/kaggle/input/tpu-getting-started/tfrecords-jpeg-224x224/test/*.tfrec',\n     \"device\": torch.device('cuda' if torch.cuda.is_available() else 'cpu'), # hardware\n     \"n_classes\": 104,\n     \"n_epochs\": 40,                                                            # number of training epochs\n     \"batch_size\": 20,                                                           # training batch size\n     \"num_prints\": 10,                                                            # number of losses to print per epoch\n     \"train_size\": 12753,                                                        # number of training data samples\n     \"check_freq\": 1,                                                            # save model if epoch is a multiple of this\n}\ncfg[\"print_freq\"] = cfg[\"train_size\"] // (cfg[\"batch_size\"] * cfg[\"num_prints\"]) + 1   #how often to print out update on training evolution","metadata":{"execution":{"iopub.status.busy":"2023-05-12T05:58:10.966264Z","iopub.execute_input":"2023-05-12T05:58:10.967086Z","iopub.status.idle":"2023-05-12T05:58:11.071275Z","shell.execute_reply.started":"2023-05-12T05:58:10.967045Z","shell.execute_reply":"2023-05-12T05:58:11.069917Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utility functions","metadata":{}},{"cell_type":"code","source":"# Utility functions:\n# ------------------\n\ndef tfrecords_to_dataframe(fp, test = False):\n    '''\n    Parse data files into rows of a dataframe.\n    \n    arguments\n    ---------\n    fp : str\n        Data files pattern.\n        \n    test : bool\n        If true, data files correspond to testing data.\n    '''\n    def parse(pb, test = False):\n        d = {'id': tf.io.FixedLenFeature([], tf.string), 'image': tf.io.FixedLenFeature([], tf.string)}\n        if not test:\n            d['class'] = tf.io.FixedLenFeature([], tf.int64)\n        return tf.io.parse_single_example(pb, d)\n\n    df = {'id': [], 'img': []} \n    if not test:\n        df['lab'] = []\n    for sample in tf.data.TFRecordDataset(glob.glob(fp)).map(lambda pb: parse(pb, test)):\n        df['id'].append(sample['id'].numpy().decode('utf-8'))\n        df['img'].append(sample['image'].numpy())\n        if not test:\n            df['lab'].append(sample['class'].numpy())\n    return pd.DataFrame(df)\n\n# ------------------------------------------------------------------------------------------------------------------------\n\ndef display_images(dataset, n, cols):\n    '''\n    Display a grid of labelled images of flowers.\n    \n    arguments\n    ---------\n    dataset : Dataset\n        Dataset containing the flower images and labels.\n        \n    n : int\n        Number of images to display.\n        \n    cols : int\n        Number of columns in the grid.\n    '''\n    rows = n // cols if n % cols == 0 else n // cols + 1\n    plt.figure(figsize = (2 * cols, 2 * rows))\n    for i in range(n):\n        plt.subplot(rows, cols, i + 1)\n        img, lab = dataset[i]\n        plt.imshow(img.permute(1, 2, 0).numpy())\n        plt.title(str(lab))\n        plt.axis('off')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-05-12T05:58:11.091054Z","iopub.execute_input":"2023-05-12T05:58:11.091497Z","iopub.status.idle":"2023-05-12T05:58:11.106576Z","shell.execute_reply.started":"2023-05-12T05:58:11.091458Z","shell.execute_reply":"2023-05-12T05:58:11.105526Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define the classes for train & valid set\n","metadata":{}},{"cell_type":"code","source":"class Trainset(Dataset):\n    '''\n    Representation of the training dataset.\n    '''\n    def __init__(self, frac = 1):\n        '''\n        arguments\n        ---------\n        frac : float\n            Fraction of data samples to keep.\n            \n            For example, if frac = 0.5, then a random sample of 50% \n            of the data is kept and the remaining 50% is discarded.\n        '''\n        super().__init__()\n        self.df = tfrecords_to_dataframe(cfg[\"train_files\"]).sample(frac = frac).reset_index(drop = True)\n        self.t1 = Lambda(lambda b: Image.open(io.BytesIO(b)))\n        self.t2 = Compose([RandomCrop(300), \n                           RandomHorizontalFlip(), \n                           ToTensor(), \n                           Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])])\n        \n    def __len__(self):\n        \"\"\"\n            Returns \n                the length of the dataframe with images\n        \"\"\"\n        return self.df.shape[0]\n    \n    def __getitem__(self, i):\n        \"\"\"\n        Returns the current item (image + label) from trainset\n        Will resize it to 300 x 641 before\n            Args\n                i: index of current item to return\n            Returns\n        \"\"\"\n        transform = Compose([self.t1, Resize(np.random.randint(300, 641)), self.t2])\n        sample = self.df.iloc[i]\n        return transform(sample['img']), sample['lab']\n\n\nclass Evalset(Dataset):\n    '''\n    Representation of the evaluation datasets.\n    '''\n    def __init__(self, frac = 1, test = False):\n        '''\n        Args\n            frac : float Fraction of data samples to keep.\n            \n            For example, if frac = 0.5, then a random sample of 50% \n            of the data is kept and the remaining 50% is discarded.\n            \n            test : bool\n                If true, this dataset contains the testing data. \n                Otherwise, this dataset contains the validation data. \n        Returns\n            none\n        '''\n        super().__init__()\n        files = cfg[\"valid_files\"] if not test else cfg[\"test_files\"]\n        self.df = tfrecords_to_dataframe(files, test).sample(frac = frac).reset_index(drop = True)\n        self.transforms = [Compose([Lambda(lambda b: Image.open(io.BytesIO(b))), \n                                    Resize(scale), \n                                    TenCrop(300), \n                                    Lambda(lambda xs: torch.stack([ToTensor()(x) for x in xs])), \n                                    Lambda(lambda xs: torch.stack([Normalize([0.485, 0.456, 0.406], \n                                                                             [0.229, 0.224, 0.225])(x) for x in xs]))])\n                           for scale in [372, 568]]\n        self.test = test\n        \n    def __len__(self):\n        \"\"\"\n        Returns\n            the valid dataset length\n        \"\"\"\n        return self.df.shape[0]\n    \n    def __getitem__(self, i):\n        \"\"\"\n        Generic evaluation class\n        Can be used either for validation set or test set\n        Args\n            i: current item\n        Returns \n            either the image and the label (if class is used for validation set) or\n            the image only (if class is used for test set)\n        \"\"\"\n        sample = self.df.iloc[i]\n        imgs = torch.stack([t(sample['img']) for t in self.transforms])\n        return imgs, sample['lab'] if not self.test else sample['id']\n","metadata":{"nteract":{"transient":{"deleting":false}},"execution":{"iopub.status.busy":"2023-05-12T05:58:11.108418Z","iopub.execute_input":"2023-05-12T05:58:11.109230Z","iopub.status.idle":"2023-05-12T05:58:11.128777Z","shell.execute_reply.started":"2023-05-12T05:58:11.109191Z","shell.execute_reply":"2023-05-12T05:58:11.127766Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model class\n","metadata":{}},{"cell_type":"code","source":"class EfficientNetB0(nn.Module):\n    '''\n    EfficientNet B0 fine-tune.\n    '''\n    def __init__(self, n_classes, learnable_modules = ('classifier.1',)):\n        '''\n        Fine tune for EfficientNetB0\n        Args\n            n_classes : int - Number of classification categories.\n            learnable_modules : tuple - Names of the modules to fine-tune.\n        Return\n            \n        '''\n        super().__init__()\n        self.efficientnet_b0 = models.efficientnet_b0(weights = 'DEFAULT')\n        self.efficientnet_b0.classifier[1] = nn.Linear(self.efficientnet_b0.classifier[1].in_features, n_classes)\n        self.efficientnet_b0.requires_grad_(False)\n        modules = dict(self.efficientnet_b0.named_modules())\n        for name in learnable_modules:\n            modules[name].requires_grad_(True)\n        \n    def forward(self, x):\n        \"\"\"\n        Forward function for the fine-tuned model\n        Args\n            x: \n        Return\n            result\n        \"\"\"\n        return F.log_softmax(self.efficientnet_b0(x), dim = 1)","metadata":{"execution":{"iopub.status.busy":"2023-05-12T05:58:11.130230Z","iopub.execute_input":"2023-05-12T05:58:11.130714Z","iopub.status.idle":"2023-05-12T05:58:11.142956Z","shell.execute_reply.started":"2023-05-12T05:58:11.130675Z","shell.execute_reply":"2023-05-12T05:58:11.142031Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare train, valid and test data","metadata":{}},{"cell_type":"code","source":"# Training, validation, and testing data:\n# ---------------------------------------\ntrain_set    = Trainset()\ntrain_loader = DataLoader(train_set, batch_size = cfg[\"batch_size\"], shuffle = True, num_workers = 2)\nvalid_loader = DataLoader(Evalset(frac = 0.20), batch_size = 1, num_workers = 2)\ntest_loader  = DataLoader(Evalset(test = True), batch_size = 1, num_workers = 2)","metadata":{"papermill":{"duration":25.462451,"end_time":"2023-03-06T18:33:12.910134","exception":false,"start_time":"2023-03-06T18:32:47.447683","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-05-12T05:58:11.144620Z","iopub.execute_input":"2023-05-12T05:58:11.145051Z","iopub.status.idle":"2023-05-12T05:58:34.784601Z","shell.execute_reply.started":"2023-05-12T05:58:11.145015Z","shell.execute_reply":"2023-05-12T05:58:34.783481Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define the optimizer","metadata":{}},{"cell_type":"code","source":"# Modelling components:\n# ---------------------\nmodel = nn.DataParallel(EfficientNetB0(n_classes = cfg[\"n_classes\"], learnable_modules = ('features.5.2', \n                                                                             'features.6', \n                                                                             'features.7', \n                                                                             'features.8', \n                                                                             'classifier')))\nmodel.to(cfg[\"device\"])\n\noptimizer = torch.optim.Adam(params = [{'params': model.module.efficientnet_b0.features[5][2].parameters()}, \n                                       {'params': model.module.efficientnet_b0.features[6].parameters()}, \n                                       {'params': model.module.efficientnet_b0.features[7].parameters()},\n                                       {'params': model.module.efficientnet_b0.features[8].parameters()},\n                                       {'params': model.module.efficientnet_b0.classifier.parameters(), 'lr': 1e-3}], \n                             lr = 1e-4, \n                             weight_decay = 1e-4)\n\nscheduler = CosineAnnealingLR(optimizer, T_max = cfg[\"n_epochs\"])\n\nloss_fn = F.nll_loss","metadata":{"papermill":{"duration":0.722297,"end_time":"2023-03-06T18:33:17.344","exception":false,"start_time":"2023-03-06T18:33:16.621703","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-05-12T05:58:38.860943Z","iopub.execute_input":"2023-05-12T05:58:38.861428Z","iopub.status.idle":"2023-05-12T05:58:39.556587Z","shell.execute_reply.started":"2023-05-12T05:58:38.861376Z","shell.execute_reply":"2023-05-12T05:58:39.555461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training loop","metadata":{}},{"cell_type":"code","source":"def train(train_losses, epoch):\n    print()\n    print(f'Epoch {epoch}:')\n    print('-' * len(f'Epoch {epoch}:'))\n    model.train() \n    for i, data in enumerate(train_loader):\n        image, label = data\n        image = image.to(cfg[\"device\"])\n        label = label.to(cfg[\"device\"])\n        # forward pass\n        outputs = model(image)\n        # calculate the loss\n        loss = loss_fn(outputs, label)\n        optimizer.zero_grad()\n        # backpropagation\n        loss.backward()\n        # optimize the weights\n        optimizer.step()\n        \n        # optionaly, print the loss\n        if i % cfg[\"print_freq\"] == 0:\n            print('Loss {}: {:.3f}'.format(i, loss.item()))\n            train_losses.append(loss.item())\n\n\ndef validate(valid_f1s):\n    model.eval()\n    valid_true_labels = []\n    valid_pred_labels = []\n    valid_running_loss = []\n    counter = 0\n    with torch.no_grad():\n        for data in valid_loader:\n            counter += 1\n            image, labels = data\n            # append true labels\n            valid_true_labels.append(labels.item())\n\n            image = image.view(-1, 3, 300, 300).to(cfg[\"device\"])\n            labels = labels.to(cfg[\"device\"])\n            # forward pass\n            outputs = model(image)\n            # calculate the accuracy\n            mean_logp = outputs.mean(dim = 0)\n            preds = torch.argmax(mean_logp).item()\n            valid_pred_labels.append(preds)\n\n    valid_f1 = f1_score(valid_true_labels, valid_pred_labels, average = 'weighted')\n    valid_f1s.append(valid_f1)\n    \n    print()\n    print('Validation F1: {:.2f}%'.format(valid_f1 * 100))\n    torch.save(model.state_dict(), f'./epoch{epoch // cfg[\"check_freq\"]}.pth')\n\n\n","metadata":{"execution":{"iopub.status.busy":"2023-05-12T05:58:39.560709Z","iopub.execute_input":"2023-05-12T05:58:39.561021Z","iopub.status.idle":"2023-05-12T05:58:39.574502Z","shell.execute_reply.started":"2023-05-12T05:58:39.560992Z","shell.execute_reply":"2023-05-12T05:58:39.573031Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Training loop:\n# --------------\ntrain_losses = []  # store train losses at each epoch       \nvalid_f1s = []  # store f1-score for validation at each epoch                                                          \nfor epoch in range(cfg[\"n_epochs\"]):\n    train(train_losses, epoch)\n    # evaluate model (with check_freq)\n    if epoch % cfg[\"check_freq\"] == 0:\n        validate(valid_f1s)\n    scheduler.step()","metadata":{"papermill":{"duration":3470.218849,"end_time":"2023-03-06T19:31:07.577597","exception":false,"start_time":"2023-03-06T18:33:17.358748","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-05-12T05:58:39.576329Z","iopub.execute_input":"2023-05-12T05:58:39.577241Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Get the optimum epoch","metadata":{}},{"cell_type":"code","source":"optimal_epoch = np.argmax(np.array(valid_f1s)) # highest validation F1 epoch / checkpoint frequency","metadata":{"papermill":{"duration":0.036756,"end_time":"2023-03-06T19:31:07.641305","exception":false,"start_time":"2023-03-06T19:31:07.604549","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Validation","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize = (12, 4))\nplt.subplot(1, 2, 1)\nplt.plot(np.arange(len(train_losses)) / cfg[\"n_epochs\"], train_losses, linewidth = 1)\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.title('Training Losses')\nplt.subplot(1, 2, 2)\nplt.plot(np.arange(len(valid_f1s)) * cfg[\"check_freq\"], valid_f1s, linewidth = 1)\nplt.vlines(optimal_epoch * cfg[\"check_freq\"], 0, valid_f1s[optimal_epoch], colors = 'black', linestyles = 'dashed', label = f'Optimal epoch ({optimal_epoch * cfg[\"check_freq\"]})')\nplt.xlabel('Epoch')\nplt.ylabel('Weighted F1')\nplt.ylim(0, 1)\nplt.title('Validation')\nplt.legend(loc = 'lower left')\nplt.savefig('plot.png')\nplt.show()","metadata":{"papermill":{"duration":0.69134,"end_time":"2023-03-06T19:31:08.358922","exception":false,"start_time":"2023-03-06T19:31:07.667582","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Use the best model","metadata":{}},{"cell_type":"code","source":"# Load the model which achieved the largest validation F1:\n# --------------------------------------------------------\nmodel = nn.DataParallel(EfficientNetB0(n_classes = cfg[\"n_classes\"], learnable_modules = ())).to(cfg[\"device\"])\nmodel.load_state_dict(torch.load(f'./epoch{optimal_epoch}.pth'))","metadata":{"papermill":{"duration":0.29888,"end_time":"2023-03-06T19:31:08.682134","exception":false,"start_time":"2023-03-06T19:31:08.383254","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"code","source":"# Submission:\n# -----------\nids = []\npreds = []\nmodel.eval()\nwith torch.no_grad():\n    for x, y in test_loader:\n        ids.append(y[0])\n        mean_logp = model(x.view(-1, 3, 300, 300).to(cfg[\"device\"])).mean(dim = 0)\n        preds.append(torch.argmax(mean_logp).item())\nsubmission = pd.DataFrame({'id': ids, 'label': preds})\nsubmission.to_csv('submission.csv', index = False)\nsubmission.head()","metadata":{"papermill":{"duration":825.091225,"end_time":"2023-03-06T19:44:53.796998","exception":false,"start_time":"2023-03-06T19:31:08.705773","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]}]}