{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"},{"sourceId":7323157,"sourceType":"datasetVersion","datasetId":3918004},{"sourceId":148233557,"sourceType":"kernelVersion"}],"dockerImageVersionId":30559,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install 'mmcv' 'mmpretrain' 'lightning' 'timm' 'segmentation-models-pytorch' 'einops' --find-links=/kaggle/input/ubc2023env --no-index\n!pip install '/kaggle/input/ubcsubmodels/torchstain-1.3.0-py3-none-any.whl'","metadata":{"execution":{"iopub.status.busy":"2023-12-31T01:00:51.676528Z","iopub.execute_input":"2023-12-31T01:00:51.676916Z","iopub.status.idle":"2023-12-31T01:01:51.32624Z","shell.execute_reply.started":"2023-12-31T01:00:51.676885Z","shell.execute_reply":"2023-12-31T01:01:51.325033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\n\nos.environ['OPENCV_IO_MAX_IMAGE_PIXELS'] = str(pow(2, 52))","metadata":{"execution":{"iopub.status.busy":"2023-12-31T01:01:51.329067Z","iopub.execute_input":"2023-12-31T01:01:51.329515Z","iopub.status.idle":"2023-12-31T01:01:51.334766Z","shell.execute_reply.started":"2023-12-31T01:01:51.329474Z","shell.execute_reply":"2023-12-31T01:01:51.333778Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import random\nfrom glob import glob\nfrom typing import Any\n\nimport albumentations as A\nimport cv2\nimport lightning as L\nimport segmentation_models_pytorch as smp\nimport torch\nimport torch.nn as nn\n# from pytorch_lightning.callbacks import ModelCheckpoint\nfrom lightning.pytorch.callbacks import ModelCheckpoint\nfrom lightning.pytorch.loggers import WandbLogger\nfrom lightning.pytorch.utilities.types import STEP_OUTPUT, OptimizerLRScheduler\n# from pytorch_lightning.loggers import WandbLogger\nfrom segmentation_models_pytorch.metrics import iou_score, get_stats\nfrom torch.utils.data import DataLoader, Dataset\nimport wandb\nimport timm\nimport torch.nn.functional as F\nfrom PIL import Image\nfrom mmpretrain.registry import MODELS\nfrom mmengine.runner import load_checkpoint","metadata":{"execution":{"iopub.status.busy":"2023-12-31T01:01:51.336238Z","iopub.execute_input":"2023-12-31T01:01:51.33699Z","iopub.status.idle":"2023-12-31T01:02:23.767112Z","shell.execute_reply.started":"2023-12-31T01:01:51.336954Z","shell.execute_reply":"2023-12-31T01:02:23.766285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SegModelStage1(L.LightningModule):\n    def __init__(self, encoder_name, num_classes=3):\n        super().__init__()\n        self.num_classes = num_classes\n        self.model = smp.Unet(encoder_name=encoder_name, encoder_weights=None, in_channels=3, classes=num_classes,\n                              activation=None)\n\n        self.valid_scores = []\n\n    def forward(self, x):\n        return self.model(x)\n\n    def training_step(self, batch: Any, batch_idx: int) -> STEP_OUTPUT:\n        x, y = batch\n        # assert torch.sum(y == 0) > 0 and torch.sum(y == 1) > 0 and torch.sum(y == 2) > 0 and torch.sum(\n        #     y == 3) > 0 and torch.sum(\n        #     y == 4) > 0, f'Gt is wrong, {torch.unique(y)}, {torch.sum(y == 0)}, {torch.sum(y == 1)}, {torch.sum(y == 2)}, {torch.sum(y == 3)}, {torch.sum(y == 4)}'\n        y_pred = self.model(x)\n        loss = F.cross_entropy(\n            y_pred,\n            y,\n            reduction='mean',\n            label_smoothing=0.1,\n            ignore_index=255)\n        self.log('train_loss', loss, prog_bar=True)\n        return loss\n\n    def validation_step(self, batch: Any, batch_idx: int) -> STEP_OUTPUT:\n        x, y = batch\n        y_pred = self.model(x)\n        loss = F.cross_entropy(\n            y_pred,\n            y,\n            reduction='mean',\n            label_smoothing=0.1,\n            ignore_index=255)\n        self.log('val_loss', loss, prog_bar=True)\n        return loss\n\n    def on_validation_epoch_end(self) -> None:\n        pass\n\n    def configure_optimizers(self) -> OptimizerLRScheduler:\n        optimizer = torch.optim.AdamW(\n            self.model.parameters(),\n            lr=1e-4,\n            eps=1e-6,\n        )\n        max_epochs = self.trainer.max_epochs\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=max_epochs, eta_min=2e-5)\n        return {\n            'optimizer': optimizer,\n            'lr_scheduler': scheduler,\n        }\n\n    def metric(self, pred, ann, only_foreground=True):\n        tp, fp, fn, tn = get_stats(output=pred, target=ann, mode='multiclass', num_classes=self.num_classes,\n                                   threshold=None)\n        if only_foreground:\n            # Ignore background class (assuming it's the first class, index 0)\n            tp, fp, fn, tn = tp[1:], fp[1:], fn[1:], tn[1:]\n        iou = iou_score(tp, fp, fn, tn, reduction='micro-imagewise')\n        return iou\n    \nclass SegModelFocus(L.LightningModule):\n    def __init__(self, encoder_name, num_classes=1):\n        super().__init__()\n        self.num_classes = num_classes\n        self.model = smp.Unet(encoder_name=encoder_name, encoder_weights=None, in_channels=3, classes=num_classes,\n                              activation=None)\n        self.loss_fn = torch.nn.BCEWithLogitsLoss()\n        self.valid_scores = []\n\n    def forward(self, x):\n        return self.model(x)","metadata":{"execution":{"iopub.status.busy":"2023-12-31T01:02:23.768737Z","iopub.execute_input":"2023-12-31T01:02:23.76909Z","iopub.status.idle":"2023-12-31T01:02:23.785535Z","shell.execute_reply.started":"2023-12-31T01:02:23.769059Z","shell.execute_reply":"2023-12-31T01:02:23.784507Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class UBCClsModel(L.LightningModule):\n    def __init__(self,model_name='convnext_base.clip_laion2b_augreg_ft_in12k_in1k_384'):\n        super().__init__()\n        self.model = timm.create_model(model_name, pretrained=False,\n                                       num_classes=6)\n#         self.model.set_grad_checkpointing()\n        # self.loss_w = torch.tensor([0.5, 0.5, 0.5, 1.0, 0.5, 0.5], dtype=torch.bfloat16)\n        self.loss_fn = nn.CrossEntropyLoss(label_smoothing=0.1)\n        self.seesaw_fn = SeesawLoss(num_classes=6)\n        self.valid_labels = []\n        self.valid_logits = []\n\n    def forward(self, x):\n        return self.model(x)\n\n    def training_step(self, batch: Any, batch_idx: int) -> STEP_OUTPUT:\n        x, y = batch\n        logits = self.model(x)\n        loss = self.loss_fn(logits, y)\n        loss += self.seesaw_fn(logits, y)\n        self.log('train/loss', loss, prog_bar=True)\n        return loss\n\n    def configure_optimizers(self) -> OptimizerLRScheduler:\n        optimizer = AdamW(\n            self.model.parameters(),\n            lr=1e-4,\n            eps=1e-6,\n        )\n        max_epochs = self.trainer.max_epochs\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=max_epochs, eta_min=2e-5)\n        return {\n            'optimizer': optimizer,\n            'lr_scheduler': scheduler,\n        }\n\n    def validation_step(self, batch: Any, batch_idx: int) -> STEP_OUTPUT:\n        x, y = batch\n        logits = self.model(x)\n        self.valid_logits.append(logits)\n        self.valid_labels.append(y)\n        return None\n\n    def on_validation_epoch_end(self) -> None:\n        self.valid_logits = torch.cat(self.valid_logits, dim=0)\n        self.valid_labels = torch.cat(self.valid_labels, dim=0)\n        loss = self.loss_fn(self.valid_logits, self.valid_labels)\n        loss += self.seesaw_fn(self.valid_logits, self.valid_labels)\n        self.log('val/loss', loss, prog_bar=True)\n        # f1 score for each class\n        self.valid_pred_labels = torch.argmax(self.valid_logits, dim=1)\n        f1_cc = f1_score(torch.where(self.valid_labels == 0, 1, 0).cpu().numpy().flatten(),\n                         torch.where(self.valid_pred_labels == 0, 1, 0).cpu().numpy().flatten())\n        f1_ec = f1_score(torch.where(self.valid_labels == 1, 1, 0).cpu().numpy().flatten(),\n                         torch.where(self.valid_pred_labels == 1, 1, 0).cpu().numpy().flatten())\n        f1_hgsc = f1_score(torch.where(self.valid_labels == 2, 1, 0).cpu().numpy().flatten(),\n                           torch.where(self.valid_pred_labels == 2, 1, 0).cpu().numpy().flatten())\n        f1_lgsc = f1_score(torch.where(self.valid_labels == 3, 1, 0).cpu().numpy().flatten(),\n                           torch.where(self.valid_pred_labels == 3, 1, 0).cpu().numpy().flatten())\n        f1_mc = f1_score(torch.where(self.valid_labels == 4, 1, 0).cpu().numpy().flatten(),\n                         torch.where(self.valid_pred_labels == 4, 1, 0).cpu().numpy().flatten())\n        self.log('val/f1_cc', f1_cc, prog_bar=True, sync_dist=True)\n        self.log('val/f1_ec', f1_ec, prog_bar=True, sync_dist=True)\n        self.log('val/f1_hgsc', f1_hgsc, prog_bar=True, sync_dist=True)\n        self.log('val/f1_lgsc', f1_lgsc, prog_bar=True, sync_dist=True)\n        self.log('val/f1_mc', f1_mc, prog_bar=True, sync_dist=True)\n\n        self.valid_logits = []\n        self.valid_labels = []\n\n        \n        \nclass UBCSwin(L.LightningModule):\n    def __init__(self, model_name='convnext_base.clip_laion2b_augreg_ft_in12k_in1k_384', lr=1e-4, lr_min=2e-5,\n                 lr_decay=0.8, img_size=1536):\n        super().__init__()\n        cfg_dict = {'type': 'ImageClassifier',\n                    'backbone': {'type': 'SwinTransformerV2', 'arch': 'base', 'img_size': img_size,\n                                 'drop_rate': 0.1, 'drop_path_rate': 0.2,\n                                 'window_size': [24, 24, 24, 12], 'pretrained_window_sizes': [12, 12, 12, 6],\n                                 'with_cp': False, },\n                    'neck': {'type': 'GlobalAveragePooling'},\n                    'head': {'type': 'LinearClsHead', 'num_classes': 6, 'in_channels': 1024,\n                             'loss': {'type': 'CrossEntropyLoss', 'loss_weight': 1.0},\n                             'topk': (1, 5)}}\n        self.model = MODELS.build(cfg_dict)\n        # self.loss_w = torch.tensor([0.5, 0.5, 0.5, 1.0, 0.5, 0.5], dtype=torch.bfloat16)\n        self.loss_fn = nn.CrossEntropyLoss(label_smoothing=0.1)  # cc, ec, hgsc, lgsc, mc, other\n        self.seesaw_fn = SeesawLoss(num_classes=6)\n        self.valid_labels = []\n        self.valid_logits = []\n        self.lr = lr\n        self.lr_min = lr_min\n        self.lr_decay = lr_decay\n        self.alpha = 0.2\n\n    def forward(self, x):\n        return self.model(x)","metadata":{"execution":{"iopub.status.busy":"2023-12-31T01:02:23.788546Z","iopub.execute_input":"2023-12-31T01:02:23.789213Z","iopub.status.idle":"2023-12-31T01:02:23.814762Z","shell.execute_reply.started":"2023-12-31T01:02:23.789183Z","shell.execute_reply":"2023-12-31T01:02:23.813708Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torchstain\nclass stain_normalizer:\n    def __init__(self):\n        stain_target = cv2.cvtColor(cv2.imread(\"/kaggle/input/ubcsubmodels/target.png\"), cv2.COLOR_BGR2RGB)\n        self.strain_normalizer = torchstain.normalizers.MacenkoNormalizer(backend='numpy')\n        self.strain_normalizer.fit(stain_target)\n\n    def normalize(self, img):\n        img = np.array(img)\n        try:\n            norm, _, _ = self.strain_normalizer.normalize(I=img, stains=True)\n            norm = norm.astype(np.uint8)\n        except:\n            norm = img\n        return Image.fromarray(norm)","metadata":{"execution":{"iopub.status.busy":"2023-12-31T01:02:23.816086Z","iopub.execute_input":"2023-12-31T01:02:23.816371Z","iopub.status.idle":"2023-12-31T01:02:23.835242Z","shell.execute_reply.started":"2023-12-31T01:02:23.816347Z","shell.execute_reply":"2023-12-31T01:02:23.834254Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nimport pandas as pd\nimport cv2\nimport numpy as np\nfrom glob import glob\nimport math\nfrom tqdm import tqdm\nimport torch\nimport albumentations as A\nimport gc\nfrom torchvision import transforms\nfrom mmpretrain.models.losses import SeesawLoss, CrossEntropyLoss\nfrom torchvision.transforms.functional import InterpolationMode\nfrom timm.data import IMAGENET_DEFAULT_MEAN, IMAGENET_DEFAULT_STD\nimport rasterio\nfrom concurrent.futures import ThreadPoolExecutor\nfrom multiprocessing import Pool\nfrom collections import deque\nm1 = SegModelFocus(encoder_name='tu-seresnextaa101d_32x8d.sw_in12k_ft_in1k_288', num_classes=1)\nweight = torch.load('/kaggle/input/ubcsubmodels/model_segFocusFinal2d5x.ckpt')\nm1.load_state_dict(weight['state_dict'])\nm1.eval()\nm1.to('cuda:1')\n# m1.half()\ndel weight\ngc.collect()\n\ntransform1 = A.Compose([\n    A.LongestMaxSize(1536 * 2),\n])\ntransform2 = A.Compose([\n    A.PadIfNeeded(None, None,\n                  pad_height_divisor=32,\n                  pad_width_divisor=32,\n                  position='top_left'),\n    A.Normalize(),\n])\n\ntransform_cam = A.Compose([\n                A.Normalize(),\n            ])\ntransform_cropwsi = A.Compose([\n                A.CenterCrop(height=50000, width=50000),\n            ])\ndef read_band_window(src_path, band, window):\n    with rasterio.open(src_path) as src:\n        return src.read(band, window=window)\nimport matplotlib.pyplot as plt\ndef get_patches_wsi(image_file, h, w, num_patches=5):\n    patches = deque(maxlen=num_patches)\n    patches_score = deque(maxlen=num_patches)\n    if h*w > img_max_size_thres:\n        h = h if h<50000 else 50000\n        w = w if w<50000 else 50000\n        with rasterio.open(image_file) as src:\n            window = rasterio.windows.Window(5000, 5000, w-5000, h-5000)\n            red = src.read(1, window=window)\n            green = src.read(2, window=window)\n            blue = src.read(3, window=window)\n        img =np.dstack((blue, green, red))\n        del window, red, green, blue\n        gc.collect()\n    else:\n        img = cv2.imread(image_file)\n        h, w = img.shape[:2]\n        if h > w*1.8:\n            img = img[:h//2, :w, :]\n        elif w > h*1.8:\n            img = img[:h, :w//2, :]\n    h, w = img.shape[:2]\n    if min(h,w)>50000:\n        img = transform_cropwsi(image=img)['image']\n        h, w = img.shape[:2]\n#     img_resized = transform1(image=img)['image']\n    h_re, w_re = h // 8, w // 8  # 8x downsample to 2.5x magnification\n    img_resized = cv2.resize(img, (w_re, h_re), interpolation=cv2.INTER_CUBIC)\n    h_re, w_re = img_resized.shape[:2]\n#     plt.imshow(img_resized)\n#     plt.show()\n    img_resized = transform2(image=img_resized)['image']\n    with torch.no_grad():\n        img_resized = torch.from_numpy(img_resized).permute(2, 0, 1).unsqueeze(0).to('cuda:1')\n        with torch.cuda.amp.autocast():\n            pred_label = m1(img_resized).sigmoid()\n        pred_label = pred_label.float().squeeze()\n        pred_label[pred_label < 0.4] = 0.0\n        del img_resized\n        pred_label = pred_label[:h_re, :w_re]\n    h, w = img.shape[:2]\n    n_h = math.ceil(h / 1536)\n    n_w = math.ceil(w / 1536)\n    scale = 8\n    scale_s = 1536 // scale\n    b_score = 1024.0\n    for i_h in range(n_h):\n        for j_w in range(n_w):\n            x1, y1 = j_w * scale_s, i_h * scale_s\n            x2 = x1 + scale_s\n            y2 = y1 + scale_s\n            if x2 > w_re:\n                x2 = w_re\n                x1 = w_re - scale_s\n            if y2 > h_re:\n                y2 = h_re\n                y1 = h_re - scale_s\n            # check if seg area large than 60%\n            x1, y1, x2, y2 = int(x1), int(y1), int(x2), int(y2)\n            p_score = torch.sum(pred_label[y1:y2, x1:x2]).item()\n#             print(f\"p score: {p_score}\")\n            if len(patches) > 0:\n                if p_score < max(patches_score):\n                    continue\n            x1, y1 = j_w * 1536, i_h * 1536\n            x2 = x1 + 1536\n            y2 = y1 + 1536\n            if x2 > w:\n                x2 = w\n                x1 = w - 1536\n            if y2 > h:\n                y2 = h\n                y1 = h - 1536\n            if len(patches) == 0:\n                if p_score > b_score:\n                    img_crop = img[y1:y2, x1:x2]\n                    img_crop = cv2.cvtColor(img_crop, cv2.COLOR_BGR2RGB)\n                    img_crop = Image.fromarray(img_crop)\n                    patches.append(img_crop)\n                    patches_score.append(p_score)\n            else:\n                if p_score > max(patches_score):\n                    img_crop = img[y1:y2, x1:x2]\n                    img_crop = cv2.cvtColor(img_crop, cv2.COLOR_BGR2RGB)\n                    img_crop = Image.fromarray(img_crop)\n                    patches.append(img_crop)\n                    patches_score.append(p_score)\n                    \n#     print('end scan')\n#     print(f\"lenght of patches: {len(patches)}\")\n    del img, pred_label\n    gc.collect()\n    torch.cuda.empty_cache()\n#     print('sart print sacores')\n#     print(f\"selected patch p scores: {patches_score}\")\n    return list(patches)\n\ntma_tta1 = transforms.Compose([\n    transforms.RandomHorizontalFlip(p=1.0),\n])\ntma_tta2 = transforms.Compose([\n    transforms.RandomVerticalFlip(p=1.0),\n    transforms.RandomRotation(45),\n    transforms.RandomGrayscale(p=1.0),\n])\ndef get_tma_img(image_file):\n    img = Image.open(image_file)\n    return [img,]\n\ncat2label ={\"CC\":0,\"EC\":1,\"HGSC\":2,\"LGSC\":3,\"MC\":4,\"Other\":5,}\nlabel2cat = {v: k for k, v in cat2label.items()}\n\ndf_test = pd.read_csv('/kaggle/input/UBC-OCEAN/test.csv')\ndf_train = pd.read_csv('/kaggle/input/UBC-OCEAN/train.csv')\ndf_train['size'] = df_train['image_width'] * df_train['image_height']\nimg_max_size = df_train['size'].max()\nimg_max_size_thres = img_max_size * 0.95\n# img_max_size_thres = 3805092100\nprint(img_max_size, 50000**2 / img_max_size)\nprint(np.sum(df_train['size']>img_max_size_thres))\n# df_train = df_train[df_train['is_tma']]\n# df_test = df_train.sort_values('size',ascending=False).head(20)\n# df_test = df_test[5:15].copy()\n# print(df_test['size'])\n# inference the wsi and tma\n# load model\nm_cls = UBCClsModel(model_name='convnext_base.clip_laion2b_augreg_ft_in12k_in1k_384')\nweight = torch.load(f'/kaggle/input/ubcsubmodels/Ex_Pseudo_Base_StainNorm_RM.ckpt')\nm_cls.load_state_dict(weight['state_dict'])\nm_cls.eval()\nm_cls.cuda()\n\nm_cls2 = UBCClsModel(model_name='convnext_large_mlp.clip_laion2b_soup_ft_in12k_in1k_384')\nweight = torch.load(f'/kaggle/input/ubcsubmodels/Ex_Pseudo_Large_StainNorm_RM.ckpt')\nm_cls2.load_state_dict(weight['state_dict'])\nm_cls2.eval()\nm_cls2.cuda()\n\nm_cls3 = UBCClsModel(model_name='eva02_large_patch14_448.mim_m38m_ft_in22k_in1k')\nweight = torch.load(f'/kaggle/input/ubcsubmodels/ExMoreOther_PseudoEVA_EVA02_448_Full.ckpt')\nm_cls3.load_state_dict(weight['state_dict'])\nm_cls3.eval()\nm_cls3.cuda()\ndel weight\ngc.collect()\n# make transform\ntrans_test_wsi = transforms.Compose([\n    transforms.ToTensor(),\n    transforms.Normalize(mean=IMAGENET_DEFAULT_MEAN, std=IMAGENET_DEFAULT_STD),\n])\ntrans_test_tma = transforms.Compose([\n    transforms.CenterCrop(size=(3072, 3072)),\n    transforms.Resize(size=(1536, 1536), interpolation=InterpolationMode.BICUBIC),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=IMAGENET_DEFAULT_MEAN, std=IMAGENET_DEFAULT_STD),\n])\nre448 = transforms.Compose([\n    transforms.Resize(size=(448, 448), interpolation=InterpolationMode.BICUBIC),\n])\npreds = []\nstainnorm = stain_normalizer()\nfor index, row in tqdm(df_test.iterrows(), total=len(df_test)):\n    imgs = None\n    gc.collect()\n    image_id = row['image_id']\n    h, w = row['image_height'], row['image_width']\n    try:\n        if max([h, w]) > 7000:\n            # WSI part\n            imgs = get_patches_wsi(f\"/kaggle/input/UBC-OCEAN/test_images/{image_id}.png\",h, w, num_patches=4)\n\n            with Pool(4) as p:\n                imgs += p.map(stainnorm.normalize, imgs)\n            imgs = [trans_test_wsi(img).unsqueeze(0).cuda() for img in imgs]\n        else:\n            # TMA part\n            imgs = get_tma_img(f\"/kaggle/input/UBC-OCEAN/test_images/{image_id}.png\")\n            with Pool(4) as p:\n                imgs += p.map(stainnorm.normalize, imgs)\n            imgs = [trans_test_tma(img).unsqueeze(0).cuda() for img in imgs]\n        with torch.no_grad():\n            imgs = torch.cat(imgs, dim=0).cuda()\n            with torch.cuda.amp.autocast():\n                pred = m_cls(imgs)\n                pred_softmax = pred.softmax(dim=1).mean(dim=0, keepdim=True)\n                # ensemble \n                pred2 = m_cls2(imgs)\n                pred_softmax2 = pred2.softmax(dim=1).mean(dim=0, keepdim=True)\n                pred3 = m_cls3(re448(imgs))\n                pred_softmax3 = pred3.softmax(dim=1).mean(dim=0, keepdim=True)\n                pred_softmax = (pred_softmax+pred_softmax2+pred_softmax3)/3\n            pred_label = pred_softmax.argmax(dim=1).squeeze().cpu().numpy()\n        cat = label2cat[pred_label.item()]\n        preds.append([image_id, cat])\n    except:\n        preds.append([image_id, 'Other'])\n\nprint('a')","metadata":{"execution":{"iopub.status.busy":"2023-12-31T01:04:32.904051Z","iopub.execute_input":"2023-12-31T01:04:32.904499Z","iopub.status.idle":"2023-12-31T01:05:39.098152Z","shell.execute_reply.started":"2023-12-31T01:04:32.904452Z","shell.execute_reply":"2023-12-31T01:05:39.09701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df = pd.DataFrame(preds, columns=['image_id','label'])\nsubmission_df.to_csv('submission.csv', index = False)\nsubmission_df","metadata":{"execution":{"iopub.status.busy":"2023-12-31T01:05:39.100119Z","iopub.execute_input":"2023-12-31T01:05:39.100478Z","iopub.status.idle":"2023-12-31T01:05:39.167015Z","shell.execute_reply.started":"2023-12-31T01:05:39.100427Z","shell.execute_reply":"2023-12-31T01:05:39.166086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}