{"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":"gpu","dataSources":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"},{"sourceId":7342249,"sourceType":"datasetVersion","datasetId":4263116}],"dockerImageVersionId":30626,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os, shutil\n\ndef pip_install(name):\n    shutil.copytree(f'/kaggle/input/{name}', f'/kaggle/working/package/{name}')\n    os.system(f'cd /kaggle/working/package/{name}; pip install . 1>/dev/null')\n    shutil.rmtree(f'/kaggle/working/package/{name}')\n\ndef setup():\n    os.system('cp -r /kaggle/input/ubc-ocean-7th/utils /kaggle/working')\n\n    print('Start installing yangdl...')\n    pip_install('ubc-ocean-7th/yangdl')\n\n    print('Start installing libvips...')\n    os.system('yes | sudo dpkg -i /kaggle/input/ubc-ocean-7th/libvips/*.deb 1>/dev/null')\n\n    print('Start installing pyvips...')\n    os.system('pip install /kaggle/input/ubc-ocean-7th/pyvips/pyvips-2.2.2-py2.py3-none-any.whl --no-index --find-links /kaggle/input/ubc-ocean-7th/pyvips 1>/dev/null')\n\n    print('Start installing einops for Perceiver...')\n    os.system('pip install /kaggle/input/ubc-ocean-7th/einops/einops-0.7.0-py3-none-any.whl 1>/dev/null')\n\n    print('Finish installing.')\n\nsetup()\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-01-06T03:24:14.490413Z","iopub.execute_input":"2024-01-06T03:24:14.491713Z","iopub.status.idle":"2024-01-06T03:25:53.411402Z","shell.execute_reply.started":"2024-01-06T03:24:14.491673Z","shell.execute_reply":"2024-01-06T03:25:53.410402Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import math\nimport os\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport numpy as np\nimport pandas as pd\nimport pyvips\nimport torch\nfrom torch.nn import functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nimport yangdl as yd\nfrom utils import (\n    CTransPath,\n    DSMIL,\n    Perceiver,\n    get_file_names,\n    rgb2gray,\n    get_biggest_component_box,\n)\n","metadata":{"execution":{"iopub.status.busy":"2024-01-06T03:25:53.413824Z","iopub.execute_input":"2024-01-06T03:25:53.414242Z","iopub.status.idle":"2024-01-06T03:25:53.421609Z","shell.execute_reply.started":"2024-01-06T03:25:53.414205Z","shell.execute_reply":"2024-01-06T03:25:53.420555Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"PATCH_SIZE = 256\nTHRESH = 0.4\n\nIMAGES_PATH = '/kaggle/input/UBC-OCEAN/test_images'\nlabels = ['CC', 'EC', 'HGSC', 'LGSC', 'MC', 'Other']\n","metadata":{"execution":{"iopub.status.busy":"2024-01-06T03:25:53.422825Z","iopub.execute_input":"2024-01-06T03:25:53.423179Z","iopub.status.idle":"2024-01-06T03:25:53.436737Z","shell.execute_reply.started":"2024-01-06T03:25:53.423146Z","shell.execute_reply":"2024-01-06T03:25:53.43579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MyModelModule(yd.ModelModule):\n    def __init__(self):\n        super().__init__()\n\n        # 1. ctranspath\n        self.ctrans = CTransPath(num_classes=0)\n        self.ctrans.load_state_dict(torch.load(f'/kaggle/input/ubc-ocean-7th/weights/ctranspath/1.pt')['model'], strict=False)\n        self.ctrans_dsmils = []\n        self.ctrans_perceivers = []\n        for fold in (1, 2, 3, 4, 5):\n            dsmil = DSMIL(\n                num_classes=5,\n                size=[768, 128, 128],\n                dropout=0.,\n            )\n            dsmil.load_state_dict(torch.load(f'/kaggle/input/ubc-ocean-7th/weights/ctrans-dsmil/{fold}.pt')['model'])\n            self.ctrans_dsmils.append(dsmil)\n            self.register_models({f'ctrans_dsmil{fold}': dsmil})\n        for fold in (1, 2, 3, 4, 5):\n            perceiver = Perceiver(\n                input_channels=768,\n                input_axis=1,\n                num_freq_bands=6,\n                max_freq=10.,\n                depth=1,\n                num_latents=1024,\n                latent_dim=768,\n                cross_heads=1,\n                latent_heads=8,\n                cross_dim_head=64,\n                latent_dim_head=64,\n                n_classes=5,\n                attn_dropout=0.2,\n                ff_dropout=0.2,\n                weight_tie_layers=True,\n                fourier_encode_data=False,\n                self_per_cross_attn=1,\n                latent_bounds=2,\n                scale=0.125,\n            )\n            perceiver.load_state_dict(torch.load(f'/kaggle/input/ubc-ocean-7th/weights/ctrans-perceiver/{fold}.pt')['model'])\n            self.ctrans_perceivers.append(perceiver)\n            self.register_models({f'ctrans_perceiver{fold}': perceiver})\n\n        # 2. vits16\n        self.vits16 = torch.load(f'/kaggle/input/ubc-ocean-7th/weights/vits16/1.pt')\n        self.vits16_dsmils = []\n        self.vits16_perceivers = []\n        for fold in (1, 2, 3, 4, 5):\n            dsmil = DSMIL(\n                num_classes=5,\n                size=[384, 128, 128],\n                dropout=0.,\n            )\n            dsmil.load_state_dict(torch.load(f'/kaggle/input/ubc-ocean-7th/weights/vits16-dsmil/{fold}.pt')['model'])\n            self.vits16_dsmils.append(dsmil)\n            self.register_models({f'vits16_dsmil{fold}': dsmil})\n        for fold in (1, 2, 3, 4, 5):\n            perceiver = Perceiver(\n                input_channels=384,\n                input_axis=1,\n                num_freq_bands=6,\n                max_freq=10.,\n                depth=1,\n                num_latents=1024,\n                latent_dim=384,\n                cross_heads=1,\n                latent_heads=8,\n                cross_dim_head=64,\n                latent_dim_head=64,\n                n_classes=5,\n                attn_dropout=0.2,\n                ff_dropout=0.2,\n                weight_tie_layers=True,\n                fourier_encode_data=False,\n                self_per_cross_attn=1,\n                latent_bounds=2,\n                scale=0.125,\n            )\n            perceiver.load_state_dict(torch.load(f'/kaggle/input/ubc-ocean-7th/weights/vits16-perceiver/{fold}.pt')['model'])\n            self.vits16_perceivers.append(perceiver)\n            self.register_models({f'vits16_perceiver{fold}': perceiver})\n\n        self.res = {'image_id': [], 'label': []}\n\n    # x (N, 3, 224, 224)\n    def gen_features(self, encoder, x):\n        features = []\n        for i in range(0, len(x), 64):\n            features.append(encoder(x[i: i + 64]))\n        features = torch.cat(features, dim=0)\n\n        return features\n\n    def predict_step(self, batch):\n        patches, image_id = batch['patches'], batch['image_id']\n        patches, image_id = patches[0], image_id[0]\n\n        self.res['image_id'].append(image_id)\n\n        if patches.shape == (0,):\n            self.res['label'].append(labels[5])\n        else:\n            probs = []\n            perceiver_probs = []\n\n            # ctranspath\n            features = self.gen_features(self.ctrans, patches)  # (N, 768)\n            for dsmil in self.ctrans_dsmils:\n                bag_logits, inst_logits, _, _ = dsmil(features)\n                inst_logits, _ = torch.max(inst_logits, dim=0)\n                bag_prob = F.softmax(bag_logits, dim=0)\n                inst_prob = F.softmax(inst_logits, dim=0)\n                probs.append((bag_prob + inst_prob) / 2)\n            for perceiver in self.ctrans_perceivers:\n                logits, _, _, _, _ = perceiver(features)\n                prob = F.sigmoid(logits[0])\n                probs.append(prob)\n                perceiver_probs.append(prob)\n\n            # vits16\n            features = self.gen_features(self.vits16, patches)  # (N, 384)\n            for dsmil in self.vits16_dsmils:\n                bag_logits, inst_logits, _, _ = dsmil(features)\n                inst_logits, _ = torch.max(inst_logits, dim=0)\n                bag_prob = F.softmax(bag_logits, dim=0)\n                inst_prob = F.softmax(inst_logits, dim=0)\n                probs.append((bag_prob + inst_prob) / 2)\n            for perceiver in self.vits16_perceivers:\n                logits, _, _, _, _ = perceiver(features)\n                prob = F.sigmoid(logits[0])\n                probs.append(prob)\n                perceiver_probs.append(prob)\n\n            probs = torch.stack(probs, dim=0).mean(dim=0)  # (5,)\n            pred = probs.argmax(dim=0).item()\n\n            perceiver_probs = torch.stack(perceiver_probs, dim=0).mean(dim=0)  # (5,)\n            if max(perceiver_probs) < THRESH:\n                pred = 5\n\n            self.res['label'].append(labels[pred])\n","metadata":{"execution":{"iopub.status.busy":"2024-01-06T03:25:53.438248Z","iopub.execute_input":"2024-01-06T03:25:53.438545Z","iopub.status.idle":"2024-01-06T03:25:53.465605Z","shell.execute_reply.started":"2024-01-06T03:25:53.43852Z","shell.execute_reply":"2024-01-06T03:25:53.46467Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MyDataset(Dataset):\n    def __init__(self):\n        super().__init__()\n\n        self.image_ids = get_file_names(IMAGES_PATH, '.png')\n        self.transform = A.Compose([\n            A.Resize(224, 224),\n            A.Normalize(mean=(0.815, 0.695, 0.808), std=(0.129, 0.147, 0.112)),\n            ToTensorV2(),\n        ])\n\n    def __len__(self):\n        return len(self.image_ids)\n    \n    def __getitem__(self, idx):\n        image_id = self.image_ids[idx]\n        image, is_tma = self.read_png(f'{IMAGES_PATH}/{image_id}.png')\n        patches = self.image2patches(image, patch_size=PATCH_SIZE, step=[256, 64][is_tma], ratio=0.25, transform=self.transform, is_tma=is_tma)\n\n        return {'patches': patches, 'image_id': image_id}\n\n    @staticmethod\n    def read_png(image_id: str):\n        image = pyvips.Image.new_from_file(image_id, access='sequential').numpy()\n        is_tma = image.shape[0] <= 5000 and image.shape[1] <= 5000\n\n        # 1. downsample\n        if is_tma:\n            resize = A.Resize(image.shape[0] // 4, image.shape[1] // 4)\n        else:\n            resize = A.Resize(image.shape[0] // 2, image.shape[1] // 2)\n        image = resize(image=image)['image']\n\n        # 2. deduplicate for WSI\n        if not is_tma:\n            resize = A.Resize(image.shape[0] // 16, image.shape[1] // 16)  # downsample for speed\n            thumbnail = resize(image=image)['image']\n            mask = rgb2gray(thumbnail) > 0\n            x0, y0, x1, y1 = get_biggest_component_box(mask)\n\n            # resize box\n            scale_h = image.shape[0] / thumbnail.shape[0]\n            scale_w = image.shape[1] / thumbnail.shape[1]\n\n            x0 = max(0, math.floor(x0 * scale_w))\n            y0 = max(0, math.floor(y0 * scale_h))\n            x1 = min(image.shape[1] - 1, math.ceil(x1 * scale_w))\n            y1 = min(image.shape[0] - 1, math.ceil(y1 * scale_h))\n            image = image[y0: y1 + 1, x0: x1 + 1]\n\n        return image, is_tma\n\n    @staticmethod\n    def image2patches(image: np.ndarray, patch_size: int, step: int, ratio: float, transform, is_tma: bool):\n        \"\"\"\n        Args:\n            image (H, W, 3)\n\n        Returns:\n            patches: (N, 256, 256, 3), np.uint8\n        \"\"\"\n\n        patches = []\n        for i in range(0, image.shape[0], step):\n            for j in range(0, image.shape[1], step):\n                patch = image[i: i + patch_size, j: j + patch_size, :]\n                if patch.shape != (patch_size, patch_size, 3):\n                    patch = np.pad(patch, ((0, patch_size - patch.shape[0]), (0, patch_size - patch.shape[1]), (0, 0)))\n\n                if is_tma:\n                    patch = transform(image=patch)['image']\n                    patches.append(patch)\n                else:\n                    patch_gray = rgb2gray(patch) # (patch_size, patch_size)\n                    patch_binary = (patch_gray <= 220) & (patch_gray > 0)\n\n                    if np.count_nonzero(patch_binary) / patch_binary.size >= ratio:\n                        patch = transform(image=patch)['image']\n                        patches.append(patch)\n\n        if len(patches) != 0:\n            patches = torch.stack(patches, dim=0)\n        else:\n            patches = torch.zeros(0, dtype=torch.uint8)\n\n        return patches\n\n\nclass MyDataModule(yd.DataModule):\n    def __init__(self):\n        super().__init__()\n\n    def predict_loader(self):\n        dataset = MyDataset()\n\n        yield DataLoader(\n            dataset,\n            batch_size=1,\n            num_workers=2,\n            shuffle=False,\n            drop_last=False,\n            pin_memory=False\n        )\n","metadata":{"execution":{"iopub.status.busy":"2024-01-06T03:25:53.467793Z","iopub.execute_input":"2024-01-06T03:25:53.468056Z","iopub.status.idle":"2024-01-06T03:25:53.489815Z","shell.execute_reply.started":"2024-01-06T03:25:53.468033Z","shell.execute_reply":"2024-01-06T03:25:53.488739Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def submit(res):\n    pd.DataFrame(res).to_csv('/kaggle/working/submission.csv', index=None)\n\nmodel_module = MyModelModule()\ndata_module = MyDataModule()\ntask_module = yd.TaskModule(model_module, data_module)\n\ntask_module.do()\nsubmit(model_module.res)\n","metadata":{"execution":{"iopub.status.busy":"2024-01-06T03:25:53.491074Z","iopub.execute_input":"2024-01-06T03:25:53.491346Z","iopub.status.idle":"2024-01-06T03:26:23.696319Z","shell.execute_reply.started":"2024-01-06T03:25:53.491322Z","shell.execute_reply":"2024-01-06T03:26:23.695539Z"},"trusted":true},"execution_count":null,"outputs":[]}]}