{"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 timm","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:28.589671Z","iopub.execute_input":"2021-06-27T17:40:28.590067Z","iopub.status.idle":"2021-06-27T17:40:28.596206Z","shell.execute_reply.started":"2021-06-27T17:40:28.590031Z","shell.execute_reply":"2021-06-27T17:40:28.594601Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport albumentations as A\nfrom albumentations.core.composition import Compose\nfrom albumentations.pytorch import ToTensorV2\nimport pytorch_lightning as pl\nimport matplotlib.pyplot as plt\nimport pandas as pd\nfrom PIL import Image\nimport cv2\nimport numpy as np\n\nimport timm\nimport torch\nimport torch.nn as nn\nimport torchvision\nfrom torch.utils.data import Dataset, DataLoader\n\n# pytorch lightning\nfrom pytorch_lightning import Trainer, seed_everything\nfrom pytorch_lightning import Callback\nfrom pytorch_lightning.loggers import CSVLogger\nfrom pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping\n\nfrom sklearn import metrics, model_selection\n%matplotlib inline\n","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:28.600454Z","iopub.execute_input":"2021-06-27T17:40:28.601273Z","iopub.status.idle":"2021-06-27T17:40:28.925800Z","shell.execute_reply.started":"2021-06-27T17:40:28.601180Z","shell.execute_reply":"2021-06-27T17:40:28.924713Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    seed = 42\n    model_name='resnet50'\n    pretrained=True\n    img_size=256\n    num_classes=5\n    lr=1e-4\n    min_lr=1e-3\n    t_max=20\n    num_epochs=10\n    batch_size=64\n    accum=1\n    n_fold=5\n    precision=16\n    device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:28.928000Z","iopub.execute_input":"2021-06-27T17:40:28.928457Z","iopub.status.idle":"2021-06-27T17:40:28.986166Z","shell.execute_reply.started":"2021-06-27T17:40:28.928412Z","shell.execute_reply":"2021-06-27T17:40:28.984830Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:28.989347Z","iopub.execute_input":"2021-06-27T17:40:28.990059Z","iopub.status.idle":"2021-06-27T17:40:29.003407Z","shell.execute_reply.started":"2021-06-27T17:40:28.990009Z","shell.execute_reply":"2021-06-27T17:40:29.002283Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('../input/cassava-leaf-disease-classification/train.csv')","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:29.007003Z","iopub.execute_input":"2021-06-27T17:40:29.007449Z","iopub.status.idle":"2021-06-27T17:40:29.047276Z","shell.execute_reply.started":"2021-06-27T17:40:29.007376Z","shell.execute_reply":"2021-06-27T17:40:29.046235Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.head()","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:29.048762Z","iopub.execute_input":"2021-06-27T17:40:29.049564Z","iopub.status.idle":"2021-06-27T17:40:29.069207Z","shell.execute_reply.started":"2021-06-27T17:40:29.049518Z","shell.execute_reply":"2021-06-27T17:40:29.067961Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.label.value_counts()","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:29.070991Z","iopub.execute_input":"2021-06-27T17:40:29.071888Z","iopub.status.idle":"2021-06-27T17:40:29.086928Z","shell.execute_reply.started":"2021-06-27T17:40:29.071837Z","shell.execute_reply":"2021-06-27T17:40:29.085541Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train, df_valid = model_selection.train_test_split(\ndf,\ntest_size=0.1,\nrandom_state=CFG.seed,\nstratify=df.label.values)\n\ndf_train = df_train.reset_index(drop=True)\ndf_valid = df_valid.reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:29.088940Z","iopub.execute_input":"2021-06-27T17:40:29.089755Z","iopub.status.idle":"2021-06-27T17:40:29.129351Z","shell.execute_reply.started":"2021-06-27T17:40:29.089651Z","shell.execute_reply":"2021-06-27T17:40:29.128141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.shape, df_valid.shape","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:29.131038Z","iopub.execute_input":"2021-06-27T17:40:29.131581Z","iopub.status.idle":"2021-06-27T17:40:29.140146Z","shell.execute_reply.started":"2021-06-27T17:40:29.131548Z","shell.execute_reply":"2021-06-27T17:40:29.138491Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_path = '../input/cassava-leaf-disease-classification/train_images'\n\ntrain_image_paths = [os.path.join(image_path, x) for x in df_train.image_id.values]\nvalid_image_paths = [os.path.join(image_path, x) for x in df_valid.image_id.values]","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:29.145430Z","iopub.execute_input":"2021-06-27T17:40:29.146006Z","iopub.status.idle":"2021-06-27T17:40:29.201209Z","shell.execute_reply.started":"2021-06-27T17:40:29.145958Z","shell.execute_reply":"2021-06-27T17:40:29.200025Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_image_paths[:5]","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:29.204160Z","iopub.execute_input":"2021-06-27T17:40:29.204659Z","iopub.status.idle":"2021-06-27T17:40:29.211321Z","shell.execute_reply.started":"2021-06-27T17:40:29.204613Z","shell.execute_reply":"2021-06-27T17:40:29.210064Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_targets = df_train.label.values\nvalid_targets = df_valid.label.values","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:29.213376Z","iopub.execute_input":"2021-06-27T17:40:29.214203Z","iopub.status.idle":"2021-06-27T17:40:29.223530Z","shell.execute_reply.started":"2021-06-27T17:40:29.214141Z","shell.execute_reply":"2021-06-27T17:40:29.222236Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CassavaLeafDataset(Dataset):\n    def __init__(self, image_paths, targets, transform=None):\n        self.image_paths = image_paths\n        self.targets = targets\n        self.transform = transform\n    \n    def __len__(self):\n        return len(self.image_paths)\n    \n    def __getitem__(self, item):\n        targets = self.targets[item]\n        image = cv2.imread(self.image_paths[item])\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        augmented = self.transform(image=image)\n        \n        image = augmented['image']\n        image = np.transpose(image, (2, 0, 1)).astype(np.float32)\n        image = image / 255.0\n\n#         image = Image.open(self.image_paths[item])\n#         image = np.array(image)\n#         augmented = self.transform(image=image)\n#         image = augmented['image']\n# #         print(image)\n#         image = np.transpose(image, (0, 1, 2)).astype(np.float32)\n#         image_tensor = torch.tensor(image)\n        return {\n            'image': image, \n            'targets': targets\n        }\n        ","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:29.225362Z","iopub.execute_input":"2021-06-27T17:40:29.225937Z","iopub.status.idle":"2021-06-27T17:40:29.239016Z","shell.execute_reply.started":"2021-06-27T17:40:29.225891Z","shell.execute_reply":"2021-06-27T17:40:29.237628Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_aug = A.Compose([\n        A.RandomResizedCrop(height=CFG.img_size, width=CFG.img_size),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.Transpose(p=0.5)\n])\n\nvalid_aug = A.Compose([\n        A.CenterCrop(height=CFG.img_size, width=CFG.img_size, p=1.0),\n        A.Resize(256, 256),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.Transpose(p=0.5)\n])","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:29.241130Z","iopub.execute_input":"2021-06-27T17:40:29.241761Z","iopub.status.idle":"2021-06-27T17:40:29.251300Z","shell.execute_reply.started":"2021-06-27T17:40:29.241716Z","shell.execute_reply":"2021-06-27T17:40:29.249875Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def get_transform(phase: str):\n#     if phase == 'train':\n#         return Compose([\n#             A.RandomResizedCrop(height=CFG.img_size, width=CFG.img_size),\n#             A.HorizontalFlip(p=0.5),\n#             A.ShiftScaleRotate(p=0.5),\n#             A.RandomBrightnessContrast(p=0.5),\n#             A.Normalize(),\n#             ToTensorV2(),\n#         ])\n#     else:\n#         return Compose([\n#             A.Resize(height=CFG.img_size, width=CFG.img_size),\n#             A.Normalize(),\n#             ToTensorV2(),\n#         ])","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:29.253377Z","iopub.execute_input":"2021-06-27T17:40:29.253904Z","iopub.status.idle":"2021-06-27T17:40:29.263852Z","shell.execute_reply.started":"2021-06-27T17:40:29.253858Z","shell.execute_reply":"2021-06-27T17:40:29.262679Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = CassavaLeafDataset(\n    image_paths=train_image_paths,\n    targets=train_targets,\n    transform=train_aug\n)\n\nvalid_dataset = CassavaLeafDataset(\n    image_paths=valid_image_paths,\n    targets=valid_targets,\n    transform=valid_aug\n)\n","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:29.267809Z","iopub.execute_input":"2021-06-27T17:40:29.268597Z","iopub.status.idle":"2021-06-27T17:40:29.275900Z","shell.execute_reply.started":"2021-06-27T17:40:29.268546Z","shell.execute_reply":"2021-06-27T17:40:29.274362Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# valid_dataset[0]['image']","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:29.277800Z","iopub.execute_input":"2021-06-27T17:40:29.278408Z","iopub.status.idle":"2021-06-27T17:40:29.285630Z","shell.execute_reply.started":"2021-06-27T17:40:29.278347Z","shell.execute_reply":"2021-06-27T17:40:29.284495Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_image(image_dict):\n    img_tensor = image_dict['image']\n    target = image_dict['targets']\n    plt.figure(figsize=(5, 5))\n    image = img_tensor/255\n    print(image.shape)\n    print(target)\n    plt.imshow(image)","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:29.287580Z","iopub.execute_input":"2021-06-27T17:40:29.288292Z","iopub.status.idle":"2021-06-27T17:40:29.296244Z","shell.execute_reply.started":"2021-06-27T17:40:29.288244Z","shell.execute_reply":"2021-06-27T17:40:29.294838Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plot_image(valid_dataset[1])","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:29.298136Z","iopub.execute_input":"2021-06-27T17:40:29.298934Z","iopub.status.idle":"2021-06-27T17:40:29.306501Z","shell.execute_reply.started":"2021-06-27T17:40:29.298886Z","shell.execute_reply":"2021-06-27T17:40:29.305385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plot_image(train_dataset[1])","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:29.308477Z","iopub.execute_input":"2021-06-27T17:40:29.308948Z","iopub.status.idle":"2021-06-27T17:40:29.317463Z","shell.execute_reply.started":"2021-06-27T17:40:29.308901Z","shell.execute_reply":"2021-06-27T17:40:29.316435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Data Loaders\n\ntrain_loader = DataLoader(train_dataset, batch_size=CFG.batch_size, num_workers=2, pin_memory=True, drop_last=True, shuffle=True)\nvalid_loader = DataLoader(valid_dataset, batch_size=CFG.batch_size, num_workers=2, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:29.320273Z","iopub.execute_input":"2021-06-27T17:40:29.321108Z","iopub.status.idle":"2021-06-27T17:40:29.328830Z","shell.execute_reply.started":"2021-06-27T17:40:29.321052Z","shell.execute_reply":"2021-06-27T17:40:29.327589Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Defining model\n# class CustomResNet(nn.Module):\n#     def __init__(self, model_name='resnet18', pretrained=False):\n#         super().__init__()\n#         self.model = timm.create_model(model_name, pretrained=pretrained)\n#         in_features = self.model.get_classifier().in_features\n#         self.model.fc = nn.Linear(in_features, CFG.num_classes)\n        \n#     def forward(self, x):\n#         x = self.model(x)\n#         return x","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:29.330657Z","iopub.execute_input":"2021-06-27T17:40:29.331241Z","iopub.status.idle":"2021-06-27T17:40:29.338837Z","shell.execute_reply.started":"2021-06-27T17:40:29.331186Z","shell.execute_reply":"2021-06-27T17:40:29.337312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#     def __init__(self):\n#         super().__init__()\n#         self.efficient_net = EfficientNet.from_name('efficientnet-b5')\n#         self.efficient_net.load_state_dict(torch.load(PRETRAINED_PATH))\n# #         self.efficient_net=EfficientNet.from_pretrained('efficientnet-b3',num_classes=CLASSES)\n#         in_features=self.efficient_net._fc.in_features\n#         self.efficient_net._fc=nn.Linear(in_features,CLASSES)\n    \n#     def forward(self,x):\n#         out=self.efficient_net(x)\n#         return out","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:29.341041Z","iopub.execute_input":"2021-06-27T17:40:29.341570Z","iopub.status.idle":"2021-06-27T17:40:29.351094Z","shell.execute_reply.started":"2021-06-27T17:40:29.341524Z","shell.execute_reply":"2021-06-27T17:40:29.349960Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CassavaLeafModel(pl.LightningModule):\n    def __init__(self):\n        super(CassavaLeafModel, self).__init__()\n        self.model = timm.create_model('resnet18', pretrained=True)\n        in_features = self.model.get_classifier().in_features\n        self.model.fc = nn.Linear(in_features, CFG.num_classes)\n        self.metric = pl.metrics.F1(num_classes=CFG.num_classes)\n        self.criterion = nn.CrossEntropyLoss()\n        self.lr = CFG.lr\n        \n    def forward(self, x, *args, **kwargs):\n        return self.model(x)\n    \n    def configure_optimizers(self):\n        self.optimizer = torch.optim.Adam(self.model.parameters(), lr=self.lr)\n        self.scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(self.optimizer, T_max= CFG.t_max, eta_min=CFG.min_lr)\n        return {'optimizer': self.optimizer, 'lr_scheduler':self.scheduler}\n    \n    def training_step(self, batch, batch_idx):\n        image = batch['image']\n        target = batch['targets']\n        output = self.model(image)\n        loss = self.criterion(output, target)\n        score = self.metric(output.argmax(1), target)\n        logs = {'train_loss': loss, 'train_f1': score, 'lr':self.optimizer.param_groups[0]['lr']}\n        self.log_dict(logs, on_step=False, on_epoch=True, prog_bar=True, logger=True)\n        return loss\n        \n    def validation_step(self, batch, batch_idx):\n        image = batch['image']\n        target = batch['targets']\n        output = self.model(image)\n        loss = self.criterion(output, target)\n        score = self.metric(output.argmax(1), target)\n        logs = {\n            'valid_loss': loss,\n            'valid_f1': score,\n        }\n        self.log_dict(logs, on_step=False, on_epoch=True, prog_bar=True, logger=True)\n        return loss","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:29.353848Z","iopub.execute_input":"2021-06-27T17:40:29.354826Z","iopub.status.idle":"2021-06-27T17:40:29.371995Z","shell.execute_reply.started":"2021-06-27T17:40:29.354779Z","shell.execute_reply":"2021-06-27T17:40:29.370536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model = CustomResNet(model_name=CFG.model_name, pretrained=CFG.pretrained)\ncassava_model = CassavaLeafModel()\n# cassava_model","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:41:01.263946Z","iopub.execute_input":"2021-06-27T17:41:01.264380Z","iopub.status.idle":"2021-06-27T17:41:07.822102Z","shell.execute_reply.started":"2021-06-27T17:41:01.264349Z","shell.execute_reply":"2021-06-27T17:41:07.821093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"logger = CSVLogger(save_dir='logs/', name=CFG.model_name)\nlogger.log_hyperparams(CFG.__dict__)\ncheckpoint_callback = ModelCheckpoint(monitor='valid_loss',\n                                    save_top_k=1,\n                                     save_last=True,\n                                     save_weights_only=True,\n                                     filename='checkpoint/{epoch:02d}-{valid_loss:.4f}-{valid_f1:.4f}',\n                                     verbose=False,\n                                     mode='min')\n\ntrainer = Trainer(max_epochs=CFG.num_epochs,\n                 gpus=1,\n                 accumulate_grad_batches=CFG.accum,\n                 precision=CFG.precision,\n                 checkpoint_callback=checkpoint_callback,\n                 logger=logger,\n                 weights_summary='top',\n)","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:41:12.376773Z","iopub.execute_input":"2021-06-27T17:41:12.377128Z","iopub.status.idle":"2021-06-27T17:41:12.398567Z","shell.execute_reply.started":"2021-06-27T17:41:12.377084Z","shell.execute_reply":"2021-06-27T17:41:12.397527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.fit(cassava_model, train_dataloader=train_loader, val_dataloaders=valid_loader)","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:41:12.658107Z","iopub.execute_input":"2021-06-27T17:41:12.658536Z","iopub.status.idle":"2021-06-27T18:26:36.625675Z","shell.execute_reply.started":"2021-06-27T17:41:12.658505Z","shell.execute_reply":"2021-06-27T18:26:36.624223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df_ = 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 test_df_.image_id.values]\n# fake targets\ntest_targets = test_df_.label.values\n\n\ntest_aug = A.Compose([\n            A.CenterCrop(256, 256, p=1.),\n            A.Resize(256, 256),\n            A.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.)\n\ntest_dataset = CassavaLeafDataset(\n    image_paths=test_image_paths,\n    targets=test_targets,\n    transform=test_aug,\n)\n\ntest_loader = DataLoader(test_dataset, batch_size=CFG.batch_size)","metadata":{"execution":{"iopub.status.busy":"2021-06-27T18:26:36.628361Z","iopub.execute_input":"2021-06-27T18:26:36.628878Z","iopub.status.idle":"2021-06-27T18:26:36.646838Z","shell.execute_reply.started":"2021-06-27T18:26:36.628824Z","shell.execute_reply":"2021-06-27T18:26:36.645683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_img = test_dataset[0]['image']\n# test_img","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:49.828331Z","iopub.status.idle":"2021-06-27T17:40:49.829391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# torch.from_numpy(test_img)","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:49.831109Z","iopub.status.idle":"2021-06-27T17:40:49.831980Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# cassava_model","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:49.833331Z","iopub.status.idle":"2021-06-27T17:40:49.834056Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"best_checkpoints = trainer.checkpoint_callback.best_model_path\n# pretrained_model = CassavaLite().load_from_checkpoint(checkpoint_path = best_checkpoints)\n# pretrained_model = pretrained_model.to(\"cuda\")\npretrained_model = CassavaLeafModel().load_from_checkpoint(checkpoint_path = best_checkpoints)\npretrained_model = pretrained_model.to('cuda')\npretrained_model.eval()\npretrained_model.freeze()\n# pretrained_model","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:49.835380Z","iopub.status.idle":"2021-06-27T17:40:49.836132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fin_out = []\nfor data in test_loader:\n    y_hat = pretrained_model(data[\"image\"].to(\"cuda\"))\n    y_hat = torch.argmax(y_hat,dim=1)\n    fin_out.extend(y_hat.cpu().detach().numpy().tolist())\ntest_df_[\"label\"] = fin_out\ntest_df_[[\"image_id\",\"label\"]].to_csv(\"submission.csv\",index=False)\ntest_df_.head()","metadata":{"execution":{"iopub.status.busy":"2021-06-27T17:40:49.837708Z","iopub.status.idle":"2021-06-27T17:40:49.838490Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}