{"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":"markdown","source":"# reference\n* This notebook help us build the general stucture of inference part.\n* [Ensemble: Resnext50_32x4d + Efficientnet = 0.903](https://www.kaggle.com/code/nanamiyvonne/ensemble-resnext50-32x4d-efficientnet-0-903/edit)\n* [[No TTA] Cassava Resnext50_32x4d Inference lb0.903](https://www.kaggle.com/piantic/no-tta-cassava-resnext50-32x4d-inference-lb0-903/output)\n* [Clean_Inference_Kernel_8xTTA_LB902](https://www.kaggle.com/underwearfitting/clean-inference-kernel-8xtta-lb902/data)","metadata":{}},{"cell_type":"code","source":"#for efficientnet\n# tf_efficientnet_b8_fold_.. are provided in zip folder\nBATCH_SIZE = 1\nimage_size = 512\nenet_type = ['tf_efficientnet_b8'] * 5\nmodel_path = ['../input/cassava-models/tf_efficientnet_b8_fold_0_7',\n              '../input/cassava-models/tf_efficientnet_b8_fold_1_7',\n              '../input/cassava-models/tf_efficientnet_b8_fold_2_7',\n              '../input/cassava-models/tf_efficientnet_b8_fold_3_7',\n              '../input/cassava-models/tf_efficientnet_b8_fold_4_7']","metadata":{"execution":{"iopub.status.busy":"2023-04-09T00:33:14.470989Z","iopub.execute_input":"2023-04-09T00:33:14.471311Z","iopub.status.idle":"2023-04-09T00:33:14.476215Z","shell.execute_reply.started":"2023-04-09T00:33:14.471279Z","shell.execute_reply":"2023-04-09T00:33:14.475357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport sys\nsys.path.append('../input/pytorch-image-models/pytorch-image-models-master')\nimport time\nimport random\nimport albumentations\nfrom contextlib import contextmanager\nfrom scipy.special import softmax\nimport numpy as np\nimport pandas as pd\nfrom sklearn.metrics import accuracy_score\nfrom tqdm.auto import tqdm\nimport cv2\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader, Dataset\nimport timm\n\nimport warnings\nwarnings.filterwarnings('ignore')\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2023-04-09T00:33:14.482401Z","iopub.execute_input":"2023-04-09T00:33:14.482694Z","iopub.status.idle":"2023-04-09T00:33:14.491850Z","shell.execute_reply.started":"2023-04-09T00:33:14.482667Z","shell.execute_reply":"2023-04-09T00:33:14.491115Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transforms_valid = albumentations.Compose([\n    albumentations.CenterCrop(image_size, image_size, p=1),\n    albumentations.Resize(image_size, image_size),\n    albumentations.Normalize()\n])","metadata":{"execution":{"iopub.status.busy":"2023-04-09T00:33:14.493548Z","iopub.execute_input":"2023-04-09T00:33:14.494046Z","iopub.status.idle":"2023-04-09T00:33:14.502195Z","shell.execute_reply.started":"2023-04-09T00:33:14.494008Z","shell.execute_reply":"2023-04-09T00:33:14.501450Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Directory settings ","metadata":{}},{"cell_type":"code","source":"OUTPUT_DIR = './'\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)\n\nTRAIN_PATH = '../input/cassava-leaf-disease-classification/train_images'\nTEST_PATH = '../input/cassava-leaf-disease-classification/test_images'\n\n#RestNet\nMODEL_DIR = '../input/cassava-resnext50-32x4d-weights/'","metadata":{"execution":{"iopub.status.busy":"2023-04-09T00:33:14.504108Z","iopub.execute_input":"2023-04-09T00:33:14.504531Z","iopub.status.idle":"2023-04-09T00:33:14.513819Z","shell.execute_reply.started":"2023-04-09T00:33:14.504494Z","shell.execute_reply":"2023-04-09T00:33:14.512939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Utils for Resnext","metadata":{}},{"cell_type":"code","source":"def get_score(y_true, y_pred):\n    return accuracy_score(y_true, y_pred)\n\n\n@contextmanager\ndef timer(name):\n    t0 = time.time()\n    LOGGER.info(f'[{name}] start')\n    yield\n    LOGGER.info(f'[{name}] done in {time.time() - t0:.0f} s.')\n\n\ndef init_logger(log_file=OUTPUT_DIR+'inference.log'):\n    from logging import getLogger, INFO, FileHandler,  Formatter,  StreamHandler\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=log_file)\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\n\nLOGGER = init_logger()\n\ndef seed_torch(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True","metadata":{"execution":{"iopub.status.busy":"2023-04-09T00:33:14.516937Z","iopub.execute_input":"2023-04-09T00:33:14.517235Z","iopub.status.idle":"2023-04-09T00:33:14.527798Z","shell.execute_reply.started":"2023-04-09T00:33:14.517210Z","shell.execute_reply":"2023-04-09T00:33:14.526810Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset for efficientnet","metadata":{}},{"cell_type":"code","source":"class CLDDataset(Dataset):\n    def __init__(self, df, mode, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.mode = mode\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, index):\n        row = self.df.loc[index]\n        image = cv2.imread(row.filepath)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\n        if self.transform is not None:\n            res = self.transform(image=image)\n            image = res['image']\n\n        image = image.astype(np.float32)\n        image = image.transpose(2, 0, 1)\n        if self.mode == 'test':\n            return torch.tensor(image).float()\n        else:\n            return torch.tensor(image).float(), torch.tensor(row.label).float()\n\n#for efficientnet\n\nclass CassvaImgClassifier(nn.Module):\n    def __init__(self, model_arch, n_class, pretrained=False):\n        super().__init__()\n        self.model = timm.create_model(model_arch, pretrained=pretrained)\n        n_features = self.model.classifier.in_features\n        # self.model.classifier = nn.Linear(n_features, n_class)\n        # '''\n        self.model.classifier = nn.Sequential(\n            nn.Dropout(0.3),\n            #nn.Linear(n_features, hidden_size,bias=True), nn.ELU(),\n            nn.Linear(n_features, n_class, bias=True)\n        )\n        # '''\n    def forward(self, x):\n        x = self.model(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-04-09T00:33:14.529165Z","iopub.execute_input":"2023-04-09T00:33:14.529779Z","iopub.status.idle":"2023-04-09T00:33:14.544448Z","shell.execute_reply.started":"2023-04-09T00:33:14.529744Z","shell.execute_reply":"2023-04-09T00:33:14.543492Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Helper functions for Resnext","metadata":{}},{"cell_type":"code","source":"def load_state(model,model_path):\n    try:  # single GPU model_file\n        model.load_state_dict(torch.load(model_path)['model'], strict=True)\n        state_dict = torch.load(model_path)['model']\n    except:  # multi GPU model_file\n        state_dict = torch.load(model_path)['model']\n        state_dict = {k[7:] if k.startswith('module.') else k: state_dict[k] for k in state_dict.keys()}\n\n    return state_dict\n\n\ndef inference(model, states, test_loader, device):\n    model.to(device)\n    tk0 = tqdm(enumerate(test_loader), total=len(test_loader))\n    probs = []\n    for i, (images) in tk0:\n        images = images.to(device)\n        avg_preds = []\n        for state in states:\n            model.load_state_dict(state)\n            model.eval()\n            with torch.no_grad():\n                y_preds = model(images)\n            avg_preds.append(y_preds.softmax(1).to('cpu').numpy())\n        avg_preds = np.mean(avg_preds, axis=0)\n        probs.append(avg_preds)\n    probs = np.concatenate(probs)\n    return probs","metadata":{"execution":{"iopub.status.busy":"2023-04-09T00:33:14.608462Z","iopub.execute_input":"2023-04-09T00:33:14.608822Z","iopub.status.idle":"2023-04-09T00:33:14.631048Z","shell.execute_reply.started":"2023-04-09T00:33:14.608786Z","shell.execute_reply":"2023-04-09T00:33:14.630392Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Helper functions for efficientnet","metadata":{}},{"cell_type":"code","source":"def inference_func(test_loader):\n    model.eval()\n    bar = tqdm(test_loader)\n\n    LOGITS = []\n    PREDS = []\n\n    with torch.no_grad():\n        for batch_idx, images in enumerate(bar):\n            x = images.to(device)\n            logits = model(x)\n            LOGITS.append(logits.cpu())\n            PREDS += [torch.softmax(logits, 1).detach().cpu()]\n        PREDS = torch.cat(PREDS).cpu().numpy()\n        LOGITS = torch.cat(LOGITS).cpu().numpy()\n    return PREDS\n\n\ndef tta_inference_func(test_loader):\n    model.eval()\n    bar = tqdm(test_loader)\n    PREDS = []\n    LOGITS = []\n\n    with torch.no_grad():\n        for batch_idx, images in enumerate(bar):\n            x = images.to(device)\n            x = torch.stack([x, x.flip(-1), x.flip(-2), x.flip(-1, -2),\n                             x.transpose(-1, -2), x.transpose(-1, -2).flip(-1),\n                             x.transpose(-1, -2).flip(-2), x.transpose(-1, -2).flip(-1, -2)], 0)\n            x = x.view(-1, 3, image_size, image_size)\n            logits = model(x)\n            logits = logits.view(BATCH_SIZE, 8, -1).mean(1)\n            PREDS += [torch.softmax(logits, 1).detach().cpu()]\n            LOGITS.append(logits.cpu())\n        PREDS = torch.cat(PREDS).cpu().numpy()\n    return PREDS","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG_b8:\n    debug=False\n    num_workers=1\n    model_name='tf_efficientnet_b8'\n    size=512\n    batch_size=32\n    seed=2020\n    target_size=5\n    target_col='label'\n    n_fold=5\n    trn_fold=[0, 1, 2, 3, 4]\n    inference=True","metadata":{"execution":{"iopub.status.busy":"2023-04-09T00:33:14.636861Z","iopub.execute_input":"2023-04-09T00:33:14.637103Z","iopub.status.idle":"2023-04-09T00:33:14.644910Z","shell.execute_reply.started":"2023-04-09T00:33:14.637079Z","shell.execute_reply":"2023-04-09T00:33:14.644048Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# CFG for Resnext\nclass CFG:\n    debug=False\n    num_workers=8\n    model_name='resnext50_32x4d'\n    size=512\n    batch_size=32\n    seed=2020\n    target_size=5\n    target_col='label'\n    n_fold=5\n    trn_fold=[0, 1, 2, 3, 4]\n    inference=True","metadata":{"execution":{"iopub.status.busy":"2023-04-09T00:33:14.646457Z","iopub.execute_input":"2023-04-09T00:33:14.646882Z","iopub.status.idle":"2023-04-09T00:33:14.654669Z","shell.execute_reply.started":"2023-04-09T00:33:14.646829Z","shell.execute_reply":"2023-04-09T00:33:14.653918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset for Resnext","metadata":{}},{"cell_type":"code","source":"\nclass TestDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.file_names = df['image_id'].values\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        file_name = self.file_names[idx]\n        file_path = f'{TEST_PATH}/{file_name}'\n        image = cv2.imread(file_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        if self.transform:\n            augmented = self.transform(image=image)\n            image = augmented['image']\n        return image","metadata":{"execution":{"iopub.status.busy":"2023-04-09T00:33:14.656325Z","iopub.execute_input":"2023-04-09T00:33:14.656693Z","iopub.status.idle":"2023-04-09T00:33:14.664839Z","shell.execute_reply.started":"2023-04-09T00:33:14.656659Z","shell.execute_reply":"2023-04-09T00:33:14.663856Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Transforms for Resnext","metadata":{}},{"cell_type":"code","source":"def get_transforms(*, data):\n    if data == 'valid':\n        return A.Compose([\n            A.Resize(CFG.size, CFG.size),\n            A.Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ])\n    \n","metadata":{"execution":{"iopub.status.busy":"2023-04-09T00:33:14.667323Z","iopub.execute_input":"2023-04-09T00:33:14.667887Z","iopub.status.idle":"2023-04-09T00:33:14.676640Z","shell.execute_reply.started":"2023-04-09T00:33:14.667851Z","shell.execute_reply":"2023-04-09T00:33:14.675836Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## ResNext Model ","metadata":{"execution":{"iopub.status.busy":"2023-04-12T13:35:22.891024Z","iopub.execute_input":"2023-04-12T13:35:22.891350Z","iopub.status.idle":"2023-04-12T13:35:22.900538Z","shell.execute_reply.started":"2023-04-12T13:35:22.891317Z","shell.execute_reply":"2023-04-12T13:35:22.898334Z"}}},{"cell_type":"code","source":"class CustomResNext(nn.Module):\n    def __init__(self, model_name='resnext50_32x4d', pretrained=False):\n        super().__init__()\n        self.model = timm.create_model(model_name, pretrained=pretrained)\n        n_features = self.model.fc.in_features\n        self.model.fc = nn.Linear(n_features, CFG.target_size)\n\n    def forward(self, x):\n        x = self.model(x)\n        return x","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## inference and submission","metadata":{}},{"cell_type":"code","source":"seed_torch(seed=CFG_b8.seed)\ntest = pd.read_csv('../input/cassava-leaf-disease-classification/sample_submission.csv')\ntest['filepath'] = test.image_id.apply(\n    lambda x: os.path.join('../input/cassava-leaf-disease-classification/test_images', f'{x}'))\nprint(test.head())","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#RestNet\nmodel_res = CustomResNext(CFG.model_name, pretrained=False)\n\nstates = [load_state(model_res,MODEL_DIR + f'{CFG.model_name}_fold{fold}.pth') for fold in CFG.trn_fold]\n\ntest_dataset = CLDDataset(test, 'test', transform=transforms_valid)\n\ntest_loader = DataLoader(test_dataset, batch_size=CFG.batch_size, shuffle=False,\n                         num_workers=CFG.num_workers, pin_memory=True)\npredictions = inference(model_res, states, test_loader, device)\n\n\ntest_dataset_efficient = CLDDataset(test, 'test', transform=transforms_valid)\ntest_loader_efficient = torch.utils.data.DataLoader(test_dataset_efficient, batch_size=BATCH_SIZE, shuffle=False,\n                                                    num_workers=1)\n# for Efficientnet\ntest_preds = []\nfor i in range(len(enet_type)):\n    model = CassvaImgClassifier(enet_type[i], n_class=5)\n    model = model.to(device)\n    model.load_state_dict(torch.load(model_path[i]))\n    test_preds += [tta_inference_func(test_loader_efficient)]\n\n# submission\npred = 0.8 * predictions + 0.2 * np.mean(test_preds, axis=0)\n# pred = np.mean(test_preds, axis=0)\ntest['label'] = softmax(pred).argmax(1)\ntest[['image_id', 'label']].to_csv(OUTPUT_DIR + '/submission.csv', index=False)\n","metadata":{"execution":{"iopub.status.busy":"2023-04-09T00:33:14.678313Z","iopub.execute_input":"2023-04-09T00:33:14.678952Z","iopub.status.idle":"2023-04-09T00:33:30.706503Z","shell.execute_reply.started":"2023-04-09T00:33:14.678916Z","shell.execute_reply":"2023-04-09T00:33:30.705623Z"},"trusted":true},"execution_count":null,"outputs":[]}]}