{"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":"# Solving cassava-leaf-disease-classification using transfer learning with convnext","metadata":{}},{"cell_type":"markdown","source":" Objective: fine tune a pretrained ConvNext tiny model and use it to classify if cassava leaf images contain any diseases,\n This will be a image classifcation problem.\n \n Data: image data as jpg files and a csv files contain all image filenames and corresponding labels.\n \n Transfer Learning: Transfer Learning is used to achieve a better result by utilise a pre-trained model usually had trained on a much larger data set with more complex tasks, and fine tune on a specific task with similar domain.","metadata":{}},{"cell_type":"markdown","source":"# Import libraries","metadata":{}},{"cell_type":"code","source":"import gc\nimport math\n\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nimport plotly.express as px\nimport plotly.graph_objs as go\n\nimport os\n\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.model_selection import train_test_split\nimport torch \n\nfrom imblearn.over_sampling import RandomOverSampler\n\n\nfrom torchvision.datasets.utils import download_url\nimport torchvision as tv\nimport torchvision.transforms as T\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom PIL import Image\n","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:58:56.196768Z","iopub.execute_input":"2023-04-14T13:58:56.197464Z","iopub.status.idle":"2023-04-14T13:59:01.604205Z","shell.execute_reply.started":"2023-04-14T13:58:56.197328Z","shell.execute_reply":"2023-04-14T13:59:01.603201Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#download_url('https://s3.amazonaws.com/fast-ai-imageclas/oxford-iiit-pet.tgz')","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:01.606030Z","iopub.execute_input":"2023-04-14T13:59:01.607599Z","iopub.status.idle":"2023-04-14T13:59:01.613479Z","shell.execute_reply.started":"2023-04-14T13:59:01.607558Z","shell.execute_reply":"2023-04-14T13:59:01.611695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import tarfile\n\n# with tarfile.open('./oxford-iiit-pet.tgz', 'r:gz') as tar:\n#     tar.extractall(path='./data')","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:01.614636Z","iopub.execute_input":"2023-04-14T13:59:01.614944Z","iopub.status.idle":"2023-04-14T13:59:01.628810Z","shell.execute_reply.started":"2023-04-14T13:59:01.614917Z","shell.execute_reply":"2023-04-14T13:59:01.627687Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Preprocessing","metadata":{}},{"cell_type":"markdown","source":"## Read the train.csv file that contains metadata","metadata":{}},{"cell_type":"code","source":"#train = pd.read_csv('train.csv').head(200)\ndata = pd.read_csv('../input/cassava-leaf-disease-classification/train.csv').head(5000)\n\ndata","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:01.632127Z","iopub.execute_input":"2023-04-14T13:59:01.632523Z","iopub.status.idle":"2023-04-14T13:59:01.685200Z","shell.execute_reply.started":"2023-04-14T13:59:01.632483Z","shell.execute_reply":"2023-04-14T13:59:01.684211Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data['label'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:01.687230Z","iopub.execute_input":"2023-04-14T13:59:01.687848Z","iopub.status.idle":"2023-04-14T13:59:01.702347Z","shell.execute_reply.started":"2023-04-14T13:59:01.687813Z","shell.execute_reply":"2023-04-14T13:59:01.701318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data.dtypes","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:01.704042Z","iopub.execute_input":"2023-04-14T13:59:01.704396Z","iopub.status.idle":"2023-04-14T13:59:01.713458Z","shell.execute_reply.started":"2023-04-14T13:59:01.704364Z","shell.execute_reply":"2023-04-14T13:59:01.712458Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label_classes = {0:\"Cassava Bacterial Blight (CBB)\",\n                 1:\"Cassava Brown Streak Disease (CBSD)\",\n                 2:\"Cassava Green Mottle (CGM)\",\n                 3:\"Cassava Mosaic Disease (CMD)\",\n                 4:\"Healthy\"\n                 }","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:01.715120Z","iopub.execute_input":"2023-04-14T13:59:01.715453Z","iopub.status.idle":"2023-04-14T13:59:01.723148Z","shell.execute_reply.started":"2023-04-14T13:59:01.715421Z","shell.execute_reply":"2023-04-14T13:59:01.722314Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class_converter = {'Cassava Bacterial Blight (CBB)': 0,\n                   'Cassava Mosaic Disease (CMD)': 3,\n                   'Cassava Brown Streak Disease (CBSD)': 1,\n                   'Cassava Green Mottle (CGM)': 2,\n                   'Healthy': 4}\n         ","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:01.726415Z","iopub.execute_input":"2023-04-14T13:59:01.726705Z","iopub.status.idle":"2023-04-14T13:59:01.733363Z","shell.execute_reply.started":"2023-04-14T13:59:01.726681Z","shell.execute_reply":"2023-04-14T13:59:01.732403Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data['class'] = data['label'].replace([0, 1, 2, 3, 4], [label_classes[0],label_classes[1],label_classes[2],label_classes[3],label_classes[4]])","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:01.735696Z","iopub.execute_input":"2023-04-14T13:59:01.736525Z","iopub.status.idle":"2023-04-14T13:59:01.746780Z","shell.execute_reply.started":"2023-04-14T13:59:01.736488Z","shell.execute_reply":"2023-04-14T13:59:01.745828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:01.751535Z","iopub.execute_input":"2023-04-14T13:59:01.752468Z","iopub.status.idle":"2023-04-14T13:59:01.764514Z","shell.execute_reply.started":"2023-04-14T13:59:01.752435Z","shell.execute_reply":"2023-04-14T13:59:01.763599Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Import the Images and build custom datasets","metadata":{}},{"cell_type":"markdown","source":"custom dataset class","metadata":{}},{"cell_type":"code","source":"#TEST_DATA_DIR = './test_images'\n# TRAIN_DATA_DIR = './train_images'\nTEST_DATA_DIR = '../input/cassava-leaf-disease-classification/test_images'\nTRAIN_DATA_DIR = '../input/cassava-leaf-disease-classification/train_images'","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:01.765944Z","iopub.execute_input":"2023-04-14T13:59:01.766563Z","iopub.status.idle":"2023-04-14T13:59:01.774624Z","shell.execute_reply.started":"2023-04-14T13:59:01.766528Z","shell.execute_reply":"2023-04-14T13:59:01.773903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.listdir(TEST_DATA_DIR)","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:01.775848Z","iopub.execute_input":"2023-04-14T13:59:01.776382Z","iopub.status.idle":"2023-04-14T13:59:01.791611Z","shell.execute_reply.started":"2023-04-14T13:59:01.776347Z","shell.execute_reply":"2023-04-14T13:59:01.790578Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Since test data folder only have 1 image, I am going to ignore this folder.","metadata":{}},{"cell_type":"code","source":"files = os.listdir(TRAIN_DATA_DIR)[:5000]\nfiles[:5]","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:01.793139Z","iopub.execute_input":"2023-04-14T13:59:01.793725Z","iopub.status.idle":"2023-04-14T13:59:02.406770Z","shell.execute_reply.started":"2023-04-14T13:59:01.793693Z","shell.execute_reply":"2023-04-14T13:59:02.405804Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"train_images folder contain all data","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader, IterableDataset","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:02.408171Z","iopub.execute_input":"2023-04-14T13:59:02.408543Z","iopub.status.idle":"2023-04-14T13:59:02.415275Z","shell.execute_reply.started":"2023-04-14T13:59:02.408508Z","shell.execute_reply":"2023-04-14T13:59:02.414053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CassavaLeafDiseaseDataset(IterableDataset):\n    def __init__(self, root, files, labels, classes, transform):\n        super(CassavaLeafDiseaseDataset).__init__()\n        self.root = root\n        self.files = files\n        self.labels = labels\n        self.classes = classes\n        self.transform = transform\n        self.start = 0\n        self.end = len(files)\n    \n    def __len__(self):\n        return len(self.files)\n\n    def __getitem__(self, i):\n        fname = self.files[i]\n        fpath = os.path.join(self.root, fname)\n        img = self.transform(Image.open(fpath).convert('RGB'))\n        class_label_int = self.labels[i]\n        return img, class_label_int\n    \n    def __iter__(self):\n        worker_info = torch.utils.data.get_worker_info()\n        if worker_info is None:  # single-process data loading, return the full iterator\n            iter_start = self.start\n            iter_end = self.end\n        else:  # in a worker process\n            # split workload\n            per_worker = int(math.ceil((self.end - self.start) / float(worker_info.num_workers)))\n            worker_id = worker_info.id\n            iter_start = self.start + worker_id * per_worker\n            iter_end = min(iter_start + per_worker, self.end)\n        return map(self.__getitem__, range(iter_start, iter_end))","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:02.416812Z","iopub.execute_input":"2023-04-14T13:59:02.418045Z","iopub.status.idle":"2023-04-14T13:59:02.429904Z","shell.execute_reply.started":"2023-04-14T13:59:02.418006Z","shell.execute_reply":"2023-04-14T13:59:02.428679Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Split dataset for training and testing","metadata":{}},{"cell_type":"markdown","source":"### Split the metadata into train data and test data","metadata":{}},{"cell_type":"code","source":"import numpy as np\nfrom sklearn.model_selection import train_test_split","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:02.432499Z","iopub.execute_input":"2023-04-14T13:59:02.432908Z","iopub.status.idle":"2023-04-14T13:59:02.443166Z","shell.execute_reply.started":"2023-04-14T13:59:02.432874Z","shell.execute_reply":"2023-04-14T13:59:02.442185Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = {}\ntrain_df, test_df = train_test_split(\n    data, stratify=data.label, train_size=0.8, random_state=0\n)","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:02.444561Z","iopub.execute_input":"2023-04-14T13:59:02.445696Z","iopub.status.idle":"2023-04-14T13:59:02.461968Z","shell.execute_reply.started":"2023-04-14T13:59:02.445659Z","shell.execute_reply":"2023-04-14T13:59:02.460835Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = train_df.reset_index(drop=True)\ntest_df = test_df.reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:02.463556Z","iopub.execute_input":"2023-04-14T13:59:02.464361Z","iopub.status.idle":"2023-04-14T13:59:02.470632Z","shell.execute_reply.started":"2023-04-14T13:59:02.464322Z","shell.execute_reply":"2023-04-14T13:59:02.469437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df, test_df","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:02.472448Z","iopub.execute_input":"2023-04-14T13:59:02.473577Z","iopub.status.idle":"2023-04-14T13:59:02.490096Z","shell.execute_reply.started":"2023-04-14T13:59:02.473535Z","shell.execute_reply":"2023-04-14T13:59:02.489018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Train dataset and test dataset serpation","metadata":{}},{"cell_type":"code","source":"datasets = {}\ndatasets['train'] = CassavaLeafDiseaseDataset(TRAIN_DATA_DIR, train_df['image_id'].to_numpy(),\n                                              train_df['label'].to_numpy(), label_classes, T.ToTensor())\ndatasets['test'] = CassavaLeafDiseaseDataset(TRAIN_DATA_DIR, test_df['image_id'].to_numpy(),\n                                             test_df['label'].to_numpy(), label_classes, T.ToTensor())","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:02.491534Z","iopub.execute_input":"2023-04-14T13:59:02.492304Z","iopub.status.idle":"2023-04-14T13:59:02.499439Z","shell.execute_reply.started":"2023-04-14T13:59:02.492267Z","shell.execute_reply":"2023-04-14T13:59:02.498697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Data loaders for train dataset and test dataset","metadata":{}},{"cell_type":"code","source":"batch_size = 16\ndataloaders = {}\ndataloaders['train'] = DataLoader(datasets['train'], batch_size, num_workers=2, pin_memory=False, shuffle=False)\ndataloaders['test'] = DataLoader(datasets['test'], batch_size, num_workers=2, pin_memory=False)","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:02.500738Z","iopub.execute_input":"2023-04-14T13:59:02.501606Z","iopub.status.idle":"2023-04-14T13:59:02.510855Z","shell.execute_reply.started":"2023-04-14T13:59:02.501541Z","shell.execute_reply":"2023-04-14T13:59:02.510114Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Image augmentation","metadata":{}},{"cell_type":"code","source":"!pip install timm","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:02.512417Z","iopub.execute_input":"2023-04-14T13:59:02.513426Z","iopub.status.idle":"2023-04-14T13:59:14.466656Z","shell.execute_reply.started":"2023-04-14T13:59:02.513386Z","shell.execute_reply":"2023-04-14T13:59:14.465179Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from timm.data import create_transform","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:14.470899Z","iopub.execute_input":"2023-04-14T13:59:14.471244Z","iopub.status.idle":"2023-04-14T13:59:15.551346Z","shell.execute_reply.started":"2023-04-14T13:59:14.471202Z","shell.execute_reply":"2023-04-14T13:59:15.550066Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_image(data):\n    img, label = data\n    plt.imshow(img.permute(1, 2, 0))\n    plt.show()\n    print('Label: ' + label_classes[label])","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:15.553990Z","iopub.execute_input":"2023-04-14T13:59:15.554888Z","iopub.status.idle":"2023-04-14T13:59:15.560920Z","shell.execute_reply.started":"2023-04-14T13:59:15.554817Z","shell.execute_reply":"2023-04-14T13:59:15.559412Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Before:","metadata":{}},{"cell_type":"code","source":"show_image(datasets['train'][3])","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:15.562876Z","iopub.execute_input":"2023-04-14T13:59:15.563578Z","iopub.status.idle":"2023-04-14T13:59:15.957225Z","shell.execute_reply.started":"2023-04-14T13:59:15.563530Z","shell.execute_reply":"2023-04-14T13:59:15.956280Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_image(datasets['test'][3])","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:15.958323Z","iopub.execute_input":"2023-04-14T13:59:15.958976Z","iopub.status.idle":"2023-04-14T13:59:16.286770Z","shell.execute_reply.started":"2023-04-14T13:59:15.958938Z","shell.execute_reply":"2023-04-14T13:59:16.285848Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Update transform","metadata":{}},{"cell_type":"markdown","source":"I have apply transforms for dataset based on timm package, these tranforms are simple but will boost the model with some oversampling cause by stratified train_test_split.","metadata":{}},{"cell_type":"code","source":"mean = 0.\nstd = 0.\ntotal_samples = 0.\nfor images, _ in dataloaders['train']:\n    batch_samples = images.size(0)\n    images = images.view(batch_samples, images.size(1), -1)\n    mean += images.mean(2).sum(0)\n    std += images.std(2).sum(0)\n    total_samples += batch_samples\n\nmean /= total_samples\nstd /= total_samples","metadata":{"execution":{"iopub.status.busy":"2023-04-14T13:59:16.288332Z","iopub.execute_input":"2023-04-14T13:59:16.289237Z","iopub.status.idle":"2023-04-14T14:00:43.356531Z","shell.execute_reply.started":"2023-04-14T13:59:16.289188Z","shell.execute_reply":"2023-04-14T14:00:43.355101Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mean","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:00:43.367261Z","iopub.execute_input":"2023-04-14T14:00:43.367616Z","iopub.status.idle":"2023-04-14T14:00:43.377655Z","shell.execute_reply.started":"2023-04-14T14:00:43.367585Z","shell.execute_reply":"2023-04-14T14:00:43.376491Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"datasets['train'].transform = create_transform(input_size=384, is_training=True, interpolation='bicubic', mean=mean, std=std)\ndatasets['test'].transform = T.Compose([T.Resize(size=400),\n                                        T.CenterCrop(size=(384, 384)),\n                                        T.ToTensor()])","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:00:43.379510Z","iopub.execute_input":"2023-04-14T14:00:43.380264Z","iopub.status.idle":"2023-04-14T14:00:43.392279Z","shell.execute_reply.started":"2023-04-14T14:00:43.380209Z","shell.execute_reply":"2023-04-14T14:00:43.391061Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"create_transform(600, is_training=False)","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:00:43.394196Z","iopub.execute_input":"2023-04-14T14:00:43.394619Z","iopub.status.idle":"2023-04-14T14:00:43.407047Z","shell.execute_reply.started":"2023-04-14T14:00:43.394583Z","shell.execute_reply":"2023-04-14T14:00:43.406067Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### After:","metadata":{}},{"cell_type":"code","source":"show_image(datasets['train'][3])","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:00:43.408810Z","iopub.execute_input":"2023-04-14T14:00:43.409286Z","iopub.status.idle":"2023-04-14T14:00:43.668974Z","shell.execute_reply.started":"2023-04-14T14:00:43.409208Z","shell.execute_reply":"2023-04-14T14:00:43.668015Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_image(datasets['test'][3])","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:00:43.670824Z","iopub.execute_input":"2023-04-14T14:00:43.671541Z","iopub.status.idle":"2023-04-14T14:00:43.925398Z","shell.execute_reply.started":"2023-04-14T14:00:43.671503Z","shell.execute_reply":"2023-04-14T14:00:43.924107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Show 40 Images from train data","metadata":{}},{"cell_type":"code","source":"from torchvision.utils import make_grid\ndef show_dataset(ds):\n    fig, ax = plt.subplots(figsize=(16, 16))\n    ax.set_xticks([]); ax.set_yticks([])\n    images = []\n    for i in range(40):\n        image, label = ds[i]\n        images.append(image)\n    ax.imshow(make_grid(images, nrow=8).permute(1, 2, 0))\n        \nshow_dataset(datasets['train'])","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:00:43.926933Z","iopub.execute_input":"2023-04-14T14:00:43.927550Z","iopub.status.idle":"2023-04-14T14:00:47.176011Z","shell.execute_reply.started":"2023-04-14T14:00:43.927508Z","shell.execute_reply":"2023-04-14T14:00:47.174691Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Show 20 Images from train samples as a Grid using matplotlib","metadata":{}},{"cell_type":"code","source":"gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:00:47.177534Z","iopub.execute_input":"2023-04-14T14:00:47.178109Z","iopub.status.idle":"2023-04-14T14:00:47.408195Z","shell.execute_reply.started":"2023-04-14T14:00:47.178071Z","shell.execute_reply":"2023-04-14T14:00:47.406885Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## GPU Utilities","metadata":{}},{"cell_type":"code","source":"device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:00:47.410078Z","iopub.execute_input":"2023-04-14T14:00:47.410497Z","iopub.status.idle":"2023-04-14T14:00:47.539752Z","shell.execute_reply.started":"2023-04-14T14:00:47.410458Z","shell.execute_reply":"2023-04-14T14:00:47.538549Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Modifying a Pretrained Model (convnext)","metadata":{}},{"cell_type":"markdown","source":"Convnext tiny is a great provide a good performance with smaller model size, shorter training time, and a small dataset. It also have similar domain with my task. Thus I will fine tune it and check it out whether it work or not.","metadata":{}},{"cell_type":"code","source":"from torchvision import models","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:00:47.541298Z","iopub.execute_input":"2023-04-14T14:00:47.541640Z","iopub.status.idle":"2023-04-14T14:00:47.552533Z","shell.execute_reply.started":"2023-04-14T14:00:47.541607Z","shell.execute_reply":"2023-04-14T14:00:47.551538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_ft = models.convnext_tiny(pretrained=True)","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:00:47.554036Z","iopub.execute_input":"2023-04-14T14:00:47.554873Z","iopub.status.idle":"2023-04-14T14:00:49.268775Z","shell.execute_reply.started":"2023-04-14T14:00:47.554837Z","shell.execute_reply":"2023-04-14T14:00:49.267671Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_ft","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:00:49.270355Z","iopub.execute_input":"2023-04-14T14:00:49.272623Z","iopub.status.idle":"2023-04-14T14:00:49.288299Z","shell.execute_reply.started":"2023-04-14T14:00:49.272579Z","shell.execute_reply":"2023-04-14T14:00:49.287229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for param in model_ft.parameters():\n    param.requires_grad = False\n\nn_inputs = 768\nn_outputs = 5\n\n\nsequential_layers = nn.Linear(n_inputs, n_outputs, bias=True)\nmodel_ft.classifier[2] = sequential_layers\n\n#model_ft.avgpool = nn.AdaptiveMaxPool2d(output_size=1)\n\nfor (param1, param2, param3) in zip(model_ft.classifier.parameters(), model_ft.features[7].parameters(), model_ft.avgpool.parameters()):\n    param1.requires_grad, param2.requires_grad, param3.requires_grad = True","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:00:49.290160Z","iopub.execute_input":"2023-04-14T14:00:49.290621Z","iopub.status.idle":"2023-04-14T14:00:49.301809Z","shell.execute_reply.started":"2023-04-14T14:00:49.290585Z","shell.execute_reply":"2023-04-14T14:00:49.300750Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_ft = model_ft.to(device)","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:00:49.303307Z","iopub.execute_input":"2023-04-14T14:00:49.303759Z","iopub.status.idle":"2023-04-14T14:00:52.956766Z","shell.execute_reply.started":"2023-04-14T14:00:49.303715Z","shell.execute_reply":"2023-04-14T14:00:52.955722Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = nn.NLLLoss()","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:00:52.958156Z","iopub.execute_input":"2023-04-14T14:00:52.959503Z","iopub.status.idle":"2023-04-14T14:00:52.964405Z","shell.execute_reply.started":"2023-04-14T14:00:52.959459Z","shell.execute_reply":"2023-04-14T14:00:52.963298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model training and validation","metadata":{}},{"cell_type":"markdown","source":"## Test if model make correct output ","metadata":{}},{"cell_type":"code","source":"test_in = torch.ones((16,3,384,384)).to(torch.device('cuda'),non_blocking=True)\ntest_out = model_ft(test_in)\nprint(test_out.shape)","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:00:52.966183Z","iopub.execute_input":"2023-04-14T14:00:52.966547Z","iopub.status.idle":"2023-04-14T14:01:00.530512Z","shell.execute_reply.started":"2023-04-14T14:00:52.966512Z","shell.execute_reply":"2023-04-14T14:01:00.528981Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Evaluate loss and accuracy before training","metadata":{}},{"cell_type":"code","source":"def evaluate(model, dataloader):\n    outputs = []\n    with torch.no_grad():\n        model.eval()\n        for batch in dataloader:\n            images, labels = batch\n            images = images.to(torch.device('cuda'),non_blocking=True)\n            labels = labels.to(torch.device('cuda'),non_blocking=True)\n            out = model(images)                    \n            loss = F.cross_entropy(out, labels)   \n        \n            _, preds = torch.max(out, dim=1)\n            acc = torch.tensor(torch.sum(preds == labels).item() / len(preds))\n            outputs.append({'test_loss': loss.detach(), 'test_acc': acc})\n\n        batch_losses = [x['test_loss'] for x in outputs]\n        epoch_loss = torch.stack(batch_losses).mean()\n        batch_accs = [x['test_acc'] for x in outputs]\n        epoch_acc = torch.stack(batch_accs).mean() \n        result = {'test_loss': epoch_loss.item(), 'test_acc': epoch_acc.item()}\n        torch.cuda.empty_cache()\n        \n    return result ","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:01:00.531999Z","iopub.execute_input":"2023-04-14T14:01:00.533097Z","iopub.status.idle":"2023-04-14T14:01:00.543064Z","shell.execute_reply.started":"2023-04-14T14:01:00.533055Z","shell.execute_reply":"2023-04-14T14:01:00.542029Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = [evaluate(model_ft, dataloaders['test'])]\nhistory","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:01:00.546582Z","iopub.execute_input":"2023-04-14T14:01:00.547141Z","iopub.status.idle":"2023-04-14T14:01:21.769966Z","shell.execute_reply.started":"2023-04-14T14:01:00.547112Z","shell.execute_reply":"2023-04-14T14:01:21.768723Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"You can see the model have no usefulness before trainning","metadata":{}},{"cell_type":"markdown","source":"## Find the best learning rate","metadata":{}},{"cell_type":"code","source":"from ignite.engine import create_supervised_trainer, create_supervised_evaluator\nfrom ignite.metrics import Loss, Accuracy\nfrom ignite.contrib.handlers import FastaiLRFinder, ProgressBar","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:01:21.771688Z","iopub.execute_input":"2023-04-14T14:01:21.772126Z","iopub.status.idle":"2023-04-14T14:01:22.127842Z","shell.execute_reply.started":"2023-04-14T14:01:21.772090Z","shell.execute_reply":"2023-04-14T14:01:22.126659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optimizer = torch.optim.Adam([{'params': model_ft.features[7].parameters()},\n                              {'params': model_ft.avgpool.parameters()},\n                              {'params': model_ft.classifier.parameters(),\n                               'lr': 1e-7 }],\n                              lr=0.001 / 10, weight_decay=0.001)\n                           ","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:01:22.129178Z","iopub.execute_input":"2023-04-14T14:01:22.129592Z","iopub.status.idle":"2023-04-14T14:01:22.138224Z","shell.execute_reply.started":"2023-04-14T14:01:22.129553Z","shell.execute_reply":"2023-04-14T14:01:22.136222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install torch-lr-finder\nfrom torch_lr_finder import LRFinder","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:01:22.139838Z","iopub.execute_input":"2023-04-14T14:01:22.140859Z","iopub.status.idle":"2023-04-14T14:01:32.843136Z","shell.execute_reply.started":"2023-04-14T14:01:22.140818Z","shell.execute_reply":"2023-04-14T14:01:32.841849Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lr_finder = LRFinder(model_ft, optimizer, criterion, device=\"cuda\")\nlr_finder.range_test(dataloaders['train'],start_lr = 1e-7, end_lr=100, num_iter=100)\nlr_finder.plot() # to inspect the loss-learning rate graph\nlr_finder.reset() # to reset the model and optimizer to their initial state","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:01:32.845389Z","iopub.execute_input":"2023-04-14T14:01:32.845846Z","iopub.status.idle":"2023-04-14T14:01:34.979419Z","shell.execute_reply.started":"2023-04-14T14:01:32.845777Z","shell.execute_reply":"2023-04-14T14:01:34.978343Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Learning rate finder provide by Fastai is quite helpful to find a good starting learning rate for your model.","metadata":{}},{"cell_type":"markdown","source":"## Finetuning the Pretrained Model","metadata":{}},{"cell_type":"code","source":"gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:01:34.981553Z","iopub.execute_input":"2023-04-14T14:01:34.981959Z","iopub.status.idle":"2023-04-14T14:01:35.201681Z","shell.execute_reply.started":"2023-04-14T14:01:34.981915Z","shell.execute_reply":"2023-04-14T14:01:35.200282Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm.notebook import tqdm","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:01:35.203186Z","iopub.execute_input":"2023-04-14T14:01:35.204199Z","iopub.status.idle":"2023-04-14T14:01:35.215363Z","shell.execute_reply.started":"2023-04-14T14:01:35.204158Z","shell.execute_reply":"2023-04-14T14:01:35.214088Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"epochs = 8\nlr = 0.001\ngrad_clip = 1\noptimizer = torch.optim.RMSprop\nweight_decay = 0.001\n#steps_per_epochs=64\naccum_iter = 10\nprint_every = 10\n#momentum = 0.9\nalpha=0.9","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:16:03.392860Z","iopub.execute_input":"2023-04-14T14:16:03.393237Z","iopub.status.idle":"2023-04-14T14:16:03.398740Z","shell.execute_reply.started":"2023-04-14T14:16:03.393204Z","shell.execute_reply":"2023-04-14T14:16:03.397805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_fine_tuning(epochs, lr, model, train_loader, test_loader, alpha,\n                      devices,accum_iter, print_every, optimizer=torch.optim.Adam):\n    \n    optimizer = optimizer([{'params': model.features[7].parameters()},\n                           {'params': model.avgpool.parameters()},\n                           {'params': model.classifier.parameters(),\n                            'lr': lr }],\n                          lr=lr / 10, alpha=alpha)\n    \n    sched = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min')\n    \n    total_hist = [[],[],[],[]]\n    \n    # batch accumulation parameter\n    accum_iter = accum_iter  \n    \n    steps = 0\n    running_loss = 0\n    running_accuracy = 0\n    print_every = print_every\n    \n    for epoch in range(epochs):\n        \n        model.train()\n        \n        for batch_idx, (images, labels) in enumerate(train_loader):\n            steps += 1\n                          \n            images,labels = images.to(device),labels.to(device)\n            \n            with torch.set_grad_enabled(True):\n                \n                y_hat = model(images)\n                loss = criterion(y_hat, labels)\n                \n                running_loss += loss.item()\n                \n                ps = torch.exp(y_hat)\n                top_p, top_class = ps.topk(1, dim=1)\n                equals = top_class == labels.view(*top_class.shape)\n                running_accuracy += torch.mean(equals.type(torch.FloatTensor)).item()\n                \n                \n                loss = loss / accum_iter \n                \n                loss.backward()\n                \n                if ((batch_idx + 1) % accum_iter == 0) or (batch_idx + 1 == len(train_loader)):\n                    optimizer.step()\n                    torch.nn.utils.clip_grad_norm_(parameters=model_ft.parameters(), max_norm=1)\n                    optimizer.zero_grad()\n                \n                \n                del images,labels,y_hat\n                torch.cuda.empty_cache()\n                \n                if steps % print_every == 0:\n                    test_loss = 0\n                    accuracy = 0\n                    with torch.no_grad():\n                        model.eval()\n                        for images2, labels2 in tqdm(test_loader):\n\n                            images2, labels2 = images2.to(device),labels2.to(device)\n\n                            y_hat2 = model(images2)\n                            batch_loss = criterion(y_hat2, labels2)\n                    \n                            test_loss += batch_loss.item()\n\n                            # Calculate accuracy\n                            ps2 = torch.exp(y_hat2)\n                            top_p2, top_class2 = ps2.topk(1, dim=1)\n                            equals2 = top_class2 == labels2.view(*top_class2.shape)\n                            accuracy += torch.mean(equals2.type(torch.FloatTensor)).item()\n                            \n                            del images2,labels2,y_hat2\n                            torch.cuda.empty_cache()\n                    \n                    train_loss = running_loss/print_every\n                    train_accuracy = running_accuracy/print_every\n                    test_loss = test_loss/len(test_loader)\n                    test_accuracy = accuracy/len(test_loader)\n                    total_hist[0].append(train_loss)\n                    total_hist[1].append(train_accuracy)\n                    total_hist[2].append(test_loss)\n                    total_hist[3].append(test_accuracy)\n                    \n                    print(f\"Epoch {epoch+1}/{epochs}.. \"\n                          f\"Train loss: {train_loss:.3f}.. \"\n                          f\"Test loss: {test_loss:.3f}.. \"\n                          f\"Test accuracy: {test_accuracy:.3f}\")\n                    running_loss = 0\n                    running_accuracy = 0\n                    model.train()\n                    \n                    sched.step(test_loss)\n                \n        torch.cuda.empty_cache()\n        gc.collect()\n        \n                                      \n    train_total_loss = np.array(total_hist[0])\n    train_total_acc = np.array(total_hist[1])\n    \n    val_total_loss = np.array(total_hist[2])   \n    val_total_acc = np.array(total_hist[3])\n    \n    \n    total = {'train_loss': train_total_loss, 'train_acc': train_total_acc,\n               'val_loss': val_total_loss, 'val_acc': val_total_acc}\n    \n    return pd.DataFrame({'Train Loss':total_hist[0], 'Train Accuracy':total_hist[1],\n                         'Val Loss':total_hist[2], 'Val Accuracy':total_hist[3]})\n","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:16:05.156698Z","iopub.execute_input":"2023-04-14T14:16:05.157267Z","iopub.status.idle":"2023-04-14T14:16:05.182452Z","shell.execute_reply.started":"2023-04-14T14:16:05.157226Z","shell.execute_reply":"2023-04-14T14:16:05.181400Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(dataloaders['train'])/10","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:16:08.907709Z","iopub.execute_input":"2023-04-14T14:16:08.908741Z","iopub.status.idle":"2023-04-14T14:16:08.915672Z","shell.execute_reply.started":"2023-04-14T14:16:08.908705Z","shell.execute_reply":"2023-04-14T14:16:08.914682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"I had try serveral hyperparameters, torch.optim.RMSprop work the best, and 8 epochs would prevent overfiting. I had tried OneCycleLR but abandon it later as it is less effective, ReduceLROnPlateau is way better.","metadata":{}},{"cell_type":"code","source":"%%time\nhistory2 = train_fine_tuning(epochs, lr, model_ft, dataloaders['train'], dataloaders['test'],\n                             weight_decay, device, accum_iter, print_every, optimizer)","metadata":{"execution":{"iopub.status.busy":"2023-04-14T14:16:10.532774Z","iopub.execute_input":"2023-04-14T14:16:10.533862Z","iopub.status.idle":"2023-04-14T15:34:28.444553Z","shell.execute_reply.started":"2023-04-14T14:16:10.533811Z","shell.execute_reply":"2023-04-14T15:34:28.443095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history2","metadata":{"execution":{"iopub.status.busy":"2023-04-14T15:40:08.569763Z","iopub.execute_input":"2023-04-14T15:40:08.570813Z","iopub.status.idle":"2023-04-14T15:40:08.586847Z","shell.execute_reply.started":"2023-04-14T15:40:08.570753Z","shell.execute_reply":"2023-04-14T15:40:08.585842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"index=history2.index.tolist()","metadata":{"execution":{"iopub.status.busy":"2023-04-14T15:40:11.027349Z","iopub.execute_input":"2023-04-14T15:40:11.027741Z","iopub.status.idle":"2023-04-14T15:40:11.033695Z","shell.execute_reply.started":"2023-04-14T15:40:11.027688Z","shell.execute_reply":"2023-04-14T15:40:11.032338Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(16,10))\nplt.title('Loss curve: ')\nplt.plot(index, history2['Train Loss'], 'b-o')\nplt.plot(index, history2['Val Loss'], 'r-o')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.legend(['Training', 'Validation'])","metadata":{"execution":{"iopub.status.busy":"2023-04-14T15:40:13.276251Z","iopub.execute_input":"2023-04-14T15:40:13.277043Z","iopub.status.idle":"2023-04-14T15:40:13.579752Z","shell.execute_reply.started":"2023-04-14T15:40:13.277004Z","shell.execute_reply":"2023-04-14T15:40:13.578812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(16,10))\nplt.title('Accuracy curve: ')\nplt.plot(index, history2['Train Accuracy'], 'b-o')\nplt.plot(index, history2['Val Accuracy'], 'r-o')\nplt.xlabel('Epoch')\nplt.ylabel('Accuracy')\nplt.legend(['Training', 'Validation'])","metadata":{"execution":{"iopub.status.busy":"2023-04-14T15:40:18.856295Z","iopub.execute_input":"2023-04-14T15:40:18.857301Z","iopub.status.idle":"2023-04-14T15:40:19.154606Z","shell.execute_reply.started":"2023-04-14T15:40:18.857253Z","shell.execute_reply":"2023-04-14T15:40:19.153681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Evaluate loss and accuracy after training","metadata":{}},{"cell_type":"code","source":"history3 = [evaluate(model_ft, dataloaders['test'])]\nhistory3","metadata":{"execution":{"iopub.status.busy":"2023-04-14T15:40:42.356896Z","iopub.execute_input":"2023-04-14T15:40:42.357277Z","iopub.status.idle":"2023-04-14T15:41:01.834309Z","shell.execute_reply.started":"2023-04-14T15:40:42.357245Z","shell.execute_reply":"2023-04-14T15:41:01.833248Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The model have improved its accuracy after training","metadata":{}},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"images,labels =[],[]\npred = []\n\nfor batch in dataloaders['test']:\n    images, labels = batch\n    images = images.to(device)\n    labels = labels.to(device)\n    pred = model_ft(images)\n    images\n    labels\n    break","metadata":{"execution":{"iopub.status.busy":"2023-04-14T15:41:43.366758Z","iopub.execute_input":"2023-04-14T15:41:43.367899Z","iopub.status.idle":"2023-04-14T15:41:44.417350Z","shell.execute_reply.started":"2023-04-14T15:41:43.367848Z","shell.execute_reply":"2023-04-14T15:41:44.416074Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred = pred.argmax(axis=1)","metadata":{"execution":{"iopub.status.busy":"2023-04-14T15:41:45.218061Z","iopub.execute_input":"2023-04-14T15:41:45.218459Z","iopub.status.idle":"2023-04-14T15:41:45.225397Z","shell.execute_reply.started":"2023-04-14T15:41:45.218421Z","shell.execute_reply":"2023-04-14T15:41:45.224143Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def display_img(img,label,prediction):\n    img,label = img.to(torch.device('cpu')),label.to(torch.device('cpu'))\n    prediction = prediction.to(torch.device('cpu'),non_blocking=True)\n    plt.imshow(img.permute(1, 2, 0))\n    plt.show()\n    print('Label: '+label_classes[int(label)])\n    print('Prediction: '+label_classes[int(prediction)])\n    ","metadata":{"execution":{"iopub.status.busy":"2023-04-14T15:41:46.587525Z","iopub.execute_input":"2023-04-14T15:41:46.587944Z","iopub.status.idle":"2023-04-14T15:41:46.594896Z","shell.execute_reply.started":"2023-04-14T15:41:46.587911Z","shell.execute_reply":"2023-04-14T15:41:46.593863Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display_img(images[2],labels[2],pred[2])","metadata":{"execution":{"iopub.status.busy":"2023-04-14T15:41:48.617850Z","iopub.execute_input":"2023-04-14T15:41:48.618963Z","iopub.status.idle":"2023-04-14T15:41:48.868865Z","shell.execute_reply.started":"2023-04-14T15:41:48.618922Z","shell.execute_reply":"2023-04-14T15:41:48.867855Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display_img(images[5],labels[5],pred[5])","metadata":{"execution":{"iopub.status.busy":"2023-04-14T15:41:53.898304Z","iopub.execute_input":"2023-04-14T15:41:53.898664Z","iopub.status.idle":"2023-04-14T15:41:54.140799Z","shell.execute_reply.started":"2023-04-14T15:41:53.898632Z","shell.execute_reply":"2023-04-14T15:41:54.139760Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display_img(images[13],labels[13],pred[13])","metadata":{"execution":{"iopub.status.busy":"2023-04-14T15:41:57.117565Z","iopub.execute_input":"2023-04-14T15:41:57.118349Z","iopub.status.idle":"2023-04-14T15:41:57.356114Z","shell.execute_reply.started":"2023-04-14T15:41:57.118309Z","shell.execute_reply":"2023-04-14T15:41:57.355019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display_img(images[9],labels[9],pred[9])","metadata":{"execution":{"iopub.status.busy":"2023-04-14T15:42:00.663014Z","iopub.execute_input":"2023-04-14T15:42:00.663419Z","iopub.status.idle":"2023-04-14T15:42:00.906856Z","shell.execute_reply.started":"2023-04-14T15:42:00.663387Z","shell.execute_reply":"2023-04-14T15:42:00.905919Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Save the model","metadata":{}},{"cell_type":"code","source":"torch.save(model_ft.state_dict(),'cassava-leaf-disease-classification.pt')","metadata":{"execution":{"iopub.status.busy":"2023-04-14T15:42:03.790818Z","iopub.execute_input":"2023-04-14T15:42:03.791193Z","iopub.status.idle":"2023-04-14T15:42:03.976919Z","shell.execute_reply.started":"2023-04-14T15:42:03.791161Z","shell.execute_reply":"2023-04-14T15:42:03.975844Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"##  Conclusion\n\nI have learned a lot in this project, not just developing transfer learning algorithms, but also things like handling imbalanced data, image augumentation, validate your model, use appropriate hyperparameters and pytorch optimizer and etc.\n\n## Future works\n\nI may try a few other CNN models, and hybrid models with LSTM. Moreover, adjusting hyperparameters, and tried few other image augumentation approaches can be helpful as well. ","metadata":{}},{"cell_type":"markdown","source":"## Reference:\n\n1. https://jovian.ai/learn/deep-learning-with-pytorch-zero-to-gans\n\n\n2. https://jovian.ai/sachin-it-ds/traffic-sign-recognition-using-pytorch-and-cnn-project\n\n\n3. https://towardsdatascience.com/demystifying-pytorchs-weightedrandomsampler-by-example-a68aceccb452","metadata":{}}]}