{"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":"code","source":"import numpy as np \nimport pandas as pd \nimport glob\nimport json\nimport os\nimport seaborn as sns\nimport cv2\nimport matplotlib.pyplot as plt","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-02-26T06:13:49.633793Z","iopub.execute_input":"2022-02-26T06:13:49.634111Z","iopub.status.idle":"2022-02-26T06:13:50.761471Z","shell.execute_reply.started":"2022-02-26T06:13:49.634024Z","shell.execute_reply":"2022-02-26T06:13:50.760716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<!-- The training set images are organized in subfolders **h22-train/images/subfolder1/subfolder2/image_id.jpg**, where subfolder1 and subfolder2 comes from the **first three and the last two digits of the image_id**. **Image_id is a result of combination between category_id and unique numbers that differentiates images within plant taxa.** -->","metadata":{}},{"cell_type":"code","source":"INPUT_BASE_FILES = glob.glob('../input/herbarium-2022-fgvc9/*')\n\ntrain_metadata_json = INPUT_BASE_FILES[0]\nsample_submission_csv = INPUT_BASE_FILES[1]\ntest_metadata_json = INPUT_BASE_FILES[2]\ntrain_images_folder = INPUT_BASE_FILES[3]\ntest_images_folder = INPUT_BASE_FILES[4]\n","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:13:50.764187Z","iopub.execute_input":"2022-02-26T06:13:50.764621Z","iopub.status.idle":"2022-02-26T06:13:50.771221Z","shell.execute_reply.started":"2022-02-26T06:13:50.764582Z","shell.execute_reply":"2022-02-26T06:13:50.770471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# EDA","metadata":{}},{"cell_type":"code","source":"with open(train_metadata_json) as json_file:\n    train_metadata = json.load(json_file)\n    \nwith open(test_metadata_json) as json_file:\n    test_metadata = json.load(json_file)","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:13:50.772566Z","iopub.execute_input":"2022-02-26T06:13:50.773051Z","iopub.status.idle":"2022-02-26T06:14:04.540367Z","shell.execute_reply.started":"2022-02-26T06:13:50.773008Z","shell.execute_reply":"2022-02-26T06:14:04.539599Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(train_metadata.keys()) # A dictionary","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:14:04.541446Z","iopub.execute_input":"2022-02-26T06:14:04.541725Z","iopub.status.idle":"2022-02-26T06:14:04.549577Z","shell.execute_reply.started":"2022-02-26T06:14:04.541691Z","shell.execute_reply":"2022-02-26T06:14:04.546909Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(test_metadata[:2]) # A list","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:14:04.552937Z","iopub.execute_input":"2022-02-26T06:14:04.553214Z","iopub.status.idle":"2022-02-26T06:14:04.561088Z","shell.execute_reply.started":"2022-02-26T06:14:04.553177Z","shell.execute_reply":"2022-02-26T06:14:04.557483Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for k,v in train_metadata.items():\n    print(f'| Key : {k}   >>  Total values  : {len(v)} ')","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:14:04.562441Z","iopub.execute_input":"2022-02-26T06:14:04.562907Z","iopub.status.idle":"2022-02-26T06:14:04.570468Z","shell.execute_reply.started":"2022-02-26T06:14:04.562851Z","shell.execute_reply":"2022-02-26T06:14:04.56959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gen = train_metadata.get('genera')\ngenera_dict = {}\nfor i in gen:\n    genera_dict[i.get('genus_id')] = i.get('genus')","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:14:04.572359Z","iopub.execute_input":"2022-02-26T06:14:04.572949Z","iopub.status.idle":"2022-02-26T06:14:04.578815Z","shell.execute_reply.started":"2022-02-26T06:14:04.572873Z","shell.execute_reply":"2022-02-26T06:14:04.578141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Sample Values of each keys .. \\n')\n\nprint('[+] Images ---\\n')\nprint(train_metadata.get('images')[0])\nprint('\\n')\nprint('[+] Annotations ---\\n')\nprint(train_metadata.get('annotations')[0])\nprint('\\n')\nprint('[+] Categories ---\\n')\nprint(train_metadata.get('categories')[0])\nprint('\\n')\nprint('[+] Genera --- \\n ')\nprint(train_metadata.get('genera')[0])\nprint('\\n')\nprint('[+] Distances --- \\n ')\nprint(train_metadata.get('distances')[0])\nprint('\\n')\nprint('[+] Institutions --- \\n ')\nprint(train_metadata.get('institutions')[0])\nprint('\\n')\nprint('[+] License --- \\n ')\nprint(train_metadata.get('license')[0])","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:14:04.580212Z","iopub.execute_input":"2022-02-26T06:14:04.580707Z","iopub.status.idle":"2022-02-26T06:14:04.593373Z","shell.execute_reply.started":"2022-02-26T06:14:04.58067Z","shell.execute_reply":"2022-02-26T06:14:04.592697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Image information\nfile_names = []\nimage_ids = []\ngenus_ids = []\ngenus_names = []\ncategory_ids = []\ninstitution_ids = []\nimage_paths = []\n\nfor i,j in zip(train_metadata.get('images'),train_metadata.get('annotations')):\n    \n    image_id_im = i.get('image_id')\n    image_id_anno = j.get('image_id')\n    \n    if image_id_im == image_id_anno:\n        file_name = i.get('file_name')\n        genus_id = j.get('genus_id')\n        category_id = j.get('category_id')\n        institution_id = j.get('institution_id')\n        \n        file_names.append(file_name)\n        image_ids.append(image_id_anno)\n        genus_ids.append(genus_id)\n        genus_names.append(genera_dict.get(genus_id))\n        category_ids.append(category_id)\n        institution_ids.append(institution_id)\n        image_paths.append(os.path.join(train_images_folder,file_name))","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:14:04.594684Z","iopub.execute_input":"2022-02-26T06:14:04.595033Z","iopub.status.idle":"2022-02-26T06:14:07.577534Z","shell.execute_reply.started":"2022-02-26T06:14:04.594999Z","shell.execute_reply":"2022-02-26T06:14:07.576818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_images_df = pd.DataFrame.from_dict({'FileNames' : file_names, 'ImageID' : image_ids, 'GenusID' : genus_ids,'GenusNames':genus_names,\n                                             'CategoryID' : category_ids,'InstitutionID' : institution_ids,'ImagePath':image_paths})","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:14:07.578752Z","iopub.execute_input":"2022-02-26T06:14:07.578991Z","iopub.status.idle":"2022-02-26T06:14:08.735029Z","shell.execute_reply.started":"2022-02-26T06:14:07.57896Z","shell.execute_reply":"2022-02-26T06:14:08.734206Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_images_df.to_csv('training_images_df.csv')","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:14:08.736447Z","iopub.execute_input":"2022-02-26T06:14:08.736721Z","iopub.status.idle":"2022-02-26T06:14:14.255862Z","shell.execute_reply.started":"2022-02-26T06:14:08.736684Z","shell.execute_reply":"2022-02-26T06:14:14.255082Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_images_df.sample(5)","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:14:14.257206Z","iopub.execute_input":"2022-02-26T06:14:14.257446Z","iopub.status.idle":"2022-02-26T06:14:14.301235Z","shell.execute_reply.started":"2022-02-26T06:14:14.257411Z","shell.execute_reply":"2022-02-26T06:14:14.300542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Check correlation between 'CategoryID','GenusID' and 'InstitutionID'\ncorr_cgi = training_images_df[['CategoryID','GenusID','InstitutionID']].corr()\ncorr_cgi","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:14:14.302398Z","iopub.execute_input":"2022-02-26T06:14:14.302632Z","iopub.status.idle":"2022-02-26T06:14:14.366723Z","shell.execute_reply.started":"2022-02-26T06:14:14.302602Z","shell.execute_reply":"2022-02-26T06:14:14.365911Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_images_df[['CategoryID','GenusID','InstitutionID']].nunique()","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:14:14.369931Z","iopub.execute_input":"2022-02-26T06:14:14.37021Z","iopub.status.idle":"2022-02-26T06:14:14.403639Z","shell.execute_reply.started":"2022-02-26T06:14:14.370171Z","shell.execute_reply":"2022-02-26T06:14:14.402886Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_images_df.sort_values(by=['CategoryID'],ascending=False).head(4)","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:14:14.40483Z","iopub.execute_input":"2022-02-26T06:14:14.406808Z","iopub.status.idle":"2022-02-26T06:14:14.519963Z","shell.execute_reply.started":"2022-02-26T06:14:14.406769Z","shell.execute_reply":"2022-02-26T06:14:14.51915Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Briet data information","metadata":{}},{"cell_type":"code","source":"print('Genus ID information')\nid,count = np.unique(genus_ids,return_counts=True)\ngenus_count_df = pd.DataFrame.from_dict({'Genus ID' : id,'Count' : count}).sort_values(by=['Count'],ascending=False)\ngenus_count_df['Count'].hist(bins=100, figsize=(18, 6), grid=True)\nplt.title('Histogram of Genus ID counts')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:14:14.521427Z","iopub.execute_input":"2022-02-26T06:14:14.521707Z","iopub.status.idle":"2022-02-26T06:14:15.081151Z","shell.execute_reply.started":"2022-02-26T06:14:14.521672Z","shell.execute_reply":"2022-02-26T06:14:15.080484Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Category ID Information') \nid,count = np.unique(category_ids,return_counts=True)\ncategory_id_df = pd.DataFrame.from_dict({'Category ID' : id,'Count' : count}).sort_values(by=['Count'],ascending=False)\ncategory_id_df['Count'].hist(bins=100, figsize=(18, 6), grid=True)\nplt.title('Histogram of Category ID counts')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:14:15.082572Z","iopub.execute_input":"2022-02-26T06:14:15.083057Z","iopub.status.idle":"2022-02-26T06:14:15.536237Z","shell.execute_reply.started":"2022-02-26T06:14:15.083018Z","shell.execute_reply":"2022-02-26T06:14:15.535504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Institution ID information')\nid,count = np.unique(institution_ids,return_counts=True)\ninstitution_id_df = pd.DataFrame.from_dict({'Institution ID' : id,'Count' : count}).sort_values(by=['Count'],ascending=False)\ninstitution_id_df['Count'].hist(bins=100, figsize=(18, 6), grid=True)\nplt.title('Histogram of Institution ID counts')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:14:15.537505Z","iopub.execute_input":"2022-02-26T06:14:15.537769Z","iopub.status.idle":"2022-02-26T06:14:16.274354Z","shell.execute_reply.started":"2022-02-26T06:14:15.537735Z","shell.execute_reply":"2022-02-26T06:14:16.273698Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Image Visualization","metadata":{}},{"cell_type":"code","source":"def visualize_data(df,show_by='Random',genus_name = None):\n    \n    if show_by == 'Genus':\n        df = df[df['GenusNames']==genus_name]\n            \n    data = df.sample(10)\n    \n    image_paths = data['ImagePath'].to_list()\n    genus_ids = data['GenusNames'].to_list()\n    category_ids = data['CategoryID'].to_list()\n    institution_ids = data['InstitutionID'].to_list()\n    \n    plt.figure(figsize=(13,13))\n    \n    for indx,im in enumerate(image_paths):\n        plt.subplot(2,5,indx+1)\n        image = cv2.imread(im)\n        plt.imshow(image[:,:,::-1])\n        plt.title(f'GeniusNames :{genus_ids[indx]},\\nCategoryID : {category_ids[indx]},\\nInstitutionID : {institution_ids[indx]}')\n        plt.axis('off')\n    plt.tight_layout()","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:14:16.275718Z","iopub.execute_input":"2022-02-26T06:14:16.27595Z","iopub.status.idle":"2022-02-26T06:14:16.283448Z","shell.execute_reply.started":"2022-02-26T06:14:16.275918Z","shell.execute_reply":"2022-02-26T06:14:16.2826Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Visualize random 10 image\nvisualize_data(training_images_df,show_by='Random')","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:14:16.284898Z","iopub.execute_input":"2022-02-26T06:14:16.285293Z","iopub.status.idle":"2022-02-26T06:14:18.397583Z","shell.execute_reply.started":"2022-02-26T06:14:16.285258Z","shell.execute_reply":"2022-02-26T06:14:18.396713Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Visualize random 10 image for a particular genus\nvisualize_data(training_images_df,show_by='Genus',genus_name='Asimina')","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:14:18.398874Z","iopub.execute_input":"2022-02-26T06:14:18.399147Z","iopub.status.idle":"2022-02-26T06:14:20.438041Z","shell.execute_reply.started":"2022-02-26T06:14:18.399114Z","shell.execute_reply":"2022-02-26T06:14:20.436368Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# MODELLING","metadata":{}},{"cell_type":"code","source":"import pytorch_lightning as pl\nfrom torch.utils.data import DataLoader\nfrom sklearn.model_selection import train_test_split\nfrom pytorch_lightning.loggers import WandbLogger\nfrom torch.utils.data import Dataset\nimport torch\nimport albumentations as A\nfrom albumentations.pytorch.transforms import ToTensorV2\nimport torch\nfrom torch.nn import functional as F\nfrom torch import nn\nimport pandas as pd\nimport torchvision\nfrom pytorch_lightning.core.lightning import LightningModule","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:14:53.488534Z","iopub.execute_input":"2022-02-26T06:14:53.48908Z","iopub.status.idle":"2022-02-26T06:14:53.495814Z","shell.execute_reply.started":"2022-02-26T06:14:53.489043Z","shell.execute_reply":"2022-02-26T06:14:53.494983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data = training_images_df #pd.read_csv('./training_images_df.csv',index_col=0,dtype=str)\n\n\nX = data['ImagePath'].to_list()\ny = data['CategoryID'].to_list()\n#y = [i-1 for i in y]\n\nX_train, X_val, y_train, y_val = train_test_split(X, y, test_size=0.33, random_state=41,stratify=y)\n\ntrain_df = pd.DataFrame.from_dict({'ImagePath' : X_train, 'CategoryID' : y_train})\nval_df = pd.DataFrame.from_dict({'ImagePath' : X_val, 'CategoryID' : y_val})","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:51:39.146883Z","iopub.execute_input":"2022-02-26T06:51:39.147434Z","iopub.status.idle":"2022-02-26T06:51:40.912801Z","shell.execute_reply.started":"2022-02-26T06:51:39.147395Z","shell.execute_reply":"2022-02-26T06:51:40.912045Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#training_images_df.to_csv('data_info.csv')","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:51:40.914329Z","iopub.execute_input":"2022-02-26T06:51:40.914577Z","iopub.status.idle":"2022-02-26T06:51:40.918827Z","shell.execute_reply.started":"2022-02-26T06:51:40.914542Z","shell.execute_reply":"2022-02-26T06:51:40.917969Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(max(y),min(y))\nprint(max(y_train),min(y_train))\nprint(max(y_val),min(y_val))","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:51:40.919988Z","iopub.execute_input":"2022-02-26T06:51:40.920256Z","iopub.status.idle":"2022-02-26T06:51:41.028193Z","shell.execute_reply.started":"2022-02-26T06:51:40.92022Z","shell.execute_reply":"2022-02-26T06:51:41.027345Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transform = A.Compose([\n    A.HorizontalFlip(p=0.3),\n    A.RandomBrightnessContrast(p=0.1),\n    A.Normalize(p=1),\n    A.Rotate(limit=30,p=0.2)\n])\n\nval_transform = A.Compose([\n    A.Normalize(p=1),\n])","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:52:02.486078Z","iopub.execute_input":"2022-02-26T06:52:02.486335Z","iopub.status.idle":"2022-02-26T06:52:02.491386Z","shell.execute_reply.started":"2022-02-26T06:52:02.486308Z","shell.execute_reply":"2022-02-26T06:52:02.490385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CDataset(Dataset):\n    def __init__(self, dataframe, transform=None,target_size=(512,512)):\n        self.transform = transform\n        self.dataframe = dataframe\n        self.image_paths = dataframe['ImagePath']\n        self.labels = dataframe['CategoryID']\n        self.target_size = target_size\n        \n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, indx):\n        \n        if torch.is_tensor(indx):\n            indx = indx.tolist()\n\n        img_name = self.image_paths[indx]\n        image = cv2.imread(img_name)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        image = cv2.resize(image, self.target_size)\n        label = np.array(self.labels[indx]).astype(int)\n        if self.transform:\n            transformed_image = self.transform(image=image)['image']\n            transformed_image,label = torch.from_numpy(transformed_image),torch.from_numpy(label)\n            transformed_image = transformed_image.permute(2, 0, 1)\n            return (transformed_image,label)\n        image = image.permute(2, 0, 1)\n        return (image,label)","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:52:19.942845Z","iopub.execute_input":"2022-02-26T06:52:19.943105Z","iopub.status.idle":"2022-02-26T06:52:20.035068Z","shell.execute_reply.started":"2022-02-26T06:52:19.943077Z","shell.execute_reply":"2022-02-26T06:52:20.033953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_data = CDataset(train_df,transform = train_transform,target_size=(224,224))\nvalidation_data = CDataset(val_df,transform = val_transform,target_size=(224,224))","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:52:20.394271Z","iopub.execute_input":"2022-02-26T06:52:20.394967Z","iopub.status.idle":"2022-02-26T06:52:20.416201Z","shell.execute_reply.started":"2022-02-26T06:52:20.394927Z","shell.execute_reply":"2022-02-26T06:52:20.415344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_batch_size = 128\ntrain_dataloader = DataLoader(training_data, batch_size=train_batch_size,shuffle=True)\nvalidation_dataloader = DataLoader(validation_data, batch_size=128,shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2022-02-26T07:20:56.996096Z","iopub.execute_input":"2022-02-26T07:20:56.996638Z","iopub.status.idle":"2022-02-26T07:20:57.00113Z","shell.execute_reply.started":"2022-02-26T07:20:56.996597Z","shell.execute_reply":"2022-02-26T07:20:57.000328Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(train_dataloader))\nprint(len(validation_dataloader))","metadata":{"execution":{"iopub.status.busy":"2022-02-26T07:22:14.6813Z","iopub.execute_input":"2022-02-26T07:22:14.681744Z","iopub.status.idle":"2022-02-26T07:22:14.687258Z","shell.execute_reply.started":"2022-02-26T07:22:14.681686Z","shell.execute_reply":"2022-02-26T07:22:14.686501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualize_from_dataloader(dl):\n    features, labels = next(iter(train_dataloader))\n    print(type(features))\n    print(type(labels))\n    print(f\"Feature batch shape: {features.size()}\")\n    print(f\"Labels batch shape: {labels.size()}\")\n    img = features[0].squeeze()\n    img = img.numpy().transpose((1, 2, 0))\n    mean = np.array([0.485, 0.456, 0.406])\n    std = np.array([0.229, 0.224, 0.225])\n    img = img * std + mean\n    label = labels[0]\n    plt.imshow(img, cmap=\"gray\")\n    plt.title(f\"Label: {label}\")\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:52:27.864358Z","iopub.execute_input":"2022-02-26T06:52:27.864616Z","iopub.status.idle":"2022-02-26T06:52:27.873003Z","shell.execute_reply.started":"2022-02-26T06:52:27.864585Z","shell.execute_reply":"2022-02-26T06:52:27.870795Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"visualize_from_dataloader(train_dataloader)","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:52:31.068779Z","iopub.execute_input":"2022-02-26T06:52:31.069032Z","iopub.status.idle":"2022-02-26T06:52:34.010214Z","shell.execute_reply.started":"2022-02-26T06:52:31.069003Z","shell.execute_reply":"2022-02-26T06:52:34.009508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install efficientnet_pytorch \n!pip install torchsummary","metadata":{"execution":{"iopub.status.busy":"2022-02-26T07:19:30.755456Z","iopub.execute_input":"2022-02-26T07:19:30.755753Z","iopub.status.idle":"2022-02-26T07:19:30.759424Z","shell.execute_reply.started":"2022-02-26T07:19:30.75572Z","shell.execute_reply":"2022-02-26T07:19:30.758424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from efficientnet_pytorch import EfficientNet\nfrom efficientnet_pytorch.utils import MemoryEfficientSwish\nfrom torchsummary import summary\nfrom torch import nn\nimport torch.nn.functional as F\nimport time\nimport copy\nfrom tqdm.autonotebook import tqdm","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:53:06.01311Z","iopub.execute_input":"2022-02-26T06:53:06.013581Z","iopub.status.idle":"2022-02-26T06:53:06.018384Z","shell.execute_reply.started":"2022-02-26T06:53:06.013546Z","shell.execute_reply":"2022-02-26T06:53:06.017396Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Net(nn.Module):\n    def __init__(self):\n        super(Net, self).__init__()\n        classes = 15505\n        self.base_model = EfficientNet.from_name(\"efficientnet-b0\",include_top=False, in_channels=3)\n        self.drop = nn.Dropout2d(p=0.2)\n        self.fc1 = nn.Linear(1280, 1280//2)\n        self.fc2 = nn.Linear(1280//2, classes)\n\n    def forward(self, x):  \n        x = self.base_model(x)\n        x = x.view(-1,1280)\n        x = F.relu(self.fc1(x))\n        x = F.dropout(x, training=self.training)\n        x = self.fc2(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:55:26.032782Z","iopub.execute_input":"2022-02-26T06:55:26.0336Z","iopub.status.idle":"2022-02-26T06:55:26.041628Z","shell.execute_reply.started":"2022-02-26T06:55:26.033548Z","shell.execute_reply":"2022-02-26T06:55:26.040721Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:55:26.591896Z","iopub.execute_input":"2022-02-26T06:55:26.592156Z","iopub.status.idle":"2022-02-26T06:55:26.596462Z","shell.execute_reply.started":"2022-02-26T06:55:26.592128Z","shell.execute_reply":"2022-02-26T06:55:26.595726Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = Net()\nmodel = model.to(device)\nlayer = 1\nfor name, param in model.named_parameters():\n    #print(f'layer : {layer}, name : {name}')\n    if layer < 210:\n        param.requires_grad = False\n    layer+=1\nsummary(model,input_size=(3,224,224))","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:55:32.935895Z","iopub.execute_input":"2022-02-26T06:55:32.936854Z","iopub.status.idle":"2022-02-26T06:55:33.804082Z","shell.execute_reply.started":"2022-02-26T06:55:32.936801Z","shell.execute_reply":"2022-02-26T06:55:33.803339Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def get_preds(model):\n#     model.eval()\n#     images, labels = next(iter(train_dataloader))\n#     images = images.to(device)\n#     labels = labels.to(device)\n#     outputs = model(images)\n#     _, preds = torch.max(outputs, 1)\n#     print(preds)\n#     return preds\n#get_preds(model)","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:55:45.272629Z","iopub.execute_input":"2022-02-26T06:55:45.272907Z","iopub.status.idle":"2022-02-26T06:55:45.278621Z","shell.execute_reply.started":"2022-02-26T06:55:45.272877Z","shell.execute_reply":"2022-02-26T06:55:45.275904Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataloaders = {'train' : train_dataloader, 'val' : validation_dataloader}\noptimizer = torch.optim.Adam(model.parameters(), lr=3e-3)\ncriterion = nn.CrossEntropyLoss()","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:55:46.980539Z","iopub.execute_input":"2022-02-26T06:55:46.981164Z","iopub.status.idle":"2022-02-26T06:55:46.99024Z","shell.execute_reply.started":"2022-02-26T06:55:46.981124Z","shell.execute_reply":"2022-02-26T06:55:46.989477Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_model(model, dataloaders, criterion, optimizer, num_epochs=25):\n    since = time.time()\n    val_acc_history = []\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_acc = 0.0\n    for epoch in range(num_epochs):\n        print('Epoch {}/{}'.format(epoch, num_epochs - 1))\n        print('-' * 10)\n        # Each epoch has a training and validation phase\n        \n        for phase in ['train', 'val']:\n            if phase == 'train':\n                model.train()  # Set model to training mode\n            else:\n                model.eval()   # Set model to evaluate mode\n            \n            running_loss = 0.0\n            running_corrects = 0\n            # Iterate over data.\n            with tqdm(dataloaders[phase],unit=\"batch\") as dl:\n                for inputs, labels in dl:\n                    inputs = inputs.to(device)\n                    labels = labels.to(device)\n                    optimizer.zero_grad()\n                    with torch.set_grad_enabled(phase == 'train'):\n                        outputs = model(inputs)\n                        loss = criterion(outputs, labels)\n                        _, preds = torch.max(outputs, 1)\n                        if phase == 'train':\n                            loss.backward()\n                            optimizer.step()\n                    running_loss += loss.item() * inputs.size(0)\n                    running_corrects += torch.sum(preds == labels.data).item()\n                    dl.set_postfix(loss=running_loss/len(dataloaders[phase].dataset), accuracy=running_corrects / len(dataloaders[phase].dataset))\n\n                epoch_loss = running_loss / len(dataloaders[phase].dataset)\n                epoch_acc = running_corrects / len(dataloaders[phase].dataset)\n\n                print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))\n\n                # deep copy the model\n                if phase == 'val' and epoch_acc > best_acc:\n                    best_acc = epoch_acc\n                    best_model_wts = copy.deepcopy(model.state_dict())\n                    torch.save(best_model_wts,'best_model.pth')\n                    \n                if phase == 'val':\n                    val_acc_history.append(epoch_acc)\n\n                del inputs, labels\n                torch.cuda.empty_cache()\n        print()\n\n    time_elapsed = time.time() - since\n    print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))\n    print('Best val Acc: {:4f}'.format(best_acc))\n\n    # load best model weights\n    model.load_state_dict(best_model_wts)\n    return model, val_acc_history","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:55:52.642542Z","iopub.execute_input":"2022-02-26T06:55:52.642982Z","iopub.status.idle":"2022-02-26T06:55:52.656712Z","shell.execute_reply.started":"2022-02-26T06:55:52.642945Z","shell.execute_reply":"2022-02-26T06:55:52.655844Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model, val_acc_history = train_model(model, dataloaders, criterion, optimizer, num_epochs=5)","metadata":{"execution":{"iopub.status.busy":"2022-02-26T07:27:43.72343Z","iopub.execute_input":"2022-02-26T07:27:43.723713Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#torch.save(model,'best_model.pth')","metadata":{"execution":{"iopub.status.busy":"2022-02-26T06:04:31.034901Z","iopub.execute_input":"2022-02-26T06:04:31.035496Z","iopub.status.idle":"2022-02-26T06:04:31.15813Z","shell.execute_reply.started":"2022-02-26T06:04:31.03546Z","shell.execute_reply":"2022-02-26T06:04:31.157219Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}