{"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":"!pip install wtfml==0.0.2\n!pip install efficientnet_pytorch","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-03-06T17:39:15.895902Z","iopub.execute_input":"2022-03-06T17:39:15.896535Z","iopub.status.idle":"2022-03-06T17:39:33.900090Z","shell.execute_reply.started":"2022-03-06T17:39:15.896415Z","shell.execute_reply":"2022-03-06T17:39:33.899245Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport os\nfrom PIL import Image\n\nfrom sklearn import model_selection\nfrom sklearn import metrics\n\nimport torch\nimport torch.nn as nn\nfrom torch.nn import functional as F\nimport torch.optim as optim\nimport efficientnet_pytorch\n\nimport albumentations as A\n\nfrom wtfml.utils import EarlyStopping\nfrom wtfml.engine import Engine\nfrom wtfml.data_loaders.image import ClassificationLoader\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2022-03-06T17:39:37.000013Z","iopub.execute_input":"2022-03-06T17:39:37.000274Z","iopub.status.idle":"2022-03-06T17:39:40.074926Z","shell.execute_reply.started":"2022-03-06T17:39:37.000247Z","shell.execute_reply":"2022-03-06T17:39:40.074155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dir = '../input/cassava-leaf-disease-classification/train_images'\ntest_dir = '../input/cassava-leaf-disease-classification/test_images'\nt = os.listdir(train_dir)\n\nt1 = os.listdir(test_dir)\nprint(len(t),len(t1),len(t)+len(t1))","metadata":{"execution":{"iopub.status.busy":"2022-03-06T17:39:45.923099Z","iopub.execute_input":"2022-03-06T17:39:45.923620Z","iopub.status.idle":"2022-03-06T17:39:46.655320Z","shell.execute_reply.started":"2022-03-06T17:39:45.923562Z","shell.execute_reply":"2022-03-06T17:39:46.654806Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(\"../input/cassava-leaf-disease-classification/train.csv\")\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2022-03-06T17:39:49.919401Z","iopub.execute_input":"2022-03-06T17:39:49.919700Z","iopub.status.idle":"2022-03-06T17:39:49.961122Z","shell.execute_reply.started":"2022-03-06T17:39:49.919674Z","shell.execute_reply":"2022-03-06T17:39:49.960621Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label_to_disease = pd.read_json((\"../input/cassava-leaf-disease-classification/label_num_to_disease_map.json\"), typ='series')\ndf['disease'] = df['label'].map(label_to_disease)\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2022-03-06T17:39:52.608448Z","iopub.execute_input":"2022-03-06T17:39:52.609246Z","iopub.status.idle":"2022-03-06T17:39:52.638005Z","shell.execute_reply.started":"2022-03-06T17:39:52.609155Z","shell.execute_reply":"2022-03-06T17:39:52.637378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''check the count of the various disease types'''\n\n#visualization imports\nimport matplotlib.pyplot as plt\nfrom matplotlib.image import imread\nimport seaborn as sns\n%matplotlib inline\n\n\nsns.countplot(df['label'])\nplt.title('Count of the various disease types in Cassava leaves')\nplt.grid()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-03-06T17:39:55.952224Z","iopub.execute_input":"2022-03-06T17:39:55.952475Z","iopub.status.idle":"2022-03-06T17:39:56.203736Z","shell.execute_reply.started":"2022-03-06T17:39:55.952449Z","shell.execute_reply":"2022-03-06T17:39:56.203243Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nfrom sklearn import model_selection\n\n\n# we create a new column called kfold and fill it with -1\ndf[\"kfold\"] = -1\n# the next step is to randomize the rows of the data\ndf = df.sample(frac=1).reset_index(drop=True)\n# fetch targets\ny = df.label.values\n# initiate the kfold class from model_selection module\nkf = model_selection.StratifiedKFold(n_splits=5)\n# fill the new kfold column\n\nfor f, (t_, v_) in enumerate(kf.split(X=df, y=y)):\n    df.loc[v_, 'kfold'] = f\n# save the new csv with kfold column\ndf.to_csv(\"train_folds.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2022-03-06T17:39:59.652818Z","iopub.execute_input":"2022-03-06T17:39:59.653211Z","iopub.status.idle":"2022-03-06T17:39:59.739554Z","shell.execute_reply.started":"2022-03-06T17:39:59.653173Z","shell.execute_reply":"2022-03-06T17:39:59.738764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train_folds =pd.read_csv(\"train_folds.csv\")\ndf_train_folds","metadata":{"execution":{"iopub.status.busy":"2022-03-06T17:40:03.691956Z","iopub.execute_input":"2022-03-06T17:40:03.692215Z","iopub.status.idle":"2022-03-06T17:40:03.727318Z","shell.execute_reply.started":"2022-03-06T17:40:03.692186Z","shell.execute_reply":"2022-03-06T17:40:03.726556Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''check the count after applying Stratified kfold'''\n\n#visualization imports\nimport matplotlib.pyplot as plt\nfrom matplotlib.image import imread\nimport seaborn as sns\n%matplotlib inline\n\n\nsns.countplot(df['kfold'])\nplt.title('Count of the various disease types in Cassava leaves')\nplt.grid()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-03-06T17:40:06.659544Z","iopub.execute_input":"2022-03-06T17:40:06.660272Z","iopub.status.idle":"2022-03-06T17:40:06.842283Z","shell.execute_reply.started":"2022-03-06T17:40:06.660222Z","shell.execute_reply":"2022-03-06T17:40:06.841503Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#done train test split and stratified folds on original train.csv file again.\n\n#reading only first 200rows for experimentation\n\ndfx = pd.read_csv('../input/cassava-leaf-disease-classification/train.csv',nrows =200)\ndf_train, df_valid = model_selection.train_test_split(\n        dfx, test_size=0.1, random_state=42, stratify=dfx.label.values\n)\nlen(df_train),len(df_valid)\n","metadata":{"execution":{"iopub.status.busy":"2022-03-06T17:40:10.455943Z","iopub.execute_input":"2022-03-06T17:40:10.456578Z","iopub.status.idle":"2022-03-06T17:40:10.472118Z","shell.execute_reply.started":"2022-03-06T17:40:10.456546Z","shell.execute_reply":"2022-03-06T17:40:10.471384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#reset the index and than drop the index\ndf_train = df_train.reset_index(drop=True)\ndf_valid = df_valid.reset_index(drop=True)\ndf_train.shape\n","metadata":{"execution":{"iopub.status.busy":"2022-03-06T17:40:14.488948Z","iopub.execute_input":"2022-03-06T17:40:14.489618Z","iopub.status.idle":"2022-03-06T17:40:14.496281Z","shell.execute_reply.started":"2022-03-06T17:40:14.489555Z","shell.execute_reply":"2022-03-06T17:40:14.495623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#join image name with path to make a list of training images\n\ntrain_images = [os.path.join(train_dir,x) for x in df_train.image_id.values]\ntrain_images[1]","metadata":{"execution":{"iopub.status.busy":"2022-03-06T17:40:17.019627Z","iopub.execute_input":"2022-03-06T17:40:17.020353Z","iopub.status.idle":"2022-03-06T17:40:17.027274Z","shell.execute_reply.started":"2022-03-06T17:40:17.020313Z","shell.execute_reply":"2022-03-06T17:40:17.026639Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_images = [os.path.join(train_dir,x) for x in df_valid.image_id.values]\nvalid_images[:5]","metadata":{"execution":{"iopub.status.busy":"2022-03-06T17:40:19.303938Z","iopub.execute_input":"2022-03-06T17:40:19.304322Z","iopub.status.idle":"2022-03-06T17:40:19.310168Z","shell.execute_reply.started":"2022-03-06T17:40:19.304283Z","shell.execute_reply":"2022-03-06T17:40:19.309626Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_targets = df_train.label.values\nvalid_targets = df_valid.label.values","metadata":{"execution":{"iopub.status.busy":"2022-03-06T17:40:21.844704Z","iopub.execute_input":"2022-03-06T17:40:21.844966Z","iopub.status.idle":"2022-03-06T17:40:21.849030Z","shell.execute_reply.started":"2022-03-06T17:40:21.844937Z","shell.execute_reply":"2022-03-06T17:40:21.848237Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_targets[1]","metadata":{"execution":{"iopub.status.busy":"2022-03-06T17:40:24.246675Z","iopub.execute_input":"2022-03-06T17:40:24.247346Z","iopub.status.idle":"2022-03-06T17:40:24.253290Z","shell.execute_reply.started":"2022-03-06T17:40:24.247298Z","shell.execute_reply":"2022-03-06T17:40:24.252816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install tez\nimport tez\nfrom tez.datasets import ImageDataset\nfrom tez.callbacks import EarlyStopping","metadata":{"execution":{"iopub.status.busy":"2022-03-06T17:40:26.406898Z","iopub.execute_input":"2022-03-06T17:40:26.407367Z","iopub.status.idle":"2022-03-06T17:40:34.427735Z","shell.execute_reply.started":"2022-03-06T17:40:26.407335Z","shell.execute_reply":"2022-03-06T17:40:34.426953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import albumentations\n\ntrain_aug = albumentations.Compose([\n            albumentations.RandomResizedCrop(256, 256),\n            albumentations.Transpose(p=0.5),\n            albumentations.HorizontalFlip(p=0.5),\n            albumentations.VerticalFlip(p=0.5),\n            albumentations.ShiftScaleRotate(p=0.5),\n            albumentations.HueSaturationValue(\n                hue_shift_limit=0.2, \n                sat_shift_limit=0.2, \n                val_shift_limit=0.2, \n                p=0.5\n            ),\n            albumentations.RandomBrightnessContrast(\n                brightness_limit=(-0.1,0.1), \n                contrast_limit=(-0.1, 0.1), \n                p=0.5\n            ),\n            albumentations.Normalize(\n                mean=[0.485, 0.456, 0.406], \n                std=[0.229, 0.224, 0.225], \n                max_pixel_value=255.0, \n                p=1.0\n            ),\n            albumentations.CoarseDropout(p=0.5),\n            albumentations.Cutout(p=0.5)], p=1.)\n  \n        \nvalid_aug = albumentations.Compose([\n            albumentations.CenterCrop(256, 256, p=1.),\n            albumentations.Resize(256, 256),\n            albumentations.Normalize(\n                mean=[0.485, 0.456, 0.406], \n                std=[0.229, 0.224, 0.225], \n                max_pixel_value=255.0, \n                p=1.0\n            )], p=1.)\n\nprint(\"hello\")","metadata":{"execution":{"iopub.status.busy":"2022-03-06T17:40:36.872052Z","iopub.execute_input":"2022-03-06T17:40:36.872757Z","iopub.status.idle":"2022-03-06T17:40:36.886450Z","shell.execute_reply.started":"2022-03-06T17:40:36.872718Z","shell.execute_reply":"2022-03-06T17:40:36.885718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ntrain_dataset = ImageDataset(\n    image_paths=train_images,\n    targets=train_targets,\n    augmentations=train_aug,\n)\n\nvalid_dataset = ImageDataset(\n    image_paths=valid_images,\n    targets=valid_targets,\n    augmentations=valid_aug,\n)","metadata":{"execution":{"iopub.status.busy":"2022-03-06T17:40:40.843137Z","iopub.execute_input":"2022-03-06T17:40:40.843779Z","iopub.status.idle":"2022-03-06T17:40:40.848210Z","shell.execute_reply.started":"2022-03-06T17:40:40.843734Z","shell.execute_reply":"2022-03-06T17:40:40.847554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#forward function should return 3 things, if we are using tez.\n#multiclass classification problem : loss == crossentropy\n\n\nclass LeafModel(tez.Model):\n    def __init__(self, num_classes,pretrained = True):\n        super().__init__()\n\n        self.convnet = torchvision.models.resnet18(pretrained= pretrained)\n\n        #changing last fc layer of resnet 18 as it gives 1000 output features and 512 input\n        #last layer was Linear and changed layer is also same but output ==num_classes\n        self.convnet.fc = nn.Linear(512, num_classes)\n        self.step_scheduler_after = \"epoch\"\n        \n    def loss(self, outputs, targets):\n        if targets is None:\n            return None\n        return nn.CrossEntropyLoss()(outputs,targets)\n\n    def monitor_metrics(self, outputs, targets):\n        if targets is None:\n            return {}\n        outputs = torch.argmax(outputs, dim=1).cpu().detach().numpy()\n        targets = targets.cpu().detach().numpy()\n        accuracy = metrics.accuracy_score(targets, outputs)\n        return {\"accuracy\": accuracy}\n\n    def fetch_optimizer(self):\n        opt = torch.optim.Adam(self.parameters(), lr=3e-4)\n        return opt\n\n    def fetch_scheduler(self):\n        sch = torch.optim.lr_scheduler.StepLR(\n            self.optimizer, step_size =0.7\n        )\n        return sch\n\n    #image and targets from dataset\n    def forward(self,image,targets =None):\n        outputs = self.convnet(image)\n        if targets is not None:\n            #calculate loss and metrics\n            loss = self.loss(outputs, targets)\n            mon_metrics=self.monitor_metrics(outputs, targets)\n            return outputs, loss, mon_metrics\n        return outputs, None, None","metadata":{"execution":{"iopub.status.busy":"2022-03-06T17:40:44.070862Z","iopub.execute_input":"2022-03-06T17:40:44.071093Z","iopub.status.idle":"2022-03-06T17:40:44.081854Z","shell.execute_reply.started":"2022-03-06T17:40:44.071069Z","shell.execute_reply":"2022-03-06T17:40:44.080988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torchvision\n\nmodel = LeafModel(num_classes = dfx.label.nunique(), pretrained = True)\nes = EarlyStopping(\n    monitor=\"valid_accuracy\", model_path=\"model.bin\", patience=3, mode=\"min\"\n)\nmodel.fit(\n    train_dataset,\n    valid_dataset=valid_dataset,\n    train_bs=32,\n    valid_bs=64,\n    device=device,\n    epochs=50,\n    callbacks=[es],\n    fp16=True,\n)\nmodel.save(\"model.bin\")","metadata":{"execution":{"iopub.status.busy":"2022-03-06T17:40:47.807096Z","iopub.execute_input":"2022-03-06T17:40:47.807782Z","iopub.status.idle":"2022-03-06T17:43:46.654058Z","shell.execute_reply.started":"2022-03-06T17:40:47.807742Z","shell.execute_reply":"2022-03-06T17:43:46.653101Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model","metadata":{"execution":{"iopub.status.busy":"2022-02-24T05:52:58.558Z","iopub.status.idle":"2022-02-24T05:52:58.558456Z","shell.execute_reply.started":"2022-02-24T05:52:58.558221Z","shell.execute_reply":"2022-02-24T05:52:58.558246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#It was easy :)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}