{"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":"## Exercise 1","metadata":{"id":"oWAprhLEhCQ4"}},{"cell_type":"code","source":"!pip install torchsummary","metadata":{"execution":{"iopub.status.busy":"2022-01-16T23:46:01.509287Z","iopub.execute_input":"2022-01-16T23:46:01.509645Z","iopub.status.idle":"2022-01-16T23:46:10.719809Z","shell.execute_reply.started":"2022-01-16T23:46:01.509545Z","shell.execute_reply":"2022-01-16T23:46:10.719018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.datasets as datasets\nfrom torch.utils.data import DataLoader, Dataset\nimport torch.optim as optim\nimport torchvision.transforms as transforms\nfrom torchsummary import summary\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport pandas as pd\nimport os\nimport tqdm\nfrom PIL import Image\nimport time\n%matplotlib inline","metadata":{"id":"_bgwk2vdeR5q","execution":{"iopub.status.busy":"2022-01-16T23:46:14.011383Z","iopub.execute_input":"2022-01-16T23:46:14.011935Z","iopub.status.idle":"2022-01-16T23:46:15.69981Z","shell.execute_reply.started":"2022-01-16T23:46:14.011892Z","shell.execute_reply":"2022-01-16T23:46:15.699076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LeNet(nn.Module):\n    def __init__(self, n_classes):\n        super(LeNet, self).__init__()\n        self.n_classes = n_classes\n        # 1x32x32 => 6x28x28\n        self.conv1 = nn.Conv2d(in_channels=1, out_channels=6, kernel_size=5, stride=1)\n        self.relu1 = nn.ReLU()\n        # 6x28x28 => 6x14x14\n        self.pool1 = nn.AvgPool2d(kernel_size=2)\n        # 6x14x14 => 16x10x10\n        self.conv2 = nn.Conv2d(in_channels=6, out_channels=16, kernel_size=5, stride=1)\n        self.relu2 = nn.ReLU()\n        # Pooling 16x10x10 => 16x5x5\n        self.pool2 = nn.AvgPool2d(kernel_size=2)\n        # Fully connected\n        self.fc1 = nn.Linear(16*5*5, 120)   # convert matrix with 16*5*5 (= 400) features to a matrix of 120 features (columns)\n        self.relu3 = nn.ReLU()\n        self.fc2 = nn.Linear(120, 84)\n        self.relu4 = nn.ReLU()\n        self.fc3 = nn.Linear(84, n_classes)\n\n    def forward(self, x):\n        output = self.relu1(self.conv1(x))\n        output = self.pool1(output)\n        output = self.relu2(self.conv2(output))\n        output = self.pool2(output)\n        #Flatten\n        output = output.view(-1, 16*5*5)\n        # Fully connected layer\n        output = self.fc1(output)\n        output = self.relu3(output)\n        output = self.fc2(output)\n        output = self.relu4(output)\n        output = self.fc3(output)\n        output = F.softmax(output, dim=1)\n        return output","metadata":{"id":"C51pAOK0ft6t","execution":{"iopub.status.busy":"2022-01-16T23:46:16.274289Z","iopub.execute_input":"2022-01-16T23:46:16.274949Z","iopub.status.idle":"2022-01-16T23:46:16.286621Z","shell.execute_reply.started":"2022-01-16T23:46:16.274906Z","shell.execute_reply":"2022-01-16T23:46:16.285902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### See architecture","metadata":{"id":"7A11GE8-oUEB"}},{"cell_type":"code","source":"model = LeNet(n_classes=10)\nprint(model)","metadata":{"id":"im1wVKPFoVnJ","outputId":"0a27182d-b34d-4561-ccd0-c945f054ef72","execution":{"iopub.status.busy":"2022-01-16T23:46:21.615232Z","iopub.execute_input":"2022-01-16T23:46:21.615989Z","iopub.status.idle":"2022-01-16T23:46:21.64265Z","shell.execute_reply.started":"2022-01-16T23:46:21.615947Z","shell.execute_reply":"2022-01-16T23:46:21.641768Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"summary(model, (1,32,32), device='cpu')","metadata":{"id":"HduNEPsjohEP","outputId":"ce7845a1-fe3b-4ad4-a1f7-5be5b9e9fbae","execution":{"iopub.status.busy":"2022-01-16T23:46:24.104436Z","iopub.execute_input":"2022-01-16T23:46:24.105052Z","iopub.status.idle":"2022-01-16T23:46:24.274465Z","shell.execute_reply.started":"2022-01-16T23:46:24.105014Z","shell.execute_reply":"2022-01-16T23:46:24.273756Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"id":"9kMpwVatpJ9c"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Exercise 2: Xây dựng kiến trúc VGG16","metadata":{"id":"1Zwk32ffqWmK"}},{"cell_type":"code","source":"class VGG16(nn.Module):\n    def __init__(self, n_classes):\n        super(VGG16, self).__init__()\n        self.n_classes = n_classes\n\n        self.block_1 = nn.Sequential(nn.Conv2d(3, 64, kernel_size=3, stride=1),\n                                    nn.ReLU(),\n                                    nn.Conv2d(64, 64, kernel_size=3, stride=1),\n                                    nn.ReLU(),\n                                    nn.MaxPool2d(2, 2))\n\n        self.block_2 = nn.Sequential(nn.Conv2d(64, 128, kernel_size=3, stride=1),\n                                     nn.ReLU(),\n                                     nn.Conv2d(128, 128, kernel_size=3, stride=1),\n                                     nn.ReLU(),\n                                     nn.MaxPool2d(2, 2))\n\n        self.block_3 = nn.Sequential(nn.Conv2d(128, 256, kernel_size=3, stride=1),\n                                    nn.ReLU(),\n                                    nn.Conv2d(256, 256, kernel_size=3, stride=1),\n                                    nn.ReLU(),\n                                    nn.Conv2d(256, 256, kernel_size=3, stride=1),\n                                    nn.ReLU(),\n                                    nn.MaxPool2d(2, 2))\n\n        self.block_4 = nn.Sequential(nn.Conv2d(256, 512, kernel_size=3, stride=1),\n                                    nn.ReLU(),\n                                    nn.Conv2d(512, 512, kernel_size=3, stride=1),\n                                    nn.ReLU(),\n                                    nn.Conv2d(512, 512, kernel_size=3, stride=1),\n                                    nn.ReLU(),\n                                    nn.MaxPool2d(2, 2))\n\n        self.block_5 = nn.Sequential(nn.Conv2d(512, 512, kernel_size=3, stride=1),\n                                nn.ReLU(),\n                                nn.Conv2d(512, 512, kernel_size=3, stride=1),\n                                nn.ReLU(),\n                                nn.Conv2d(512, 512, kernel_size=3, stride=1),\n                                nn.ReLU(),\n                                # max pooling (kernel_size, stride)\n                                nn.MaxPool2d(2))\n\n        self.avg = nn.AdaptiveAvgPool2d((7,7))\n\n        self.classifier = nn.Sequential(nn.Linear(7*7*512, 4096),\n                        nn.ReLU(),\n                        nn.Dropout(p=0.5),\n                        nn.Linear(4096, 4096),\n                        nn.ReLU(),\n                        nn.Dropout(p=0.5),\n                        nn.Linear(4096, n_classes))\n\n        for m in self.modules():\n            if isinstance(m, torch.torch.nn.Conv2d) or isinstance(m, torch.torch.nn.Linear):\n                torch.nn.init.kaiming_uniform_(m.weight, mode='fan_in', nonlinearity='relu')\n                if m.bias is not None:\n                    m.bias.detach().zero_()\n\n    def forward(self, x):\n        x = self.block_1(x)\n        x = self.block_2(x)\n        x = self.block_3(x)\n        x = self.block_4(x)\n        x = self.block_5(x)\n        x = self.avg(x)\n        # Flatten\n        x = x.view(-1, 7*7*512)\n        x = self.classifier(x)\n        # output = F.softmax(self.fc3(x), dim=-1)\n        return x","metadata":{"id":"EoiruWDfqe0z","execution":{"iopub.status.busy":"2022-01-16T23:46:28.399876Z","iopub.execute_input":"2022-01-16T23:46:28.400286Z","iopub.status.idle":"2022-01-16T23:46:28.417276Z","shell.execute_reply.started":"2022-01-16T23:46:28.400248Z","shell.execute_reply":"2022-01-16T23:46:28.416579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### See architecture VGG16","metadata":{"id":"peytcU8Fwx-T"}},{"cell_type":"code","source":"net = VGG16(n_classes=1000)\nprint(net)","metadata":{"id":"PBkkUiL3w0sx","outputId":"1d24da95-85e3-4ed3-e8b4-39d44715343f","execution":{"iopub.status.busy":"2022-01-16T23:46:34.393241Z","iopub.execute_input":"2022-01-16T23:46:34.393717Z","iopub.status.idle":"2022-01-16T23:46:36.395231Z","shell.execute_reply.started":"2022-01-16T23:46:34.39368Z","shell.execute_reply":"2022-01-16T23:46:36.394457Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"summary(net, (3,224,224), device='cpu')","metadata":{"id":"zeVmuV-5xewD","outputId":"2b7ae534-4748-4ee2-bbbc-1cf463ecb8fc","execution":{"iopub.status.busy":"2022-01-16T23:46:37.336064Z","iopub.execute_input":"2022-01-16T23:46:37.336722Z","iopub.status.idle":"2022-01-16T23:46:38.023183Z","shell.execute_reply.started":"2022-01-16T23:46:37.336677Z","shell.execute_reply":"2022-01-16T23:46:38.02245Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Prepare dataset","metadata":{"id":"i-ZWOmDB7uox"}},{"cell_type":"code","source":"from google.colab import files\nfiles.upload()         # expire any previous token(s) and upload recreated token","metadata":{"id":"pjxlEP8T7y6n","outputId":"5f1b07d7-c70b-49e3-c065-904bdf803e6e"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install --upgrade --force-reinstall --no-deps kaggle","metadata":{"id":"sNEj6Y9mofDy","outputId":"f8aeac77-f332-4d82-beee-2c1c2d58409d"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir ~/.kaggle","metadata":{"id":"XSYMqidQ8tfW"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!cp kaggle.json ~/.kaggle/","metadata":{"id":"rz7adBf99Nps"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!chmod 600 ~/.kaggle/kaggle.json","metadata":{"id":"Abu9RWQx9QeG"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!kaggle competitions download -c cassava-leaf-disease-classification","metadata":{"id":"T-dyWa4SoYim","outputId":"3c90a15c-e0e6-4c29-a928-a24ee03718e2"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!unzip /content/cassava-leaf-disease-classification.zip","metadata":{"id":"oXy5n3Tx9WW3","outputId":"8ec968ab-a04c-4df0-a56a-c3f4329ba286"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Read file train.csv\ndf = pd.read_csv(\"../input/cassava-leaf-disease-classification/train.csv\")\ndf.head()","metadata":{"id":"opNCPoyh_PRF","outputId":"4d22c42d-0cfd-4d8d-caaa-ec879d237f5b","execution":{"iopub.status.busy":"2022-01-16T23:46:48.671967Z","iopub.execute_input":"2022-01-16T23:46:48.672249Z","iopub.status.idle":"2022-01-16T23:46:48.715478Z","shell.execute_reply.started":"2022-01-16T23:46:48.672219Z","shell.execute_reply":"2022-01-16T23:46:48.714699Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['path'] = df['image_id'].map(lambda x: os.path.join('../input/cassava-leaf-disease-classification/','train_images',x))","metadata":{"id":"t0VaurE2_18i","execution":{"iopub.status.busy":"2022-01-16T23:46:48.817574Z","iopub.execute_input":"2022-01-16T23:46:48.817899Z","iopub.status.idle":"2022-01-16T23:46:48.877239Z","shell.execute_reply.started":"2022-01-16T23:46:48.817869Z","shell.execute_reply":"2022-01-16T23:46:48.876601Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.drop(columns=['image_id'],inplace=True)\ndf = df.reset_index(drop=True)\ndf.head(10)","metadata":{"id":"1q_3Fe9KAGkQ","outputId":"6ef2b5b0-e3b9-4c5b-a374-3ff37ea73b8c","execution":{"iopub.status.busy":"2022-01-16T23:46:51.092537Z","iopub.execute_input":"2022-01-16T23:46:51.092822Z","iopub.status.idle":"2022-01-16T23:46:51.109474Z","shell.execute_reply.started":"2022-01-16T23:46:51.092769Z","shell.execute_reply":"2022-01-16T23:46:51.108742Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## EDA","metadata":{"id":"M0Yu7rQPAVam"}},{"cell_type":"code","source":"df.label.value_counts()","metadata":{"id":"GNvgpGX0AWbF","outputId":"1fee256c-2491-4387-b0a7-fcf3d113f54e","execution":{"iopub.status.busy":"2022-01-16T23:46:53.693471Z","iopub.execute_input":"2022-01-16T23:46:53.693905Z","iopub.status.idle":"2022-01-16T23:46:53.702768Z","shell.execute_reply.started":"2022-01-16T23:46:53.693865Z","shell.execute_reply":"2022-01-16T23:46:53.702081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.plot.hist(df.label,figsize=(10,5))","metadata":{"id":"2LAdlbDGEuDJ","outputId":"2d71394f-8bf9-401f-c81f-3a463a67dcaa","execution":{"iopub.status.busy":"2022-01-16T23:46:54.54743Z","iopub.execute_input":"2022-01-16T23:46:54.547957Z","iopub.status.idle":"2022-01-16T23:46:54.833403Z","shell.execute_reply.started":"2022-01-16T23:46:54.547915Z","shell.execute_reply":"2022-01-16T23:46:54.832639Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import seaborn as sns","metadata":{"id":"E7ajqcU7NO_R","execution":{"iopub.status.busy":"2022-01-16T23:46:56.669649Z","iopub.execute_input":"2022-01-16T23:46:56.670279Z","iopub.status.idle":"2022-01-16T23:46:57.407511Z","shell.execute_reply.started":"2022-01-16T23:46:56.670237Z","shell.execute_reply":"2022-01-16T23:46:57.40681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(20,8))\nsns.countplot(data=df, x='label')","metadata":{"id":"n128oxjwNLLN","outputId":"f7c6fe85-dad4-4c57-cee9-c7e899e2b4df","execution":{"iopub.status.busy":"2022-01-16T23:46:57.408921Z","iopub.execute_input":"2022-01-16T23:46:57.409386Z","iopub.status.idle":"2022-01-16T23:46:57.613723Z","shell.execute_reply.started":"2022-01-16T23:46:57.409355Z","shell.execute_reply":"2022-01-16T23:46:57.613095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import json\nwith open(\"../input/cassava-leaf-disease-classification/label_num_to_disease_map.json\") as fn:\n    print(json.loads(fn.read()))","metadata":{"id":"zWAsTnqJEwkU","outputId":"5d4a29d4-a75a-4a33-9103-898970d33c28","execution":{"iopub.status.busy":"2022-01-16T23:46:59.489139Z","iopub.execute_input":"2022-01-16T23:46:59.489693Z","iopub.status.idle":"2022-01-16T23:46:59.503052Z","shell.execute_reply.started":"2022-01-16T23:46:59.489654Z","shell.execute_reply":"2022-01-16T23:46:59.501521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(df))","metadata":{"id":"n1FVUfgSFCDf","outputId":"5bb4cafd-d31b-4320-baf7-890fcb76adcc","execution":{"iopub.status.busy":"2022-01-16T23:47:00.762324Z","iopub.execute_input":"2022-01-16T23:47:00.763102Z","iopub.status.idle":"2022-01-16T23:47:00.768195Z","shell.execute_reply.started":"2022-01-16T23:47:00.763061Z","shell.execute_reply":"2022-01-16T23:47:00.767167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img = Image.open(df['path'][1])\nw,b = img.size\nprint(w,b)","metadata":{"id":"Pj7ck0VhFeET","outputId":"57ca0ec2-16b2-4160-f636-f77c83640079","execution":{"iopub.status.busy":"2022-01-16T23:47:01.730777Z","iopub.execute_input":"2022-01-16T23:47:01.731619Z","iopub.status.idle":"2022-01-16T23:47:01.754449Z","shell.execute_reply.started":"2022-01-16T23:47:01.731563Z","shell.execute_reply":"2022-01-16T23:47:01.753536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Split dataset\nfrom sklearn.model_selection import train_test_split\n\ndf_train, df_valid = train_test_split( df, test_size =0.15, random_state = 42, stratify=df.label.values)\ndf_train = df_train.reset_index(drop=True)\ndf_valid = df_valid.reset_index(drop=True)","metadata":{"id":"2bgSKpfpFjx7","execution":{"iopub.status.busy":"2022-01-16T23:47:02.916718Z","iopub.execute_input":"2022-01-16T23:47:02.917308Z","iopub.status.idle":"2022-01-16T23:47:03.074628Z","shell.execute_reply.started":"2022-01-16T23:47:02.917268Z","shell.execute_reply":"2022-01-16T23:47:03.073908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(df_valid))","metadata":{"id":"R5X-LRIbF7fF","outputId":"9fbdaf5a-6e49-4e8c-fc0a-f51a67e8498d","execution":{"iopub.status.busy":"2022-01-16T23:47:04.299235Z","iopub.execute_input":"2022-01-16T23:47:04.300096Z","iopub.status.idle":"2022-01-16T23:47:04.307334Z","shell.execute_reply.started":"2022-01-16T23:47:04.300049Z","shell.execute_reply":"2022-01-16T23:47:04.306445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"id":"uuaMbRcgGAy2"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Create custom datasets for cassava","metadata":{"id":"QPmSOIPuGLvM"}},{"cell_type":"code","source":"# Data Augumentation\nclass ImageTransform():\n    def __init__(self, input_size):\n        self.data_transform = {\n            'train': transforms.Compose([\n                transforms.RandomResizedCrop(input_size),\n                transforms.RandomHorizontalFlip(),\n                transforms.ToTensor(),\n                transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n            ]),\n            'test': transforms.Compose([\n                    transforms.Resize(input_size),\n                    transforms.CenterCrop(input_size),\n                    transforms.ToTensor(),\n                    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n            ])\n        }\n        \n    def __call__(self, img, phase='train'):\n        return self.data_transform[phase](img)","metadata":{"id":"PAMNqSzzH0-j","execution":{"iopub.status.busy":"2022-01-16T23:47:10.713456Z","iopub.execute_input":"2022-01-16T23:47:10.713901Z","iopub.status.idle":"2022-01-16T23:47:10.721013Z","shell.execute_reply.started":"2022-01-16T23:47:10.713859Z","shell.execute_reply":"2022-01-16T23:47:10.720011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"classes = ['Cassava Bacterial Blight (CBB)',  'Cassava Brown Streak Disease (CBSD)',  'Cassava Green Mottle (CGM)', \n           'Cassava Mosaic Disease (CMD)', 'Healthy']","metadata":{"id":"TJ0XSlCcRhig","execution":{"iopub.status.busy":"2022-01-16T23:47:11.673946Z","iopub.execute_input":"2022-01-16T23:47:11.674511Z","iopub.status.idle":"2022-01-16T23:47:11.678489Z","shell.execute_reply.started":"2022-01-16T23:47:11.67447Z","shell.execute_reply":"2022-01-16T23:47:11.677417Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"idx_to_class = {i:j for i, j in enumerate(classes)}\nclass_to_idx = {value:key for key,value in idx_to_class.items()}","metadata":{"id":"8Vb8w9aKSTQx","execution":{"iopub.status.busy":"2022-01-16T23:47:12.453457Z","iopub.execute_input":"2022-01-16T23:47:12.453736Z","iopub.status.idle":"2022-01-16T23:47:12.458639Z","shell.execute_reply.started":"2022-01-16T23:47:12.453704Z","shell.execute_reply":"2022-01-16T23:47:12.457747Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(class_to_idx.items())","metadata":{"id":"rWi6XbMaSYrK","outputId":"1ec975c9-ac65-4b72-f4cd-12d3e1f6d21e","execution":{"iopub.status.busy":"2022-01-16T23:47:13.696628Z","iopub.execute_input":"2022-01-16T23:47:13.697063Z","iopub.status.idle":"2022-01-16T23:47:13.707754Z","shell.execute_reply.started":"2022-01-16T23:47:13.697009Z","shell.execute_reply":"2022-01-16T23:47:13.70676Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Cassava(Dataset):\n    def __init__(self, img_path, transform=None, phase='train'):\n        super(Cassava, self).__init__()\n        self.img_path = img_path\n        self.transform = transform\n        self.phase = phase\n\n    def __len__(self):\n        return len(self.img_path)\n\n    def __getitem__(self, idx):\n        img_file_path = self.img_path.iloc[idx,1]\n        img = Image.open(img_file_path)\n\n        img_transformed = self.transform(img, self.phase)\n        \n        label = self.img_path.iloc[idx, 0]\n            \n        return img_transformed, label","metadata":{"id":"jlCSZtJGGO6g","execution":{"iopub.status.busy":"2022-01-16T23:47:13.731385Z","iopub.execute_input":"2022-01-16T23:47:13.731841Z","iopub.status.idle":"2022-01-16T23:47:13.739266Z","shell.execute_reply.started":"2022-01-16T23:47:13.731805Z","shell.execute_reply":"2022-01-16T23:47:13.738432Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Resize to this size\ninput_size = 224","metadata":{"id":"DoEIx6bYS5dw","execution":{"iopub.status.busy":"2022-01-16T23:47:14.533685Z","iopub.execute_input":"2022-01-16T23:47:14.534455Z","iopub.status.idle":"2022-01-16T23:47:14.538658Z","shell.execute_reply.started":"2022-01-16T23:47:14.534417Z","shell.execute_reply":"2022-01-16T23:47:14.537917Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds = Cassava(df_train, transform=ImageTransform(input_size), phase='train')\nvalid_ds = Cassava(df_valid, transform=ImageTransform(input_size), phase='test')","metadata":{"id":"qPDsJFzcRn2I","execution":{"iopub.status.busy":"2022-01-16T23:47:17.004882Z","iopub.execute_input":"2022-01-16T23:47:17.005498Z","iopub.status.idle":"2022-01-16T23:47:17.010478Z","shell.execute_reply.started":"2022-01-16T23:47:17.005459Z","shell.execute_reply":"2022-01-16T23:47:17.00945Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds.__len__()","metadata":{"id":"pbEMTyKhTWOd","outputId":"cf09a146-fde7-49d6-8219-92ce0443dc54","execution":{"iopub.status.busy":"2022-01-16T23:47:17.714217Z","iopub.execute_input":"2022-01-16T23:47:17.71497Z","iopub.status.idle":"2022-01-16T23:47:17.720658Z","shell.execute_reply.started":"2022-01-16T23:47:17.71493Z","shell.execute_reply":"2022-01-16T23:47:17.71994Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Test datasets\nindex = 0\nimg, label = train_ds.__getitem__(index)\nprint(img.size())\nprint(label)","metadata":{"id":"tLAdYhCTS4W_","outputId":"2bc19f47-194e-43b5-8549-93f836bca997","execution":{"iopub.status.busy":"2022-01-16T23:47:18.943232Z","iopub.execute_input":"2022-01-16T23:47:18.943815Z","iopub.status.idle":"2022-01-16T23:47:18.984907Z","shell.execute_reply.started":"2022-01-16T23:47:18.943757Z","shell.execute_reply":"2022-01-16T23:47:18.98404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_image_tensor(image_tensor,label):\n    print(image_tensor.shape)\n    plt.figure(figsize=(10,10))\n    plt.title(f'Class: {label}')\n    plt.imshow(img.permute(1,2,0))\n    plt.axis('off')","metadata":{"id":"Hh_2AHHhS97n","execution":{"iopub.status.busy":"2022-01-16T23:47:20.731008Z","iopub.execute_input":"2022-01-16T23:47:20.731618Z","iopub.status.idle":"2022-01-16T23:47:20.73732Z","shell.execute_reply.started":"2022-01-16T23:47:20.73157Z","shell.execute_reply":"2022-01-16T23:47:20.736182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_image_tensor(img, label)","metadata":{"id":"IFYsbC0YUgff","outputId":"cb87fa16-2fff-4a5f-bf81-7ad0110a9709","execution":{"iopub.status.busy":"2022-01-16T23:47:21.819154Z","iopub.execute_input":"2022-01-16T23:47:21.8197Z","iopub.status.idle":"2022-01-16T23:47:22.161208Z","shell.execute_reply.started":"2022-01-16T23:47:21.819663Z","shell.execute_reply":"2022-01-16T23:47:22.160567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')","metadata":{"id":"ihM4xI_aUkhp","execution":{"iopub.status.busy":"2022-01-16T23:47:24.051551Z","iopub.execute_input":"2022-01-16T23:47:24.051834Z","iopub.status.idle":"2022-01-16T23:47:24.096035Z","shell.execute_reply.started":"2022-01-16T23:47:24.051801Z","shell.execute_reply":"2022-01-16T23:47:24.094754Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 32\n\ntrain_dataloader = DataLoader(train_ds, batch_size,shuffle=True)\nvalid_dataloader = DataLoader(valid_ds, batch_size, shuffle=False)\n\ndataloader_dict = {\"train\": train_dataloader, 'test': valid_dataloader}","metadata":{"id":"eJQRH7W4lwgs","execution":{"iopub.status.busy":"2022-01-16T23:47:27.534196Z","iopub.execute_input":"2022-01-16T23:47:27.534738Z","iopub.status.idle":"2022-01-16T23:47:27.539114Z","shell.execute_reply.started":"2022-01-16T23:47:27.534697Z","shell.execute_reply":"2022-01-16T23:47:27.538409Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_iteration = iter(dataloader_dict['train'])\ninputs, labels = next(batch_iteration)","metadata":{"id":"zhshPRiPl7iM","execution":{"iopub.status.busy":"2022-01-16T23:47:28.846471Z","iopub.execute_input":"2022-01-16T23:47:28.847024Z","iopub.status.idle":"2022-01-16T23:47:29.579395Z","shell.execute_reply.started":"2022-01-16T23:47:28.846983Z","shell.execute_reply":"2022-01-16T23:47:29.578681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(inputs.size())\nprint(labels)","metadata":{"id":"e8xsy7ESmCkP","outputId":"811d5329-eb2f-45c4-d236-995fb9e21bbe","execution":{"iopub.status.busy":"2022-01-16T23:47:30.135832Z","iopub.execute_input":"2022-01-16T23:47:30.136514Z","iopub.status.idle":"2022-01-16T23:47:30.141625Z","shell.execute_reply.started":"2022-01-16T23:47:30.136471Z","shell.execute_reply":"2022-01-16T23:47:30.140833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchvision.utils import make_grid","metadata":{"id":"ewZAmwg2mK5w","execution":{"iopub.status.busy":"2022-01-16T23:47:30.329869Z","iopub.execute_input":"2022-01-16T23:47:30.330523Z","iopub.status.idle":"2022-01-16T23:47:30.334099Z","shell.execute_reply.started":"2022-01-16T23:47:30.330487Z","shell.execute_reply":"2022-01-16T23:47:30.333293Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for img, label in train_dataloader:\n    fig, ax = plt.subplots(figsize=(10,8))\n    ax.set_xticks([])\n    ax.set_yticks([])\n    ax.imshow(make_grid(img, 6).permute(1,2,0))\n    break","metadata":{"id":"gMZ8qNnbmE5V","outputId":"fb992977-53fb-4a91-d22e-fb0e87d5b551","execution":{"iopub.status.busy":"2022-01-16T23:47:31.588606Z","iopub.execute_input":"2022-01-16T23:47:31.589373Z","iopub.status.idle":"2022-01-16T23:47:32.86526Z","shell.execute_reply.started":"2022-01-16T23:47:31.589319Z","shell.execute_reply":"2022-01-16T23:47:32.864629Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.tensorboard import SummaryWriter","metadata":{"id":"5tewKAkmm44d","execution":{"iopub.status.busy":"2022-01-16T23:47:36.445674Z","iopub.execute_input":"2022-01-16T23:47:36.446214Z","iopub.status.idle":"2022-01-16T23:47:36.7315Z","shell.execute_reply.started":"2022-01-16T23:47:36.446172Z","shell.execute_reply":"2022-01-16T23:47:36.730696Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"writer = SummaryWriter('runs/vgg16/cassava')","metadata":{"id":"_6n1QmXBnCbp","execution":{"iopub.status.busy":"2022-01-16T23:47:46.178983Z","iopub.execute_input":"2022-01-16T23:47:46.179503Z","iopub.status.idle":"2022-01-16T23:47:49.820328Z","shell.execute_reply.started":"2022-01-16T23:47:46.179464Z","shell.execute_reply":"2022-01-16T23:47:49.819584Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training with VGG16","metadata":{"id":"MQuGoOIvLkN1"}},{"cell_type":"code","source":"def validate(model, data_loader, criterion, device):\n    model.eval()\n    running_loss = 0\n\n    for features, targets in data_loader:\n        features = features.to(device)\n        targets = targets.to(device)\n\n        # forward pass and record loss\n        outputs = model(features)\n        loss = criterion(outputs, targets)\n        running_loss += loss.item()\n    \n    epoch_loss = running_loss / len(data_loader)\n\n    return model, epoch_loss","metadata":{"execution":{"iopub.status.busy":"2022-01-16T23:48:05.95623Z","iopub.execute_input":"2022-01-16T23:48:05.956833Z","iopub.status.idle":"2022-01-16T23:48:05.962567Z","shell.execute_reply.started":"2022-01-16T23:48:05.956788Z","shell.execute_reply":"2022-01-16T23:48:05.961556Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_model(model, num_epochs, train_loader,\n                valid_loader,optimizer,\n                device, logging_interval=50,\n                scheduler=None,\n                scheduler_on='valid_acc'):\n\n    start_time = time.time()\n    minibatch_loss_list, train_acc_list, valid_acc_list = [], [], []\n    valid_loss_min = 1.\n    \n    for epoch in range(num_epochs):\n\n        model.train()\n        for batch_idx, (features, targets) in enumerate(train_loader):\n\n            features = features.to(device)\n            targets = targets.to(device)\n\n            # ## FORWARD AND BACK PROP\n            logits = model(features)\n            loss = torch.nn.functional.cross_entropy(logits, targets)\n            optimizer.zero_grad()\n\n            loss.backward()\n\n            # ## UPDATE MODEL PARAMETERS\n            optimizer.step()\n\n            # ## LOGGING\n            minibatch_loss_list.append(loss.item())\n            if not batch_idx % logging_interval:\n                print(f'Epoch: {epoch+1:03d}/{num_epochs:03d} '\n                      f'| Batch {batch_idx:04d}/{len(train_loader):04d} '\n                      f'| Loss: {loss:.4f}')\n                writer.add_scalar(\"Loss/Epoch\", loss, epoch)\n                if loss < valid_loss_min:\n                    valid_loss_min = loss\n                    checkpoint = {'model': model,\n                                  'epoch' : epoch+1,\n                                  'model_state_dict': model.state_dict(),\n                                  'loss': loss,\n                                  'optimizer' : optimizer.state_dict()}\n                    torch.save(checkpoint, 'model.pth')\n        \n        model.eval()\n        with torch.no_grad():  # save memory during inference\n            train_acc = compute_accuracy(model, train_loader, device=device)\n            valid_acc = compute_accuracy(model, valid_loader, device=device)\n            _, valid_loss = validate(model, valid_loader, criterion, device=device)\n            # Write valid loss to tensorboard\n            writer.add_scalar(\"Valid_loss/Epoch\", valid_loss, epoch)\n            print(f'Epoch: {epoch+1:03d}/{num_epochs:03d} '\n                  f'| Train: {train_acc :.2f}% '\n                  f'| Validation: {valid_acc :.2f}%')\n            train_acc_list.append(train_acc.item())\n            valid_acc_list.append(valid_acc.item())\n            # Write to tensorboard\n            writer.add_scalar(\"Train_acc\", train_acc, epoch)\n            writer.add_scalar(\"Valid_acc\", valid_acc, epoch)\n\n        writer.flush()\n        writer.close()\n        elapsed = (time.time() - start_time)/60\n        print(f'Time elapsed: {elapsed:.2f} min')\n        \n        if scheduler is not None:\n\n            if scheduler_on == 'valid_acc':\n                scheduler.step(valid_acc_list[-1])\n            elif scheduler_on == 'minibatch_loss':\n                scheduler.step(minibatch_loss_list[-1])\n            else:\n                raise ValueError(f'Invalid `scheduler_on` choice.')\n        \n\n    elapsed = (time.time() - start_time)/60\n    print(f'Total Training Time: {elapsed:.2f} min')\n    print('-'*50)\n    # print(f'Test accuracy {test_acc :.2f}%')\n\n    return minibatch_loss_list, train_acc_list, valid_acc_list","metadata":{"id":"b4M2zqibmJKO","execution":{"iopub.status.busy":"2022-01-17T03:12:05.532054Z","iopub.execute_input":"2022-01-17T03:12:05.53254Z","iopub.status.idle":"2022-01-17T03:12:05.548765Z","shell.execute_reply.started":"2022-01-17T03:12:05.532497Z","shell.execute_reply":"2022-01-17T03:12:05.547843Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def compute_accuracy(model, data_loader, device):\n\n    with torch.no_grad():\n\n        correct_pred, num_examples = 0, 0\n\n        for i, (features, targets) in enumerate(data_loader):\n\n            features = features.to(device)\n            targets = targets.float().to(device)\n\n            logits = model(features)\n            _, predicted_labels = torch.max(logits, 1)\n\n            num_examples += targets.size(0)\n            correct_pred += (predicted_labels == targets).sum()\n    return correct_pred.float()/num_examples * 100\n","metadata":{"id":"HTDRkf9bSG4V","execution":{"iopub.status.busy":"2022-01-16T23:49:26.733121Z","iopub.execute_input":"2022-01-16T23:49:26.733578Z","iopub.status.idle":"2022-01-16T23:49:26.739477Z","shell.execute_reply.started":"2022-01-16T23:49:26.73354Z","shell.execute_reply":"2022-01-16T23:49:26.738803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Load checkpoint","metadata":{"id":"5vEt6eBeSIzN"}},{"cell_type":"code","source":"def load_checkpoint(model, optimizer, model_path):\n    checkpoint = torch.load(PATH)\n    model.load_state_dict(checkpoint['model_state_dict'])\n    optimizer.load_state_dict(checkpoint['optimizer'])\n    epoch = checkpoint['epoch']\n    loss = checkpoint['loss']\n    print(\"Loading checkpoint: {}\".format(epoch))\n    print(\"Resume training !!!\")\n    # Resume training\n    model.train()\n    return model","metadata":{"id":"YVTWmGesSK3u","execution":{"iopub.status.busy":"2022-01-16T23:49:28.768471Z","iopub.execute_input":"2022-01-16T23:49:28.769163Z","iopub.status.idle":"2022-01-16T23:49:28.774384Z","shell.execute_reply.started":"2022-01-16T23:49:28.769113Z","shell.execute_reply":"2022-01-16T23:49:28.773588Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(dataloader_dict['train'])","metadata":{"id":"pOpef7y6xJrS","outputId":"b4954bbe-22f3-4678-9e23-1fb2894f3360","execution":{"iopub.status.busy":"2022-01-16T23:49:31.829555Z","iopub.execute_input":"2022-01-16T23:49:31.829833Z","iopub.status.idle":"2022-01-16T23:49:31.834924Z","shell.execute_reply.started":"2022-01-16T23:49:31.8298Z","shell.execute_reply":"2022-01-16T23:49:31.834238Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NUM_EPOCHS = 20","metadata":{"id":"csFLX-PlSPTx","execution":{"iopub.status.busy":"2022-01-16T23:49:34.962716Z","iopub.execute_input":"2022-01-16T23:49:34.963397Z","iopub.status.idle":"2022-01-16T23:49:34.967293Z","shell.execute_reply.started":"2022-01-16T23:49:34.963352Z","shell.execute_reply":"2022-01-16T23:49:34.966419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss().to(device)\nmodel = VGG16(n_classes=5).to(device)\noptimizer = optim.SGD(model.parameters(), lr=0.01, momentum=0.9)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer,\n                                                       factor=0.1,\n                                                       mode='max',\n                                                       verbose=True)","metadata":{"id":"WvxhPaRFuAs2","execution":{"iopub.status.busy":"2022-01-17T03:13:31.75142Z","iopub.execute_input":"2022-01-17T03:13:31.751691Z","iopub.status.idle":"2022-01-17T03:13:33.84668Z","shell.execute_reply.started":"2022-01-17T03:13:31.751654Z","shell.execute_reply":"2022-01-17T03:13:33.845891Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NUM_EPOCHS = 5\nminibatch_loss_list, train_acc_list, valid_acc_list = train_model(\n    model=model,\n    num_epochs=NUM_EPOCHS,\n    train_loader=dataloader_dict['train'],\n    valid_loader=dataloader_dict['test'],\n    optimizer=optimizer,\n    device=device,\n    scheduler=scheduler,\n    scheduler_on='valid_acc',\n    logging_interval=500)","metadata":{"execution":{"iopub.status.busy":"2022-01-17T03:13:38.731437Z","iopub.execute_input":"2022-01-17T03:13:38.732227Z","iopub.status.idle":"2022-01-17T03:24:51.532613Z","shell.execute_reply.started":"2022-01-17T03:13:38.732186Z","shell.execute_reply":"2022-01-17T03:24:51.530997Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"minibatch_loss_list, train_acc_list, valid_acc_list = train_model(\n    model=model,\n    num_epochs=NUM_EPOCHS,\n    train_loader=dataloader_dict['train'],\n    valid_loader=dataloader_dict['test'],\n    optimizer=optimizer,\n    device=device,\n    scheduler=scheduler,\n    scheduler_on='valid_acc',\n    logging_interval=500)","metadata":{"id":"qS-YUkz9xKwg","outputId":"aa2d4c0b-a3cd-4b42-c4e1-3a2a74167787","execution":{"iopub.status.busy":"2022-01-16T23:49:39.052807Z","iopub.execute_input":"2022-01-16T23:49:39.053073Z","iopub.status.idle":"2022-01-17T03:04:28.22594Z","shell.execute_reply.started":"2022-01-16T23:49:39.053044Z","shell.execute_reply":"2022-01-17T03:04:28.224103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_accuracy(train_acc_list, valid_acc_list, results_dir):\n\n    num_epochs = len(train_acc_list)\n\n    plt.plot(np.arange(1, num_epochs+1),\n             train_acc_list, label='Training')\n    plt.plot(np.arange(1, num_epochs+1),\n             valid_acc_list, label='Validation')\n\n    plt.xlabel('Epoch')\n    plt.ylabel('Accuracy')\n    plt.legend()\n\n    plt.tight_layout()\n\n    if results_dir is not None:\n        image_path = os.path.join(\n            results_dir, 'plot_acc_training_validation.pdf')\n        plt.savefig(image_path)","metadata":{"execution":{"iopub.status.busy":"2022-01-17T03:04:43.529732Z","iopub.execute_input":"2022-01-17T03:04:43.530002Z","iopub.status.idle":"2022-01-17T03:04:43.536107Z","shell.execute_reply.started":"2022-01-17T03:04:43.529968Z","shell.execute_reply":"2022-01-17T03:04:43.535135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%reload_ext tensorboard\n%tensorboard --logdir runs/vgg16/cassava","metadata":{"id":"UjANmm37SWCB","execution":{"iopub.status.busy":"2022-01-17T03:05:28.287169Z","iopub.execute_input":"2022-01-17T03:05:28.28806Z","iopub.status.idle":"2022-01-17T03:05:31.370245Z","shell.execute_reply.started":"2022-01-17T03:05:28.288004Z","shell.execute_reply":"2022-01-17T03:05:31.369443Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_accuracy(train_acc_list=train_acc_list,\n              valid_acc_list=valid_acc_list,\n              results_dir=None)\nplt.ylim([60, 100])\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-01-17T03:06:39.405118Z","iopub.execute_input":"2022-01-17T03:06:39.406253Z","iopub.status.idle":"2022-01-17T03:06:39.672955Z","shell.execute_reply.started":"2022-01-17T03:06:39.406206Z","shell.execute_reply":"2022-01-17T03:06:39.67223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Show examples","metadata":{"id":"jtpXOevuSbuS"}},{"cell_type":"code","source":"def show_examples(model, data_loader, unnormalizer=None, class_dict=None):\n    for batch_idx, (features, targets) in enumerate(data_loader):\n        with torch.no_grad():\n            features = features\n            targets = targets\n            logits = model(features)\n            predictions = torch.argmax(logits, dim=1)\n        break\n    \n    fig, axes = plt.subplots(nrows=3, ncols=5, sharex=True, sharey=True)\n    if unnormalizer is not None:\n        for idx in range(features.shape[0]):\n            features[idx] = unnormalizer(features[idx])\n    nhwc_img = features.permute(0, 2, 3, 1)\n    \n    if nhwc_img.shape[-1] == 1:\n        nhw_img = np.squeeze(nhwc_img.numpy(), axis=3)\n        \n        for idx, ax in enumerate(axes.ravel()):\n            ax.imshow(nhw_img[idx], cmap='binary')\n            if class_dict is not None:\n                ax.title.set_text(f'P: {class_dict[predictions[idx].item()]}'\n                                  f'\\nT: {class_dict[targets[idx].item()]}')\n            else:\n                ax.title.set_text(f'P: {predictions[idx]} | T: {targets[idx]}')\n            ax.axison = False\n    else:\n\n        for idx, ax in enumerate(axes.ravel()):\n            ax.imshow(nhwc_img[idx])\n            if class_dict is not None:\n                ax.title.set_text(f'P: {class_dict[predictions[idx].item()]}'\n                                  f'\\nT: {class_dict[targets[idx].item()]}')\n            else:\n                ax.title.set_text(f'P: {predictions[idx]} | T: {targets[idx]}')\n            ax.axison = False\n            \n    plt.tight_layout()\n    plt.show()","metadata":{"id":"4ZkKHXsayOOx","execution":{"iopub.status.busy":"2022-01-17T03:06:53.866703Z","iopub.execute_input":"2022-01-17T03:06:53.866984Z","iopub.status.idle":"2022-01-17T03:06:53.878651Z","shell.execute_reply.started":"2022-01-17T03:06:53.866951Z","shell.execute_reply":"2022-01-17T03:06:53.877672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Plot confusion matrix","metadata":{"id":"39cQKZIFSfnu"}},{"cell_type":"code","source":"from itertools import product","metadata":{"execution":{"iopub.status.busy":"2022-01-17T03:06:59.9832Z","iopub.execute_input":"2022-01-17T03:06:59.984039Z","iopub.status.idle":"2022-01-17T03:06:59.988174Z","shell.execute_reply.started":"2022-01-17T03:06:59.983994Z","shell.execute_reply":"2022-01-17T03:06:59.987194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_confusion_matrix(conf_mat,\n                          hide_spines=False,\n                          hide_ticks=False,\n                          figsize=None,\n                          cmap=None,\n                          colorbar=False,\n                          show_absolute=True,\n                          show_normed=False,\n                          class_names=None):\n\n    if not (show_absolute or show_normed):\n        raise AssertionError('Both show_absolute and show_normed are False')\n    if class_names is not None and len(class_names) != len(conf_mat):\n        raise AssertionError('len(class_names) should be equal to number of'\n                             'classes in the dataset')\n\n    total_samples = conf_mat.sum(axis=1)[:, np.newaxis]\n    normed_conf_mat = conf_mat.astype('float') / total_samples\n\n    fig, ax = plt.subplots(figsize=figsize)\n    ax.grid(False)\n    if cmap is None:\n        cmap = plt.cm.Blues\n\n    if figsize is None:\n        figsize = (len(conf_mat)*1.25, len(conf_mat)*1.25)\n\n    if show_normed:\n        matshow = ax.matshow(normed_conf_mat, cmap=cmap)\n    else:\n        matshow = ax.matshow(conf_mat, cmap=cmap)\n\n    if colorbar:\n        fig.colorbar(matshow)\n\n    for i in range(conf_mat.shape[0]):\n        for j in range(conf_mat.shape[1]):\n            cell_text = \"\"\n            if show_absolute:\n                cell_text += format(conf_mat[i, j], 'd')\n                if show_normed:\n                    cell_text += \"\\n\" + '('\n                    cell_text += format(normed_conf_mat[i, j], '.2f') + ')'\n            else:\n                cell_text += format(normed_conf_mat[i, j], '.2f')\n            ax.text(x=j,\n                    y=i,\n                    s=cell_text,\n                    va='center',\n                    ha='center',\n                    color=\"white\" if normed_conf_mat[i, j] > 0.5 else \"black\")\n    \n    if class_names is not None:\n        tick_marks = np.arange(len(class_names))\n        plt.xticks(tick_marks, class_names, rotation=90)\n        plt.yticks(tick_marks, class_names)\n        \n    if hide_spines:\n        ax.spines['right'].set_visible(False)\n        ax.spines['top'].set_visible(False)\n        ax.spines['left'].set_visible(False)\n        ax.spines['bottom'].set_visible(False)\n    ax.yaxis.set_ticks_position('left')\n    ax.xaxis.set_ticks_position('bottom')\n    if hide_ticks:\n        ax.axes.get_yaxis().set_ticks([])\n        ax.axes.get_xaxis().set_ticks([])\n\n    plt.xlabel('Predicted label')\n    plt.ylabel('True label')\n    return fig, ax","metadata":{"id":"ptzXuxsvSjDr","execution":{"iopub.status.busy":"2022-01-17T03:07:01.246601Z","iopub.execute_input":"2022-01-17T03:07:01.246907Z","iopub.status.idle":"2022-01-17T03:07:01.262481Z","shell.execute_reply.started":"2022-01-17T03:07:01.246872Z","shell.execute_reply":"2022-01-17T03:07:01.261622Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Compute confusion matrix","metadata":{"id":"AD1Bx3YNSlML"}},{"cell_type":"code","source":"def compute_confusion_matrix(model, data_loader, device):\n    all_targets, all_predictions = [], []\n    with torch.no_grad():\n        for i, (features, targets) in enumerate(data_loader):\n            features = features.to(device)\n            targets = targets\n            logits = model(features)\n            _, predicted_labels = torch.max(logits, 1)\n            all_targets.extend(targets.to('cpu'))\n            all_predictions.extend(predicted_labels.to('cpu'))\n            \n    all_predictions = all_predictions\n    all_predictions = np.array(all_predictions)\n    all_targets = np.array(all_targets)\n    \n    class_labels = np.unique(np.concatenate((all_targets, all_predictions)))\n    if class_labels.shape[0] == 1:\n        if class_labels[0] != 0:\n            class_labels = np.array([0, class_labels[0]])\n        else:\n            class_labels = np.array([class_labels[0], 1])\n            \n    n_labels = class_labels.shape[0]\n    lst = []\n    z = list(zip(all_targets, all_predictions))\n    for combi in product(class_labels, repeat=2):\n        lst.append(z.count(combi))\n    mat = np.asarray(lst)[:, None].reshape(n_labels, n_labels)\n    return mat","metadata":{"id":"J9pHrl0ZSpTj","execution":{"iopub.status.busy":"2022-01-17T03:07:03.297264Z","iopub.execute_input":"2022-01-17T03:07:03.299109Z","iopub.status.idle":"2022-01-17T03:07:03.315819Z","shell.execute_reply.started":"2022-01-17T03:07:03.297763Z","shell.execute_reply":"2022-01-17T03:07:03.314937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Unnormalized data","metadata":{"id":"EHUt4tBxSrY-"}},{"cell_type":"code","source":"class UnNormalize(object):\n    def __init__(self, mean, std):\n        self.mean = mean\n        self.std = std\n\n    def __call__(self, tensor):\n        \"\"\"\n        Parameters:\n        ------------\n        tensor (Tensor): Tensor image of size (C, H, W) to be normalized.\n        \n        Returns:\n        ------------\n        Tensor: Normalized image.\n        \"\"\"\n        for t, m, s in zip(tensor, self.mean, self.std):\n            t.mul_(s).add_(m)\n        return tensor\n","metadata":{"id":"_cmlCfuySt9l","execution":{"iopub.status.busy":"2022-01-17T03:07:05.007586Z","iopub.execute_input":"2022-01-17T03:07:05.00787Z","iopub.status.idle":"2022-01-17T03:07:05.014035Z","shell.execute_reply.started":"2022-01-17T03:07:05.007838Z","shell.execute_reply":"2022-01-17T03:07:05.012465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.cpu()\nunnormalizer = UnNormalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))\n\n\nshow_examples(model=model, data_loader=dataloader_dict['test'], unnormalizer=unnormalizer, class_dict=idx_to_class)","metadata":{"id":"IdQOy65aSzRM","execution":{"iopub.status.busy":"2022-01-17T03:25:04.811859Z","iopub.execute_input":"2022-01-17T03:25:04.81231Z","iopub.status.idle":"2022-01-17T03:25:15.309422Z","shell.execute_reply.started":"2022-01-17T03:25:04.812275Z","shell.execute_reply":"2022-01-17T03:25:15.308655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Load model","metadata":{"id":"wmnFnYgjSvuW"}},{"cell_type":"code","source":"def load_model(model, model_path):\n    print(\"Loading model\")\n    load_weights = torch.load(model_path,  map_location={\"cuda:0\": \"cpu\"})\n    model.load_state_dict(load_weights)\n    print(\"Loading successful!\")\n    return model","metadata":{"id":"bhCtuc8NS3f2","execution":{"iopub.status.busy":"2022-01-17T03:10:19.491973Z","iopub.execute_input":"2022-01-17T03:10:19.492742Z","iopub.status.idle":"2022-01-17T03:10:19.498011Z","shell.execute_reply.started":"2022-01-17T03:10:19.492701Z","shell.execute_reply":"2022-01-17T03:10:19.497169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Predict","metadata":{"id":"bS3nIoGIS5Ow"}},{"cell_type":"code","source":"class Predictor():\n    def __init__(self, class_index):\n        self.class_index = class_index\n        \n    def predict_max(self, output):\n        max_id = np.argmax(output.detach().numpy())\n        predicted_label = self.class_index[max_id]\n        return predicted_label","metadata":{"id":"-dj2NtdKS7NE","execution":{"iopub.status.busy":"2022-01-17T03:07:40.702369Z","iopub.execute_input":"2022-01-17T03:07:40.703052Z","iopub.status.idle":"2022-01-17T03:07:40.707833Z","shell.execute_reply.started":"2022-01-17T03:07:40.70301Z","shell.execute_reply":"2022-01-17T03:07:40.707103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictor = Predictor(classes)","metadata":{"execution":{"iopub.status.busy":"2022-01-17T03:07:40.738816Z","iopub.execute_input":"2022-01-17T03:07:40.739235Z","iopub.status.idle":"2022-01-17T03:07:40.743592Z","shell.execute_reply.started":"2022-01-17T03:07:40.739203Z","shell.execute_reply":"2022-01-17T03:07:40.742437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict(img):\n    model.eval()\n    \n    # model = load_model(model, 'model.pth')\n    \n\n    transform = ImageTransform(input_size)\n    img = transform(img, phase=\"test\")\n    img = img.unsqueeze(0)\n    \n    output = model(img)\n    response = predictor.predict_max(output)\n    \n    return response","metadata":{"id":"-_Yi13IoS9FP","execution":{"iopub.status.busy":"2022-01-17T03:07:40.888159Z","iopub.execute_input":"2022-01-17T03:07:40.888462Z","iopub.status.idle":"2022-01-17T03:07:40.893102Z","shell.execute_reply.started":"2022-01-17T03:07:40.88843Z","shell.execute_reply":"2022-01-17T03:07:40.892334Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_img = Image.open('../input/cassava-leaf-disease-classification/test_images/2216849948.jpg').convert('RGB')\nplt.imshow(test_img)\nplt.show()\nprint('\\tPredicted image:' + \" \" + predict(test_img))","metadata":{"id":"6ftfyoJNS_Ua","execution":{"iopub.status.busy":"2022-01-17T03:24:51.5386Z","iopub.status.idle":"2022-01-17T03:24:51.54004Z","shell.execute_reply.started":"2022-01-17T03:24:51.539795Z","shell.execute_reply":"2022-01-17T03:24:51.539825Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Training with LeNet5","metadata":{}},{"cell_type":"code","source":"net = LeNet(n_classes=5)\nnet.conv1 = nn.Sequential(nn.Conv2d(3, 6, 5, 1))","metadata":{"execution":{"iopub.status.busy":"2022-01-16T12:19:12.245527Z","iopub.execute_input":"2022-01-16T12:19:12.246239Z","iopub.status.idle":"2022-01-16T12:19:12.277109Z","shell.execute_reply.started":"2022-01-16T12:19:12.246198Z","shell.execute_reply":"2022-01-16T12:19:12.276353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"INPUT_SIZE = 32","metadata":{"execution":{"iopub.status.busy":"2022-01-16T12:19:14.417207Z","iopub.execute_input":"2022-01-16T12:19:14.417459Z","iopub.status.idle":"2022-01-16T12:19:14.421061Z","shell.execute_reply.started":"2022-01-16T12:19:14.417431Z","shell.execute_reply":"2022-01-16T12:19:14.420222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds = Cassava(df_train, transform=ImageTransform(INPUT_SIZE), phase='train')\nvalid_ds = Cassava(df_valid, transform=ImageTransform(INPUT_SIZE), phase='test')","metadata":{"execution":{"iopub.status.busy":"2022-01-16T12:19:17.392829Z","iopub.execute_input":"2022-01-16T12:19:17.393308Z","iopub.status.idle":"2022-01-16T12:19:17.398711Z","shell.execute_reply.started":"2022-01-16T12:19:17.39327Z","shell.execute_reply":"2022-01-16T12:19:17.397455Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 32\n\ntrain_dataloader = DataLoader(train_ds, batch_size,shuffle=True)\nvalid_dataloader = DataLoader(valid_ds, batch_size, shuffle=False)\n\ndataloader_dict = {\"train\": train_dataloader, 'test': valid_dataloader}","metadata":{"execution":{"iopub.status.busy":"2022-01-16T12:19:18.725547Z","iopub.execute_input":"2022-01-16T12:19:18.726245Z","iopub.status.idle":"2022-01-16T12:19:18.731165Z","shell.execute_reply.started":"2022-01-16T12:19:18.726207Z","shell.execute_reply":"2022-01-16T12:19:18.7301Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"net.to(device)\ncriterion = nn.CrossEntropyLoss().to(device)\noptimizer = optim.SGD(model.parameters(), lr=0.001, momentum=0.9)","metadata":{"execution":{"iopub.status.busy":"2022-01-16T12:19:21.236377Z","iopub.execute_input":"2022-01-16T12:19:21.236738Z","iopub.status.idle":"2022-01-16T12:19:22.617151Z","shell.execute_reply.started":"2022-01-16T12:19:21.236702Z","shell.execute_reply":"2022-01-16T12:19:22.616364Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"net","metadata":{"execution":{"iopub.status.busy":"2022-01-16T12:19:23.901397Z","iopub.execute_input":"2022-01-16T12:19:23.901963Z","iopub.status.idle":"2022-01-16T12:19:23.907164Z","shell.execute_reply.started":"2022-01-16T12:19:23.901925Z","shell.execute_reply":"2022-01-16T12:19:23.906464Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"writer = SummaryWriter('runs/lenet/cassava')","metadata":{"execution":{"iopub.status.busy":"2022-01-16T12:19:27.195226Z","iopub.execute_input":"2022-01-16T12:19:27.195792Z","iopub.status.idle":"2022-01-16T12:19:27.200793Z","shell.execute_reply.started":"2022-01-16T12:19:27.195746Z","shell.execute_reply":"2022-01-16T12:19:27.200036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NUM_EPOCHS = 30","metadata":{"execution":{"iopub.status.busy":"2022-01-16T12:19:29.334318Z","iopub.execute_input":"2022-01-16T12:19:29.335091Z","iopub.status.idle":"2022-01-16T12:19:29.339384Z","shell.execute_reply.started":"2022-01-16T12:19:29.335039Z","shell.execute_reply":"2022-01-16T12:19:29.338456Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"minibatch_loss_list, train_acc_list, valid_acc_list = train_model(\n    model=net,\n    num_epochs=NUM_EPOCHS,\n    train_loader=dataloader_dict['train'],\n    valid_loader=dataloader_dict['test'],\n    optimizer=optimizer,\n    device=device,\n    scheduler=scheduler,\n    scheduler_on='valid_acc',\n    logging_interval=500)","metadata":{"execution":{"iopub.status.busy":"2022-01-16T12:25:16.847239Z","iopub.execute_input":"2022-01-16T12:25:16.847862Z","iopub.status.idle":"2022-01-16T12:59:45.428004Z","shell.execute_reply.started":"2022-01-16T12:25:16.847823Z","shell.execute_reply":"2022-01-16T12:59:45.426692Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}