{"cells":[{"metadata":{},"cell_type":"markdown","source":"Log\n- v7 add vit"},{"metadata":{"execution":{"iopub.execute_input":"2021-02-10T09:09:14.632054Z","iopub.status.busy":"2021-02-10T09:09:14.631274Z","iopub.status.idle":"2021-02-10T09:09:14.63426Z","shell.execute_reply":"2021-02-10T09:09:14.633639Z"},"papermill":{"duration":0.023111,"end_time":"2021-02-10T09:09:14.634353","exception":false,"start_time":"2021-02-10T09:09:14.611242","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"#for efficientnet\nBATCH_SIZE = 1\nimage_size = 512\nenet_type = 'tf_efficientnet_b5_ns'\nb4_model_path = ['../input/cassavaext/ExtEff/baseline_cld_fold0_epoch8_tf_efficientnet_b4_ns_512.pth', \n              '../input/cassavaext/ExtEff/baseline_cld_fold1_epoch9_tf_efficientnet_b4_ns_512.pth', \n              '../input/cassavaext/ExtEff/baseline_cld_fold2_epoch9_tf_efficientnet_b4_ns_512.pth',\n              '../input/cassavaext/ExtEff/baseline_cld_fold3_epoch5_tf_efficientnet_b4_ns_512.pth',\n              '../input/cassavaext/ExtEff/baseline_cld_fold4_epoch11_tf_efficientnet_b4_ns_512.pth']","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-02-10T09:09:14.672356Z","iopub.status.busy":"2021-02-10T09:09:14.671798Z","iopub.status.idle":"2021-02-10T09:09:18.403692Z","shell.execute_reply":"2021-02-10T09:09:18.402651Z"},"papermill":{"duration":3.75519,"end_time":"2021-02-10T09:09:18.4038","exception":false,"start_time":"2021-02-10T09:09:14.64861","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# Library\n# ====================================================\nimport sys\nsys.path.append('../input/pytorch-image-models/pytorch-image-models-master')\nsys.path.append('../input/vision-transformer-pytorch/VisionTransformer-Pytorch')\nimport os\nimport math\nimport time\nimport random\nimport shutil\nimport albumentations\nfrom pathlib import Path\nfrom contextlib import contextmanager\nfrom collections import defaultdict, Counter\n\nimport scipy as sp\nfrom scipy.special import softmax\nimport numpy as np\nimport pandas as pd\n\nfrom sklearn import preprocessing\nfrom sklearn.metrics import accuracy_score\nfrom sklearn.model_selection import StratifiedKFold\n\nfrom tqdm import tqdm\nfrom functools import partial\n\nimport cv2\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim import Adam, SGD\nimport torchvision.models as models\nfrom torch.nn.parameter import Parameter\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, CosineAnnealingLR, ReduceLROnPlateau\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport timm\nfrom vision_transformer_pytorch import VisionTransformer\n\nimport warnings \nwarnings.filterwarnings('ignore')\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-02-10T09:09:18.439812Z","iopub.status.busy":"2021-02-10T09:09:18.438089Z","iopub.status.idle":"2021-02-10T09:09:18.440402Z","shell.execute_reply":"2021-02-10T09:09:18.440817Z"},"papermill":{"duration":0.021446,"end_time":"2021-02-10T09:09:18.44092","exception":false,"start_time":"2021-02-10T09:09:18.419474","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"#Transform for efficientnet\ntransforms_valid = albumentations.Compose([\n    albumentations.CenterCrop(image_size, image_size, p=1),\n    albumentations.Resize(image_size, image_size),\n    albumentations.Normalize()\n])","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.013666,"end_time":"2021-02-10T09:09:18.46881","exception":false,"start_time":"2021-02-10T09:09:18.455144","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# Directory settings"},{"metadata":{"execution":{"iopub.execute_input":"2021-02-10T09:09:18.503102Z","iopub.status.busy":"2021-02-10T09:09:18.50153Z","iopub.status.idle":"2021-02-10T09:09:18.503769Z","shell.execute_reply":"2021-02-10T09:09:18.504183Z"},"papermill":{"duration":0.021349,"end_time":"2021-02-10T09:09:18.50428","exception":false,"start_time":"2021-02-10T09:09:18.482931","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# Directory settings for Resnext\n# ====================================================\nimport os\n\nOUTPUT_DIR = './'\nMODEL_DIR = '../input/cassavaext/ExtRes/'\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'","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.013805,"end_time":"2021-02-10T09:09:18.532237","exception":false,"start_time":"2021-02-10T09:09:18.518432","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# CFG"},{"metadata":{"execution":{"iopub.execute_input":"2021-02-10T09:09:18.566887Z","iopub.status.busy":"2021-02-10T09:09:18.565251Z","iopub.status.idle":"2021-02-10T09:09:18.567552Z","shell.execute_reply":"2021-02-10T09:09:18.567949Z"},"papermill":{"duration":0.021778,"end_time":"2021-02-10T09:09:18.568049","exception":false,"start_time":"2021-02-10T09:09:18.546271","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# CFG for Resnext\n# ====================================================\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\n    epochs = [8]","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.01412,"end_time":"2021-02-10T09:09:18.596368","exception":false,"start_time":"2021-02-10T09:09:18.582248","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# Utils"},{"metadata":{"execution":{"iopub.execute_input":"2021-02-10T09:09:18.635544Z","iopub.status.busy":"2021-02-10T09:09:18.634885Z","iopub.status.idle":"2021-02-10T09:09:18.639939Z","shell.execute_reply":"2021-02-10T09:09:18.639514Z"},"papermill":{"duration":0.029473,"end_time":"2021-02-10T09:09:18.640017","exception":false,"start_time":"2021-02-10T09:09:18.610544","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# Utils for Resnext\n# ====================================================\ndef 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\n#LOGGER = init_logger()\n\ndef seed_everything(seed):\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\n    torch.backends.cudnn.benchmark = True\n\nseed_everything(seed=CFG.seed)","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.013912,"end_time":"2021-02-10T09:09:18.668744","exception":false,"start_time":"2021-02-10T09:09:18.654832","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# Data Loading"},{"metadata":{"execution":{"iopub.execute_input":"2021-02-10T09:09:18.707537Z","iopub.status.busy":"2021-02-10T09:09:18.707018Z","iopub.status.idle":"2021-02-10T09:09:18.718821Z","shell.execute_reply":"2021-02-10T09:09:18.718189Z"},"papermill":{"duration":0.034112,"end_time":"2021-02-10T09:09:18.718908","exception":false,"start_time":"2021-02-10T09:09:18.684796","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"test = pd.read_csv('../input/cassava-leaf-disease-classification/sample_submission.csv')\ntest['filepath'] = test.image_id.apply(lambda x: os.path.join('../input/cassava-leaf-disease-classification/test_images', f'{x}'))\n#test.head()","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.015431,"end_time":"2021-02-10T09:09:18.750209","exception":false,"start_time":"2021-02-10T09:09:18.734778","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# Dataset"},{"metadata":{"execution":{"iopub.execute_input":"2021-02-10T09:09:18.789492Z","iopub.status.busy":"2021-02-10T09:09:18.788823Z","iopub.status.idle":"2021-02-10T09:09:18.795587Z","shell.execute_reply":"2021-02-10T09:09:18.795185Z"},"papermill":{"duration":0.029703,"end_time":"2021-02-10T09:09:18.795668","exception":false,"start_time":"2021-02-10T09:09:18.765965","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# Dataset for Resnext\n# ====================================================\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","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-02-10T09:09:18.836606Z","iopub.status.busy":"2021-02-10T09:09:18.835926Z","iopub.status.idle":"2021-02-10T09:09:18.838965Z","shell.execute_reply":"2021-02-10T09:09:18.838535Z"},"papermill":{"duration":0.029168,"end_time":"2021-02-10T09:09:18.839051","exception":false,"start_time":"2021-02-10T09:09:18.809883","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# Dataset for efficientnet\n# ====================================================\nclass 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#test = pd.read_csv('../input/cassava-leaf-disease-classification/sample_submission.csv')\n#test_dataset = CLDDataset(test, 'test', transform=transforms_valid)\n#test_loader = torch.utils.data.DataLoader(test_dataset, batch_size=batch_size, shuffle=False,  num_workers=4)","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-02-10T09:09:18.876073Z","iopub.status.busy":"2021-02-10T09:09:18.875413Z","iopub.status.idle":"2021-02-10T09:09:18.878061Z","shell.execute_reply":"2021-02-10T09:09:18.878456Z"},"papermill":{"duration":0.023867,"end_time":"2021-02-10T09:09:18.878558","exception":false,"start_time":"2021-02-10T09:09:18.854691","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"#for efficientnet\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,  num_workers=4)","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-02-10T09:09:18.914104Z","iopub.status.busy":"2021-02-10T09:09:18.913242Z","iopub.status.idle":"2021-02-10T09:09:18.919676Z","shell.execute_reply":"2021-02-10T09:09:18.918838Z"},"papermill":{"duration":0.025565,"end_time":"2021-02-10T09:09:18.919795","exception":false,"start_time":"2021-02-10T09:09:18.89423","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# Transforms for Resnext\n# ====================================================\ndef 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        ])","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.015526,"end_time":"2021-02-10T09:09:18.952392","exception":false,"start_time":"2021-02-10T09:09:18.936866","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# MODELS"},{"metadata":{"execution":{"iopub.execute_input":"2021-02-10T09:09:18.98868Z","iopub.status.busy":"2021-02-10T09:09:18.987988Z","iopub.status.idle":"2021-02-10T09:09:18.99052Z","shell.execute_reply":"2021-02-10T09:09:18.990935Z"},"papermill":{"duration":0.023595,"end_time":"2021-02-10T09:09:18.991028","exception":false,"start_time":"2021-02-10T09:09:18.967433","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# ResNext101 Model\n# ====================================================\nclass ResNext101(nn.Module):\n    def __init__(self, model_name='resnext101_32x8d', 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, 5)\n\n    def forward(self, x):\n        return self.model(x)","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-02-10T09:09:19.02858Z","iopub.status.busy":"2021-02-10T09:09:19.027976Z","iopub.status.idle":"2021-02-10T09:09:19.61425Z","shell.execute_reply":"2021-02-10T09:09:19.613788Z"},"papermill":{"duration":0.608251,"end_time":"2021-02-10T09:09:19.614359","exception":false,"start_time":"2021-02-10T09:09:19.006108","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# EfficientNetb5 Model\n# ====================================================\nclass EffnetB5(nn.Module):\n    def __init__(self, transfer_model = timm.create_model(model_name='tf_efficientnet_b5_ns', \n                                                          pretrained=False)):\n        super().__init__()\n        self.model = transfer_model\n        self.model.classifier = nn.Linear(transfer_model.classifier.in_features, 5)\n        \n    def forward(self,x):\n        return self.model(x)","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-02-10T09:09:19.651405Z","iopub.status.busy":"2021-02-10T09:09:19.650771Z","iopub.status.idle":"2021-02-10T09:09:19.654744Z","shell.execute_reply":"2021-02-10T09:09:19.654298Z"},"papermill":{"duration":0.025148,"end_time":"2021-02-10T09:09:19.654894","exception":false,"start_time":"2021-02-10T09:09:19.629746","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# ResNext50 Model\n# ====================================================\nclass ResNext50(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","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-02-10T09:09:19.693203Z","iopub.status.busy":"2021-02-10T09:09:19.692565Z","iopub.status.idle":"2021-02-10T09:09:19.696079Z","shell.execute_reply":"2021-02-10T09:09:19.696603Z"},"papermill":{"duration":0.026223,"end_time":"2021-02-10T09:09:19.696712","exception":false,"start_time":"2021-02-10T09:09:19.670489","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# EfficientNetB4 Model\n# ====================================================\nclass EffnetB4(nn.Module):\n    def __init__(self, backbone, out_dim, pretrained=False):\n        super().__init__()\n        self.enet = timm.create_model(backbone, pretrained=pretrained)\n        in_ch = self.enet.classifier.in_features\n        self.myfc = nn.Linear(in_ch, out_dim)\n        self.enet.classifier = nn.Identity()\n\n    def forward(self, x):\n        x = self.enet(x)\n        x = self.myfc(x)\n        return x","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.015195,"end_time":"2021-02-10T09:09:19.727099","exception":false,"start_time":"2021-02-10T09:09:19.711904","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# Helper functions"},{"metadata":{"execution":{"iopub.execute_input":"2021-02-10T09:09:19.769745Z","iopub.status.busy":"2021-02-10T09:09:19.768585Z","iopub.status.idle":"2021-02-10T09:09:19.771425Z","shell.execute_reply":"2021-02-10T09:09:19.771011Z"},"papermill":{"duration":0.028797,"end_time":"2021-02-10T09:09:19.77151","exception":false,"start_time":"2021-02-10T09:09:19.742713","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# Helper functions for Resnext\n# ====================================================\ndef load_state(model_path, this_model):\n    model = this_model\n    try:  # single GPU model_file\n        model.load_state_dict(torch.load(model_path), strict=True)\n        state_dict = torch.load(model_path)\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(test_loader, total=len(test_loader))\n    probs = []\n    for i, (images) in enumerate(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","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-02-10T09:09:19.818545Z","iopub.status.busy":"2021-02-10T09:09:19.817811Z","iopub.status.idle":"2021-02-10T09:09:19.820471Z","shell.execute_reply":"2021-02-10T09:09:19.82005Z"},"papermill":{"duration":0.03353,"end_time":"2021-02-10T09:09:19.820556","exception":false,"start_time":"2021-02-10T09:09:19.787026","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# Helper functions for efficientnet\n# ====================================================\ndef 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\ndef tta_inference_func(test_loader):\n    model.eval()\n    pbar = tqdm(test_loader, total=len(test_loader))\n    PREDS = []\n    LOGITS = []\n\n    with torch.no_grad():\n        for batch_idx, images in enumerate(pbar):\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\n        PREDS = torch.cat(PREDS).cpu().numpy()\n    pbar.set_description(f'Model: {model.__class__.__name__}') \n    return PREDS","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.015177,"end_time":"2021-02-10T09:09:19.851351","exception":false,"start_time":"2021-02-10T09:09:19.836174","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# inference and Submit"},{"metadata":{"execution":{"iopub.execute_input":"2021-02-10T09:09:19.891185Z","iopub.status.busy":"2021-02-10T09:09:19.890368Z","iopub.status.idle":"2021-02-10T09:09:42.98824Z","shell.execute_reply":"2021-02-10T09:09:42.988853Z"},"papermill":{"duration":23.121772,"end_time":"2021-02-10T09:09:42.989016","exception":false,"start_time":"2021-02-10T09:09:19.867244","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# inference\n# ====================================================\n\nres101_preds = []\neffb5_preds = []\ntest_dataset = TestDataset(test, transform=get_transforms(data='valid'))\ntest_loader = DataLoader(test_dataset, batch_size=CFG.batch_size, shuffle=False, \n                             num_workers=CFG.num_workers, pin_memory=True)\nfor epoch in CFG.epochs:\n    #for Resnext101\n    model = ResNext101().to(device)\n    states = [load_state(f'../input/cassavamodels/resnext101_32x8d/resnext101_32x8d_fold_{fold}_{epoch}.pt', model) for fold in CFG.trn_fold]\n    res101_preds += [inference(model, states, test_loader, device)]\n    \n    for fold in range(5):\n        #for EfficientnetB5\n        model = EffnetB5().to(device)\n        model.load_state_dict(torch.load(f'../input/cassavamodels/tf_efficientnet_b5_ns/tf_efficientnet_b5_ns_fold_{fold}_{epoch}.pt'))\n        effb5_preds += [tta_inference_func(test_loader_efficient)]\n        \n#for Resnext50\nmodel = ResNext50(CFG.model_name, pretrained=False)\nstates = [load_state(MODEL_DIR+f'{CFG.model_name}_fold{fold}.pth', model) for fold in CFG.trn_fold]\nres50_mean = inference(model, states, test_loader, device)\n \neffb4_preds = []\nfor fold in range(5):\n    #for EfficientnetB4\n    model = EffnetB4('tf_efficientnet_b4_ns', out_dim=5).to(device)\n    model.load_state_dict(torch.load(b4_model_path[fold]))\n    effb4_preds += [tta_inference_func(test_loader_efficient)]","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# VIT"},{"metadata":{"trusted":true},"cell_type":"code","source":"def get_vit_transforms():\n    return A.Compose([\n            A.RandomResizedCrop(384, 384),\n            A.Transpose(p=0.5),\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.HueSaturationValue(hue_shift_limit=0.2, sat_shift_limit=0.2, val_shift_limit=0.2, p=0.5),\n            A.RandomBrightnessContrast(brightness_limit=(-0.1,0.1), contrast_limit=(-0.1, 0.1), p=0.5),\n            A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225], max_pixel_value=255.0, p=1.0),\n            ToTensorV2(p=1.0),\n        ], p=1.)\n\nclass VitNet(nn.Module):\n    def __init__(self, model_name='vit_base_patch16_384', pretrained=False):\n        super().__init__()\n        self.model = timm.create_model(model_name, pretrained=pretrained)\n        self.model.head = nn.Linear(self.model.head.in_features, 5)\n            \n    def forward(self, x):\n        x = self.model(x)\n        return x\n    \ndef inference_one_epoch(model, data_loader):\n    model.eval()\n\n    image_preds_all = []\n    \n    pbar = tqdm(enumerate(data_loader), total=len(data_loader))\n    for step, (imgs) in pbar:\n        imgs = imgs.to(device).float()\n        \n        image_preds = model(imgs)   #output = model(input)\n        image_preds_all += [torch.softmax(image_preds, 1).detach().cpu().numpy()]\n        \n    \n    image_preds_all = np.concatenate(image_preds_all, axis=0)\n    return image_preds_all\n\ndef vit_tta_inference_func(test_loader):\n    model.eval()\n    pbar = tqdm(test_loader, total=len(test_loader))\n    PREDS = []\n    LOGITS = []\n\n    with torch.no_grad():\n        for batch_idx, images in enumerate(pbar):\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, 384, 384)\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\n        PREDS = torch.cat(PREDS).cpu().numpy()\n    pbar.set_description(f'Model: {model.__class__.__name__}') \n    return PREDS\n\n\nvit_dataset = TestDataset(test, transform=get_vit_transforms())\nvit_loader = DataLoader(vit_dataset, batch_size=CFG.batch_size, shuffle=False, \n                             num_workers=CFG.num_workers, pin_memory=False)\n\nvit_preds = []\nfor fold in range(5):\n    model = VitNet().to(device)\n    model.load_state_dict(torch.load(f'../input/cassavamodels/VIT_B16_384/vit_base_patch16_384_fold_{fold}_best.pt'))\n    vit_preds += [vit_tta_inference_func(vit_loader)] ","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-02-10T09:09:43.047511Z","iopub.status.busy":"2021-02-10T09:09:43.046684Z","iopub.status.idle":"2021-02-10T09:09:43.276457Z","shell.execute_reply":"2021-02-10T09:09:43.275655Z"},"papermill":{"duration":0.262172,"end_time":"2021-02-10T09:09:43.276558","exception":false,"start_time":"2021-02-10T09:09:43.014386","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"effb5_mean = np.mean(effb5_preds, axis=0)\nres101_mean = np.mean(res101_preds, axis=0)\neffb4_mean = np.mean(effb4_preds, axis=0)\nvit_mean = np.mean(vit_preds, axis=0)\npreds = 0.2*res50_mean + 0.2*effb4_mean + 0.2*res101_mean + 0.2*effb5_mean + 0.2*vit_mean \nprint(f'effb5_mean:{effb5_mean}\\\n        \\nres101_mean:{res101_mean}\\\n        \\neffb4_mean:{effb4_mean}\\\n        \\nres50_mean:{res50_mean}\\\n        \\nvit_mean:{vit_mean}\\\n        \\npreds:{preds}')\ntest['label'] = softmax(preds).argmax(1)\ntest[['image_id', 'label']].to_csv(OUTPUT_DIR+'submission.csv', index=False)\ntest.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}