{"cells":[{"metadata":{"trusted":true},"cell_type":"code","source":"!pip install efficientnet_pytorch","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"tez_path = \"../input/tez-lib/\"\nimport sys\nsys.path.append(tez_path)\n\nout_dir = \"models\"\n!mkdir -p $out_dir\n\nout1_dir = \"snapshot_models\"\n!mkdir -p $out1_dir","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import os\nimport argparse\nimport albumentations\nimport matplotlib.pyplot as plt\nimport pandas as pd\nfrom efficientnet_pytorch import EfficientNet\n\nimport tez \nfrom tez.model.model import Model\nfrom tez.datasets import ImageDataset\nfrom tez.callbacks import EarlyStopping, Callback\n\nimport torch\nimport torch.nn as nn\nfrom torch.nn import functional as F\n\nimport torchvision\n\nfrom sklearn import metrics, model_selection \n\n%matplotlib inline","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"train_dfx = pd.read_csv(\"../input/cassava-leaf-disease-classification/train.csv\")\ntrain_dfx.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_dfx.label.value_counts()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_set,valid_set= model_selection.train_test_split(\n                            train_dfx,\n                            test_size=0.1,\n                            random_state=42,\n                            stratify=train_dfx.label.values)\ntrain_set=train_set.reset_index(drop=True)\nvalid_set=valid_set.reset_index(drop=True)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_set.shape\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"valid_set.shape","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"image_path = \"../input/cassava-leaf-disease-classification/train_images\"\n\ntrain_img_path = [os.path.join(image_path,x) for x in train_set.image_id.values]\nvalid_img_path = [os.path.join(image_path,x) for x in valid_set.image_id.values]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_img_path[:10]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_target = train_set.label.values\nvalid_target = valid_set.label.values\n\n\ntrain_dataset = ImageDataset(\n    image_paths = train_img_path,\n    targets=train_target,\n    augmentations=None,\n    \n)\n\nvalid_dataset = ImageDataset(\n    image_paths = valid_img_path,\n    targets= valid_target,\n    augmentations=None,\n    \n)\n\ndef img_plot(image_dict):\n    img_tensor = image_dict[\"image\"]\n    target = image_dict[\"targets\"]\n    plt.figure(figsize=(5,5))\n    image = img_tensor.permute(1,2,0)/255\n    plt.imshow(image)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_dataset[0]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"img_plot(train_dataset[50])","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":""},{"metadata":{"trusted":true},"cell_type":"code","source":"train_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.3\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.2),\n            #albumentations.Cutout(p=0.5)\n            ], 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.0)\n\ntrain_dataset = ImageDataset(\n    image_paths = train_img_path,\n    targets=train_target,\n    augmentations=train_aug,\n    \n)\n\nvalid_dataset = ImageDataset(\n    image_paths = valid_img_path,\n    targets= valid_target,\n    augmentations=valid_aug,\n    \n)\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"img_plot(train_dataset[50])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class LeafModel(Model):\n    def __init__(self,num_classes):\n        super().__init__()\n        self.effnet= EfficientNet.from_pretrained('efficientnet-b3')\n        self.dropout1 = nn.Dropout(0.2)\n        self.lin1 = nn.Linear(1536, 256)\n        self.bn1=nn.BatchNorm1d(256)\n        self.dropout2 = nn.Dropout(0.1)\n        self.lin2 = nn.Linear(256,num_classes)\n        self.step_scheduler_after = \"epoch\"\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.CosineAnnealingWarmRestarts(\n            self.optimizer, T_0=10, T_mult=1, eta_min=1e-6, last_epoch=-1\n        )\n        return sch    \n        \n    def forward(self,image,targets=None):\n        batch_size, _, _, _ = image.shape\n        x = self.effnet.extract_features(image)\n        x = F.adaptive_avg_pool2d(x, 1).reshape(batch_size, -1)\n        \n        x = F.relu((self.lin1(self.dropout1(x))))\n        x = self.dropout2(self.bn1(x))\n        outputs = F.softmax(self.lin2(x))\n        \n        if targets is not None:\n            loss = nn.CrossEntropyLoss()(outputs, targets) \n            metrics = self.monitor_metrics(outputs, targets)\n            return outputs, loss, metrics\n        return outputs, None, None\n        ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model = LeafModel(num_classes=train_dfx.label.nunique())\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"print(model)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class ModelSave(Callback):\n    def __init__(self, model_path):\n        self.model_path = model_path\n        self.model_temp_path = model_path\n        self.counter = 0;\n    \n    def on_epoch_end(self, model):\n        if(self.counter >= 3):\n            self.model_temp_path = self.model_path + 'epoch' + f'{self.counter}' + '.bin' \n            model.save(self.model_temp_path)\n            #print('saving model ----->>  \"./SnapShot/model_\" + f'{self.counter}')\n        self.counter += 1\n        \n       ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"es = EarlyStopping(\n        monitor=\"valid_loss\",\n        model_path='model.bin',\n        patience=4,\n        mode=\"min\",\n    )\n\nms = ModelSave(model_path = './snapshot_models/')\n\nmodel.fit(\n        train_dataset,\n        valid_dataset=valid_dataset,\n        train_bs=32,\n        valid_bs=64,\n        device=\"cuda\",\n        epochs=15,\n        callbacks=[es, ms],\n    )\n#model.save(\"./models/model.bin\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"test_dfx = pd.read_csv(\"../input/cassava-leaf-disease-classification/sample_submission.csv\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"test_img_path = \"../input/cassava-leaf-disease-classification/test_images\"\ntest_img = [os.path.join(test_img_path,x) for x in test_dfx.image_id.values]\ntest_img","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"\"\"\"test_target = test_dfx.label.values\ntest_aug = albumentations.Compose(\n    [\n        albumentations.RandomResizedCrop(256,256),\n        albumentations.Resize(256,256),\n        albumentations.Transpose(p=0.5),\n        albumentations.HorizontalFlip(p=0.5),\n        albumentations.VerticalFlip(p=0.5),\n    ]\n)\ntest_dataset = ImageDataset(\n    image_paths = test_img,\n    targets=test_target,\n    augmentations=test_aug,\n)\n\npreds = model.predict(test_dataset, batch_size=64, n_jobs=1)\nfinal_preds=None\nfor p in preds:\n    if final_preds is None:\n        final_preds=p\n    else:\n        final_preds= np.vstack((final_preds,p))\nfinal_preds=final_preds.argmax(axis=1)\ntest_dfx.label= final_preds\ntest_dfx.to_csv(\"submission.csv\",index=False)\n    \n\"\"\"","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}