{"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":"# Plant Pathology 2021 - FGVC8","metadata":{}},{"cell_type":"markdown","source":"## Introduction\nIn this notebook, we will train two architectures to detect several diseases on plants using leaf images. The dataset was extracted from the **Plant Pathology 2021 - FGVC8** Kaggle competition. The first architecture will be a simple CNN model from scratch to have make sure the data is well structured for training. The second model will be a pre-trained model that will be finetuned on our dataset.","metadata":{}},{"cell_type":"markdown","source":"## Table of Content\n### [1. Dataset](#Dataset)\n### &nbsp;&nbsp;&nbsp;&nbsp; [1.1. EDA](#EDA)\n### &nbsp;&nbsp;&nbsp;&nbsp; [1.2. Pytorch Dataset](#Pytorch-Dataset)\n### &nbsp;&nbsp;&nbsp;&nbsp; [1.3. Dataset Overview](#Dataset-Overview)\n### &nbsp;&nbsp;&nbsp;&nbsp; [1.4. Data Loader](#Data-Loader)\n### [2. Focal Loss](#Focal-Loss)\n### [3. Custom Model Architecture](#Custom-Model-Architecture)\n### &nbsp;&nbsp;&nbsp;&nbsp; [3.1. Design](#Design)\n### &nbsp;&nbsp;&nbsp;&nbsp; [3.2. Training](#Training)\n### &nbsp;&nbsp;&nbsp;&nbsp; [3.3. Evaluation](#Training)\n### [4. Transfer Learning](#Transfer-Learning)\n### &nbsp;&nbsp;&nbsp;&nbsp; [4.1  Loading Model](#Loading-Model)\n### &nbsp;&nbsp;&nbsp;&nbsp; [4.2  Training](#Training)\n### &nbsp;&nbsp;&nbsp;&nbsp; [4.3  Evaluation](#Training)","metadata":{}},{"cell_type":"markdown","source":"## Libraries & Setup\nWe will install `wandb` which is a very useful MLOps tool. This tool is mainly used to: \n- Track the model's evolution during training and validation through their website\n- Remotely save checkpoints of the model after each epoch in order to be able to *continue where you left off*. This helps especially in our case as kaggle wipes the working directory between sessions.","metadata":{}},{"cell_type":"code","source":"! pip install wandb > /kaggle/working/log.txt","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-01-21T08:03:06.632868Z","iopub.execute_input":"2023-01-21T08:03:06.633276Z","iopub.status.idle":"2023-01-21T08:03:16.150253Z","shell.execute_reply.started":"2023-01-21T08:03:06.633238Z","shell.execute_reply":"2023-01-21T08:03:16.149055Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# File system\nimport os\n\n# Typing\nfrom typing import Callable\nfrom typing import Union\n\n# PyTorch\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset\nfrom torch.nn.functional import one_hot\nfrom torch.utils.data import DataLoader\nfrom torch.utils.data import random_split\nfrom torch.optim import SGD\nfrom torch.optim.lr_scheduler import StepLR\n\n# Torchvision\nfrom torchvision.io import read_image\nfrom torchvision.transforms import Lambda\nfrom torchvision.transforms import Resize, Compose, ToTensor, RandomVerticalFlip, RandomHorizontalFlip\nfrom torchvision.models import resnet50\n\n# Torchmetrics\nfrom torchmetrics import Accuracy\n\n# Sklearn\nfrom sklearn.preprocessing import LabelEncoder\n\n# Numpy/Pandas\nimport pandas as pd\nimport numpy as np\n\n# DataViz\nimport matplotlib.pyplot as plt\n\n# Progress bar\nfrom tqdm.notebook import tqdm\n\n# MLOps\nimport wandb","metadata":{"execution":{"iopub.status.busy":"2023-01-21T08:03:22.569097Z","iopub.execute_input":"2023-01-21T08:03:22.570095Z","iopub.status.idle":"2023-01-21T08:03:27.223445Z","shell.execute_reply.started":"2023-01-21T08:03:22.570042Z","shell.execute_reply":"2023-01-21T08:03:27.222244Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset\n### EDA\n\nThe dataset folder is structured as follows,\n```\n.\n├── plant-pathology-2021-fgvc8\n│   ├── train_images\n│   ├── test_images\n│   ├── sample_submission.csv\n└── └── train.csv\n```\nIntuitively, the first thing to do is take a look at a sample of the train images along with their labels.","metadata":{}},{"cell_type":"code","source":"sample_size = 6\nnrows = 2\nncols = 3\n\nsample = pd.read_csv('/kaggle/input/plant-pathology-2021-fgvc8/train.csv').sample(sample_size)\n\nfig, axs = plt.subplots(nrows = nrows, ncols = ncols, squeeze = False, figsize = (15, 10))\nfor i, (img_path, label) in enumerate(sample.values):\n    img = plt.imread(os.path.join('/kaggle/input/plant-pathology-2021-fgvc8/train_images', img_path))\n    axs[i//3, i%3].imshow(img, aspect = 'auto')\n    axs[i//3, i%3].axis('off')\n    axs[i//3, i%3].set_title(label, fontsize = 16)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-01-21T08:03:29.399557Z","iopub.execute_input":"2023-01-21T08:03:29.399928Z","iopub.status.idle":"2023-01-21T08:03:37.806220Z","shell.execute_reply.started":"2023-01-21T08:03:29.399897Z","shell.execute_reply":"2023-01-21T08:03:37.804885Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's delve deeper into our dataset and figure out some more information.","metadata":{}},{"cell_type":"code","source":"%%bash\n# Using bash mainly because it is faster or simpler than python for os manipulation\necho 'Number of train images'\nls /kaggle/input/plant-pathology-2021-fgvc8/train_images | wc -l\necho\necho 'Number of test images'\nls /kaggle/input/plant-pathology-2021-fgvc8/test_images | wc -l\necho\necho 'Top 2 largest files'\ndu -a /kaggle/input/plant-pathology-2021-fgvc8/train_images/* | sort -n -r | head -n 2\necho ''\necho 'Top 2 smallest files'\ndu -a /kaggle/input/plant-pathology-2021-fgvc8/train_images/* | sort -n -r | tail -n 2","metadata":{"execution":{"iopub.status.busy":"2023-01-21T08:04:59.371040Z","iopub.execute_input":"2023-01-21T08:04:59.371470Z","iopub.status.idle":"2023-01-21T08:06:30.612821Z","shell.execute_reply.started":"2023-01-21T08:04:59.371437Z","shell.execute_reply":"2023-01-21T08:06:30.611316Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"largest_img = plt.imread('/kaggle/input/plant-pathology-2021-fgvc8/train_images/cd3a1d64e6806eb5.jpg')\nsmallest_img = plt.imread('/kaggle/input/plant-pathology-2021-fgvc8/train_images/feb0b0208b5f89ae.jpg')\n\nprint('{:<25}{:<10}'.format('Largest image shape', str(largest_img.shape)))\nprint('{:<25}{:<10}'.format('Smallest image shape', str(smallest_img.shape)))","metadata":{"execution":{"iopub.status.busy":"2023-01-21T08:06:30.615986Z","iopub.execute_input":"2023-01-21T08:06:30.618729Z","iopub.status.idle":"2023-01-21T08:06:31.086398Z","shell.execute_reply.started":"2023-01-21T08:06:30.618690Z","shell.execute_reply":"2023-01-21T08:06:31.085428Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"As expected, images are relatively large, as the smallest has 4,478,976 pixels. Images might require a resize to fit GPU storage.\n\nNow, let's look at our labels and their distribution.","metadata":{}},{"cell_type":"code","source":"labels = pd.read_csv('/kaggle/input/plant-pathology-2021-fgvc8/train.csv', usecols = [1])\n\nplt.style.use('seaborn')\n\nfig, (ax1, ax2) = plt.subplots(1, 2, figsize = (17, 8))\n\ncounts = labels.value_counts()\nindex = [t[0] for t in list(counts.index.values)]\n\nax1.bar(x = index[:7], height = counts.values[:7])\nax2.bar(x = index[7:], height = counts.values[7:])\n\nax1.tick_params(axis = 'x', rotation = 15)\nax2.tick_params(axis = 'x', rotation = 15)\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-01-21T08:06:39.226793Z","iopub.execute_input":"2023-01-21T08:06:39.227180Z","iopub.status.idle":"2023-01-21T08:06:39.580217Z","shell.execute_reply.started":"2023-01-21T08:06:39.227143Z","shell.execute_reply":"2023-01-21T08:06:39.579220Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We notice a huge imbalance in the dataset as 5 classes count less than 20% of the most frequent class. This represents a major problem as the model is less likely to learn to identify less frequent classes. One of the solutions to counter the imbalance problem is using Focal loss. We will go further into details about this type of loss later in the notebook.","metadata":{}},{"cell_type":"markdown","source":"## Focal Loss\n\nAs we mentioned earlier, their is huge imbalance in our dataset and using the classic cross-entropy loss will have the model training be more biased towards dominant classes. In other words, less frequent classes will have a more significant loss. \nThe focal loss is a finetuned cross-entropy loss that puts emphasis on classes with a high loss and give less importance to classes with low loss. This *regularisation* term pushes the model to focus on learning rare classes and put less emphasis on dominant classes.\n$$\n\\large\n\\text{FL}(p_t) = -(1-p_t)^\\gamma\\log(p_t)\n$$\nwith $p_t$ being the probability of object being a certain class.\nHere is an implementation of the Focal Loss as a PyTorch `Module`:","metadata":{}},{"cell_type":"code","source":"class FocalLoss(nn.Module):\n    def __init__(self, alpha: Union[float, torch.Tensor] = 1, gamma: float = 2.):\n        super(FocalLoss, self).__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.epsilon = 1e-12  # prevent training from Nan-loss error \n    \n    def forward(self, logits, target):\n        probs = torch.sigmoid(logits)\n        one_subtract_probs = 1.0 - probs\n        # add epsilon\n        probs_new = probs + self.epsilon\n        one_subtract_probs_new = one_subtract_probs + self.epsilon\n        # calculate focal loss\n        log_pt =  target * torch.log(probs_new) + (1.0 - target) * torch.log(one_subtract_probs_new)\n        pt = torch.exp(log_pt)\n        focal_loss = -1.0 * (self.alpha * (1 - pt) ** self.gamma) * log_pt\n        return torch.mean(focal_loss)","metadata":{"execution":{"iopub.status.busy":"2023-01-21T08:20:02.826710Z","iopub.execute_input":"2023-01-21T08:20:02.827063Z","iopub.status.idle":"2023-01-21T08:20:02.834903Z","shell.execute_reply.started":"2023-01-21T08:20:02.827034Z","shell.execute_reply":"2023-01-21T08:20:02.833968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## PyTorch Dataset\n\nLet's start preparing and structuring our dataset. First, we define a `PlantDataset` class that inherits from PyTorch's `Dataset` class. ","metadata":{}},{"cell_type":"code","source":"class PlantDataset(Dataset):\n    def __init__(self, annotations_file: str, img_dir: str, transform: Callable = None, target_transform: Callable = None):\n        self.img_labels = pd.read_csv(annotations_file)\n        self.n_classes = self.img_labels['labels'].nunique()\n        self.img_dir = img_dir\n        self.transform = transform\n        self.target_transform = target_transform\n        self.class_freq = self.img_labels['labels'].value_counts()\n        \n        # Label encoding target\n        encoder = LabelEncoder()\n        self.img_labels['labels_enc'] = encoder.fit_transform(self.img_labels['labels'].values)\n\n    def __len__(self):\n        return len(self.img_labels)\n\n    def __getitem__(self, idx):\n        img_path = os.path.join(self.img_dir, self.img_labels.iloc[idx, 0])\n        image = read_image(img_path)*1. # We need image tensors to be tensors of floats\n        label = self.img_labels.iloc[idx, -1]\n        if self.transform:\n            image = self.transform(image)\n        if self.target_transform:\n            label = self.target_transform(label)\n        return image, label","metadata":{"execution":{"iopub.status.busy":"2023-01-21T08:20:20.478651Z","iopub.execute_input":"2023-01-21T08:20:20.479044Z","iopub.status.idle":"2023-01-21T08:20:20.490279Z","shell.execute_reply.started":"2023-01-21T08:20:20.479012Z","shell.execute_reply":"2023-01-21T08:20:20.489128Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = PlantDataset(\n    annotations_file = '/kaggle/input/plant-pathology-2021-fgvc8/train.csv',\n    img_dir = '/kaggle/input/plant-pathology-2021-fgvc8/train_images',\n    transform = Compose([\n        Resize((512, 512)), # We need to rescale to avoid GPU memory saturation and make sure the model trains in a reasonable amount of time\n        RandomVerticalFlip(.2), # Data augumentation (1)\n        RandomHorizontalFlip(.2), # Data augumentation (2)\n    ]), \n    target_transform = Lambda(lambda y: torch.zeros(12, dtype = torch.float)\n                                             .scatter_(dim = 0, index = torch.tensor(y), value = 1) # One-hot-encoding target\n    )\n)\n\nprint('Constructed dataset of type', dataset.__class__.__name__)\nprint('Dataset contains', len(dataset), 'annotated images')","metadata":{"execution":{"iopub.status.busy":"2023-01-21T08:21:08.755746Z","iopub.execute_input":"2023-01-21T08:21:08.756094Z","iopub.status.idle":"2023-01-21T08:21:08.789896Z","shell.execute_reply.started":"2023-01-21T08:21:08.756065Z","shell.execute_reply":"2023-01-21T08:21:08.788877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"To prepare the dataset for training, we need to split it into train and validation subsets.\n\n**Note** - A sub-dataset is defined to accelerate model training and evaluation to make sure it works instead of having to wait hours each time.","metadata":{}},{"cell_type":"code","source":"subset_size = 100\nsub_dataset = torch.utils.data.Subset(dataset, list(range(subset_size)))","metadata":{"execution":{"iopub.status.busy":"2023-01-21T08:21:28.182901Z","iopub.execute_input":"2023-01-21T08:21:28.183268Z","iopub.status.idle":"2023-01-21T08:21:28.188819Z","shell.execute_reply.started":"2023-01-21T08:21:28.183236Z","shell.execute_reply":"2023-01-21T08:21:28.187728Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"When doing sanity checks, replace `dataset` with `sub_dataset` for quick model training and evaluation","metadata":{}},{"cell_type":"code","source":"# Train and Validation splits\ntrain_split = int(.95*len(dataset)) # Train is 95%\nval_split = len(dataset) - train_split # Validation is the rest (5%)\ntrain_dataset, val_dataset = random_split(dataset, [train_split, val_split])\n\nprint('Dataset has been split into train and validation subsets.\\n')\nprint('='*35)\nprint('{:<20}{:<15}'.format('Subset', 'Size (images)'))\nprint('-'*35)\nprint('{:<20}{:<15}'.format('Train', len(train_dataset)))\nprint('{:<20}{:<15}'.format('Validation', len(val_dataset)))\nprint('='*35)","metadata":{"execution":{"iopub.status.busy":"2023-01-21T08:22:12.974880Z","iopub.execute_input":"2023-01-21T08:22:12.975247Z","iopub.status.idle":"2023-01-21T08:22:12.985845Z","shell.execute_reply.started":"2023-01-21T08:22:12.975191Z","shell.execute_reply":"2023-01-21T08:22:12.984244Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"To properly load the dataset into our GPU and later use it to fit the model, we will create a `DeviceDataLoader` class that wraps around PyTorch's `DataLoader` class.","metadata":{}},{"cell_type":"code","source":"train_dataloader = DataLoader(train_dataset, batch_size = 8, shuffle = True)\nval_dataloader = DataLoader(val_dataset, batch_size = 8, shuffle = True)\n\nbatch = next(iter(train_dataloader))\nprint('Successfully loaded a batch of {} images and {} labels'.format(len(batch[0]), len(batch[1])))","metadata":{"execution":{"iopub.status.busy":"2023-01-21T08:22:24.616439Z","iopub.execute_input":"2023-01-21T08:22:24.616817Z","iopub.status.idle":"2023-01-21T08:22:27.474491Z","shell.execute_reply.started":"2023-01-21T08:22:24.616787Z","shell.execute_reply":"2023-01-21T08:22:27.473257Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def to_device(data, device): # Move data to a device ('CPU' or 'GPU')\n    if isinstance(data, (list, tuple)):\n        return [to_device(x, device) for x in data] # If data is a list of tensors\n    return data.to(device, non_blocking = True) # If data is a single tensor\n\nclass DeviceDataLoader(): # Wraps a dataloader to move data to a device\n    def __init__(self, dl, device):\n        self.dl = dl\n        self.device = device\n    def __iter__(self):\n        for b in self.dl: # For each batch\n            yield to_device(b, self.device) # Return but doesn't stop the for loop (NB: Garabage collection happens automatically)\n    def __len__(self):\n        return (len(self.dl))\n\nget_default_device = lambda : torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\n\ntrain_dataloader = DeviceDataLoader(train_dataloader, device = get_default_device())\nval_dataloader = DeviceDataLoader(val_dataloader, device = get_default_device())","metadata":{"execution":{"iopub.status.busy":"2023-01-21T08:23:41.696676Z","iopub.execute_input":"2023-01-21T08:23:41.697061Z","iopub.status.idle":"2023-01-21T08:23:41.704893Z","shell.execute_reply.started":"2023-01-21T08:23:41.697033Z","shell.execute_reply":"2023-01-21T08:23:41.703772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Custom Model Architecture\n### Design\n\nFirst, we will construct a custom model architecture. The architecture is pretty simple,\n\n$$\n\\fbox{Conv} \\xrightarrow{\\text{ReLU}} \\fbox{MaxPool} \\rightarrow \\fbox{Dropout} \\rightarrow \\fbox{Conv} \\xrightarrow{\\text{ReLU}} \\fbox{MaxPool} \\rightarrow \\fbox{Dropout} \\rightarrow \\fbox{Conv} \\xrightarrow{\\text{ReLU}} \\fbox{MaxPool} \\rightarrow \\fbox{Dropout} \\rightarrow \\fbox{Conv} \\xrightarrow{\\text{ReLU}} \\fbox{MaxPool} \\rightarrow \\fbox{Dropout} \\rightarrow \\fbox{FC} \\xrightarrow{\\text{ReLU}} \\fbox{Dropout} \\rightarrow \\fbox{FC}\n$$\n\nWe also define our feed forward method as well as a fit method which will be called to train the model.","metadata":{}},{"cell_type":"code","source":"class ScratchModel(nn.Module):\n    def __init__(self, num_classes: int):\n        super(ScratchModel, self).__init__()\n        self.conv1 = nn.Sequential(\n            nn.Conv2d(in_channels = 3, out_channels = 16, kernel_size = (5, 5)), # out_channels = number of filters\n            nn.ReLU(),\n            nn.MaxPool2d(kernel_size = (2, 2)),\n            nn.Dropout(.25)\n        )\n        self.conv2 = nn.Sequential(\n            nn.Conv2d(in_channels = 16, out_channels = 32, kernel_size = (5, 5)), # out_channels = number of filters\n            nn.ReLU(),\n            nn.MaxPool2d(kernel_size = (2, 2)),\n            nn.Dropout(.25)\n        )\n        self.conv3 = nn.Sequential(\n            nn.Conv2d(in_channels = 32, out_channels = 64, kernel_size = (5, 5)), # out_channels = number of filters\n            nn.ReLU(),\n            nn.MaxPool2d(kernel_size = (2, 2)),\n            nn.Dropout(.25)\n        )\n        self.conv4 = nn.Sequential(\n            nn.Conv2d(in_channels = 64, out_channels = 64, kernel_size = (5, 5)), # out_channels = number of filters\n            nn.ReLU(),\n            nn.MaxPool2d(kernel_size = (2, 2)),\n            nn.Dropout(.25)\n        )\n        self.fc1 = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(in_features = 64*4*4, out_features = 128), # TODO: Write formula\n            nn.ReLU(),\n            nn.Dropout(.5)\n        )\n        self.fc2 = nn.Sequential(\n            nn.Linear(in_features = 128, out_features = 64),\n            nn.ReLU(),\n            nn.Dropout(.5)\n        )\n        self.fc3 = nn.Sequential(\n            nn.Linear(64, num_classes)\n        )\n    \n    def forward(self, x):\n        out = self.conv1(x)\n        out = self.conv2(out)\n        out = self.conv3(out)\n        # out = self.conv4(out)\n        print(out.shape)\n        out = self.fc1(out)\n        out = self.fc2(out)\n        out = self.fc3(out)\n        return out\n    \n    def train_one_epoch(self, optimizer: torch.optim.Optimizer, loss_fn: Callable, metric_fn: Callable, train_loader: DataLoader, epoch_idx: int):\n        # init\n        running_loss = 0.\n        running_metric = 0.\n        avg_loss = 0.\n        avg_metric = 0.\n        \n        with tqdm(train_loader, unit = ' batch', colour = '#ffc933') as tepoch:\n            tepoch.set_description(f'Training (Epoch {epoch_idx})')\n            for i, (inputs, labels) in enumerate(tepoch):\n                # Reset optimiser gradients to zero\n                optimizer.zero_grad()\n                # Forward propagation\n                outputs = self(inputs)\n                # Compute the loss\n                loss = loss_fn(outputs, labels)\n                # Backward propagation\n                loss.backward()\n                # Update weights\n                optimizer.step()\n                # Gather data and report\n                running_loss += loss.item()\n                running_metric += metric_fn(outputs, labels)\n                if i % 100 == 99: # Each 100 batches\n                    avg_loss = running_loss / 100 # average loss per 100 batch\n                    avg_metric = running_metric / 100 # average accuracy per 100 batch\n                    tepoch.set_postfix_str('Average loss = {:.4f}\\tAverage accuracy = {:.4f}'.format(avg_loss, avg_metric))\n                    running_loss = 0.\n                    running_metric = 0.\n                    \n    def fit(self, optimizer: torch.optim.Optimizer, lr_scheduler, loss_fn: Callable, metric_fn: Callable, train_loader: DataLoader, val_loader: DataLoader, n_epochs: int):\n        for epoch in range(1, n_epochs + 1):        \n            # Training (Gradient tracking is on)\n            self.train()\n            avg_loss = self.train_one_epoch(optimizer, loss_fn, metric_fn, train_loader, epoch_idx = epoch)\n\n            # Evaluation (Gradient tracking is off)\n            self.eval() # Does not set gradient tracking off (only deactivates Dropout, BatchNorm, ...)\n            with torch.no_grad():\n                running_vloss = 0\n                running_vmetric = 0.\n                avg_vloss = 0.\n                avg_vmetric = 0\n                \n                with tqdm(val_loader, unit = ' batch', colour = '#20b5f2') as vtepoch:\n                    vtepoch.set_description(f'Evaluation (Epoch {epoch})')\n                    for vbatch in vtepoch:\n                        vinputs, vlabels = vbatch\n                        voutputs = self(vinputs)\n                        vloss = loss_fn(voutputs, vlabels)\n                        running_vloss += vloss\n                        running_vmetric += metric_fn(voutputs, vlabels)\n                    avg_vloss = running_vloss / len(val_loader)\n                    avg_vmetric = running_vmetric / len(val_loader)\n                    vtepoch.set_postfix_str('Average loss = {:.4f}\\t Average accuracy = {:.4f}'.format(avg_vloss, avg_vmetric))\n\n            # Update learning rate\n            lr_scheduler.step()","metadata":{"execution":{"iopub.status.busy":"2023-01-21T08:49:03.694095Z","iopub.execute_input":"2023-01-21T08:49:03.694520Z","iopub.status.idle":"2023-01-21T08:49:03.717948Z","shell.execute_reply.started":"2023-01-21T08:49:03.694487Z","shell.execute_reply":"2023-01-21T08:49:03.716914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We will also use **accuracy** to track the model's performance during training.","metadata":{}},{"cell_type":"code","source":"def accuracy(outputs, labels):\n    _, preds = torch.max(outputs, dim = 1)\n    _, labels = torch.max(labels, dim = 1)\n    return torch.sum(preds == labels).item() / len(preds)","metadata":{"execution":{"iopub.status.busy":"2023-01-21T08:42:40.038410Z","iopub.execute_input":"2023-01-21T08:42:40.038783Z","iopub.status.idle":"2023-01-21T08:42:40.046132Z","shell.execute_reply.started":"2023-01-21T08:42:40.038753Z","shell.execute_reply":"2023-01-21T08:42:40.044945Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's gather everything back and define our optimiser.","metadata":{}},{"cell_type":"code","source":"model = ScratchModel(num_classes = 12).to(get_default_device())\noptimizer = torch.optim.SGD(model.parameters(), lr = 0.001, momentum = 0.9)\nlr_scheduler = StepLR(optimizer, step_size = 10, gamma = .5) # Cut learning rate in half every 30 epochs\nloss_fn = nn.CrossEntropyLoss()\nmetric_fn = accuracy","metadata":{"execution":{"iopub.status.busy":"2023-01-21T08:49:14.887948Z","iopub.execute_input":"2023-01-21T08:49:14.888314Z","iopub.status.idle":"2023-01-21T08:49:14.906273Z","shell.execute_reply.started":"2023-01-21T08:49:14.888283Z","shell.execute_reply":"2023-01-21T08:49:14.905355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Training","metadata":{}},{"cell_type":"code","source":"model.fit(optimizer, lr_scheduler, loss_fn, metric_fn, train_dataloader, val_dataloader, n_epochs = 1)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Transfer Learning\n\n### Loading Model\nNow we will use a pretrained model `ResNet-50` to classify plant images. The last layer is changed to fit our dataset (a vector of 12 values each being a probability that an image belong to a certain class).\n\nSince this is a relatively heavy model and a single epoch is likely ot take hours, we will integrate `wandb`.","metadata":{}},{"cell_type":"code","source":"run = wandb.init(project = 'plant-pathology-2021', name = 'ResNet50', entity = 'nizarmasmoudi', reinit = True)","metadata":{"execution":{"iopub.status.busy":"2023-01-21T08:30:46.352311Z","iopub.execute_input":"2023-01-21T08:30:46.352998Z","iopub.status.idle":"2023-01-21T08:31:22.188077Z","shell.execute_reply.started":"2023-01-21T08:30:46.352964Z","shell.execute_reply":"2023-01-21T08:31:22.187112Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Similar to what we did earlier, we define our model with different helpful method. Additionally, we define a `push_to_wandb` method and a `load_from_wandb` in order to push/load the model to/from Weight & Biases ","metadata":{}},{"cell_type":"code","source":"class ResNet50(nn.Module):\n    '''\n    ResNet50 model.\n    '''\n    def __init__(self):\n        super(ResNet50, self).__init__()\n        self.resnet50 = resnet50(pretrained = True)\n        self.resnet50.fc = nn.Linear(self.resnet50.fc.in_features, 12) # No need for softmax because nn.CrossEntropyLoss computes softmax\n        \n    def forward(self, x: torch.Tensor):\n        x = self.resnet50(x)\n        return x\n    \n    def train_one_epoch(self, optimizer: torch.optim.Optimizer, loss_fn: Callable, metric_fn: Callable, train_loader: DataLoader, epoch_idx: int):\n        # init\n        running_loss = 0.\n        running_metric = 0.\n        avg_loss = 0.\n        avg_metric = 0.\n        \n        with tqdm(train_loader, unit = ' batch', colour = '#ffc933') as tepoch:\n            tepoch.set_description(f'Training (Epoch {epoch_idx})')\n            for i, (inputs, labels) in enumerate(tepoch):\n                # Reset optimiser gradients to zero\n                optimizer.zero_grad()\n                # Forward propagation\n                outputs = self(inputs)\n                # Compute the loss\n                loss = loss_fn(outputs, labels)\n                # Backward propagation\n                loss.backward()\n                # Update weights\n                optimizer.step()\n                # Gather data and report\n                running_loss += loss.item()\n                running_metric += metric_fn(outputs, labels)\n                if i % 100 == 99: # Each 100 batches\n                    avg_loss = running_loss / 100 # average loss per 100 batch\n                    avg_metric = running_metric / 100 # average accuracy per 100 batch\n                    tepoch.set_postfix_str('Average loss = {:.4f}\\tAverage accuracy = {:.4f}'.format(avg_loss, avg_metric))\n                    wandb.log({'Training loss': avg_loss, 'Training accuracy': avg_metric, 'Batch': i + 1, 'Epoch': epoch_idx})\n                    running_loss = 0.\n                    running_metric = 0.\n                    \n    def fit(self, optimizer: torch.optim.Optimizer, lr_scheduler, loss_fn: Callable, metric_fn: Callable, train_loader: DataLoader, val_loader: DataLoader, n_epochs: int):\n        for epoch in range(1, n_epochs + 1):        \n            # Training (Gradient tracking is on)\n            self.train()\n            avg_loss = self.train_one_epoch(optimizer, loss_fn, metric_fn, train_loader, epoch_idx = epoch)\n\n            # Evaluation (Gradient tracking is off)\n            self.eval() # Does not set gradient tracking off (only deactivates Dropout, Batchnorm, ...)\n            with torch.no_grad():\n                running_vloss = 0\n                running_vmetric = 0.\n                avg_vloss = 0.\n                avg_vmetric = 0\n                \n                with tqdm(val_loader, unit = ' batch', colour = '#20b5f2') as vtepoch:\n                    vtepoch.set_description(f'Evaluation (Epoch {epoch})')\n                    for vbatch in vtepoch:\n                        vinputs, vlabels = vbatch\n                        voutputs = self(vinputs)\n                        vloss = loss_fn(voutputs, vlabels)\n                        running_vloss += vloss\n                        running_vmetric += metric_fn(voutputs, vlabels)\n                    avg_vloss = running_vloss / len(val_loader)\n                    avg_vmetric = running_vmetric / len(val_loader)\n                    vtepoch.set_postfix_str('Average loss = {:.4f}\\t Average accuracy = {:.4f}'.format(avg_vloss, avg_vmetric))\n                    wandb.log({'Validation loss': avg_vloss, 'Validation accuracy': avg_vmetric, 'Epoch': epoch})\n\n            # Update learning rate\n            lr_scheduler.step()\n            \n    def evaluate(self, val_loader, loss_fn, metric_fn):\n        self.eval() # Does not set gradient tracking off (only deactivates Dropout, Batchnorm, ...)\n        with torch.no_grad():\n            running_vloss = 0\n            running_vmetric = 0.\n            avg_vloss = 0.\n            avg_vmetric = 0\n                \n            with tqdm(val_loader, unit = ' batch', colour = '#20b5f2') as vtepoch:\n                vtepoch.set_description(f'Evaluation')\n                for vbatch in vtepoch:\n                    vinputs, vlabels = vbatch\n                    voutputs = self(vinputs)\n                    vloss = loss_fn(voutputs, vlabels)\n                    running_vloss += vloss\n                    running_vmetric += metric_fn(voutputs, vlabels)\n                avg_vloss = running_vloss / len(val_loader)\n                avg_vmetric = running_vmetric / len(val_loader)\n                vtepoch.set_postfix_str('Average loss = {:.4f}\\t Average accuracy = {:.4f}'.format(avg_vloss, avg_vmetric))\n        print('Average loss = {:.4f}\\t Average accuracy = {:.4f}'.format(avg_vloss, avg_vmetric))\n            \n    def push_to_wandb(self, desc: str, wandb_run, pth_file: str = '/kaggle/working/model/model.pth', end_run: bool = False):\n        if not os.path.exists(os.path.dirname(pth_file)):\n            os.makedirs(os.path.dirname(pth_file))\n        torch.save(self.state_dict(), pth_file)\n\n        artifact = wandb.Artifact('model', type = 'model', description = desc)\n        artifact.add_file(pth_file)\n        run.log_artifact(artifact)\n        if end_run:\n            run.finish()\n        \n    def load_from_wandb(self, name: str):\n        artifact = run.use_artifact(name, type = 'model')\n        artifact_dir = artifact.download()\n        self.load_state_dict(torch.load(os.path.join(artifact_dir, 'model.pth')))","metadata":{"execution":{"iopub.status.busy":"2023-01-21T08:33:51.041632Z","iopub.execute_input":"2023-01-21T08:33:51.042021Z","iopub.status.idle":"2023-01-21T08:33:51.066510Z","shell.execute_reply.started":"2023-01-21T08:33:51.041988Z","shell.execute_reply":"2023-01-21T08:33:51.065537Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def accuracy(outputs, labels):\n    _, preds = torch.max(outputs, dim = 1)\n    _, labels = torch.max(labels, dim = 1)\n    return torch.sum(preds == labels).item() / len(preds)","metadata":{"execution":{"iopub.status.busy":"2023-01-21T08:33:53.558916Z","iopub.execute_input":"2023-01-21T08:33:53.559340Z","iopub.status.idle":"2023-01-21T08:33:53.567662Z","shell.execute_reply.started":"2023-01-21T08:33:53.559306Z","shell.execute_reply":"2023-01-21T08:33:53.566661Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = ResNet50().to(get_default_device())\noptimizer = torch.optim.SGD(model.parameters(), lr = 0.001, momentum = 0.9)\nlr_scheduler = StepLR(optimizer, step_size = 10, gamma = .5) # Cut learning rate in half every 30 epochs\nloss_fn = FocalLoss()\nmetric_fn = accuracy","metadata":{"execution":{"iopub.status.busy":"2023-01-21T08:33:57.122089Z","iopub.execute_input":"2023-01-21T08:33:57.122476Z","iopub.status.idle":"2023-01-21T08:34:02.184593Z","shell.execute_reply.started":"2023-01-21T08:33:57.122445Z","shell.execute_reply":"2023-01-21T08:34:02.183441Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Training\nAs mentioned earlier,  training this model might takes multiple runs distributed on different days. We will, therefore, omit the training process in this notebook and skip to evaluating the model. However, to train the model, we use the `fit` method as follows,\n```\nmodel.fit(optimizer, lr_scheduler, loss_fn, metric_fn, train_dataloader, val_dataloader, n_epochs = 3)\n```\nand then, push to model to W&B, using the pre-defined method,\n```\nmodel.push_to_wandb(wandb_run = run, desc = 'Some description')\n```","metadata":{}},{"cell_type":"markdown","source":"### Evaluation\nIt's time to evaluate the model. Let's, first, load our last checkpoint from W&B and run the evaluation method.","metadata":{}},{"cell_type":"code","source":"model = ResNet50()\nmodel.load_from_wandb(name = 'nizarmasmoudi/plant-pathology-2021/model:v2')\nmodel = model.to(get_default_device())","metadata":{"execution":{"iopub.status.busy":"2023-01-21T08:35:11.036158Z","iopub.execute_input":"2023-01-21T08:35:11.036918Z","iopub.status.idle":"2023-01-21T08:35:13.887560Z","shell.execute_reply.started":"2023-01-21T08:35:11.036874Z","shell.execute_reply":"2023-01-21T08:35:13.886514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.evaluate(val_dataloader, loss_fn, metric_fn)","metadata":{"execution":{"iopub.status.busy":"2023-01-21T08:35:21.658722Z","iopub.execute_input":"2023-01-21T08:35:21.659109Z","iopub.status.idle":"2023-01-21T08:41:21.807312Z","shell.execute_reply.started":"2023-01-21T08:35:21.659076Z","shell.execute_reply":"2023-01-21T08:41:21.806269Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The model manages to get relatively good results on validation subset. The test subset, however, is omitted by Kaggle and we'll only be able to get results once we submit the notebook.","metadata":{}}]}