{"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 torch\nimport matplotlib.pyplot as plt \nimport torch.nn as nn\nimport torchvision \nimport albumentations as A\nimport cv2\nfrom albumentations import pytorch as ATorch\nimport os\nfrom sklearn import metrics , model_selection\nfrom torch.utils.data import Dataset\nfrom PIL import Image\nfrom torchvision import transforms\nfrom torch.utils.data import DataLoader\nimport copy\nfrom torchvision import transforms\nfrom datetime import datetime\nimport time\n%matplotlib inline","metadata":{"execution":{"iopub.status.busy":"2022-12-16T20:24:15.907882Z","iopub.execute_input":"2022-12-16T20:24:15.908233Z","iopub.status.idle":"2022-12-16T20:24:15.918484Z","shell.execute_reply.started":"2022-12-16T20:24:15.908202Z","shell.execute_reply":"2022-12-16T20:24:15.917512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/cassava-leaf-disease-classification/train.csv')\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2022-12-16T20:24:15.923310Z","iopub.execute_input":"2022-12-16T20:24:15.924192Z","iopub.status.idle":"2022-12-16T20:24:15.954854Z","shell.execute_reply.started":"2022-12-16T20:24:15.924157Z","shell.execute_reply":"2022-12-16T20:24:15.953920Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train , df_val = model_selection.train_test_split(df , test_size =0.1 , random_state = 42,stratify=df.label.values)\ndf_train = df_train.reset_index(drop=True)\ndf_val = df_val.reset_index(drop=True)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-12-16T20:24:15.956775Z","iopub.execute_input":"2022-12-16T20:24:15.957270Z","iopub.status.idle":"2022-12-16T20:24:15.973673Z","shell.execute_reply.started":"2022-12-16T20:24:15.957212Z","shell.execute_reply":"2022-12-16T20:24:15.972759Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train['image_id'] = df_train['image_id'].apply(lambda x : '/kaggle/input/cassava-leaf-disease-classification/train_images/'+x )\ndf_val['image_id'] = df_val['image_id'].apply(lambda x : '/kaggle/input/cassava-leaf-disease-classification/train_images/'+x )","metadata":{"execution":{"iopub.status.busy":"2022-12-16T20:24:15.976562Z","iopub.execute_input":"2022-12-16T20:24:15.976876Z","iopub.status.idle":"2022-12-16T20:24:15.989902Z","shell.execute_reply.started":"2022-12-16T20:24:15.976846Z","shell.execute_reply":"2022-12-16T20:24:15.988977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_transforms():\n    return A.Compose(\n        [\n            A.Resize(224, 224),            \n            A.Rotate(limit=30, border_mode=cv2.BORDER_REPLICATE, p=0.5),\n            A.HorizontalFlip(p=0.5),\n            A.RandomBrightnessContrast(p=0.2),\n            A.Blur(p=0.25),\n            A.Normalize(\n                mean=[0.485, 0.456, 0.406], \n                std=[0.229, 0.224, 0.225], \n                p=1.0\n            ),\n            ATorch.transforms.ToTensorV2(p=1.0),\n        ],\n        p=1.0\n    )\ndef val_transforms():\n    return A.Compose(\n        [\n            A.Resize(224, 224),\n            A.Normalize(\n                mean=[0.485, 0.456, 0.406], \n                std=[0.229, 0.224, 0.225], \n                p=1.0\n            ),\n            ATorch.transforms.ToTensorV2(p=1.0),\n        ],\n        p=1.0\n    )","metadata":{"execution":{"iopub.status.busy":"2022-12-16T20:24:15.992888Z","iopub.execute_input":"2022-12-16T20:24:15.993706Z","iopub.status.idle":"2022-12-16T20:24:16.005413Z","shell.execute_reply.started":"2022-12-16T20:24:15.993669Z","shell.execute_reply":"2022-12-16T20:24:16.004532Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LeafModel(nn.Module):\n    def __init__(self,num_classes, pretrained = True):\n        super().__init__()\n        self.effnet =  torchvision.models.efficientnet_b0(pretrained = pretrained)\n        self.effnet.classifier= nn.Linear(1280,num_classes)\n    def forward(self, image , targets=None):\n        outputs = self.effnet(image)\n        return outputs","metadata":{"execution":{"iopub.status.busy":"2022-12-16T20:24:16.006982Z","iopub.execute_input":"2022-12-16T20:24:16.007471Z","iopub.status.idle":"2022-12-16T20:24:16.016601Z","shell.execute_reply.started":"2022-12-16T20:24:16.007424Z","shell.execute_reply":"2022-12-16T20:24:16.015593Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model =  LeafModel(num_classes = df.label.nunique() , pretrained = False ).to(torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\"))","metadata":{"execution":{"iopub.status.busy":"2022-12-16T20:24:16.018690Z","iopub.execute_input":"2022-12-16T20:24:16.019094Z","iopub.status.idle":"2022-12-16T20:24:16.150670Z","shell.execute_reply.started":"2022-12-16T20:24:16.019061Z","shell.execute_reply":"2022-12-16T20:24:16.149642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CassavaLeafDataset(Dataset):\n    \n    def __init__(self, metadata_csv, transform = None):\n        super(CassavaLeafDataset, self).__init__()\n        \n        self.df = metadata_csv \n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        img_path = self.df['image_id'].iloc[index]\n        class_label = torch.tensor(int(self.df['label'].iloc[index]))\n        image = cv2.imread(img_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)    \n        if self.transform is not None:\n            image = self.transform(image=image)['image']\n            \n        return (image, class_label)","metadata":{"execution":{"iopub.status.busy":"2022-12-16T20:24:16.151940Z","iopub.execute_input":"2022-12-16T20:24:16.152403Z","iopub.status.idle":"2022-12-16T20:24:16.160429Z","shell.execute_reply.started":"2022-12-16T20:24:16.152366Z","shell.execute_reply":"2022-12-16T20:24:16.159465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = CassavaLeafDataset(df_train, train_transforms())\nval_dataset = CassavaLeafDataset(df_val, val_transforms())","metadata":{"execution":{"iopub.status.busy":"2022-12-16T20:24:16.163352Z","iopub.execute_input":"2022-12-16T20:24:16.164059Z","iopub.status.idle":"2022-12-16T20:24:16.182112Z","shell.execute_reply.started":"2022-12-16T20:24:16.164024Z","shell.execute_reply":"2022-12-16T20:24:16.181281Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_SIZE = 16\nNUM_OF_CLASSES = 5\nepochs = 20\nMODEL_SAVE_PATH = \"best_model.torch\"\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\ntrain_loader = DataLoader(dataset = train_dataset, batch_size = BATCH_SIZE, shuffle = True)\nval_loader = DataLoader(dataset = val_dataset, batch_size = BATCH_SIZE, shuffle = False)","metadata":{"execution":{"iopub.status.busy":"2022-12-16T20:24:16.183462Z","iopub.execute_input":"2022-12-16T20:24:16.183888Z","iopub.status.idle":"2022-12-16T20:24:16.194947Z","shell.execute_reply.started":"2022-12-16T20:24:16.183850Z","shell.execute_reply":"2022-12-16T20:24:16.194138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=0.001)\nexp_lr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.1)","metadata":{"execution":{"iopub.status.busy":"2022-12-16T20:24:16.196448Z","iopub.execute_input":"2022-12-16T20:24:16.197179Z","iopub.status.idle":"2022-12-16T20:24:16.206208Z","shell.execute_reply.started":"2022-12-16T20:24:16.197139Z","shell.execute_reply":"2022-12-16T20:24:16.205286Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Trainer:\n    def __init__(self, model, device, score, loss, optimizer = None):\n        self.model = model\n        self.device = device\n        self.score = score\n        self.loss = loss\n        self.optimizer = optimizer\n        \n    def run(self, train_dataloader, val_dataloader, epochs, save_path):\n        best_score = 0\n        \n        for epoch in range(epochs):\n            train_loss, train_score, train_time = self.train_epoch(train_dataloader)\n            val_loss, val_score, valid_time = self.val_epoch(val_dataloader)\n            \n            print(\n                f\"Epoch {epoch+1}\",\n                f\"Train Loss: {train_loss:.3f}, Train Accuracy: {train_score:.3f}, Time: {train_time} sec.\",\n                f\"Validation Loss: {val_loss:.3f}, Validation Accuracy: {val_score:.3f}, Time: {valid_time} sec.\",\n                f\"------------------------------\",\n                sep=\"\\n\",\n            )\n            \n            if best_score < val_score:\n                best_score = val_score\n                print(\"updating weights\")\n                self.save_model(save_path)\n    def train_epoch(self, dataloader):\n        self.model.train()\n        t = time.time()\n        train_loss, train_accuracy = 0, 0\n        \n        for batch, (X, y) in enumerate(dataloader):\n            X, y = X.to(self.device), y.to(self.device)\n            pred = self.model(X)\n            loss = self.loss(pred, y)\n            accuracy = self.score(pred.detach().cpu().numpy(), y.detach().cpu().numpy())\n            train_accuracy, train_loss = update_metrics(train_accuracy, accuracy, train_loss, loss, batch)\n\n            self.optimizer.zero_grad()\n            loss.backward()\n            self.optimizer.step()\n            \n        return train_loss, train_accuracy, int(time.time() - t)\n    def val_epoch(self, dataloader):\n        self.model.eval()\n        t = time.time()\n        val_loss, val_accuracy = 0, 0\n\n        with torch.no_grad():\n            for batch, (X, y) in enumerate(dataloader):\n                X, y = X.to(self.device), y.to(self.device)\n                pred = self.model(X)\n                loss = self.loss(pred, y)\n                accuracy = self.score(pred, y)\n                val_accuracy, val_loss = update_metrics(val_accuracy, accuracy, val_loss, loss, batch)\n                \n        return val_loss, val_accuracy, int(time.time() - t)\n\n    \n    def save_model(self, save_path):\n        torch.save(\n            {\n                \"model_state_dict\": self.model.state_dict(),\n            },\n            save_path,\n        )","metadata":{"execution":{"iopub.status.busy":"2022-12-16T20:24:16.209726Z","iopub.execute_input":"2022-12-16T20:24:16.210000Z","iopub.status.idle":"2022-12-16T20:24:16.225783Z","shell.execute_reply.started":"2022-12-16T20:24:16.209976Z","shell.execute_reply":"2022-12-16T20:24:16.224901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def accuracy_fn(y_pred, y):\n    return (y_pred.argmax(1) == y).sum().item() / y.shape[0]","metadata":{"execution":{"iopub.status.busy":"2022-12-16T20:24:16.227398Z","iopub.execute_input":"2022-12-16T20:24:16.227973Z","iopub.status.idle":"2022-12-16T20:24:16.251250Z","shell.execute_reply.started":"2022-12-16T20:24:16.227935Z","shell.execute_reply":"2022-12-16T20:24:16.250279Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def update_metrics(mean_score, score, mean_loss, loss, step):\n    mean_score = (mean_score * step + score)/(step+1)\n    mean_loss = (mean_loss * step + loss.detach().cpu().item())/(step+1)\n    return mean_score, mean_loss","metadata":{"execution":{"iopub.status.busy":"2022-12-16T20:24:16.254396Z","iopub.execute_input":"2022-12-16T20:24:16.254682Z","iopub.status.idle":"2022-12-16T20:24:16.264337Z","shell.execute_reply.started":"2022-12-16T20:24:16.254657Z","shell.execute_reply":"2022-12-16T20:24:16.263392Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer = Trainer(model, device, accuracy_fn, criterion, optimizer)\ntrainer.run(train_loader, val_loader, epochs, MODEL_SAVE_PATH)","metadata":{"execution":{"iopub.status.busy":"2022-12-16T20:24:16.265751Z","iopub.execute_input":"2022-12-16T20:24:16.266306Z","iopub.status.idle":"2022-12-16T21:06:26.434049Z","shell.execute_reply.started":"2022-12-16T20:24:16.266216Z","shell.execute_reply":"2022-12-16T21:06:26.431990Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_model = LeafModel(num_classes = df.label.nunique() ,pretrained = False).to(torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\"))\ncheck_point = torch.load('/kaggle/working/best_model.torch')\nsub_model.load_state_dict(check_point['model_state_dict'])","metadata":{"execution":{"iopub.status.busy":"2022-12-16T21:06:26.435147Z","iopub.status.idle":"2022-12-16T21:06:26.436054Z","shell.execute_reply.started":"2022-12-16T21:06:26.435787Z","shell.execute_reply":"2022-12-16T21:06:26.435814Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df = pd.DataFrame()\ntest_dir = '/kaggle/input/cassava-leaf-disease-classification/test_images'\nsubmission_df['image_id'] = pd.Series('/kaggle/input/cassava-leaf-disease-classification/test_images/2216849948.jpg')\nsubmission_df['label'] = 0\nsubmission_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-12-16T21:06:26.438043Z","iopub.status.idle":"2022-12-16T21:06:26.438775Z","shell.execute_reply.started":"2022-12-16T21:06:26.438476Z","shell.execute_reply":"2022-12-16T21:06:26.438504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 1\ndef test_transforms():\n    return A.Compose(\n        [\n            A.Resize(224, 224),\n            A.Normalize(\n                mean=[0.485, 0.456, 0.406], \n                std=[0.229, 0.224, 0.225], \n                p=1.0\n            ),\n            ATorch.transforms.ToTensorV2(p=1.0),\n        ],\n        p=1.0\n    )\ntest_dataset = val_dataset = CassavaLeafDataset(submission_df, test_transforms())\ntest_loader = DataLoader(dataset = test_dataset, batch_size = batch_size, shuffle = False)","metadata":{"execution":{"iopub.status.busy":"2022-12-16T21:06:26.440384Z","iopub.status.idle":"2022-12-16T21:06:26.440855Z","shell.execute_reply.started":"2022-12-16T21:06:26.440611Z","shell.execute_reply":"2022-12-16T21:06:26.440634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def inference_one_epoch(model, data_loader, device):\n    model.eval()\n\n    image_preds_all = []\n    \n    for batch, (X, y) in enumerate(data_loader):\n        X, y = X.to(device), y.to(device)\n        preds = model(X)\n        image_preds_all += [torch.argmax(preds, 1).detach().cpu().numpy()]\n        \n    \n    image_preds_all = np.concatenate(image_preds_all, axis=0)\n    return image_preds_all","metadata":{"execution":{"iopub.status.busy":"2022-12-16T21:06:26.442339Z","iopub.status.idle":"2022-12-16T21:06:26.443605Z","shell.execute_reply.started":"2022-12-16T21:06:26.443346Z","shell.execute_reply":"2022-12-16T21:06:26.443371Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_preds = inference_one_epoch (sub_model ,test_loader, device )\nprint(sub_preds)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df['label'] = sub_preds\nsubmission_df['image_id'] = submission_df['image_id'].apply(lambda x : x[-14:])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df.to_csv('submission.csv', index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}