{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.12"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":13836,"databundleVersionId":1718836,"sourceType":"competition"},{"sourceId":7177978,"sourceType":"datasetVersion","datasetId":4148359},{"sourceId":7195349,"sourceType":"datasetVersion","datasetId":4161246},{"sourceId":7325702,"sourceType":"datasetVersion","datasetId":4251941}],"dockerImageVersionId":30616,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Cassava Leaf Disease Detection\n- Jan Burian\n- https://www.kaggle.com/competitions/cassava-leaf-disease-classification","metadata":{}},{"cell_type":"markdown","source":"## Modules import","metadata":{}},{"cell_type":"code","source":"import numpy as np\nfrom pathlib import Path \nimport os\nimport math\nimport pandas as pd\nimport json \nimport matplotlib.pyplot as plt\nfrom PIL import Image\n\nimport wandb\n\nfrom sklearn.model_selection import train_test_split \n\nimport torch\nimport torchvision\nimport torchvision.transforms as transforms\nfrom torchvision.transforms import ToTensor, Resize\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision.io import read_image\nimport torch.nn.functional as F\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\n\nfrom albumentations.pytorch import ToTensorV2\nimport albumentations as A\n\nfrom datetime import datetime","metadata":{"execution":{"iopub.status.busy":"2024-01-02T15:56:30.611010Z","iopub.execute_input":"2024-01-02T15:56:30.611404Z","iopub.status.idle":"2024-01-02T15:56:40.845465Z","shell.execute_reply.started":"2024-01-02T15:56:30.611374Z","shell.execute_reply":"2024-01-02T15:56:40.844459Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !pip install timm\n# !pip install -U albumentations\n# !pip install wandb","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Preparing data","metadata":{}},{"cell_type":"code","source":"data_directory = Path('/kaggle/input/cassava-leaf-disease-classification/')\nBASE_directory = os.path.join(data_directory)\n!pwd # current working directory\nprint(BASE_directory)","metadata":{"execution":{"iopub.status.busy":"2024-01-02T15:56:43.948933Z","iopub.execute_input":"2024-01-02T15:56:43.949989Z","iopub.status.idle":"2024-01-02T15:56:44.906025Z","shell.execute_reply.started":"2024-01-02T15:56:43.949954Z","shell.execute_reply":"2024-01-02T15:56:44.904828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_csv = pd.read_csv(os.path.join(BASE_directory, 'train.csv'))\ntrain_csv.head()","metadata":{"execution":{"iopub.status.busy":"2024-01-02T15:56:46.203519Z","iopub.execute_input":"2024-01-02T15:56:46.204554Z","iopub.status.idle":"2024-01-02T15:56:46.269012Z","shell.execute_reply.started":"2024-01-02T15:56:46.204503Z","shell.execute_reply":"2024-01-02T15:56:46.267952Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training data analysis","metadata":{}},{"cell_type":"code","source":"def get_mapping_dictionary() -> dict:\n    num_to_disease_map = open(os.path.join(BASE_directory, 'label_num_to_disease_map.json'))\n    num_to_disease_map_dict = json.load(num_to_disease_map)\n    num_to_disease_map_dict = {int(key):num_to_disease_map_dict[key] for key in num_to_disease_map_dict}\n    \n    return num_to_disease_map_dict\n\n\ndef pair_indices_and_class_names(class_distribution_dict: dict) -> dict:\n    num_to_disease_map_dict = get_mapping_dictionary()\n    \n#     print(num_to_disease_map_dict)\n    res_dict = {}\n    \n    for key in class_distribution_dict.keys():\n        if key in num_to_disease_map_dict:\n            new_key = num_to_disease_map_dict[key]\n            res_dict[new_key] = class_distribution_dict[key]\n    \n#     print(class_distribution_dict)\n    return res_dict","metadata":{"execution":{"iopub.status.busy":"2024-01-02T15:56:47.893156Z","iopub.execute_input":"2024-01-02T15:56:47.894012Z","iopub.status.idle":"2024-01-02T15:56:47.900973Z","shell.execute_reply.started":"2024-01-02T15:56:47.893976Z","shell.execute_reply":"2024-01-02T15:56:47.899749Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualize_class_distributions_training_data(mapped_class_distribution_dict: dict):\n    fig, ax = plt.subplots()\n\n    classes = list(mapped_class_distribution_dict.keys())\n    counts = list(mapped_class_distribution_dict.values())\n    bar_colors = ['tab:red', 'tab:blue', 'tab:orange', 'tab:purple', 'tab:green']\n\n    ax.bar(classes, counts, color=bar_colors)\n\n    for i, count in enumerate(counts):\n        ax.text(classes[i], count + 10, str(count), ha='center', va='bottom')\n\n    ax.set_xticks(classes)  # Use set_xticks to set the exact tick positions\n    ax.set_xticklabels(classes, rotation=90)\n    ax.set_ylabel('Number of training samples')\n    ax.set_xlabel('Classes')\n    ax.set_title('Number of samples in training data')\n\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-01-02T15:56:49.908840Z","iopub.execute_input":"2024-01-02T15:56:49.909587Z","iopub.status.idle":"2024-01-02T15:56:49.916658Z","shell.execute_reply.started":"2024-01-02T15:56:49.909552Z","shell.execute_reply":"2024-01-02T15:56:49.915756Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Class distribution in training data\nclass_distribution_dict = {}\nfor i in range(len(train_csv)): \n#     print(train_csv.loc[i, \"label\"])\n    key = train_csv.loc[i, \"label\"]\n    \n    if key not in class_distribution_dict:\n        class_distribution_dict[key] = 1\n        \n    else:\n        class_distribution_dict[key] += 1\n        \nprint(class_distribution_dict)\nmapped_class_distribution_dict = pair_indices_and_class_names(class_distribution_dict)\nclasses_list = list(mapped_class_distribution_dict.keys())\nprint(mapped_class_distribution_dict)\nvisualize_class_distributions_training_data(mapped_class_distribution_dict)","metadata":{"execution":{"iopub.status.busy":"2024-01-02T15:56:51.616704Z","iopub.execute_input":"2024-01-02T15:56:51.617528Z","iopub.status.idle":"2024-01-02T15:56:52.294118Z","shell.execute_reply.started":"2024-01-02T15:56:51.617493Z","shell.execute_reply":"2024-01-02T15:56:52.293227Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualizing data","metadata":{}},{"cell_type":"code","source":"train_images_dir = os.path.join(BASE_directory, 'train_images')\nmapping_dictionary = get_mapping_dictionary()\n\ndef plot_n_training_images(n_images: int, train_csv: pd.DataFrame, number_classes: int):\n    for j in range(number_classes):\n        filtered_df = train_csv[train_csv['label'] == j]\n        random_rows = filtered_df.sample(n_images)\n\n        fig, axes = plt.subplots(1, n_images, figsize=(30, 5))  \n\n        for i in range(len(random_rows)):\n            image_id = random_rows.iloc[i]['image_id']\n            label = random_rows.iloc[i]['label']\n\n            class_name = mapping_dictionary[label]\n\n            img_path = os.path.join(train_images_dir, image_id)\n            img = Image.open(img_path)\n\n            axes[i].imshow(img)\n            axes[i].set_title(f'{class_name} ({image_id})')\n            axes[i].axis('off')\n\n        plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-01-02T15:56:53.406886Z","iopub.execute_input":"2024-01-02T15:56:53.407791Z","iopub.status.idle":"2024-01-02T15:56:53.416627Z","shell.execute_reply.started":"2024-01-02T15:56:53.407752Z","shell.execute_reply":"2024-01-02T15:56:53.415611Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_classes = len(mapped_class_distribution_dict)\nplot_n_training_images(5, train_csv, 5)","metadata":{"execution":{"iopub.status.busy":"2024-01-02T15:56:55.491621Z","iopub.execute_input":"2024-01-02T15:56:55.491982Z","iopub.status.idle":"2024-01-02T15:57:01.306676Z","shell.execute_reply.started":"2024-01-02T15:56:55.491954Z","shell.execute_reply":"2024-01-02T15:57:01.305336Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_class_representative_images(n_images: int, train_csv: pd.DataFrame, number_classes: int):\n    fig, axes = plt.subplots(1, number_classes, figsize=(30, 5))\n    for j in range(number_classes):\n        filtered_df = train_csv[train_csv['label'] == j]\n        random_rows = filtered_df.sample(n_images)\n\n        \n        image_id = random_rows.iloc[0]['image_id']\n        label = random_rows.iloc[0]['label']\n\n        class_name = mapping_dictionary[label]\n\n        img_path = os.path.join(train_images_dir, image_id)\n        img = Image.open(img_path)\n\n        axes[j].imshow(img)\n        axes[j].set_title(f'{class_name} ({image_id})')\n        axes[j].axis('off')\n\n    plt.show()\n        \nplot_class_representative_images(1, train_csv, 5)","metadata":{"execution":{"iopub.status.busy":"2024-01-02T15:57:01.308225Z","iopub.execute_input":"2024-01-02T15:57:01.308537Z","iopub.status.idle":"2024-01-02T15:57:02.556742Z","shell.execute_reply.started":"2024-01-02T15:57:01.308511Z","shell.execute_reply":"2024-01-02T15:57:02.555857Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"# https://colab.research.google.com/github/pytorch/tutorials/blob/gh-pages/_downloads/f498e3bcd9b6159ecfb1a07d6551287d/data_loading_tutorial.ipynb#scrollTo=amZmMAzuGgTu\n# https://pytorch.org/tutorials/beginner/data_loading_tutorial.html\n# https://pytorch.org/tutorials/beginner/basics/data_tutorial.html\n\nclass CassavaLeafDataset(Dataset):\n    def __init__(self, image_ids: list, labels: list, image_dir: str, dimension=(224, 224), transform=None):\n        self.image_ids = image_ids\n        self.labels = labels\n        self.image_dir = image_dir\n        self.dimension = dimension\n        self.transform = transform\n    \n    # returns the length\n    def __len__(self):\n        return len(self.image_ids)\n    \n    # returns the image and label for that index\n    def __getitem__(self, idx):\n        img_path = os.path.join(self.image_dir, self.image_ids[idx])\n        img = Image.open(img_path)\n        label = int(self.labels[idx])\n\n        if self.transform:\n            img_transformed = self.transform(image=np.array(img)) \n            img = img_transformed['image']\n\n        return img, label","metadata":{"execution":{"iopub.status.busy":"2024-01-02T15:57:02.558132Z","iopub.execute_input":"2024-01-02T15:57:02.558476Z","iopub.status.idle":"2024-01-02T15:57:02.568032Z","shell.execute_reply.started":"2024-01-02T15:57:02.558446Z","shell.execute_reply":"2024-01-02T15:57:02.567228Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Augmentations","metadata":{}},{"cell_type":"code","source":"# Augmentations \nresize_dimension = (224, 224) # ImageNet size\ntransform = A.Compose([\n    A.Resize(width = resize_dimension[0], height = resize_dimension[1]),\n    A.HorizontalFlip(p=0.5),\n    A.Blur(blur_limit=(3, 7), p=0.3),\n    A.VerticalFlip(p=0.5),\n    A.RandomRotate90(p=0.5),\n    A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), \n    ToTensorV2(),\n])\n\ntransform_test = A.Compose([\n    A.Resize(width = resize_dimension[0], height = resize_dimension[1]),\n    A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n    ToTensorV2(),\n])","metadata":{"execution":{"iopub.status.busy":"2024-01-02T15:57:02.569738Z","iopub.execute_input":"2024-01-02T15:57:02.570030Z","iopub.status.idle":"2024-01-02T15:57:02.590104Z","shell.execute_reply.started":"2024-01-02T15:57:02.570005Z","shell.execute_reply":"2024-01-02T15:57:02.589175Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Parameters\nbatch_size = 64 # number of samples for training\nnum_workers = 1\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Splitting dataset into 2 parts\npath_to_csv_file = os.path.join(BASE_directory, 'train.csv')\npath_to_image_directory = os.path.join(BASE_directory, 'train_images')\nprint(path_to_image_directory)\n\ntrain_csv = pd.read_csv(path_to_csv_file)\n\n# # Dividing data to train, validation and test subsets (90 % + 10 % + 10 %)\n# x_train, x_test, y_train, y_test = train_test_split(train_csv['image_id'], \n#                                                     train_csv['label'], \n#                                                     test_size=0.1, \n#                                                     random_state=1)\n# #                                                     stratify=train_csv['label']) # Test data\n                                                    \n\n# # Dividing again obtained train data to get validation data and final train data\n# x_train, x_val, y_train, y_val = train_test_split(x_train, \n#                                                   y_train, \n#                                                   test_size=(1/9), \n#                                                   random_state=1)\n# #                                                   stratify=y_train) # (1/9) x 0.9 = 0.1 # Train and validation data\n#                                                               # stratify parameter to represent data distribution\n    \n# Dividing data to train, validation and validation subsets (90 % + 10 %)\nx_train, x_val, y_train, y_val = train_test_split(train_csv['image_id'], \n                                                    train_csv['label'], \n                                                    test_size=0.1, \n                                                    random_state=1)\n#                                                     stratify=train_csv['label']) # Test data\n\nprint(len(x_train))\nprint(len(x_val))\n# print(len(x_test))\n\n\n# Getting train, val and test datasets\ntrain_dataset = CassavaLeafDataset(\n    image_ids=x_train.values,\n    labels=y_train.values,\n    image_dir=path_to_image_directory,\n    dimension=resize_dimension,\n    transform=transform\n)\n\n\nval_dataset = CassavaLeafDataset(\n    image_ids=x_val.values,\n    labels=y_val.values,\n    image_dir=path_to_image_directory,\n    dimension=resize_dimension,\n    transform=transform\n)\n\n# test_dataset = CassavaLeafDataset(\n#     image_ids=x_test.values,\n#     labels=y_test.values,\n#     image_dir=path_to_image_directory,\n#     dimension=resize_dimension,\n#     transform=transform_test\n# )","metadata":{"execution":{"iopub.status.busy":"2024-01-02T15:57:03.111287Z","iopub.execute_input":"2024-01-02T15:57:03.112443Z","iopub.status.idle":"2024-01-02T15:57:03.203308Z","shell.execute_reply.started":"2024-01-02T15:57:03.112396Z","shell.execute_reply":"2024-01-02T15:57:03.202347Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(type(x_train.values[0]))\nprint(x_train.values[0])\nprint(type(y_train.values[0]))","metadata":{"execution":{"iopub.status.busy":"2024-01-02T15:57:14.323807Z","iopub.execute_input":"2024-01-02T15:57:14.324899Z","iopub.status.idle":"2024-01-02T15:57:14.330738Z","shell.execute_reply.started":"2024-01-02T15:57:14.324855Z","shell.execute_reply":"2024-01-02T15:57:14.329445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## DataLoader","metadata":{}},{"cell_type":"code","source":"train_loader = DataLoader(\n    train_dataset,\n    batch_size=batch_size,\n    num_workers=num_workers,\n    shuffle=True,\n)\n\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=batch_size,\n    num_workers=num_workers,\n    shuffle=False\n)\n\n\n# test_loader = DataLoader(\n#     test_dataset,\n#     batch_size=batch_size,\n#     num_workers=num_workers,\n#     shuffle=False\n# )\n\n\n# loaders = {'train': train_loader, 'val': val_loader, 'test': test_loader}\n\n# print(len(x_train))\n# print(len(y_train))\n# print(len(x_val))\n# print(len(y_val))\n# print(len(x_test))\n# print(len(y_test))","metadata":{"execution":{"iopub.status.busy":"2024-01-02T15:57:16.215554Z","iopub.execute_input":"2024-01-02T15:57:16.216518Z","iopub.status.idle":"2024-01-02T15:57:16.222443Z","shell.execute_reply.started":"2024-01-02T15:57:16.216482Z","shell.execute_reply":"2024-01-02T15:57:16.221368Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Dataset size","metadata":{}},{"cell_type":"code","source":"print(f\"Train set size: {len(train_loader.dataset)} images.\")\nprint(f\"Validation set size: {len(val_loader.dataset)} images.\")\n# print(f\"Test set size: {len(test_loader.dataset)} images.\\n\")\n\n# print(f\"Dataset size: {np.sum([len(train_loader.dataset), len(val_loader.dataset), len(test_loader.dataset)])} images.\")","metadata":{"execution":{"iopub.status.busy":"2024-01-02T15:57:18.401895Z","iopub.execute_input":"2024-01-02T15:57:18.402631Z","iopub.status.idle":"2024-01-02T15:57:18.407692Z","shell.execute_reply.started":"2024-01-02T15:57:18.402598Z","shell.execute_reply":"2024-01-02T15:57:18.406592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualizing training data\nfrom train_loader","metadata":{}},{"cell_type":"code","source":"classes = list(get_mapping_dictionary().values())\n\n# functions to show an image\ndef imshow(img):\n    img = img / 2 + 0.5     # denormalization\n    npimg = img.numpy()\n    plt.imshow(np.transpose(npimg))\n    plt.show()\n\n\n# get some random training images\ndataiter = iter(train_loader)\nimages, labels = next(dataiter)\n# print(next(dataiter))\n\n# show images\nimshow(torchvision.utils.make_grid(images))\n\n# print labels\nprint(' '.join('%5s' % classes[labels[j]] for j in range(batch_size)))","metadata":{"execution":{"iopub.status.busy":"2024-01-02T15:57:20.351417Z","iopub.execute_input":"2024-01-02T15:57:20.352003Z","iopub.status.idle":"2024-01-02T15:57:22.888773Z","shell.execute_reply.started":"2024-01-02T15:57:20.351968Z","shell.execute_reply":"2024-01-02T15:57:22.887803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Usage of a pretrained model","metadata":{}},{"cell_type":"code","source":"from torchvision import models","metadata":{"execution":{"iopub.status.busy":"2024-01-02T15:57:23.071159Z","iopub.execute_input":"2024-01-02T15:57:23.071515Z","iopub.status.idle":"2024-01-02T15:57:23.076452Z","shell.execute_reply.started":"2024-01-02T15:57:23.071485Z","shell.execute_reply":"2024-01-02T15:57:23.075450Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# dir(models) # available models and weights from torchvision","metadata":{"execution":{"iopub.status.busy":"2024-01-02T08:18:58.346863Z","iopub.execute_input":"2024-01-02T08:18:58.347871Z","iopub.status.idle":"2024-01-02T08:18:58.351946Z","shell.execute_reply.started":"2024-01-02T08:18:58.347835Z","shell.execute_reply":"2024-01-02T08:18:58.350986Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(torch.cuda.is_available()) # availability of GPU","metadata":{"execution":{"iopub.status.busy":"2024-01-02T11:46:01.835467Z","iopub.execute_input":"2024-01-02T11:46:01.835852Z","iopub.status.idle":"2024-01-02T11:46:01.840849Z","shell.execute_reply.started":"2024-01-02T11:46:01.835821Z","shell.execute_reply":"2024-01-02T11:46:01.839930Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_actual_timestamp():\n    current_time = datetime.now()\n    time_string = current_time.strftime(\"%Y%m%d%H%M%S\")\n    \n    return time_string","metadata":{"execution":{"iopub.status.busy":"2024-01-02T15:57:28.862066Z","iopub.execute_input":"2024-01-02T15:57:28.862738Z","iopub.status.idle":"2024-01-02T15:57:28.867486Z","shell.execute_reply.started":"2024-01-02T15:57:28.862702Z","shell.execute_reply":"2024-01-02T15:57:28.866473Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config = {\n    \"learning_rate\": 1e-3,\n    \"batch_size\": batch_size, \n    \"architecture\": \"\",\n    \"dataset\": \"Cassava leaf disease\",\n    \"num_epochs\": 20,\n    \"dropout\": 0,\n    \"device\": device,\n    \"timestamp\": create_actual_timestamp(),\n    }","metadata":{"execution":{"iopub.status.busy":"2024-01-02T15:57:30.736613Z","iopub.execute_input":"2024-01-02T15:57:30.737336Z","iopub.status.idle":"2024-01-02T15:57:30.742000Z","shell.execute_reply.started":"2024-01-02T15:57:30.737292Z","shell.execute_reply":"2024-01-02T15:57:30.741109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model initialization","metadata":{}},{"cell_type":"code","source":"#pretrained_model = models.resnet50(weights='ResNet50_Weights.DEFAULT')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ResNet-50\n# class ImageClassificationModel(nn.Module):\n#     def __init__(self, num_classes, dropout_prob, model_path):\n#         super(ImageClassificationModel, self).__init__()\n        \n#         # Load a ResNet model\n#         self.model = models.resnet50(pretrained=False)\n    \n#         #     model.load_state_dict(weights, strict=False)\n#         self.model.load_state_dict(torch.load(model_path))\n        \n#         # Modify the final fully connected layer to match the number of classes in your dataset\n# #         self.pretrained_model.fc = nn.Sequential(\n# #             nn.Linear(in_features, 512),\n# #             nn.ReLU(),\n# #             nn.Dropout(p=dropout_prob),\n# #             nn.Linear(512, num_classes)\n# #         )\n\n#         in_features = self.model.fc.in_features\n#         self.model.fc = nn.Linear(in_features, num_classes)\n        \n\n#     def forward(self, x):\n#         return self.model(x)","metadata":{"execution":{"iopub.status.busy":"2024-01-02T13:22:59.089073Z","iopub.execute_input":"2024-01-02T13:22:59.089450Z","iopub.status.idle":"2024-01-02T13:22:59.096670Z","shell.execute_reply.started":"2024-01-02T13:22:59.089408Z","shell.execute_reply":"2024-01-02T13:22:59.095759Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ResNet-18\nclass ImageClassificationModel(nn.Module):\n    def __init__(self, num_classes, dropout_prob, model_path):\n        super(ImageClassificationModel, self).__init__()\n        \n        # Load a ResNet model\n        self.model = models.resnet18(pretrained=False)\n    \n        #     model.load_state_dict(weights, strict=False)\n        self.model.load_state_dict(torch.load(model_path))\n        \n        # Modify the final fully connected layer to match the number of classes in your dataset\n#         self.pretrained_model.fc = nn.Sequential(\n#             nn.Linear(in_features, 512),\n#             nn.ReLU(),\n#             nn.Dropout(p=dropout_prob),\n#             nn.Linear(512, num_classes)\n#         )\n\n        in_features = self.model.fc.in_features\n        self.model.fc = nn.Linear(in_features, num_classes)\n        \n\n    def forward(self, x):\n        return self.model(x)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# DenseNet\n# class ImageClassificationModel(nn.Module):\n#     def __init__(self, num_classes, dropout_prob, model_path):\n#         super(ImageClassificationModel, self).__init__()\n        \n#         # Load a pretrained DenseNet model\n# #         self.pretrained_model = models.densenet121(pretrained=False)\n#         self.model = models.densenet121()\n#         state_dict = torch.load(model_path)\n        \n#         for key in list(state_dict.keys()):\n#             state_dict[key.replace('.1.', '1.'). replace('.2.', '2.')] = state_dict.pop(key)\n            \n#         self.model.load_state_dict(state_dict)\n        \n# #         self.pretrained_model.classifier = nn.Sequential(\n# #             nn.Linear(in_features, 512),\n# #             nn.ReLU(),\n# #             nn.Dropout(p=dropout_prob),\n# #             nn.Linear(512, num_classes)\n# #         )\n\n#         # Modify the classifier part of the model\n#         in_features = self.model.classifier.in_features\n#         self.model.classifier = nn.Linear(in_features, num_classes)\n    \n#     def forward(self, x):\n#         return self.model(x)","metadata":{"execution":{"iopub.status.busy":"2024-01-02T16:09:57.695834Z","iopub.execute_input":"2024-01-02T16:09:57.696208Z","iopub.status.idle":"2024-01-02T16:09:57.703659Z","shell.execute_reply.started":"2024-01-02T16:09:57.696181Z","shell.execute_reply":"2024-01-02T16:09:57.702595Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_classification_model(train_loader, val_loader, num_epochs, learning_rate, dropout, model_path):\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    \n    # Define the model\n    num_classes = 5\n#     weights = torch.load(model_path)\n    model = ImageClassificationModel(num_classes, dropout, model_path).to(device)\n\n    # Define the loss function, optimizer and scheduler\n    loss_function = nn.CrossEntropyLoss() # Cross entropy loss function\n    optimizer = torch.optim.Adam(model.parameters(), lr=config[\"learning_rate\"]) # Adam optimizer\n#     scheduler = lr_scheduler.StepLR(optimizer, step_size=500, gamma=0.1)\n\n    n_steps_per_epoch = math.ceil(len(train_loader.dataset) / config[\"batch_size\"])\n    \n#     train_losses = []\n    \n    # Training loop\n    for epoch in range(config[\"num_epochs\"]):\n        model.train()\n        epoch_loss = 0.0\n        avg_epoch_loss = 0.0\n        for step, (images, labels) in enumerate(train_loader):\n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images)\n            \n            train_loss = loss_function(outputs, labels)\n            \n            optimizer.zero_grad()\n            train_loss.backward()\n            optimizer.step()\n            \n            epoch_loss += train_loss.item()\n            \n            metrics = {\"train/train_loss\": train_loss, \n                       \"train/epoch\": (step + 1 + (n_steps_per_epoch * epoch)) / n_steps_per_epoch}\n            \n            # Log learning rate\n#             current_lr = optimizer.param_groups[0]['lr']\n#             metrics[\"train/learning_rate\"] = current_lr\n                \n#             scheduler.step() # learning rate update\n            \n        avg_epoch_loss = epoch_loss / n_steps_per_epoch\n#         train_losses.append(avg_epoch_loss)\n#         wandb.log({\"train/avg_epoch_loss\": avg_epoch_loss})\n            \n\n        # Validation loop\n        model.eval()\n        val_loss = 0\n        with torch.no_grad():\n            correct = 0\n            for i, (images, labels) in enumerate(val_loader):\n                images, labels = images.float().to(device), labels.to(device)\n                \n                # Forward pass\n                outputs = model(images)\n                val_loss += loss_function(outputs, labels) * labels.size(0)\n                \n                # Computation of the accuracy \n                _, predicted = torch.max(outputs.data, 1)\n                correct += (predicted == labels).sum().item()\n\n        accuracy = correct / len(val_loader.dataset)\n        val_loss = val_loss / len(val_loader.dataset)\n\n        # Log train and validation metrics\n        val_metrics = {\"val/val_loss\": val_loss, \n                       \"val/val_accuracy\": accuracy}\n        \n        print(f\"Epoch [{epoch+1}/{num_epochs}], Train Loss: {avg_epoch_loss:.4f}, Validation Loss: {val_loss:4f}, Accuracy: {accuracy:.4f}\")\n        \n    print('Training finished.')\n    return model","metadata":{"execution":{"iopub.status.busy":"2024-01-02T15:57:38.107320Z","iopub.execute_input":"2024-01-02T15:57:38.108165Z","iopub.status.idle":"2024-01-02T15:57:38.121142Z","shell.execute_reply.started":"2024-01-02T15:57:38.108129Z","shell.execute_reply":"2024-01-02T15:57:38.120235Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def test_classification_model(model, test_loader):\n#     device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    \n#     model.eval()\n#     correct = 0\n#     total = 0\n    \n#     with torch.no_grad():\n#         for inputs, labels in test_loader:\n#             inputs, labels = inputs.to(device), labels.to(device)\n#             outputs = model(inputs)\n#             _, predicted = torch.max(outputs.data, 1)\n#             total += labels.size(0)\n#             correct += (predicted == labels).sum().item()\n\n#     accuracy = correct / total\n#     print(f'Test Accuracy: {accuracy:.4f}')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def visualize_test_predictions(model, test_loader, class_labels):\n#     device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    \n#     mean = np.array([0.485, 0.456, 0.406])\n#     std = np.array([0.229, 0.224, 0.225])\n\n#     model.eval()\n#     correct_predictions = 0\n#     total_samples = 0\n\n#     with torch.no_grad():\n#         for images, labels in test_loader:\n#             images, labels = images.to(device), labels.to(device)\n\n#             # Forward pass\n#             outputs = model(images)\n#             _, predicted = torch.max(outputs.data, 1)\n\n#             # Convert tensors to NumPy arrays\n#             images_np = images.cpu().numpy()\n#             labels_np = labels.cpu().numpy()\n#             predicted_np = predicted.cpu().numpy()\n\n#             # Visualize a few images with their predicted labels\n#             for i in range(len(images_np)):\n#                 img = np.transpose(images_np[i], (1, 2, 0))\n#                 img_unnormalized = (img * std) + mean # Unnormalize\n#                 plt.imshow(img_unnormalized)  # Assuming images are in (C, H, W) format\n# #                 plt.imshow(images_np[i])\n#                 true_label = class_labels[labels_np[i]]\n#                 predicted_label = class_labels[predicted_np[i]]\n\n#                 plt.title(f\"True: {true_label}, Predicted: {predicted_label}\")\n#                 plt.show()\n\n#                 # Update accuracy counts\n#                 total_samples += 1\n#                 correct_predictions += 1 if labels_np[i] == predicted_np[i] else 0\n\n#             # Only visualize a few batches for brevity\n#             break\n\n#     # Calculate overall accuracy\n#     overall_accuracy = correct_predictions / total_samples\n#     print(f\"Overall Accuracy: {overall_accuracy:.4f}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# trained_model = train_classification_model(train_loader, \n#                                            val_loader, \n#                                            num_epochs=config[\"num_epochs\"], \n#                                            learning_rate=config[\"learning_rate\"], \n#                                            dropout=config[\"dropout\"], \n#                                            model_path = \"/kaggle/input/resnet-50-model/resnet50.pth\")\n\n# trained_model = train_classification_model(train_loader, \n#                                            val_loader, \n#                                            num_epochs=config[\"num_epochs\"], \n#                                            learning_rate=config[\"learning_rate\"], \n#                                            dropout=config[\"dropout\"], \n#                                            model_path = \"/kaggle/input/densenet-121/densenet121-a639ec97.pth\")\n\ntrained_model = train_classification_model(train_loader, \n                                           val_loader, \n                                           num_epochs=config[\"num_epochs\"], \n                                           learning_rate=config[\"learning_rate\"], \n                                           dropout=config[\"dropout\"], \n                                           model_path = \"/kaggle/input/resnet-18/resnet18.pth\")","metadata":{"execution":{"iopub.status.busy":"2024-01-02T16:10:03.026044Z","iopub.execute_input":"2024-01-02T16:10:03.026423Z","iopub.status.idle":"2024-01-02T16:19:20.011561Z","shell.execute_reply.started":"2024-01-02T16:10:03.026394Z","shell.execute_reply":"2024-01-02T16:19:20.010235Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test_classification_model(trained_model, test_loader)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Visualizing trained model on test set","metadata":{}},{"cell_type":"code","source":"# visualize_test_predictions(trained_model, test_loader, classes_list)","metadata":{"scrolled":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Save trained model ","metadata":{}},{"cell_type":"code","source":"trained_models_directory = Path(\"/kaggle/working/\")\noutput_directory = os.path.join(trained_models_directory)\nprint(os.path.exists(output_directory))\n# torch.save(trained_model.state_dict(), os.path.join(output_directory, \"model_resnet_50_dropout_0.pth\"))\n# torch.save(trained_model.state_dict(), os.path.join(output_directory, \"model_densenet_121.pth\"))\ntorch.save(trained_model.state_dict(), os.path.join(output_directory, \"model_resnet_18.pth\"))","metadata":{"execution":{"iopub.status.busy":"2024-01-02T16:19:33.523856Z","iopub.execute_input":"2024-01-02T16:19:33.524674Z","iopub.status.idle":"2024-01-02T16:19:33.641403Z","shell.execute_reply.started":"2024-01-02T16:19:33.524640Z","shell.execute_reply":"2024-01-02T16:19:33.640457Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"def do_inference(img_path: str, model, transform, classes: list):\n#     model.eval()\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    img = Image.open(img_path)\n    \n    img_transformed = transform(image=np.array(img))\n    img_tensor = img_transformed[\"image\"]\n\n    # Add batch dimension to the image\n    img_tensor = img_tensor.unsqueeze(0)\n    \n    with torch.no_grad():\n        output = model(img_tensor)\n\n    # Get the predicted class index\n    _, predicted_class = torch.max(output, 1)\n\n    plt.imshow(img)\n    plt.title(f\"Predicted class: {classes[predicted_class]}.\")\n    class_idx = predicted_class.item()\n    \n    return class_idx","metadata":{"execution":{"iopub.status.busy":"2024-01-02T16:19:36.103223Z","iopub.execute_input":"2024-01-02T16:19:36.103591Z","iopub.status.idle":"2024-01-02T16:19:36.110660Z","shell.execute_reply.started":"2024-01-02T16:19:36.103562Z","shell.execute_reply":"2024-01-02T16:19:36.109585Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Implementing inference","metadata":{}},{"cell_type":"code","source":"def change_state_dict_format(weights_path):\n    ckpt = torch.load(weights_path)\n    state_dict = {k.replace(\"model.\", \"\"): v for k, v in ckpt.items()}\n        \n    return state_dict","metadata":{"execution":{"iopub.status.busy":"2024-01-02T16:19:39.197546Z","iopub.execute_input":"2024-01-02T16:19:39.198154Z","iopub.status.idle":"2024-01-02T16:19:39.203211Z","shell.execute_reply.started":"2024-01-02T16:19:39.198121Z","shell.execute_reply":"2024-01-02T16:19:39.202280Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Paths to files\n# cassava_weights_path = os.path.join(output_directory, \"model_resnet_50_dropout_0.pth\")\n# cassava_weights_path = os.path.join(output_directory, \"model_densenet_121.pth\")\ncassava_weights_path = os.path.join(output_directory, \"model_resnet_18.pth\")\ntest_images_path = os.path.join(BASE_directory, \"test_images\")\n\n# Model setup\n# model = models.resnet50(weights = None)\n# model = models.densenet121(weights = None)\nmodel = models.resnet18(weights = None)\n\n# Change number of neurons in fully connected layer\nnum_classes = len(classes_list)  # Number of cassava disease classes\n\n# ResNet\nin_features = model.fc.in_features\nmodel.fc = nn.Linear(in_features, num_classes)\n\n# DenseNet\n# in_features = model.classifier.in_features\n# model.classifier = nn.Linear(in_features, num_classes)\n\ncassava_weights = change_state_dict_format(cassava_weights_path)\n\nmodel.load_state_dict(cassava_weights)\n\nmodel.eval()","metadata":{"execution":{"iopub.status.busy":"2024-01-02T16:19:40.846263Z","iopub.execute_input":"2024-01-02T16:19:40.846628Z","iopub.status.idle":"2024-01-02T16:19:41.160266Z","shell.execute_reply.started":"2024-01-02T16:19:40.846599Z","shell.execute_reply":"2024-01-02T16:19:41.159368Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Augmentations\nresize_dimension = (224, 224) # ImageNet size\ntransform = A.Compose([\n    A.Resize(width = resize_dimension[0], height = resize_dimension[1]),\n    A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n    ToTensorV2(),\n])\n\n    \n# Submission dataframe\nsubmission_df = pd.DataFrame(columns=[\"image_id\", \"label\"])\n\nfor image_name in os.listdir(test_images_path):\n    image_path = os.path.join(test_images_path, image_name)\n    predicted_class = do_inference(image_path, model, transform, classes_list)\n\n    # Create a dictionary with \"image_id\" and \"label\"\n    entry = {\"image_id\": image_name, \"label\": predicted_class}\n\n    # Append the entry to the DataFrame\n    submission_df = submission_df._append(entry, ignore_index=True)","metadata":{"execution":{"iopub.status.busy":"2024-01-02T16:19:54.672628Z","iopub.execute_input":"2024-01-02T16:19:54.673467Z","iopub.status.idle":"2024-01-02T16:19:55.595886Z","shell.execute_reply.started":"2024-01-02T16:19:54.673433Z","shell.execute_reply":"2024-01-02T16:19:55.595034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df","metadata":{"execution":{"iopub.status.busy":"2024-01-02T16:19:59.230587Z","iopub.execute_input":"2024-01-02T16:19:59.230948Z","iopub.status.idle":"2024-01-02T16:19:59.241097Z","shell.execute_reply.started":"2024-01-02T16:19:59.230919Z","shell.execute_reply":"2024-01-02T16:19:59.239745Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2024-01-02T16:20:01.349745Z","iopub.execute_input":"2024-01-02T16:20:01.350694Z","iopub.status.idle":"2024-01-02T16:20:01.364047Z","shell.execute_reply.started":"2024-01-02T16:20:01.350655Z","shell.execute_reply":"2024-01-02T16:20:01.363105Z"},"trusted":true},"execution_count":null,"outputs":[]}]}