{"cells":[{"metadata":{"_uuid":"152add5c-c486-48c5-bfb4-5ff621e34390","_cell_guid":"a689793b-4f18-4708-ab12-2d7e966ea095","trusted":true},"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n# for dirname, _, filenames in os.walk('/kaggle/input/'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"dae53868-1ad0-4457-b0f5-a1b4da3878b4","_cell_guid":"95cbbd13-9e0d-4efa-8e62-0bc30ddb24d4","trusted":true},"cell_type":"code","source":"!pip install -U albumentations\n!pip install torchsummary\n# !pip install torch-lr-finder","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d5891d74-6611-4622-abf6-84b8e6d1bdb9","_cell_guid":"16f89a6e-5ee9-480b-97f4-540db90a42f6","trusted":true},"cell_type":"code","source":"#imports\nimport torch\nimport torch.nn as nn\nimport torchvision\nfrom torchvision import models,transforms\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.hub\nfrom torch.cuda import amp\nimport torch.nn.functional as F\nimport time\n# from torch_lr_finder import LRFinder\n\nfrom collections import Counter\nimport matplotlib.gridspec as gridspec\nimport cv2\nfrom torchsummary import summary\nfrom sklearn.metrics import precision_score,f1_score,recall_score\n\n\nfrom torch.optim.lr_scheduler import MultiStepLR\nimport math\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom collections import Counter\nfrom pylab import rcParams\nimport json\nfrom sklearn.model_selection import train_test_split\nfrom albumentations import (Compose,OneOf,RandomBrightnessContrast,\n                            RandomGamma,ShiftScaleRotate,HorizontalFlip,\n                            Rotate,FancyPCA,RandomCrop,RandomBrightnessContrast,ToGray,\n                            MultiplicativeNoise,Resize,VerticalFlip,Normalize,\n                               ChannelShuffle )\nrcParams['figure.figsize'] = 5, 10","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"805d8350-2467-42bc-866e-8de1add583de","_cell_guid":"b9f65249-4aef-48a4-adc1-66832fcdfef0","trusted":true},"cell_type":"markdown","source":"## PATHS"},{"metadata":{"_uuid":"c930550f-5c6f-4c14-aaee-b6bae1f05968","_cell_guid":"2d1dc3f0-d923-4020-b5b7-2fa42b83635f","trusted":true},"cell_type":"code","source":"data_path = \"/kaggle/input/cassava-leaf-disease-classification/train_images/\"","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"a3e79052-190a-46ce-b892-dba39b8acf2e","_cell_guid":"e5ca3f0d-d80a-45fd-a5d8-081c300bf96f","trusted":true},"cell_type":"code","source":"csv_path = \"/kaggle/input/cassava-leaf-disease-classification/train.csv\"\ndf = pd.read_csv(csv_path)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"b91ae7b7-ee34-430b-9b23-51c1280f5d77","_cell_guid":"b6e6bf63-d4ba-408e-b64b-dd59add5a256","trusted":true},"cell_type":"markdown","source":"## Training data visualization"},{"metadata":{"_uuid":"27509ba1-c96d-4e7f-a3d9-d34c04a40750","_cell_guid":"01452f3f-e676-4d33-8571-5c18e680af1c","trusted":true},"cell_type":"code","source":"f = open('/kaggle/input/cassava-leaf-disease-classification/label_num_to_disease_map.json')\nidx_to_cls_mapping = json.load(f)\nidx_to_cls_mapping","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"1f9df1f0-ad91-4b4f-ad1b-f24e9fc3f0dc","_cell_guid":"8337ece5-eaa6-439b-b3c8-61936f269045","trusted":true},"cell_type":"code","source":"df.label.value_counts(ascending=False,sort=False).plot(kind='bar',figsize=(10,10))\nCounter(df.label)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"2e45ca97-8731-49d8-a814-e73e92acfcfc","_cell_guid":"dccf867d-6446-49e6-936d-95cfb72312ff","trusted":true},"cell_type":"markdown","source":"## Visualize Images in a Grid"},{"metadata":{"_uuid":"4deb2412-b5c5-4e52-89f5-7c589ac5097e","_cell_guid":"5120c477-92e7-41a7-8bae-55b98a3c10ce","trusted":true},"cell_type":"code","source":"label0_imageids = (df.image_id[df.label==0]).tolist()\nlabel1_imageids = (df.image_id[df.label==1]).tolist()\nlabel2_imageids = (df.image_id[df.label==2]).tolist()\nlabel3_imageids = (df.image_id[df.label==3]).tolist()\nlabel4_imageids = (df.image_id[df.label==4]).tolist()","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"b4fa66ab-795a-4e0e-a58e-eb0994b88024","_cell_guid":"f1353d87-9fb0-459b-9ced-7396a7f2b4a5","trusted":true},"cell_type":"code","source":"def get_images(image_id_list,label,datapath=data_path):\n    plt.figure(figsize=(20,20))\n    for i in range(25):\n        image = Image.open(data_path+image_id_list[i])\n        plt.subplot(5,5,i+1)\n        plt.text(400,0,s=(label,image_id_list[i]))\n        plt.imshow(image)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"06aa2895-fe32-4eaa-92af-c414a54ce0be","_cell_guid":"49f7765f-f657-40d8-a34d-b8fd76640e59","trusted":true},"cell_type":"markdown","source":"## Class 0 images - Cassava Bacterial Blight (CBB)"},{"metadata":{"_uuid":"63901e1c-ca9d-4f56-8e49-66ac2e2357ed","_cell_guid":"92484c6d-5310-4de9-a040-41fb0ec672f5","trusted":true},"cell_type":"code","source":"get_images(label0_imageids,0)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"0b0ce425-efa5-4b10-897e-c61fe33b0f63","_cell_guid":"316d2078-d7fe-4dc2-b6d7-34960941b7f9","trusted":true},"cell_type":"markdown","source":"## Class 1 images - Cassava Brown Streak Disease (CBSD)"},{"metadata":{"_uuid":"b913ea69-fabf-4a11-b389-b32bd8cbe3b8","_cell_guid":"d94fa658-8825-408c-8392-2adbb69a167e","trusted":true},"cell_type":"code","source":"get_images(label1_imageids,1)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"4d9d7946-f505-417a-ab93-2ac0652b2f2b","_cell_guid":"caf07e07-ee68-4eed-be49-2cd6a77256b0","trusted":true},"cell_type":"markdown","source":"## Class 2 images - Cassava Green Mottle (CGM)"},{"metadata":{"_uuid":"19e6f1a2-40b9-4036-9d60-120281e33f7d","_cell_guid":"13f67c98-d0cc-41aa-862a-48330ee3e0d1","trusted":true},"cell_type":"code","source":"get_images(label2_imageids,2)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"77d9bc6c-521d-40ce-afa8-bb25413a63c1","_cell_guid":"92480920-9841-48f4-9569-7a7459022997","trusted":true},"cell_type":"markdown","source":"## Class 3 images - Cassava Mosaic Disease (CMD)"},{"metadata":{"_uuid":"5c60bfb0-b319-4a59-b462-1970c0ad7f86","_cell_guid":"fb773a0e-495d-465c-b1f2-d34a9e68b77e","trusted":true},"cell_type":"code","source":"get_images(label3_imageids,3)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"6fbd869d-e316-4315-a72e-5f4014986fb0","_cell_guid":"5aa982d4-72b2-4b1f-b090-b975cfaaa2cc","trusted":true},"cell_type":"markdown","source":"## Class 4 images - Healthy"},{"metadata":{"_uuid":"42f7cb19-5260-46d5-b8dc-2a83c62d3881","_cell_guid":"c6ee328e-6251-4ef4-a08b-81b594626305","trusted":true},"cell_type":"code","source":"get_images(label4_imageids,4)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"28ac1f45-9d45-4c3f-a51a-2b253f7aad1a","_cell_guid":"177d2864-e0d7-405a-8fd4-6176b9adc73e","trusted":true},"cell_type":"markdown","source":"### Data split \n1. Using stratified splitting"},{"metadata":{"_uuid":"c2338d16-97ab-4174-b0dd-70599f7837a2","_cell_guid":"d726af49-c4da-4dc0-b7ca-b133c72e3c0e","trusted":true},"cell_type":"code","source":"Xtrain,Xval,ytrain,yval = train_test_split((df.image_id).tolist(),(df.label).tolist(),shuffle=True,random_state=42,stratify=df.label,test_size=0.33)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"a6c18984-ef8f-403d-b71c-4857fc12fd7d","_cell_guid":"7814c376-3434-457f-83d7-d13fd2802540","trusted":true},"cell_type":"code","source":"assert len(Xtrain) == len(ytrain)\nassert len(Xval) == len(yval)\nprint(f'[INFO] Training images - {len(Xtrain)}')\nprint(f'[INFO] Validation images - {len(Xval)}')","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"c673e6e6-fc42-443c-9ead-d2b70b40e3fe","_cell_guid":"d9095b65-42cf-40bc-bba6-89bf086cf045","trusted":true},"cell_type":"markdown","source":"# **Visualize the plots**"},{"metadata":{"_uuid":"ed7b33f7-2ac3-4765-8945-44edf7b4fbd3","_cell_guid":"b1944865-16f1-4d90-85dd-c0c978cfda78","trusted":true},"cell_type":"code","source":"def plot(splots,title='Training Data'):\n    for p in splots.patches:\n        splots.annotate(format(p.get_height(),'.1f'),\n                           (p.get_x() + p.get_width()/2. , p.get_height()),\n                           ha= 'center', va = 'center',\n                           xytext = (0,9),\n                           textcoords = 'offset points')\n\n    plt.xlabel('Classes');\n    plt.ylabel('Count');\n    plt.title(title);","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"a47f6cdf-0b91-4f29-a2d5-775087250473","_cell_guid":"68d3cc87-8ce9-47d8-b043-53a0e82720e8","trusted":true},"cell_type":"code","source":"#training images\ntraining_counts = Counter(ytrain)\ntrain_plot = sns.barplot(x=list(training_counts.keys()),y=list(training_counts.values()))\n\nfor p in train_plot.patches:\n    train_plot.annotate(format(p.get_height(),'.1f'),\n                       (p.get_x() + p.get_width()/2. , p.get_height()),\n                       ha= 'center', va = 'center',\n                       xytext = (0,9),\n                       textcoords = 'offset points')\n\nplt.xlabel('Classes');\nplt.ylabel('Count');\nplt.title('Training Data');","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"5a16c6b2-8da3-494d-8ef0-b176008dc51a","_cell_guid":"c6a8115a-f3de-4647-a96c-e3801c07c0b7","trusted":true},"cell_type":"code","source":"#validation images\nvalidation_counts = Counter(yval)\nvalidation_plot = sns.barplot(x=list(validation_counts.keys()),y=list(validation_counts.values())) \n\nplots = [train_plot,validation_plot]\n\nfor p in validation_plot.patches:\n    validation_plot.annotate(format(p.get_height(),'.1f'),\n                       (p.get_x() + p.get_width()/2. , p.get_height()),\n                       ha= 'center', va = 'center',\n                       xytext = (0,9),\n                       textcoords = 'offset points')\n\nplt.xlabel('Classes');\nplt.ylabel('Count');\nplt.title('Validation Data');","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"074c5b36-18ad-4f68-a47d-a5f86e5d3187","_cell_guid":"d610be2c-fa92-42ff-b5c4-73ec8141f839","trusted":true},"cell_type":"markdown","source":"# Generating Train and validation excel sheets"},{"metadata":{"_uuid":"99571505-3436-46a9-83ee-b69e4aed9816","_cell_guid":"1ceabcf8-5875-4bc4-b9f0-54d51a639239","trusted":true},"cell_type":"code","source":"#generate train and validation excel sheet\nfor i in range(len(Xtrain)):\n    Xtrain[i] = data_path + Xtrain[i]","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"cbecbe28-529f-487e-8075-23026c103fea","_cell_guid":"9ad31c62-1c64-49b9-bdc8-75da5e2e1141","trusted":true},"cell_type":"code","source":"for i in range(len(Xval)):\n    Xval[i] = data_path + Xval[i]","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"484a0aac-ea9d-42c5-93a2-5f63ec64f976","_cell_guid":"97c59bee-aa50-46a4-b95b-687ef6a1148d","trusted":true},"cell_type":"code","source":"#create the dictionary of training and validation images and save them as excel sheet\ndef create_dict(X,y):\n    assert len(X) == len(y)\n    \n    xy_dict = {}\n    for i in range(len(X)):\n        xy_dict[X[i]] = y[i]\n    return xy_dict\n\ndef save_excel(dicts,name,save_in,idx_to_cls_mapping):\n    df = pd.DataFrame.from_dict(data = {'Path': [i for i in dicts.keys()],\n                                       'Label': [i for i in dicts.values()]},\n                                        orient = 'columns')\n    df['Name'] = pd.Series(data=[idx_to_cls_mapping[str(i)] for i in train_dict.values()])\n    \n    writer = pd.ExcelWriter(save_in+name+'.xlsx')\n    df.to_excel(excel_writer=writer,sheet_name=str(name),index=False)\n    writer.save()","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"49b3a64c-2997-4067-b9b5-391c4f74222c","_cell_guid":"685029e7-d357-411c-bbd4-81737ceeb482","trusted":true},"cell_type":"code","source":"train_dict = create_dict(Xtrain,ytrain)\nval_dict = create_dict(Xval,yval)\n\nf = open('/kaggle/input/cassava-leaf-disease-classification/label_num_to_disease_map.json')\nmappings = json.load(f)\n\nsave_excel(train_dict,name = 'train',save_in='/kaggle/working/',idx_to_cls_mapping = mappings)\nsave_excel(val_dict,name = 'val', save_in='/kaggle/working/',idx_to_cls_mapping = mappings)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"6954ca6d-d1ed-4bee-b21e-b3336bfd5de2","_cell_guid":"1aaa5a01-5ec7-4b3d-8008-3d27530bd8c4","trusted":true},"cell_type":"markdown","source":"# Configurations"},{"metadata":{"_uuid":"d9a84bd4-ca8f-4fce-9059-be5dc9b87762","_cell_guid":"4d68b286-ea40-46a7-9949-b07b7bd8dc4f","trusted":true},"cell_type":"code","source":"class Config():\n    def __init__(self,lr,checkpoint_name,epochs,w_decay,classify,train_batch_size,valid_batch_size):\n        self.DEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n        self.LR = lr\n        self.TRAIN_BATCH = train_batch_size\n        self.VALID_BATCH = valid_batch_size\n        \n        #model related\n        self.CHECKPOINT_NAME = checkpoint_name\n        self.MODEL_SAVE_PATH = '/kaggle/working/'\n        self.EPOCHS = epochs\n        self.W_DECAY = w_decay\n        self.CLASSIFY = classify","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"6d9b76e8-7c62-49a9-8ff7-a4f923c39fb9","_cell_guid":"8a83fb9f-f03f-4ebf-ac5e-f69bcfd580ac","trusted":true},"cell_type":"markdown","source":"## Augmentations Visualization"},{"metadata":{"_uuid":"86ed2aa6-9427-41f0-a237-1d4b122bdaaf","_cell_guid":"090eb278-e13e-4bf9-ac05-140cf31a4e08","trusted":true},"cell_type":"code","source":"aug = Compose([\n               HorizontalFlip(p=0.8),\n               VerticalFlip(p=0.8),\n               ShiftScaleRotate(shift_limit=(0.02),\n                                scale_limit=(0.2),\n                                rotate_limit=30,p=0.8),\n                RandomCrop(224,224,p=0.8),\n                MultiplicativeNoise(multiplier=[0.8, 1], elementwise=True, per_channel=True, p=0.8),\n                    ])","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"3fc6038a-122e-4af9-a512-6d68ab10d087","_cell_guid":"c1ed209f-8fee-4783-b0bd-5eb3100796d2","trusted":true},"cell_type":"code","source":"df = pd.read_excel('/kaggle/input/casava/train.xlsx')\npath = np.array(df.Path)\nlabels = np.array(df.Label)\nname = np.array(df.Name)\n\nimages = []\nclasses = []\nfor i in np.random.randint(14000,size=32):\n    image = Image.open(path[i]).convert('RGB')\n    image = image.resize(size=(224,224))\n    image = np.array(image,np.uint8)\n    aug_image = aug(image=image)['image']\n\n    images.append(image)\n    classes.append(labels[i])\n\n    images.append(aug_image)\n    classes.append(labels[i])\n    \n    \n# gs1 = gridspec.GridSpec(4,8)\n# gs1.update(wspace=0.5, hspace=1) # set the spacing between axes. \n# plt.figure(figsize=(20,20))\n\n# for i in range(32):\n#     ax = plt.subplot(gs1[i])\n#     ax.imshow(images[i])\n    \n# print(images[0].min(),images[0].max())\n# print(aug_image.min(),aug_image.max())\nfig, axes = plt.subplots(nrows=8, ncols=8, figsize = (21,16))\n\nfor i in range(64):\n    ax = axes.flat[i]\n    ax.imshow(images[i])\n    ax.set_title(f'Class - {classes[i]}')\n    ax.set_xticks([])\n    ax.set_yticks([])","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"e7ec9a0f-bf10-4d6a-865a-b0b9e0af7a2d","_cell_guid":"705b328b-01d1-40db-9304-fb82037e7ffa","trusted":true},"cell_type":"code","source":"# #training loader\n# test_train_dataset = CasavaDataset(excel_path='/kaggle/input/casava/train.xlsx',augmentations=aug)\n# test_valid_dataset = CasavaDataset(excel_path='/kaggle/input/casava/val.xlsx',augmentations=None)\n\n# #train dataloader\n# test_train_loader = torch.utils.data.DataLoader(test_train_dataset,32,True)\n\n# #validation dataloader\n# test_valid_loader = torch.utils.data.DataLoader(test_valid_dataset,16,True)\n\n# im,lab,path = next(iter(test_train_loader))\n# print(f'[INFO] Training image details - Shape - {im.shape} min,max - {im.min(),im.max()}')\n\n# im_val, lab_val, path_val = next(iter(test_valid_loader))\n# print(f'[INFO] Validation image details - Shape - {im_val.shape} min,max - {im_val.min(),im_val.max()}')","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"4cb9dcfc-52ac-4533-8806-f1099ef6d203","_cell_guid":"c730df2c-1956-4e61-b271-1fbffe2041cb","trusted":true},"cell_type":"markdown","source":"# Model Exploration"},{"metadata":{"_uuid":"261f208e-ea49-4f82-b30a-f8e9d7d71515","_cell_guid":"27ca2c68-7a7a-43ac-b3f8-20dcd5bb3ae0","trusted":true},"cell_type":"code","source":"#SeNet\n# model = torch.hub.load(\n#     'moskomule/senet.pytorch',\n#     'se_resnet50',\n#     pretrained=True,)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"0a5cbb60-b1c8-46f4-8aae-1aa0f25c33c7","_cell_guid":"d4fa56bc-8d8c-4fc8-b6f3-2f8ea7309303","trusted":true},"cell_type":"code","source":"model = torchvision.models.resnet50(pretrained=True)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"22aa5ac8-7d59-4fa5-96fd-ba166ed56cfb","_cell_guid":"3d57c541-a378-43bf-8a7f-c6aa8b0aeb03","trusted":true,"collapsed":true},"cell_type":"code","source":"model","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"cc66429f-3303-4b4d-8195-b52b8befd347","_cell_guid":"7be1f0e6-847f-475c-b828-ab2da6d4ea0f","trusted":true},"cell_type":"code","source":"model.fc = nn.Sequential(\n                         nn.Dropout(p=0.8),\n                         nn.Linear(2048, 512,bias=False),\n                         nn.BatchNorm1d(512))","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"b75c4132-0c8c-44c5-bef5-d53fde53b646","_cell_guid":"1b0895ad-3bbd-4d9a-86f1-c5638838253f","trusted":true},"cell_type":"code","source":"class Modified_Model(nn.Module):\n    def __init__(self,mymodel,num_classes,classify=False):\n        super(Modified_Model,self).__init__()\n        self.model = mymodel\n        self.logits = nn.Linear(512,num_classes)\n        self.classify = classify\n    \n    def forward(self,x):\n        x = self.model(x)\n        \n        if self.classify:\n            x = self.logits(x)     #return class score\n            return x\n        else:\n            return x        #return embeddings","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"6c76b067-954f-4790-8e29-cae0ed521fe0","_cell_guid":"88968839-bff0-49fd-a1eb-6d5dd4830786","trusted":true},"cell_type":"code","source":"mod = Modified_Model(model,num_classes=10,classify=True)\nfor p in mod.parameters():\n    p.requires_grad =False","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"bf812b63-79c8-45fa-903b-c575dc91e78b","_cell_guid":"d590c967-5597-4aa6-a735-eb162ead73de","trusted":true},"cell_type":"code","source":"mod","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"ef22f95e-d109-45a6-abf8-c80e5b204493","_cell_guid":"b36a4810-ba7f-4a8e-9f21-c4c7b59ee70a","trusted":true},"cell_type":"code","source":"#unfreeze few layers\nfor p in mod.logits.parameters():\n    p.requires_grad = True\n    \n#fc layer\nfor p in mod.model.fc.parameters():\n    p.requires_grad = True\n    \n#last SE layer\n# for p in mod.model.layer4[2].parameters():\n#     p.requires_grad = False","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"b84de9d5-77c3-49d5-ac65-565f66294859","_cell_guid":"994a69db-07f7-45a2-b931-e2bb881deebf","trusted":true},"cell_type":"code","source":"for p in mod.parameters():\n    print(p.requires_grad)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"7d767846-080e-4d2e-88cd-2cf6cab43174","_cell_guid":"3746ca79-b6a6-439c-b9cb-64b4802d073d","trusted":true},"cell_type":"code","source":"summary(mod.cpu(),(3,160,160),32,device='cpu')","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"f63e3e40-abe9-4526-bd97-16cb0c73e0a9","_cell_guid":"2e05c21a-33b8-4dcd-aca1-ac65770e4563","trusted":true},"cell_type":"markdown","source":"# Optimizer + MultiStep LR"},{"metadata":{"_uuid":"ff73d744-298c-4d8d-8422-96b57fe4cd08","_cell_guid":"834029aa-1bbd-4e02-a1e0-b1a92697d6c2","trusted":true},"cell_type":"code","source":"def separate_bn(model):\n    all_params = model.parameters()\n    paras_only_bn = []\n    for pname,p in model.named_parameters():\n        if pname.find('bn') >= 0:\n            paras_only_bn.append(p)\n    paras_only_bn_id = list(map(id,paras_only_bn))\n    paras_wo_bn = list(filter(lambda p: id(p) not in paras_only_bn_id,all_params))\n    return paras_only_bn, paras_wo_bn","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d70da514-b11b-4c85-88ff-7dcfdfe31624","_cell_guid":"12c27433-c934-462d-9e95-3554b8671e1e","trusted":true},"cell_type":"code","source":"#separate bn and non bn layers\nwith_bn,without_bn = separate_bn(mod)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"211f69fd-7e65-4b17-9a77-8b5b5f947e30","_cell_guid":"6b5fa207-1811-42da-86fe-3c0828add609","trusted":true},"cell_type":"markdown","source":"# Loss Function"},{"metadata":{"_uuid":"beb08b30-3cc2-41ad-9e1b-355d3753e019","_cell_guid":"c1bcfbda-4c41-415b-b67b-6e0e8d2f14f3","trusted":true},"cell_type":"code","source":"class FocalLoss(nn.Module):\n    def __init__(self,alpha= 1,gamma= 1,reduce= True,cls_weights=None):\n        super(FocalLoss,self).__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.reduce = reduce\n        self.cls_weights = cls_weights\n        \n    def forward(self,logits,labels):\n        CE_loss = F.cross_entropy(logits,labels,weight=self.cls_weights,reduction = \"none\")\n        \n        pt = torch.exp(-CE_loss)\n        F_loss = self.alpha * (1-pt)**self.gamma * CE_loss\n\n        if self.reduce:\n            return torch.mean(F_loss)\n        else:\n            return F_loss","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"98c94082-3af7-4770-9c42-189150260e0e","_cell_guid":"8fffff2b-7543-4e3f-8e49-6f43021333f2","trusted":true},"cell_type":"markdown","source":"# Dataloader"},{"metadata":{"_uuid":"d59326a1-7081-4911-a048-e83504a6294d","_cell_guid":"aa088c64-0952-4445-a6d4-72cfabdc523a","trusted":true},"cell_type":"code","source":"#dataloader\nclass CasavaDataset(Dataset):\n    def __init__(self,excel_path,augmentations=None):\n        self.df = pd.read_excel(excel_path)\n        self.colums = self.df.columns #['Path','Label','Name']\n        self.image_paths = self.df['Path'].to_numpy()\n        self.targets = self.df['Label']\n        self.names = self.df['Name']\n        self.classes = self.targets.unique()\n        self.augmentations = augmentations\n        \n        self.cls_to_idx = dict(zip(self.names,self.targets))\n        self.idx_to_cls = dict(zip(self.targets,self.names))\n        \n    def __len__(self):\n            return len(self.image_paths)\n        \n    def __getitem__(self,index):\n        image_paths = self.image_paths[index]\n        image = Image.open(image_paths).convert('RGB')\n        image = image.resize(size=(224,224))\n        image_arr = np.array(image,dtype=np.uint8)\n        \n        if self.augmentations:\n            image = self.augmentations(image=image_arr)['image']\n            \n        \n        image = transforms.ToTensor()(image)\n        \n        #normalization\n        image = transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                                    std = [0.229, 0.224, 0.225])(image)\n            \n\n        labels = self.targets[index]\n\n        return image,labels,image_paths","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"5090a86e-b61a-4ec3-a948-a9d67a1ac298","_cell_guid":"8aeec6d1-c47f-4c9d-9942-770c5718a414","trusted":true},"cell_type":"code","source":"#dataset initialization\ntrain_dataset = CasavaDataset(excel_path='/kaggle/input/casava/train.xlsx',augmentations=aug)\nvalid_dataset = CasavaDataset(excel_path='/kaggle/input/casava/val.xlsx',augmentations=None)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"013466dc-acb7-40bd-8dc1-23e0babb3100","_cell_guid":"ea1657ba-0d0f-41ac-ad19-edc56e708915","trusted":true},"cell_type":"markdown","source":"# Class weights for weighted Sampling"},{"metadata":{"_uuid":"940efb70-5457-4b8a-bcde-f6509b44d8d3","_cell_guid":"d0769212-e223-4df8-8512-145603f951be","trusted":true},"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\ndef weighted_sampling(y,effective_weights=False):\n    \n    samples_per_cls = np.array([len(np.where(y==t)[0]) for t in np.unique(y)])\n    \n    if effective_weights: #from Class balance loss based on Effective number of samples paper\n        effective_num = 1.0 - np.power(cfg.BETA,samples_per_cls)\n        weights = (1.0-cfg.BETA)/np.array(effective_num)\n        weights = weights / np.sum(weights) * len(samples_per_cls)\n        \n    else:\n         weights = 1./(samples_per_cls)\n    \n    samples_weights = torch.from_numpy(np.array([weights[t] for t in y]))\n    \n    #define a sampler\n    sampler = torch.utils.data.WeightedRandomSampler(samples_weights.type('torch.DoubleTensor'),len(samples_weights))\n    \n    return weights,samples_per_cls,samples_weights,sampler","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"350ca302-0817-416d-88ee-6788ffbfaa99","_cell_guid":"79e6d203-ba67-429c-8e87-2eccdf91f8a6","trusted":true},"cell_type":"code","source":"weights,samples_per_class,samples_weights,weighted_sampler = weighted_sampling(train_dataset.targets)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"e66d3baa-b92c-43af-96b4-a52acff98adc","_cell_guid":"e91edbee-4951-4253-a2c2-dd640e2a39c9","trusted":true},"cell_type":"code","source":"print(f'[INFO] Samples per class is - {samples_per_class}')\nprint(f'[INFO] Weights per calss is - {weights}')\nCounter(samples_weights.numpy())","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"06aa84b6-9e0e-4b40-8f70-c29861bd3607","_cell_guid":"8f2b50e5-3d94-47ca-8057-0ac75475d565","trusted":true},"cell_type":"code","source":"# im,lab,path = next(iter(train_loader))\n# print(f'[INFO] Training image details - Shape - {im.shape} min,max - {im.min(),im.max()}')\n\n# im_val, lab_val, path_val = next(iter(valid_loader))\n# print(f'[INFO] Validation image details - Shape - {im_val.shape} min,max - {im_val.min(),im_val.max()}')\n\n# for i in range(5):\n#     print(f'[INFO] Class {i} - {len(lab[lab==i])}')","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"460c17fd-467a-4f14-afd7-82025ce7c9e8","_cell_guid":"dd0039df-f432-49e6-8e0a-2160f65f36f0","trusted":true},"cell_type":"markdown","source":"# Mixed precision Training"},{"metadata":{"_uuid":"6bc12eb2-f977-45b4-bbe5-f1e7593cb935","_cell_guid":"977997ca-4fd2-4dbc-9405-43752a67d356","trusted":true},"cell_type":"code","source":"#initialize grad saclar for mixed precision training\nscaler = amp.GradScaler()","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d23ad8ba-dd86-40e1-b65e-1a73464419a4","_cell_guid":"4e6a212a-8dc4-4449-8c9d-f5bf1cebdd46","trusted":true},"cell_type":"markdown","source":"# Logger\n1. Define functions for calculating accuracy,batch timings, f1-score etc"},{"metadata":{"_uuid":"6f513f96-a2b6-455d-8f86-e3b7464e8065","_cell_guid":"031b35d0-1d4a-4614-bf4d-9b2b5ddb7889","trusted":true},"cell_type":"code","source":"class Logger(object):\n    def __init__(self,mode,length,calculate_mean=False):\n        self.mode = mode\n        self.length = length\n        self.calculate_mean = calculate_mean\n        \n        if self.calculate_mean:\n            self.fn = lambda x,i: x/(i+1)\n        else:\n            self.fn = lambda x,i : x\n            \n    def __call__(self,loss,metrics,i,lr):\n        track_str = track_str = '\\r{} | {:5d}/{:<5d}| '.format(self.mode, i + 1, self.length)\n        loss_str = 'loss: {:9.4f} | '.format(self.fn(loss, i))\n        metric_str = ' | '.join('{}: {:9.4f}'.format(k, self.fn(v, i)) for k, v in metrics.items())\n        print(track_str + loss_str + metric_str + '|' + str(lr)  + '   ', end='')\n        if i + 1 == self.length:\n            print('')     \n            \n\nclass BatchTimer(object):\n    \"\"\"Batch timing class.\n    Use this class for tracking training and testing time/rate per batch or per sample.\n    \n    Keyword Arguments:\n        rate {bool} -- Whether to report a rate (batches or samples per second) or a time (seconds\n            per batch or sample). (default: {True})\n        per_sample {bool} -- Whether to report times or rates per sample or per batch.\n            (default: {True})\n    \"\"\"\n\n    def __init__(self, rate=True, per_sample=True):\n        self.start = time.time()\n        self.end = None\n        self.rate = rate\n        self.per_sample = per_sample\n\n    def __call__(self, y_pred, y):\n        self.end = time.time()\n        elapsed = self.end - self.start\n        self.start = self.end\n        self.end = None\n\n        if self.per_sample:\n            elapsed /= len(y_pred)\n        if self.rate:\n            elapsed = 1 / elapsed\n\n        return torch.tensor(elapsed)\n    \nclass Accuracy(object):\n    \n    def __init__(self):\n        pass\n    \n    def __call__(self,logits,y):\n        _,preds = torch.max(logits,1)\n        return (preds==y).float().mean().detach().cpu()\n\n    \nclass F1_score(object):\n    def __init__(self):\n        pass\n    \n    def __call__(self,logits,y):\n        _,y_preds = torch.max(logits,1)\n        fscore = f1_score(y_preds.cpu().numpy(),y.cpu().numpy(),average='weighted')\n        return fscore","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"67d281d4-84db-45fa-9ee2-f4aa19fbe303","_cell_guid":"5d4ef571-491b-4a23-8258-1b89923159ff","trusted":true},"cell_type":"code","source":"metrics = {'Accuracy': Accuracy(),\n           'F1 Score': F1_score(),\n           'Batch Time': BatchTimer()\n           }","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d0ee46b3-906d-4afe-98bf-2145929360b7","_cell_guid":"a02f6633-7cab-4b88-ac94-5f95063dd16d","trusted":true},"cell_type":"markdown","source":"# Training and validation Loop"},{"metadata":{"_uuid":"9dca45f1-2fec-47e7-9edd-0dd024997fd3","_cell_guid":"0bf02dae-34e4-4bca-980c-3e3675fc286d","trusted":true},"cell_type":"code","source":"def log(loss,output,gt,batch_metrics,metrics,batch_idx,logger,lr):\n    \n    metrics_batch = {}\n    \n    for metric_name, metric_fn in batch_metrics.items():\n        metrics_batch[metric_name] = metric_fn(output,gt)\n\n        metrics[metric_name] = metrics.get(metric_name,0) + metrics_batch[metric_name]\n\n    logger(loss,metrics,batch_idx,lr)\n\n    \ndef train(loaders, model, optimizer,scheduler, criterion,cfg,grad_scaler,batch_metrics):\n    \"\"\"returns trained model\"\"\"\n    \n    # initialize tracker for minimum validation loss\n    valid_loss_min = np.Inf \n    \n    count = 0\n\n    print(f'[INFO] Starting training...')\n    \n    print(f'[INFO] Training the model on - {cfg.DEVICE}')\n    for epoch in range(1, cfg.EPOCHS+1):\n        t0 = time.time()\n        # initialize variables to monitor training and validation loss\n        train_loss = 0.0\n        valid_loss = 0.0\n        train_metrics = {}\n        val_metrics = {}\n\n\n        count = 0\n\n        ###################\n        # train the model #\n        ###################\n        \n        #init logging\n        train_logger = Logger('Train',length=len(loaders['train']),calculate_mean=True)\n        \n        model.train()\n        for batch_idx, (data, target,train_image_paths) in enumerate(loaders['train']):\n\n            # move to GPU\n            if cfg.DEVICE:\n                data, target = data.cuda(), target.cuda()\n            \n            \n            with amp.autocast():\n                output = model(data)\n                loss = criterion(output,target)\n                \n                optimizer.zero_grad()\n                grad_scaler.scale(loss).backward()\n                grad_scaler.step(optimizer)\n                grad_scaler.update()\n        \n            ## find the loss and update the model parameters accordingly\n            ## record the average training loss, using something like\n#             batch_loss = loss\n#             train_loss += batch_loss\n            \n            train_loss = train_loss + ((1 / (batch_idx + 1)) * (loss.data - train_loss))\n            \n            #send the data to the logger\n#             lr = [param_groups['lr'] for param_groups in optimizer.param_groups][0]\n#             log(train_loss,output,target,batch_metrics,train_metrics,batch_idx,train_logger,lr)\n        \n        scheduler.step()\n            \n        ######################    \n        # validate the model #\n        ######################\n        \n        #init logging\n#         val_logger = Logger('Valid',length=len(loaders['valid']),calculate_mean=True)\n\n        model.eval()\n        with torch.no_grad():\n            for batch_idx, (data, target,valid_image_paths) in enumerate(loaders['valid']):\n                t1 = time.time()\n                # move to GPU\n                if cfg.DEVICE:\n                    data, target = data.cuda(), target.cuda()\n\n                output = model(data)\n                loss = criterion(output,target)\n                \n                valid_loss = valid_loss + ((1 / (batch_idx + 1)) * (loss.data - valid_loss))\n\n                ## update the average validation loss\n#                 valid_loss += loss.detach().item()\n#                 lr = [param_groups['lr'] for param_groups in optimizer.param_groups][0]\n#                 log(valid_loss,output,target,batch_metrics,val_metrics,batch_idx,val_logger,lr)\n        \n            print('Epoch: {} \\tTraining Loss: {:.6f} \\tValidation Loss: {:.6f}'.format(\n                epoch, \n                train_loss,\n                valid_loss\n                ))\n\n            ## TODO: save the model if validation loss has decreased\n            if valid_loss <= valid_loss_min:\n                print(f'Validation loss decreased {valid_loss_min :.4f} --> {valid_loss :.4f}. Saving the mode..')\n                torch.save(model,cfg.MODEL_SAVE_PATH+cfg.CHECKPOINT_NAME)\n                valid_loss_min = valid_loss\n    \n    # return trained model\n    \n    return model","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"6ba61b4d-03b7-44ba-bad5-5e05f1d85cda","_cell_guid":"30c65ee3-a274-430f-b61b-a8d9a2b26667","trusted":true},"cell_type":"code","source":"#define configurations\ncfg = Config(lr=0.0305,\n             checkpoint_name='/Resnet50_1.pt',\n             epochs = 1,\n             w_decay=0.01,\n             classify=True,\n             train_batch_size=512,\n             valid_batch_size=128)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"430db26a-c8cc-443c-9307-e76636e151d4","_cell_guid":"a3669ca0-9f44-4b05-92cd-c432cdd472b7","trusted":true},"cell_type":"code","source":"#optimizers and schedulers\noptimizer = torch.optim.AdamW([{'params': filter(lambda p: p.requires_grad, with_bn)},\n                               {'params': filter(lambda p: p.requires_grad, without_bn),\n                                'weight_decay': cfg.W_DECAY}],\n                                 lr = cfg.LR)\nscheduler = MultiStepLR(optimizer=optimizer,milestones=[2,4,90])\n\n#loss functions\n# criterion = FocalLoss(alpha = 0.25,gamma = 2)\ncriterion = nn.CrossEntropyLoss()","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"1376e232-3b2b-43fc-ad4b-fb917a4c7fa8","_cell_guid":"9ee3eb18-335a-410e-bd65-b5b5d1747e49","trusted":true},"cell_type":"code","source":"#weighted sampling on training data\ntrain_loader = torch.utils.data.DataLoader(train_dataset,batch_size = cfg.TRAIN_BATCH,\n                                               shuffle = False,num_workers = 0,sampler = weighted_sampler)\n\n#validation loader\nvalid_loader = torch.utils.data.DataLoader(valid_dataset,batch_size = cfg.VALID_BATCH,shuffle=True,num_workers=0)\n\n#define dataloaders\nloaders = {'train': train_loader,\n            'valid': valid_loader}","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# #Learing rate finder\n# lr_finder = LRFinder(mod,optimizer,criterion,device=device)\n# lr_finder.range_test(train_loader,end_lr=100,num_iter=100)\n# lr_finder.plot()\n# lr_finder.reset()","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"62fb80e6-dd08-4852-91ba-18322c766be2","_cell_guid":"a44da2b8-8efa-4ab1-aa28-773fa98d9a93","trusted":true},"cell_type":"code","source":"model = train(loaders = loaders,\n              model = mod.to(cfg.DEVICE),\n              optimizer = optimizer,\n              scheduler = scheduler,\n              criterion = criterion,\n              cfg = cfg,\n              grad_scaler = scaler,\n              batch_metrics = metrics)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"2c009174-a05e-404e-bc53-9d1458c5074e","_cell_guid":"201e405a-4c24-41cd-a0ff-f3e7e961bb87","trusted":true},"cell_type":"code","source":"!nvidia-smi","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"0ef192f0-e081-4e0e-a6db-1af39bf56e60","_cell_guid":"b42bf839-f317-4e47-9301-ec82d73aba68","trusted":true},"cell_type":"code","source":"print(torch.cuda.memory_summary())","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"83da8143-d280-43f0-8d5b-f68ddaa49703","_cell_guid":"f3f8c3dc-16a4-4454-9704-cefb9d1834c1","trusted":true},"cell_type":"code","source":"torch.cuda.empty_cache()\ntorch.cuda.reset_max_memory_allocated()","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"0cad46e7-0f39-4404-94e9-66ecca0df8ce","_cell_guid":"7a982773-047c-4970-a285-b35d70a2015c","trusted":true},"cell_type":"code","source":"model = torch.load(\"/kaggle/working/Senet1.pt\")","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"77a24faa-cfb3-4274-9693-ea7ecf0855f9","_cell_guid":"7a62dc8d-5e59-46e0-b00a-e83f994339b0","trusted":true},"cell_type":"code","source":"test_image = Image.open(\"/kaggle/input/cassava-leaf-disease-classification/test_images/2216849948.jpg\").convert('RGB')\noriginal_image = test_image.copy()\ntest_image = np.array(test_image)\ntest_image = transforms.ToTensor()(test_image)\ntest_image = transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                                    std = [0.229, 0.224, 0.225])(test_image)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"bc0526bd-185a-4848-8a3c-f19e5779b606","_cell_guid":"c2efdb02-9dd7-4e68-a20d-d9c31b621abc","trusted":true},"cell_type":"code","source":"state_dict = torch.load(\"/kaggle/working/Resnet50_1.pt\")\nmod.load_state_dict(state_dict)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"c11f986b-31c0-4500-b817-141ea1b62b16","_cell_guid":"9dbb7609-f021-4355-bc96-648e7913cca5","trusted":true},"cell_type":"code","source":"logits = mod(test_image.unsqueeze(0).cuda())\nprint(logits.shape)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"4ecdefbf-3633-4557-bb3d-3f329225ba01","_cell_guid":"f2dc1734-6878-48a8-8ffc-11dc57897e10","trusted":true},"cell_type":"code","source":"probs,preds = torch.topk(F.softmax(logits,dim=1),1)\npreds = preds.detach().cpu().numpy().flatten()\nprobs = probs.detach().cpu().numpy()\nprint(f'[INFO] Predictions - {preds}')\nprint(f'[INFO] Probabilities - {probs}')","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"378a43a6-abac-44bb-942e-327dcbb082ac","_cell_guid":"1840eda0-eb37-4ecc-887c-3266137d39e7","trusted":true},"cell_type":"code","source":"original_image","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"8ae41909-c829-49a7-af65-a3076666acfe","_cell_guid":"a53e800b-7afb-4956-9035-c35e98cb14fb","trusted":true},"cell_type":"markdown","source":"# Submission"},{"metadata":{"_uuid":"48b28cdc-d163-45ad-aea2-964e6d18b14a","_cell_guid":"ebb78cf5-0849-4e97-936b-6a77060f5f76","trusted":true},"cell_type":"code","source":"test_image_id = \"2216849948.jpg\"\nmy_submission = pd.DataFrame({'image_id':test_image_id,\n                             'label': preds})\nmy_submission.to_csv('submission.csv',index = False)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}