{"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 tez\n!pip install efficientnet-pytorch","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-05-09T07:28:49.796782Z","iopub.execute_input":"2022-05-09T07:28:49.797343Z","iopub.status.idle":"2022-05-09T07:33:49.848655Z","shell.execute_reply.started":"2022-05-09T07:28:49.797292Z","shell.execute_reply":"2022-05-09T07:33:49.847457Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tez_path = '../input/tez-lib/'\neffnet_path = '../input/efficientnet-pytorch/'\nimport sys\nsys.path.append(tez_path)\nsys.path.append(effnet_path)","metadata":{"execution":{"iopub.status.busy":"2022-05-09T07:33:49.885764Z","iopub.execute_input":"2022-05-09T07:33:49.886062Z","iopub.status.idle":"2022-05-09T07:33:49.894035Z","shell.execute_reply.started":"2022-05-09T07:33:49.886034Z","shell.execute_reply":"2022-05-09T07:33:49.893107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport albumentations\nimport pandas as pd\nimport numpy as np\n\nimport tez\nfrom tez.datasets import ImageDataset\n\nimport torch\nimport torch.nn as nn\nfrom torch.nn import functional as F\n\nfrom efficientnet_pytorch import EfficientNet","metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","execution":{"iopub.status.busy":"2022-05-09T07:35:21.586507Z","iopub.execute_input":"2022-05-09T07:35:21.586885Z","iopub.status.idle":"2022-05-09T07:35:21.592582Z","shell.execute_reply.started":"2022-05-09T07:35:21.586853Z","shell.execute_reply":"2022-05-09T07:35:21.591553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LeafModel(tez.Model):\n    def __init__(self, num_classes):\n        super().__init__()\n\n        self.effnet = EfficientNet.from_name(\"efficientnet-b4\")\n        self.dropout = nn.Dropout(0.1)\n        self.out = nn.Linear(1792, num_classes)\n        self.step_scheduler_after = \"epoch\"\n\n    def forward(self, image, targets=None):\n        batch_size, _, _, _ = image.shape\n\n        x = self.effnet.extract_features(image)\n        x = F.adaptive_avg_pool2d(x, 1).reshape(batch_size, -1)\n        outputs = self.out(self.dropout(x))\n        return outputs, None, None","metadata":{"execution":{"iopub.status.busy":"2022-05-09T07:35:30.753393Z","iopub.execute_input":"2022-05-09T07:35:30.753743Z","iopub.status.idle":"2022-05-09T07:35:30.763414Z","shell.execute_reply.started":"2022-05-09T07:35:30.753714Z","shell.execute_reply":"2022-05-09T07:35:30.762242Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# augmentations taken from: https://www.kaggle.com/khyeh0719/pytorch-efficientnet-baseline-inference-tta\ntest_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.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], p=1.)","metadata":{"execution":{"iopub.status.busy":"2022-05-09T07:35:36.539305Z","iopub.execute_input":"2022-05-09T07:35:36.539698Z","iopub.status.idle":"2022-05-09T07:35:36.550347Z","shell.execute_reply.started":"2022-05-09T07:35:36.539666Z","shell.execute_reply":"2022-05-09T07:35:36.54909Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import cv2\nimport numpy as np\nimport torch\nfrom PIL import Image, ImageFile\n\n\nImageFile.LOAD_TRUNCATED_IMAGES = True\n\n\nclass ImageDataset:\n    def __init__(\n        self,\n        image_paths,\n        targets,\n        augmentations=None,\n        backend=\"pil\",\n        channel_first=True,\n        grayscale=False,\n    ):\n        \"\"\"\n        :param image_paths: list of paths to images\n        :param targets: numpy array\n        :param augmentations: albumentations augmentations\n        \"\"\"\n        self.image_paths = image_paths\n        self.targets = targets\n        self.augmentations = augmentations\n        self.backend = backend\n        self.channel_first = channel_first\n        self.grayscale = grayscale\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, item):\n        targets = self.targets[item]\n        if self.backend == \"pil\":\n            image = Image.open(self.image_paths[item])\n            image = np.array(image)\n            if self.augmentations is not None:\n                augmented = self.augmentations(image=image)\n                image = augmented[\"image\"]\n        elif self.backend == \"cv2\":\n            if self.grayscale is False:\n                image = cv2.imread(self.image_paths[item])\n                image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n            else:\n                image = cv2.imread(self.image_paths[item], cv2.IMREAD_GRAYSCALE)\n            if self.augmentations is not None:\n                augmented = self.augmentations(image=image)\n                image = augmented[\"image\"]\n        else:\n            raise Exception(\"Backend not implemented\")\n        if self.channel_first is True and self.grayscale is False:\n            image = np.transpose(image, (2, 0, 1)).astype(np.float32)\n\n        image_tensor = torch.tensor(image)\n        if self.grayscale:\n            image_tensor = image_tensor.unsqueeze(0)\n        return {\n            \"image\": image_tensor,\n            \"targets\": torch.tensor(targets),\n        }","metadata":{"execution":{"iopub.status.busy":"2022-05-09T07:35:48.058727Z","iopub.execute_input":"2022-05-09T07:35:48.059419Z","iopub.status.idle":"2022-05-09T07:35:48.078578Z","shell.execute_reply.started":"2022-05-09T07:35:48.059365Z","shell.execute_reply":"2022-05-09T07:35:48.07765Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dfx = pd.read_csv(\"../input/cassava-leaf-disease-classification/sample_submission.csv\")\nimage_path = \"../input/cassava-leaf-disease-classification/test_images/\"\ntest_image_paths = [os.path.join(image_path, x) for x in dfx.image_id.values]\n# fake targets\ntest_targets = dfx.label.values\ntest_dataset = ImageDataset(\n    image_paths=test_image_paths,\n    targets=test_targets,\n    augmentations= test_aug,\n)","metadata":{"execution":{"iopub.status.busy":"2022-05-09T07:41:46.884888Z","iopub.execute_input":"2022-05-09T07:41:46.885267Z","iopub.status.idle":"2022-05-09T07:41:46.907076Z","shell.execute_reply.started":"2022-05-09T07:41:46.885232Z","shell.execute_reply":"2022-05-09T07:41:46.906263Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dfx = pd.read_csv(\"../input/cassava-leaf-disease-classification/train.csv\")\nmodel = LeafModel(num_classes=train_dfx.label.nunique())\nmodel.load(\"../input/leafmodel/model.bin\")","metadata":{"execution":{"iopub.status.busy":"2022-05-09T07:41:53.968613Z","iopub.execute_input":"2022-05-09T07:41:53.968998Z","iopub.status.idle":"2022-05-09T07:41:54.303557Z","shell.execute_reply.started":"2022-05-09T07:41:53.968965Z","shell.execute_reply":"2022-05-09T07:41:54.301074Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# run inference 5 times\nfinal_preds = None\nfor j in range(5):\n    preds = model.predict(test_dataset, batch_size=32, n_jobs=-1, device=\"cuda\")\n    temp_preds = None\n    for p in preds:\n        if temp_preds is None:\n            temp_preds = p\n        else:\n            temp_preds = np.vstack((temp_preds, p))\n    if final_preds is None:\n        final_preds = temp_preds\n    else:\n        final_preds += temp_preds\nfinal_preds /= 5","metadata":{"execution":{"iopub.status.busy":"2022-05-09T07:42:05.712696Z","iopub.execute_input":"2022-05-09T07:42:05.713057Z","iopub.status.idle":"2022-05-09T07:42:05.736448Z","shell.execute_reply.started":"2022-05-09T07:42:05.713027Z","shell.execute_reply":"2022-05-09T07:42:05.734644Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"final_preds = final_preds.argmax(axis=1)","metadata":{"execution":{"iopub.status.busy":"2022-05-09T07:28:45.23777Z","iopub.status.idle":"2022-05-09T07:28:45.238437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dfx.label = final_preds","metadata":{"execution":{"iopub.status.busy":"2022-05-09T07:28:45.239495Z","iopub.status.idle":"2022-05-09T07:28:45.240227Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dfx.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2022-05-09T07:28:45.241334Z","iopub.status.idle":"2022-05-09T07:28:45.24196Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}