{"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":"# Training pipline with ViT using Pytorch \nThis is a pipeline on training with ViT using PyTorch. If anyone finds any improvement, please comment in the notebook!\n\nReferences:\n\n[paper](https://arxiv.org/abs/2010.11929)\n\n[Github](https://github.com/rwightman/pytorch-image-models)","metadata":{}},{"cell_type":"markdown","source":"# Install Timm","metadata":{}},{"cell_type":"code","source":"# !/opt/conda/bin/python3.7 -m pip install --upgrade pip\n# !c ../input/timm031py3noneanywhl/timm-0.3.1-py3-none-any.whl","metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","execution":{"iopub.status.busy":"2022-11-19T10:38:55.881052Z","iopub.execute_input":"2022-11-19T10:38:55.881378Z","iopub.status.idle":"2022-11-19T10:38:55.888348Z","shell.execute_reply.started":"2022-11-19T10:38:55.881345Z","shell.execute_reply":"2022-11-19T10:38:55.886733Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append('../input/pytorchimagemodels/pytorch-image-models-main')\nimport timm","metadata":{"execution":{"iopub.status.busy":"2022-11-19T10:38:55.890176Z","iopub.execute_input":"2022-11-19T10:38:55.890744Z","iopub.status.idle":"2022-11-19T10:38:59.466542Z","shell.execute_reply.started":"2022-11-19T10:38:55.890701Z","shell.execute_reply":"2022-11-19T10:38:59.465475Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Print GPU info","metadata":{}},{"cell_type":"code","source":"# gpu_info = !nvidia-smi\n# gpu_info = '\\n'.join(gpu_info)\n# if gpu_info.find('failed') >= 0:\n#     print('Select the Runtime > \"Change runtime type\" menu to enable a GPU accelerator, ')\n#     print('and then re-execute this cell.')\n# else:\n#     print(gpu_info)","metadata":{"execution":{"iopub.status.busy":"2022-11-19T10:38:59.468656Z","iopub.execute_input":"2022-11-19T10:38:59.468970Z","iopub.status.idle":"2022-11-19T10:38:59.475041Z","shell.execute_reply.started":"2022-11-19T10:38:59.468917Z","shell.execute_reply":"2022-11-19T10:38:59.472943Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Import 3rdparty**","metadata":{}},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nimport random\nimport cv2\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader, Dataset\nimport torch.nn.functional as F\nfrom torch.cuda.amp import autocast, GradScaler\n\nimport timm\nfrom matplotlib import pyplot as plt\nfrom sklearn.model_selection import StratifiedKFold\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm import tqdm","metadata":{"execution":{"iopub.status.busy":"2022-11-19T10:38:59.478584Z","iopub.execute_input":"2022-11-19T10:38:59.478962Z","iopub.status.idle":"2022-11-19T10:39:00.983512Z","shell.execute_reply.started":"2022-11-19T10:38:59.478901Z","shell.execute_reply":"2022-11-19T10:39:00.982641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def 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","metadata":{"execution":{"iopub.status.busy":"2022-11-19T10:39:00.985110Z","iopub.execute_input":"2022-11-19T10:39:00.985548Z","iopub.status.idle":"2022-11-19T10:39:00.992447Z","shell.execute_reply.started":"2022-11-19T10:39:00.985504Z","shell.execute_reply":"2022-11-19T10:39:00.991318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Global config**","metadata":{}},{"cell_type":"code","source":"class Config:\n    seed = 42\n    data_dir = '../input/cassava-leaf-disease-classification/'\n    train_data_dir = data_dir + 'test_images/'\n    train_csv_path = data_dir + 'test.csv'\n    arch = 'maxxvit_rmlp_nano_rw_256' ## model name\n    device = 'cuda'\n    debug = True                 ##\n    \n    image_size = 256    \n    train_batch_size = 16\n    val_batch_size = 32\n    epochs = 10                 ## total train epochs\n    freeze_bn_epochs = 5        ## freeze bn weights before epochs\n    \n    lr=1e-4                     ## init learning rate\n    min_lr = 1e-6               ## min learning rate\n    weight_decay = 1e-6\n    num_workers = 4\n    num_splits = 5             ## numbers splits\n    num_classes = 5            ## numbers classes\n    T_0 = 10\n    T_mult = 1\n    accum_iter = 2\n    verbose_step = 1\n    tta = 3\n    \n    criterion = 'LabelSmoothingCrossEntropy' ## CrossEntropy, LabelSmoothingCrossEntropy\n    label_smoothing = 0.3\n    \n    train_id = [0,1,2,3,4]","metadata":{"execution":{"iopub.status.busy":"2022-11-19T10:39:00.994087Z","iopub.execute_input":"2022-11-19T10:39:00.994770Z","iopub.status.idle":"2022-11-19T10:39:01.003849Z","shell.execute_reply.started":"2022-11-19T10:39:00.994730Z","shell.execute_reply":"2022-11-19T10:39:01.003003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Load Image**","metadata":{}},{"cell_type":"code","source":"def load_image(image_path):\n    img = cv2.imread(image_path)\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    return img","metadata":{"execution":{"iopub.status.busy":"2022-11-19T10:39:01.005474Z","iopub.execute_input":"2022-11-19T10:39:01.005959Z","iopub.status.idle":"2022-11-19T10:39:01.019242Z","shell.execute_reply.started":"2022-11-19T10:39:01.005887Z","shell.execute_reply":"2022-11-19T10:39:01.018391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **CassavaDataset**","metadata":{}},{"cell_type":"code","source":"class CassavaDataset(Dataset):\n    def __init__(self, data_dir, df, transforms=None, output_label=True):\n        self.data_dir = data_dir\n        self.df = df\n        self.transforms = transforms\n        self.output_label = output_label\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, index):\n        image_infos = self.df.iloc[index]\n        image_path = self.data_dir + image_infos.image_id\n\n        image = load_image(image_path)\n\n        if image is None:\n            raise FileNotFoundError(image_path)\n\n        ### augment\n        if self.transforms is not None:\n            image = self.transforms(image=image)['image']\n        else:\n            image = torch.from_numpy(image)\n\n        if self.output_label:\n            return image, image_infos.label\n        else:\n            return image","metadata":{"execution":{"iopub.status.busy":"2022-11-19T10:39:01.020299Z","iopub.execute_input":"2022-11-19T10:39:01.020533Z","iopub.status.idle":"2022-11-19T10:39:01.030402Z","shell.execute_reply.started":"2022-11-19T10:39:01.020510Z","shell.execute_reply":"2022-11-19T10:39:01.029379Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **CassavaClassifier**","metadata":{}},{"cell_type":"code","source":"class CassavaClassifier(nn.Module):\n    def __init__(self, model_arch, num_classes, pretrained=False):\n        super().__init__()\n        self.model = timm.create_model(model_arch, pretrained=pretrained)\n        ### vit\n\n        \n        '''\n        self.model.classifier = nn.Sequential(\n            nn.Dropout(0.3),\n            #nn.Linear(num_features, hidden_size,bias=True), nn.ELU(),\n            nn.Linear(num_features, num_classes, bias=True)\n        )\n        '''\n    def forward(self, x):\n        x = self.model(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-11-19T10:39:01.031742Z","iopub.execute_input":"2022-11-19T10:39:01.032482Z","iopub.status.idle":"2022-11-19T10:39:01.039906Z","shell.execute_reply.started":"2022-11-19T10:39:01.032443Z","shell.execute_reply":"2022-11-19T10:39:01.039005Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Train and Val transforms**","metadata":{}},{"cell_type":"code","source":"def get_train_transforms(CFG):\n    return A.Compose([\n            A.RandomResizedCrop(height=CFG.image_size, width=CFG.image_size, p=0.5),\n            A.Transpose(p=0.5),\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.RandomRotate90(p=0.5),\n            A.ShiftScaleRotate(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.CenterCrop(CFG.image_size, CFG.image_size),\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            A.CoarseDropout(p=0.5),\n            A.Cutout(p=0.5),\n            ToTensorV2(),\n        ],p=1.0)\n\ndef get_val_transforms(cfg):\n    return A.Compose([\n            A.CenterCrop(CFG.image_size, CFG.image_size, p=0.5),\n            A.Resize(CFG.image_size, CFG.image_size),\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(),\n        ],p=1.0)","metadata":{"execution":{"iopub.status.busy":"2022-11-19T10:39:01.041199Z","iopub.execute_input":"2022-11-19T10:39:01.041782Z","iopub.status.idle":"2022-11-19T10:39:01.054292Z","shell.execute_reply.started":"2022-11-19T10:39:01.041745Z","shell.execute_reply":"2022-11-19T10:39:01.053416Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_inference_transforms(CFG):\n    return A.Compose([\n            A.RandomResizedCrop(CFG.image_size, CFG.image_size),\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.0)","metadata":{"execution":{"iopub.status.busy":"2022-11-19T10:39:01.055491Z","iopub.execute_input":"2022-11-19T10:39:01.056034Z","iopub.status.idle":"2022-11-19T10:39:01.064228Z","shell.execute_reply.started":"2022-11-19T10:39:01.055998Z","shell.execute_reply":"2022-11-19T10:39:01.063473Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Train and Val data loader**","metadata":{}},{"cell_type":"markdown","source":"# **Train one epoch**","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(epoch,model,loss_fn,optimizer,train_loader,device,scheduler=None,schd_batch_update=False):\n    model.train()\n    lr = optimizer.state_dict()['param_groups'][0]['lr']   ##不同\n    \n    running_loss = None\n    pbar = tqdm(enumerate(train_loader),total=len(train_loader))\n    for step,(images,targets) in pbar:\n        images = images.to(device).float()\n        targets = targets.to(device).long()\n        \n        with autocast():\n            preds = model(images)\n            loss = loss_fn(preds,targets)\n        \n            scaler.scale(loss).backward()\n            if running_loss is None:\n                running_loss = loss.item()\n            else:\n                running_loss = running_loss* 0.99 + loss.item()*0.01\n                \n            if ((step + 1) % CFG.accum_iter == 0) or ((step + 1) == len(train_loader)):\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n                \n                if scheduler is not None and schd_batch_update:\n                    scheduler.step()\n            if ((step + 1) % CFG.accum_iter == 0) or ((step + 1) == len(train_loader)):\n                description = f'Train epoch {epoch} loss: {running_loss:.5f}'\n                pbar.set_description(description)\n                \n    if scheduler is not None and schd_batch_update:\n        scheduler.step()","metadata":{"execution":{"iopub.status.busy":"2022-11-19T10:39:01.067586Z","iopub.execute_input":"2022-11-19T10:39:01.068124Z","iopub.status.idle":"2022-11-19T10:39:01.079867Z","shell.execute_reply.started":"2022-11-19T10:39:01.068085Z","shell.execute_reply":"2022-11-19T10:39:01.079331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_dataloader(CFG, df, train_idx):\n    df_train = df.loc[train_idx,:].reset_index(drop=True)\n#     df_val = df.loc[val_idx,:].reset_index(drop=True)\n\n    train_dataset = CassavaDataset(\n        CFG.train_data_dir,\n        df_train,\n        transforms=get_train_transforms(CFG), \n        output_label=False)\n\n#     val_dataset = CassavaDataset(\n#         CFG.train_data_dir,\n#         df_val,\n#         transforms=get_val_transforms(CFG), \n#         output_label=True)\n\n    train_loader = torch.utils.data.DataLoader(\n        train_dataset,\n        batch_size=CFG.train_batch_size,\n        pin_memory=False,\n#         drop_last=False,\n        shuffle=False,        \n        num_workers=CFG.num_workers,\n        #sampler=BalanceClassSampler(labels=train_['label'].values, mode=\"downsampling\")\n    )\n\n#     val_loader = torch.utils.data.DataLoader(\n#         val_dataset, \n#         batch_size=CFG.val_batch_size,\n#         num_workers=CFG.num_workers,\n#         shuffle=False,\n#         pin_memory=False,\n#     )\n    \n#     return train_loader, val_loader\n    return train_loader","metadata":{"execution":{"iopub.status.busy":"2022-11-19T10:39:01.081120Z","iopub.execute_input":"2022-11-19T10:39:01.081658Z","iopub.status.idle":"2022-11-19T10:39:01.089461Z","shell.execute_reply.started":"2022-11-19T10:39:01.081559Z","shell.execute_reply":"2022-11-19T10:39:01.088768Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Valid one epoch**","metadata":{}},{"cell_type":"code","source":"def valid_one_epoch(epoch,model,loss_fn,val_loader,device,scheduler=None,schd_loss_update=False):\n    model.eval()\n    \n    loss_sum = 0\n    sample_num = 0\n    preds_all = []\n    targets_all = []\n    scores = []\n    \n    pbar = tqdm(enumerate(val_loader),total=len(val_loader))\n    for step,(images,targets) in pbar:\n        images = images.to(device).float()\n        targets = targets.to(device).long()\n        preds = model(images)\n            \n        preds_all += [torch.argmax(preds,1).detach().cpu().numpy()]\n        targets_all += [targets.detach().cpu().numpy()]\n\n        loss = loss_fn(preds,targets)\n        loss_sum += loss.item()*targets.shape[0]\n        sample_num += targets.shape[0]\n           \n        if ((step + 1) % CFG.accum_iter == 0) or ((step + 1) == len(train_loader)):\n            description = f'Val epoch {epoch} loss: {loss_sum/sample_num:.5f}'\n            pbar.set_description(description)\n            \n    preds_all = np.concatenate(preds_all)\n    targets_all = np.concatenate(targets_all)\n    accuracy = (preds_all == targets_all).mean()\n    print(f'Validation multi-class accuracy = {accuracy:.5f}')\n    \n    if scheduler is not None:\n        if schd_loss_update:\n            scheduler.step(loss_sum/sample_num)\n        else:\n            scheduler.step()\n    \n    return accuracy","metadata":{"execution":{"iopub.status.busy":"2022-11-19T10:39:01.090564Z","iopub.execute_input":"2022-11-19T10:39:01.091115Z","iopub.status.idle":"2022-11-19T10:39:01.103760Z","shell.execute_reply.started":"2022-11-19T10:39:01.091080Z","shell.execute_reply":"2022-11-19T10:39:01.103231Z"},"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    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","metadata":{"execution":{"iopub.status.busy":"2022-11-19T10:39:01.104977Z","iopub.execute_input":"2022-11-19T10:39:01.105491Z","iopub.status.idle":"2022-11-19T10:39:01.116556Z","shell.execute_reply.started":"2022-11-19T10:39:01.105448Z","shell.execute_reply":"2022-11-19T10:39:01.115860Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Freeze bn weights**","metadata":{}},{"cell_type":"code","source":"################ freeze bn \ndef freeze_batchnorm_stats(net):\n    try:\n        for m in net.modules():\n            if isinstance(m,nn.BatchNorm2d) or isinstance(m,nn.LayerNorm):\n                m.eval()\n    except ValuError:\n        print('error with batchnorm2d or layernorm')\n        return","metadata":{"execution":{"iopub.status.busy":"2022-11-19T10:39:01.117540Z","iopub.execute_input":"2022-11-19T10:39:01.118045Z","iopub.status.idle":"2022-11-19T10:39:01.126392Z","shell.execute_reply.started":"2022-11-19T10:39:01.118009Z","shell.execute_reply":"2022-11-19T10:39:01.125799Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Label Smoothing Cross Entropy Loss**","metadata":{}},{"cell_type":"code","source":"class LabelSmoothingCrossEntropy(nn.Module):\n    \"\"\"\n    NLL loss with label smoothing.\n    \"\"\"\n    def __init__(self, smoothing=0.1):\n        \"\"\"\n        Constructor for the LabelSmoothing module.\n        :param smoothing: label smoothing factor\n        \"\"\"\n        super(LabelSmoothingCrossEntropy, self).__init__()\n        assert smoothing < 1.0\n        self.smoothing = smoothing\n        self.confidence = 1. - smoothing\n\n    def forward(self, x, target):\n        logprobs = F.log_softmax(x, dim=-1)\n        nll_loss = -logprobs.gather(dim=-1, index=target.unsqueeze(1))\n        nll_loss = nll_loss.squeeze(1)\n        smooth_loss = -logprobs.mean(dim=-1)\n        loss = self.confidence * nll_loss + self.smoothing * smooth_loss\n        return loss.mean()","metadata":{"execution":{"iopub.status.busy":"2022-11-19T10:39:01.127410Z","iopub.execute_input":"2022-11-19T10:39:01.127914Z","iopub.status.idle":"2022-11-19T10:39:01.138994Z","shell.execute_reply.started":"2022-11-19T10:39:01.127878Z","shell.execute_reply":"2022-11-19T10:39:01.138409Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"asd=['../input/maxxvit-ckpt/maxxvit_rmlp_nano_rw_256_fold0_best_se_false.ckpt',\n     '../input/maxxvit-ckpt/maxxvit_rmlp_nano_rw_256_fold1_best_se_false.ckpt',\n     '../input/maxxvit-ckpt/maxxvit_rmlp_nano_rw_256_fold2_best_se_false.ckpt',\n     '../input/maxxvit-ckpt/maxxvit_rmlp_nano_rw_256_fold3_best_se_false.ckpt',\n     '../input/maxxvit-ckpt/maxxvit_rmlp_nano_rw_256_fold4_best_se_false.ckpt']","metadata":{"execution":{"iopub.status.busy":"2022-11-19T10:39:01.140079Z","iopub.execute_input":"2022-11-19T10:39:01.140571Z","iopub.status.idle":"2022-11-19T10:39:01.148671Z","shell.execute_reply.started":"2022-11-19T10:39:01.140529Z","shell.execute_reply":"2022-11-19T10:39:01.148084Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv('../input/cassava-leaf-disease-classification/train.csv')\ntrain.head()","metadata":{"execution":{"iopub.status.busy":"2022-11-19T10:39:01.150784Z","iopub.execute_input":"2022-11-19T10:39:01.151039Z","iopub.status.idle":"2022-11-19T10:39:01.193442Z","shell.execute_reply.started":"2022-11-19T10:39:01.151014Z","shell.execute_reply":"2022-11-19T10:39:01.192714Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Main Loop**","metadata":{}},{"cell_type":"code","source":"if __name__ == '__main__':\n     # for training only, need nightly build pytorch\n    CFG = Config\n    seed_everything(CFG.seed)\n    \n    folds = StratifiedKFold(n_splits=CFG.seed).split(np.arange(train.shape[0]), train.label.values)\n    tst_preds = []\n    for fold, (trn_idx, val_idx) in enumerate(folds):\n        # we'll train fold 0 first\n        if fold > 0:\n            break \n\n#         print('Inference fold {} started'.format(fold))\n\n#         valid_ = train.loc[val_idx,:].reset_index(drop=True)\n#         valid_ds = CassavaDataset(valid_, '../input/cassava-leaf-disease-classification/train_images/', transforms=get_inference_transforms(), output_label=False)\n        \n        test = pd.DataFrame()\n        test['image_id'] = list(os.listdir('../input/cassava-leaf-disease-classification/test_images/'))\n#         test_ds = CassavaDataset(test, '../input/cassava-leaf-disease-classification/test_images/', transforms=get_inference_transforms(CFG), output_label=False)\n        \n#         val_loader = torch.utils.data.DataLoader(\n#             valid_ds, \n#             batch_size=CFG['valid_bs'],\n#             num_workers=CFG['num_workers'],\n#             shuffle=False,\n#             pin_memory=False,\n#         )\n        \n#         tst_loader = torch.utils.data.DataLoader(\n#             test_ds, \n#             batch_size=CFG.val_batch_size,\n#             num_workers=CFG.num_workers,\n#             shuffle=False,\n#             pin_memory=False,\n#         )\n        \n        tst_loader=load_dataloader(CFG, test,np.array(list(test.index)))\n        \n\n        device = torch.device(CFG.device)\n#         model = CassvaImgClassifier(CFG.arch, train.label.nunique()).to(device)\n        model = CassavaClassifier(CFG.arch, train.label.nunique(), pretrained=False).to(device)\n#         val_preds = []\n\n        tst_preds_max = []\n        #for epoch in range(CFG['epochs']-3):\n        for i in asd:\n            model.load_state_dict(torch.load(i,map_location=device))\n        \n            with torch.no_grad():\n                for _ in range(CFG.tta):\n                    tst_preds_max+=[inference_one_epoch(model, tst_loader, device)]\n        \n        tst_preds_max = np.mean(tst_preds_max, axis=0)\n        tst_preds += [tst_preds_max]\n        \n#         for i, epoch in enumerate(CFG['used_epochs']):    \n#             model.load_state_dict(torch.load('../input/pytorch-efficientnet-baseline-train-amp-aug/{}_fold_{}_{}'.format(CFG['model_arch'], fold, epoch)))\n            \n#             with torch.no_grad():\n#                 for _ in range(CFG['tta']):\n# #                     val_preds += [CFG['weights'][i]/sum(CFG['weights'])/CFG['tta']*inference_one_epoch(model, val_loader, device)]\n#                     tst_preds += [CFG['weights'][i]/sum(CFG['weights'])/CFG['tta']*inference_one_epoch(model, tst_loader, device)]\n\n# #         val_preds = np.mean(val_preds, axis=0) \n#         tst_preds = np.mean(tst_preds, axis=0) \n        \n# #         print('fold {} validation loss = {:.5f}'.format(fold, log_loss(valid_.label.values, val_preds)))\n# #         print('fold {} validation accuracy = {:.5f}'.format(fold, (valid_.label.values==np.argmax(val_preds, axis=1)).mean()))\n        \n        del model\n        torch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2022-11-19T10:39:01.196238Z","iopub.execute_input":"2022-11-19T10:39:01.196482Z","iopub.status.idle":"2022-11-19T10:39:15.881950Z","shell.execute_reply.started":"2022-11-19T10:39:01.196456Z","shell.execute_reply":"2022-11-19T10:39:15.881131Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tst_preds_max","metadata":{"execution":{"iopub.status.busy":"2022-11-19T10:39:15.883657Z","iopub.execute_input":"2022-11-19T10:39:15.884033Z","iopub.status.idle":"2022-11-19T10:39:15.902634Z","shell.execute_reply.started":"2022-11-19T10:39:15.883991Z","shell.execute_reply":"2022-11-19T10:39:15.901640Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test['label'] = np.argmax(np.mean(tst_preds,axis=0), axis=1)\ntest.head()","metadata":{"execution":{"iopub.status.busy":"2022-11-19T10:49:04.706741Z","iopub.execute_input":"2022-11-19T10:49:04.707104Z","iopub.status.idle":"2022-11-19T10:49:04.721698Z","shell.execute_reply.started":"2022-11-19T10:49:04.707070Z","shell.execute_reply":"2022-11-19T10:49:04.720717Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-11-19T10:39:16.208389Z","iopub.status.idle":"2022-11-19T10:39:16.209154Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Please Upvote if you liked the kernel ! Cheers.\n\nReferences :\nhttps://www.kaggle.com/khyeh0719/pytorch-efficientnet-baseline-train-amp-aug. \nPlease Upvote too.","metadata":{}}]}