{"cells":[{"metadata":{"trusted":true},"cell_type":"code","source":"import os\nimport gc\nimport cv2\nimport time\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\n\nimport torch\nimport torch.nn as nn\nimport albumentations as A\n\nfrom sklearn.model_selection import StratifiedKFold","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_df=pd.read_csv('../input/cassava-leaf-disease-classification/train.csv')\ntrain_df.head()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Loading Pretrained Model"},{"metadata":{"trusted":true},"cell_type":"code","source":"device=torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\npretrained=torch.hub.load('pytorch/vision:v0.6.0', 'resnext50_32x4d', pretrained=True)\n","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# configuration"},{"metadata":{"trusted":true},"cell_type":"code","source":"CFG = {\n    \"IMG_SIZE\": 512,\n    \"BATCH_SIZE\": 16,\n    \"IMG_FOLDER\": \"../input/cassava-leaf-disease-classification/train_images\",\n    \"EPOCHS\": 15,\n    \"NUM_FOLDS\": 5,\n    \"device\": device\n}","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Helper Functions"},{"metadata":{"trusted":true},"cell_type":"code","source":"def read_image(imgpath):\n    img=cv2.imread(imgpath)\n    img=cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    return img\n\ndef train_augmentation():\n    transform=A.Compose([\n        A.RandomResizedCrop(CFG[\"IMG_SIZE\"], CFG[\"IMG_SIZE\"], p=1.0,\n                            scale=[0.75, 0.95], ratio=[0.8, 1.33]),\n        A.RandomBrightnessContrast(brightness_limit=[-0.1, 0.1], contrast_limit=[-0.1, 0.1]),\n        A.HorizontalFlip(),\n        A.VerticalFlip(),\n        A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225))\n    ])\n    return transform\n\ndef val_augmentation():\n    transform=A.Compose([\n        A.Resize(CFG[\"IMG_SIZE\"], CFG[\"IMG_SIZE\"], p=1.0),\n        A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225))\n    ])\n    return transform","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Dataset"},{"metadata":{"trusted":true},"cell_type":"code","source":"class CassavaDataset(torch.utils.data.Dataset):\n    def __init__(self, df, img_folder, augmentation=None):\n        super().__init__()\n        self.df=df\n        self.img_folder=img_folder\n        self.augmentation=augmentation\n    def __len__(self):\n        return len(self.df)\n    def __getitem__(self, idx):\n        row=self.df.iloc[idx]\n        image_id=row.image_id\n        label=row.label\n        ytrue=torch.zeros(5)\n        \n        image_path=os.path.join(self.img_folder, image_id)\n        img=read_image(image_path)\n        if self.augmentation:\n            img=self.augmentation(image=img)['image']\n        img=torch.tensor(img).permute(2, 1, 0)\n        ytrue[label]=1.0\n        return (img, ytrue)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Model"},{"metadata":{"trusted":true},"cell_type":"code","source":"class CassavaModel(nn.Module):\n    def __init__(self, _backbone):\n        super(CassavaModel, self).__init__()\n        self._backbone=_backbone\n        self._backbone.fc=nn.Linear(in_features=_backbone.fc.in_features, \n                                     out_features=5,\n                                     bias=True)      \n        \n    def forward(self, x):\n        x=self._backbone(x)\n        return x","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"fold_map={}\ntrain_transform=train_augmentation()\nval_transform=val_augmentation()\n\nskf=StratifiedKFold(n_splits=CFG[\"NUM_FOLDS\"], shuffle=True, random_state=42)\nfor idx, (train_index, val_index) in enumerate(skf.split(train_df.image_id, train_df.label)):\n    fold_train_df=train_df.iloc[train_index].copy()\n    fold_val_df=train_df.iloc[val_index].copy()    \n    fold_map[idx]=(fold_train_df, fold_val_df)\n","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Train Model"},{"metadata":{"trusted":true},"cell_type":"code","source":"def train_epoch(train_dataloader, model, optim, criterion, device):\n    train_loss=0.0\n    start_time=time.time()\n    for (img, label) in train_dataloader:\n        img=img.to(device)\n        label=label.to(device)\n        optim.zero_grad()\n        \n        yout=model(img)\n        loss_=criterion(yout, label)\n        loss_.backward()\n        optim.step()\n        train_loss+=loss_.item()\n        \n        del img\n        del label\n        gc.collect()\n    end_time=time.time()\n    print('TrainLoss:', train_loss)\n    print('Train Epoch Time:')\n    print( (end_time-start_time)/60 )\n    return train_loss/len(train_dataloader)\n        \n\ndef val_epoch(val_dataloader, model, criterion, device):\n    val_loss=0.0\n    start_time=time.time()\n    with torch.no_grad():\n        for (img, label) in val_dataloader:\n            img=img.to(device)\n            label=label.to(device)\n            yout=model(img)\n            \n            loss_=criterion(yout, label)\n            val_loss+=loss_.item()\n            del img\n            del label\n            gc.collect()\n    end_time=time.time()\n    print('Val Loss:', val_loss)\n    print('Time Taken for validation')\n    print( (end_time-start_time)/60 )\n    return val_loss/len(val_dataloader)\n            \n\ndef train_model(train_dataloader, val_dataloder, model, optim, criterion, device):\n    best_loss=None\n    for i in range(CFG[\"EPOCHS\"]):\n        train_loss=train_epoch(train_dataloader, model, optim, criterion, device)\n        val_loss=val_epoch(val_dataloader, model, criterion, device)\n        \n        if (best_loss is None) or (best_loss > val_loss):\n            best_loss=val_loss\n            torch.save(model.state_dict(), \"model.pth\")\n        print('Epoch:{} | TrainLoss: {:.4f} | Val Loss: {:.4f}'.format(i+1, train_loss, val_loss))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model=CassavaModel(pretrained).to(CFG[\"device\"])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"for fold_id in range(CFG[\"NUM_FOLDS\"]):\n    if fold_id==1:\n        break\n    fold_train_df=fold_map[fold_id][0]\n    fold_val_df=fold_map[fold_id][1]\n    \n    train_dataset=CassavaDataset(fold_train_df, CFG[\"IMG_FOLDER\"], train_transform)\n    val_dataset=CassavaDataset(fold_val_df, CFG[\"IMG_FOLDER\"], val_transform)\n    \n    train_dataloader=torch.utils.data.DataLoader(train_dataset, num_workers=4, \n                                                 shuffle=True, \n                                                 batch_size=CFG[\"BATCH_SIZE\"],\n                                                 pin_memory=True\n                                                )\n    val_dataloader=torch.utils.data.DataLoader(val_dataset, num_workers=4,\n                                               shuffle=False,\n                                               batch_size=CFG[\"BATCH_SIZE\"],\n                                               pin_memory=True\n                                              )\n    \n    \n    \n    #Model\n    model=CassavaModel(pretrained).to(CFG[\"device\"])\n    optim=torch.optim.Adam(model.parameters())\n    criterion=nn.BCEWithLogitsLoss()\n    \n    train_model(train_dataloader, val_dataloader, model, optim, criterion, CFG[\"device\"])","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Inference"},{"metadata":{"trusted":true},"cell_type":"code","source":"test_folder=\"../input/cassava-leaf-disease-classification/test_images\"\nclass TestDataset(torch.utils.data.Dataset):\n    def __init__(self, augmentation):\n        self.test_images=os.listdir(test_folder)\n        self.augmentation=augmentation\n    def __len__(self):\n        return len(self.test_images)\n    def __getitem__(self, idx):\n        image_id=self.test_images[idx]\n        image_path=os.path.join(test_folder, image_id)\n        img=read_image(image_path)\n        if self.augmentation:\n            img=self.augmentation(image=img)['image']\n        img=torch.tensor(img).permute(2, 1, 0)\n        return (image_id, img)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model.load_state_dict(torch.load('model.pth'))\ndevice=CFG[\"device\"]\nmodel=model.to(device)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"pred=nn.Softmax(dim=1)\ntest_dataset=TestDataset(val_transform)\ntest_dataloader=torch.utils.data.DataLoader(test_dataset, \n                                            shuffle=False,\n                                            pin_memory=True,\n                                            batch_size=16,\n                                            num_workers=4)\n\nsubmission_data=[]\nwith torch.no_grad():\n    for (image_id,img) in test_dataloader:\n        img=img.to(device)\n        yout=model(img)\n        yout=pred(yout)\n        ypred=torch.argmax(yout, dim=1).cpu().numpy()\n        \n        for i in range(img.shape[0]):\n            submission_data.append({\n                'image_id': image_id[i],\n                'label': ypred[i]\n            })","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"submission_df=pd.DataFrame.from_dict(submission_data)\nsubmission_df.to_csv('submission.csv', index=False)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"submission_df.head()","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}