{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":13836,"databundleVersionId":1718836,"sourceType":"competition"}],"dockerImageVersionId":30528,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np \nimport pandas as pd\n\nimport os\nimport random\nimport matplotlib.pyplot as plt","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-09-01T21:16:18.698172Z","iopub.execute_input":"2024-09-01T21:16:18.698503Z","iopub.status.idle":"2024-09-01T21:16:18.709523Z","shell.execute_reply.started":"2024-09-01T21:16:18.698478Z","shell.execute_reply":"2024-09-01T21:16:18.708738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Understanding Dataset","metadata":{}},{"cell_type":"code","source":"input_path = '/kaggle/input/cassava-leaf-disease-classification/'","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:16:24.330419Z","iopub.execute_input":"2024-09-01T21:16:24.331243Z","iopub.status.idle":"2024-09-01T21:16:24.335037Z","shell.execute_reply.started":"2024-09-01T21:16:24.331211Z","shell.execute_reply":"2024-09-01T21:16:24.333987Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv(input_path+'train.csv')\nprint(train_df.shape)\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:16:27.647430Z","iopub.execute_input":"2024-09-01T21:16:27.647787Z","iopub.status.idle":"2024-09-01T21:16:27.689715Z","shell.execute_reply.started":"2024-09-01T21:16:27.647756Z","shell.execute_reply":"2024-09-01T21:16:27.688680Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Number of Unique values under columns\nprint(train_df.nunique())\n\n# Unique values of Column 'label'\nprint(train_df['label'].unique())","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:16:31.685828Z","iopub.execute_input":"2024-09-01T21:16:31.686620Z","iopub.status.idle":"2024-09-01T21:16:31.703537Z","shell.execute_reply.started":"2024-09-01T21:16:31.686588Z","shell.execute_reply":"2024-09-01T21:16:31.702538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Visualize the distribution of labels in the dataset\ndf = train_df['label'].value_counts() # normalize=True\nax = df.plot(kind='bar',figsize=(10,4))\n#ax.set_ylim(0,1)\nax.set_title('Distribution of target labels')\nfor container in ax.containers:\n    ax.bar_label(container, fmt=lambda x: f'{x:0.0f} ({x/df.sum():0.3f})')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:16:34.584150Z","iopub.execute_input":"2024-09-01T21:16:34.585073Z","iopub.status.idle":"2024-09-01T21:16:34.887699Z","shell.execute_reply.started":"2024-09-01T21:16:34.585029Z","shell.execute_reply":"2024-09-01T21:16:34.886771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We can see there is an imbalance in the distribution of target label. Notably, label `3` has twelve times more samples compared to label `0`, and about six times more than the rest of the labels.\n\nWe have to use data augmentation and sampling techniques to improve the class balance during training","metadata":{}},{"cell_type":"code","source":"health_df = pd.read_json(input_path+'label_num_to_disease_map.json',typ='series')\ndisplay(health_df)","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:16:38.632187Z","iopub.execute_input":"2024-09-01T21:16:38.632557Z","iopub.status.idle":"2024-09-01T21:16:38.648254Z","shell.execute_reply.started":"2024-09-01T21:16:38.632528Z","shell.execute_reply":"2024-09-01T21:16:38.646541Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Including the corresponding 'disease' and 'img_path' for the dataset\nhealth_map = dict(health_df)\ntrain_df['disease'] = train_df['label'].map(health_map)\ntrain_df['img_path'] = input_path + 'train_images/' + train_df['image_id']\n\nprint(train_df.shape)\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:16:42.124114Z","iopub.execute_input":"2024-09-01T21:16:42.124471Z","iopub.status.idle":"2024-09-01T21:16:42.144531Z","shell.execute_reply.started":"2024-09-01T21:16:42.124441Z","shell.execute_reply":"2024-09-01T21:16:42.143581Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_img_dir = input_path+'train_images/'\nimg = 0 \nfor _,_,files in os.walk(train_img_dir):\n    img += len(files)\nprint('Total number of available images: '+str(img))","metadata":{"execution":{"iopub.status.busy":"2024-09-01T20:49:47.982689Z","iopub.execute_input":"2024-09-01T20:49:47.983637Z","iopub.status.idle":"2024-09-01T20:50:20.811413Z","shell.execute_reply.started":"2024-09-01T20:49:47.983582Z","shell.execute_reply":"2024-09-01T20:50:20.810052Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Building Training and Validation datasets","metadata":{}},{"cell_type":"code","source":"from tqdm import tqdm\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import f1_score\n\nimport torch\nfrom torch import nn\nfrom torch.optim import Adam\nfrom torch.optim.lr_scheduler import ExponentialLR\nfrom torch.utils.data import Dataset, WeightedRandomSampler, BatchSampler, DataLoader\n\nimport torchvision.transforms.v2 as transforms\nimport torchvision.transforms.functional as func\nimport torchvision.models as models\nimport torchvision.io as io","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:16:52.224667Z","iopub.execute_input":"2024-09-01T21:16:52.225747Z","iopub.status.idle":"2024-09-01T21:16:52.231440Z","shell.execute_reply.started":"2024-09-01T21:16:52.225713Z","shell.execute_reply":"2024-09-01T21:16:52.230297Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Config:\n    SEED = 42\n    VAL_SPLIT = 0.2\n    IMAGE_SIZE = [512, 512]\n    BATCH_SIZE = 16\n    EPOCHS = 50\n    VERBOSE = 1\n    CLASSES = ['0', '1', '2', '3', '4']\n    #AUTOTUNE = tf.data.AUTOTUNE\n\nconfig = Config()","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:16:55.266582Z","iopub.execute_input":"2024-09-01T21:16:55.266941Z","iopub.status.idle":"2024-09-01T21:16:55.272160Z","shell.execute_reply.started":"2024-09-01T21:16:55.266912Z","shell.execute_reply":"2024-09-01T21:16:55.271162Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Split with stratification on labels, i.e., split labels in equal proportions\ntrain_data, valid_data = train_test_split(\n    train_df, \n    test_size    = config.VAL_SPLIT, \n    random_state = config.SEED,\n    stratify     = train_df['label']\n)\n\ntrain_data.shape, valid_data.shape","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:17:00.386412Z","iopub.execute_input":"2024-09-01T21:17:00.386810Z","iopub.status.idle":"2024-09-01T21:17:00.408972Z","shell.execute_reply.started":"2024-09-01T21:17:00.386779Z","shell.execute_reply":"2024-09-01T21:17:00.408074Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Add 'label_counts' for sampling data later, and reset_index after split\ncounts_df = train_data['label'].value_counts()\ncounts_map = dict(counts_df)\ntrain_data['label_counts'] = train_data['label'].map(counts_map)\ntrain_data.reset_index(drop=True, inplace=True)\ntrain_data.head()","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:17:03.726200Z","iopub.execute_input":"2024-09-01T21:17:03.726870Z","iopub.status.idle":"2024-09-01T21:17:03.742964Z","shell.execute_reply.started":"2024-09-01T21:17:03.726839Z","shell.execute_reply":"2024-09-01T21:17:03.742024Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Add 'label_counts' for sampling data later, and reset_index after split\ncounts_df = valid_data['label'].value_counts()\ncounts_map = dict(counts_df)\nvalid_data['label_counts'] = valid_data['label'].map(counts_map)\nvalid_data.reset_index(drop=True, inplace=True)\nvalid_data.sample(n=5)","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:17:06.772367Z","iopub.execute_input":"2024-09-01T21:17:06.772723Z","iopub.status.idle":"2024-09-01T21:17:06.788469Z","shell.execute_reply.started":"2024-09-01T21:17:06.772692Z","shell.execute_reply":"2024-09-01T21:17:06.787564Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualize training and validation splits","metadata":{}},{"cell_type":"code","source":"# Visualize the distribution of labels in training and validation splits\nax = pd.concat([\n        train_data['label'].value_counts(normalize=True),\n        valid_data['label'].value_counts(normalize=True)],\n        axis=1\n     ).plot(kind='bar',figsize=(10,3))\n\nax.set_ylim(0,1)\nax.set_title('Ratios of labels in Train & Test splits')\nax.legend(labels=['train_splits','valid_splits'])\nfor container in ax.containers:\n    ax.bar_label(container, fmt='%.2f')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-09-01T20:51:53.883093Z","iopub.execute_input":"2024-09-01T20:51:53.883535Z","iopub.status.idle":"2024-09-01T20:51:54.260245Z","shell.execute_reply.started":"2024-09-01T20:51:53.883504Z","shell.execute_reply":"2024-09-01T20:51:54.259023Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"decoder = transforms.Compose([\n    transforms.ToImageTensor(),\n    transforms.ConvertDtype(torch.float32)\n])\n\npaths  = train_data.img_path.tolist() # very useful to pick random paths later on","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:17:13.879566Z","iopub.execute_input":"2024-09-01T21:17:13.880303Z","iopub.status.idle":"2024-09-01T21:17:13.885799Z","shell.execute_reply.started":"2024-09-01T21:17:13.880272Z","shell.execute_reply":"2024-09-01T21:17:13.884846Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Visualize orginal dataset images\n# np.random.randint(len(train_data.index))\nimage = io.read_image(random.choice(paths))\nimage = decoder(image)\ndisplay(image.shape)\nplt.imshow(func.to_pil_image(image))\n# plt.title(train_data[]\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:17:21.157461Z","iopub.execute_input":"2024-09-01T21:17:21.157809Z","iopub.status.idle":"2024-09-01T21:17:21.649809Z","shell.execute_reply.started":"2024-09-01T21:17:21.157779Z","shell.execute_reply":"2024-09-01T21:17:21.648777Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"decoder = transforms.Compose([\n    transforms.Resize(config.IMAGE_SIZE),\n    transforms.ToImageTensor(),\n    transforms.ConvertDtype(torch.float32)\n])\n\naugmenter = transforms.Compose([\n    transforms.RandomHorizontalFlip(0.5),\n    transforms.ColorJitter(brightness=0.3, contrast=0.3),\n    transforms.RandomRotation(0.2*2*np.pi)\n])","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:17:25.750821Z","iopub.execute_input":"2024-09-01T21:17:25.751200Z","iopub.status.idle":"2024-09-01T21:17:25.757656Z","shell.execute_reply.started":"2024-09-01T21:17:25.751171Z","shell.execute_reply":"2024-09-01T21:17:25.756412Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NUM_FIGS = 3\nfig, ax = plt.subplots(NUM_FIGS, 2, sharex=True, sharey=True, figsize=(8,8))\n\nfor i in range(NUM_FIGS):\n    path = random.choice(paths)\n    image = io.read_image(path)\n    # Original\n    image = decoder(image)\n    ax[i][0].imshow(func.to_pil_image(image))\n    ax[i][0].set_title('Original', fontsize=8)\n    # Augmented\n    image = augmenter(image)\n    ax[i][1].imshow(func.to_pil_image(image))\n    ax[i][1].set_title('Augmented', fontsize=8)\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-09-01T20:52:59.404894Z","iopub.execute_input":"2024-09-01T20:52:59.405817Z","iopub.status.idle":"2024-09-01T20:53:01.238346Z","shell.execute_reply.started":"2024-09-01T20:52:59.405779Z","shell.execute_reply":"2024-09-01T20:53:01.237113Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Custom Dataloader and sampler for having a balanced Training dataset ","metadata":{}},{"cell_type":"code","source":"class Cassava_Dataset(Dataset):\n    def __init__(self, paths, labels, augment=False):\n        self.paths = paths\n        self.labels = labels\n        self.augment = augment\n        \n    def __len__(self):\n        return len(self.labels)\n    \n    def __getitem__(self, index):\n        image = io.read_image(self.paths.iloc[index])\n        image = decoder(image)\n        if self.augment:\n            image = augmenter(image)\n        label = self.labels.iloc[index]\n        return image, label","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:17:36.855060Z","iopub.execute_input":"2024-09-01T21:17:36.855732Z","iopub.status.idle":"2024-09-01T21:17:36.862089Z","shell.execute_reply.started":"2024-09-01T21:17:36.855698Z","shell.execute_reply":"2024-09-01T21:17:36.861153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Train and Valid datasets\ntrain_dataset = Cassava_Dataset(train_data['img_path'], train_data['label'], augment=True)\nvalid_dataset = Cassava_Dataset(valid_data['img_path'], valid_data['label'])\n\n# Training dataset\nrand_index = np.random.randint(len(train_dataset))\nimages, labels = train_dataset[rand_index]\nimages.shape, labels.shape","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:17:45.221066Z","iopub.execute_input":"2024-09-01T21:17:45.222015Z","iopub.status.idle":"2024-09-01T21:17:45.253818Z","shell.execute_reply.started":"2024-09-01T21:17:45.221980Z","shell.execute_reply":"2024-09-01T21:17:45.252755Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Visualize few images and labels after building dataset\nfig, ax = plt.subplots(3, 3, sharex=True, sharey=True, figsize=(10,10))\n\nfor i in range(3):\n    for j in range(3):\n        rand_index = np.random.randint(len(train_dataset))\n        image, label = train_dataset[rand_index]\n        label = int(label)\n        #label = int(labels[i*3+j])\n        imgs = image.permute(1,2,0)\n        ax[i][j].imshow(imgs)\n        ax[i][j].set_title(health_map[label], fontsize=8)\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-09-01T20:53:40.751026Z","iopub.execute_input":"2024-09-01T20:53:40.751521Z","iopub.status.idle":"2024-09-01T20:53:44.012835Z","shell.execute_reply.started":"2024-09-01T20:53:40.751483Z","shell.execute_reply":"2024-09-01T20:53:44.011401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CustomSampler = WeightedRandomSampler(\n    torch.FloatTensor(list(1.0 /train_data['label_counts'])), \n    # weightage based on class count\n    num_samples = len(train_dataset), \n    replacement = True\n)\n\nTrainBatchSampler = BatchSampler(\n    sampler = CustomSampler,\n    batch_size = config.BATCH_SIZE,\n    drop_last = False\n)\n\n# No ValidBatchSampler, the validation set should reflect the real data distribution\n# https://datascience.stackexchange.com/questions/76056/for-imbalanced-classification-should-the-validation-dataset-be-balanced","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:17:58.878868Z","iopub.execute_input":"2024-09-01T21:17:58.879540Z","iopub.status.idle":"2024-09-01T21:17:58.889802Z","shell.execute_reply.started":"2024-09-01T21:17:58.879509Z","shell.execute_reply":"2024-09-01T21:17:58.888852Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Check class distribution while using CustomSampler\nsample = list(CustomSampler)\ncounts = {'0':0, '1':0, '2':0, '3':0, '4':0}\nfor i in sample:\n    label = train_data['label'][i]\n    counts[str(label)] += 1\n\nplt.bar(range(len(counts)), list(counts.values()), tick_label=list(counts.keys()))\nplt.title('Sample Class distribution when using CustomSampler (Training)')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:18:02.713138Z","iopub.execute_input":"2024-09-01T21:18:02.713572Z","iopub.status.idle":"2024-09-01T21:18:03.168615Z","shell.execute_reply.started":"2024-09-01T21:18:02.713532Z","shell.execute_reply":"2024-09-01T21:18:03.167543Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Check class distribution while using TrainBatchSampler\nbatches = list(TrainBatchSampler)\ndist = {'0':0, '1':0, '2':0, '3':0, '4':0} # distribution\nfor batch in batches:    \n    for id in batch:\n        label = train_data['label'][id]\n        dist[str(label)]+=1\n\nprint(\"Batch samples: \",sum(dist.values()))\nplt.bar(range(len(dist)), list(dist.values()), tick_label=list(dist.keys()))\nplt.title('Sample Class distribution when using TrainBatchSampler (Training)')\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:18:12.325793Z","iopub.execute_input":"2024-09-01T21:18:12.326517Z","iopub.status.idle":"2024-09-01T21:18:12.789490Z","shell.execute_reply.started":"2024-09-01T21:18:12.326483Z","shell.execute_reply":"2024-09-01T21:18:12.788624Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Train and valid. dataloader\ntrain_loader = DataLoader(train_dataset, batch_sampler = TrainBatchSampler)\nvalid_loader = DataLoader(valid_dataset, config.BATCH_SIZE, drop_last = False, shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:18:20.622169Z","iopub.execute_input":"2024-09-01T21:18:20.622792Z","iopub.status.idle":"2024-09-01T21:18:20.627342Z","shell.execute_reply.started":"2024-09-01T21:18:20.622762Z","shell.execute_reply":"2024-09-01T21:18:20.626471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Transfer Learning using ResNet50 model","metadata":{}},{"cell_type":"code","source":"active_device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nprint(active_device)","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:18:24.625067Z","iopub.execute_input":"2024-09-01T21:18:24.625920Z","iopub.status.idle":"2024-09-01T21:18:24.656402Z","shell.execute_reply.started":"2024-09-01T21:18:24.625887Z","shell.execute_reply":"2024-09-01T21:18:24.655477Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We use ResNet50 network with custom last layers. This architecture is applied following certain transformations to resize and normalize our input image to get to the same input dimensions used while training the ResNet50.","metadata":{}},{"cell_type":"code","source":"class CassavaResNet50(nn.Module):\n    def __init__(self):\n        super().__init__()\n        \n        # Batch Norm 2d\n        self.bn = nn.BatchNorm2d(3)\n        \n        # Custom input top layer\n        self.top_layer = nn.Sequential(\n            transforms.Resize((224, 224)),\n            transforms.ToImageTensor(),\n            transforms.ConvertDtype(torch.float32),\n            transforms.Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225]\n            ) \n        )\n\n        # Load ResNet-50 and remove the last fc layer\n        self.resnet50 = models.resnet50(pretrained=True)\n        self.resnet50 = nn.Sequential(*list(self.resnet50.children())[:-2])\n\n        # Global Average Pooling 2D\n        self.global_avg_pool = nn.AdaptiveAvgPool2d(1)\n\n        # Custom fully connected layers\n        self.fc1 = nn.Linear(2048, 8)  # 2048 is the output feature size of ResNet-50\n        self.relu1 = nn.ReLU()\n        self.fc2 = nn.Linear(8, 5)\n        self.softmax = nn.Softmax(dim=1)\n\n    def forward(self, x):\n        # Forward pass through the layers\n        x = self.bn(x)\n        x = self.top_layer(x)\n        x = self.resnet50(x)\n        x = self.global_avg_pool(x)\n        x = x.view(x.size(0), -1)\n        x = self.fc1(x)\n        x = self.relu1(x)\n        x = self.fc2(x)\n        x = self.softmax(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:18:28.486644Z","iopub.execute_input":"2024-09-01T21:18:28.487502Z","iopub.status.idle":"2024-09-01T21:18:28.497581Z","shell.execute_reply.started":"2024-09-01T21:18:28.487468Z","shell.execute_reply":"2024-09-01T21:18:28.496674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Remember, we have changed the distribution of train dataset compared to the test dataset when we sampled with WeightedRandomSampler to tackle class imbalance. Unfortunately, due to this the training and validation metrics are not based on the same data distributions due to which a meaningful comparison of metrics cannnot be possible.\n\nTherfore, we employ metrics that are more suitable for a dataset with class imbalance like F1 score, confusion matrix. We then leave the validation dataset without tackling class imbalance, assuming that the initial train dataset provided before the train_test_split was applied actually reflects the data distribution from the environement we intend to deploy this model at.\n\nWe go ahead with F1 score for comparing metrics between train and validation datasets, even if the metric for the comeptition is Accuracy.","metadata":{}},{"cell_type":"code","source":"def train_model(train_loader, valid_loader, model, active_device):\n    clf = model.to(active_device)\n    optimizer = Adam(clf.parameters(), lr=1e-3, eps=0.001)\n    scheduler = ExponentialLR(optimizer, gamma=0.9)\n    loss_fn = nn.CrossEntropyLoss()\n    \n    tqdm._instances.clear()\n    train_losses = []  # List to store training losses per epoch\n    valid_losses = []  # List to store validation losses per epoch\n    train_accuracies = []  # List to store training accuracies per epoch\n    valid_accuracies = []  # List to store validation accuracies per epoch\n    train_F1 = []  # List to store training F1 scores per epoch\n    valid_F1 = []  # List to store validation F1 scores per epoch\n    \n    for epoch in range(config.EPOCHS):\n        train_loss = 0\n        valid_loss = 0\n        # for accuracy\n        correct_train = 0\n        total_train = 0\n        correct_valid = 0\n        total_valid = 0\n        # for F1 score\n        train_labels = []\n        valid_labels = []\n        train_preds  = []\n        valid_preds  = []\n        \n        scheduler.step()\n        model.train()\n        train_loop = tqdm(\n            enumerate(train_loader), \n            desc=f'Train Epoch {epoch + 1}/{config.EPOCHS}', \n            ncols=150, \n            leave=True\n        )\n        for itr, (images, labels) in train_loop:\n            # ForwardPass\n            images, labels = images.to(active_device), labels.to(active_device)\n            pred_labels = model(images)\n            # Loss\n            loss = loss_fn(pred_labels, labels) \n            train_loss += loss.item()\n            train_loop.set_postfix(curr_train_loss = \"{:.8f}\".format(train_loss / (itr+1)))\n            \n            # Accuracy\n            _, preds = torch.max(pred_labels, 1)\n            total_train += labels.size(0)\n            correct_train += (preds == labels).sum().item()\n            # F1 score\n            train_labels.extend(labels.cpu())\n            train_preds.extend(preds.cpu())\n            # Backpass\n            optimizer.zero_grad()\n            loss.backward()\n            optimizer.step()\n        \n        # train loss, accuracy and F1 score\n        train_accuracy = 100 * correct_train / total_train \n        train_f1 = f1_score(train_labels, train_preds, average='weighted', zero_division=1.0)\n        \n        train_losses.append(train_loss / len(train_loader))\n        train_accuracies.append(train_accuracy)\n        train_F1.append(train_f1)\n        \n        valid_loop = tqdm(\n            enumerate(valid_loader), \n            desc=f'Valid Epoch {epoch + 1}/{config.EPOCHS}', \n            ncols=150, \n            leave=True\n        )\n        with torch.no_grad():\n            model.eval()\n            for itr, (images, labels) in valid_loop:\n                # ForwardPass\n                images_val, labels_val = images.to(active_device), labels.to(active_device)\n                pred_labels_val = model(images_val)\n                # Loss\n                loss_val = loss_fn(pred_labels_val, labels_val)\n                valid_loss += loss_val.item()\n                valid_loop.set_postfix(\n                    val_loss = \"{:.8f}\".format(valid_loss / (itr+1))\n                )\n        \n                # Accuracy\n                _, preds_val = torch.max(pred_labels_val, 1)\n                total_valid += labels_val.size(0)\n                correct_valid += (preds_val == labels_val).sum().item()\n                # F1 score\n                valid_labels.extend(labels_val.cpu())\n                valid_preds.extend(preds_val.cpu())\n                \n        \n        # valid loss, accuracy and F1 score\n        valid_accuracy = 100 * correct_valid / total_valid\n        valid_f1 = f1_score(valid_labels, valid_preds, average='weighted', zero_division=1.0)\n        \n        valid_losses.append(valid_loss / len(valid_loader))\n        valid_accuracies.append(valid_accuracy)\n        valid_F1.append(valid_f1)\n        \n        tqdm.write(f\"Train Loss: {train_losses[-1]:.4f} | Train Acc: {train_accuracies[-1]:.2f}%\")\n        tqdm.write(f\"Valid Loss: {valid_losses[-1]:.4f} | Valid Acc: {valid_accuracies[-1]:.2f}%\")\n        tqdm.write(\"=\"*50)\n    \n    return train_losses, valid_losses, train_accuracies, valid_accuracies, train_F1, valid_F1","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:50:40.549873Z","iopub.execute_input":"2024-09-01T21:50:40.550285Z","iopub.status.idle":"2024-09-01T21:50:40.571651Z","shell.execute_reply.started":"2024-09-01T21:50:40.550254Z","shell.execute_reply":"2024-09-01T21:50:40.570553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Training\nmodel = CassavaResNet50()\ntrain_losses, valid_losses, train_accuracies, valid_accuracies, train_F1, valid_F1 = train_model(train_loader, valid_loader, model, active_device)","metadata":{"execution":{"iopub.status.busy":"2024-09-01T21:50:49.834194Z","iopub.execute_input":"2024-09-01T21:50:49.834898Z","iopub.status.idle":"2024-09-02T02:57:21.625695Z","shell.execute_reply.started":"2024-09-01T21:50:49.834866Z","shell.execute_reply":"2024-09-02T02:57:21.624755Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Function to visualize loss and accuracy over epochs\ndef visualize_metrics(train_losses, valid_losses, train_accuracies, valid_accuracies, train_F1, valid_F1):\n    plt.figure(figsize=(12, 12))\n    \n    # Loss plot\n    plt.subplot(3, 1, 1)\n    plt.plot(train_losses, label='Training Loss', marker='o')\n    plt.plot(valid_losses, label='Validation Loss', marker='o')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.title('Loss Over Epochs')\n    plt.legend()\n    \n    # Accuracy plot\n    plt.subplot(3, 1, 2)\n    plt.plot(train_accuracies, label='Training Accuracy', marker='o')\n    plt.plot(valid_accuracies, label='Validation Accuracy', marker='o')\n    plt.xlabel('Epoch')\n    plt.ylabel('Accuracy (%)')\n    plt.title('Accuracy Over Epochs')\n    plt.legend()\n    \n    \n    # Accuracy plot\n    plt.subplot(3, 1, 3)\n    plt.plot(train_F1, label='Training F1 score', marker='o')\n    plt.plot(valid_F1, label='Validation F1 score', marker='o')\n    plt.xlabel('Epoch')\n    plt.ylabel('F1 score')\n    plt.title('F1 score Over Epochs')\n    plt.legend()\n    \n    plt.tight_layout()\n    plt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-09-02T02:58:58.492847Z","iopub.execute_input":"2024-09-02T02:58:58.493334Z","iopub.status.idle":"2024-09-02T02:58:58.502232Z","shell.execute_reply.started":"2024-09-02T02:58:58.493300Z","shell.execute_reply":"2024-09-02T02:58:58.501343Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Visualization\nvisualize_metrics(train_losses, valid_losses, train_accuracies, valid_accuracies, train_F1, valid_F1)","metadata":{"execution":{"iopub.status.busy":"2024-09-02T02:59:03.028468Z","iopub.execute_input":"2024-09-02T02:59:03.029042Z","iopub.status.idle":"2024-09-02T02:59:03.860691Z","shell.execute_reply.started":"2024-09-02T02:59:03.029012Z","shell.execute_reply":"2024-09-02T02:59:03.859895Z"},"trusted":true},"execution_count":null,"outputs":[]}]}