{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","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":6774553,"sourceType":"datasetVersion","datasetId":3898019},{"sourceId":7120688,"sourceType":"datasetVersion","datasetId":4107067},{"sourceId":219386,"sourceType":"modelInstanceVersion","modelInstanceId":187094,"modelId":209179}],"dockerImageVersionId":30805,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!ls /kaggle/input/pyvips-python-and-deb-package-gpu\n# intall the deb packages\n!yes | dpkg -i --force-depends /kaggle/input/pyvips-python-and-deb-package-gpu/linux_packages/archives/*.deb\n# install the python wrapper\n!pip install pyvips -f /kaggle/input/pyvips-python-and-deb-package-gpu/python_packages/ --no-index\n!pip list | grep pyvips","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-15T13:02:32.637607Z","iopub.execute_input":"2025-01-15T13:02:32.63787Z","iopub.status.idle":"2025-01-15T13:04:06.256659Z","shell.execute_reply.started":"2025-01-15T13:02:32.637824Z","shell.execute_reply":"2025-01-15T13:04:06.255741Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pyvips","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-15T13:04:06.261646Z","iopub.execute_input":"2025-01-15T13:04:06.26194Z","iopub.status.idle":"2025-01-15T13:04:06.510988Z","shell.execute_reply.started":"2025-01-15T13:04:06.261913Z","shell.execute_reply":"2025-01-15T13:04:06.510415Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import albumentations as A","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nimport random\nimport cv2\nfrom PIL import Image, ImageFile\nimport gc \nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchvision import transforms\nimport torch.nn.functional as F\nimport timm\n\nImageFile.LOAD_TRUNCATED_IMAGES = True\nImage.MAX_IMAGE_PIXELS = None\nos.environ['VIPS_CONCURRENCY'] = '4'\nos.environ['VIPS_DISC_THRESHOLD'] = '15gb'\n\n# Paths and constants\ndata_dir = \"/kaggle/input/UBC-OCEAN/\"  # Update with your data directory\ntest_images_dir = os.path.join(data_dir, \"test_images\")\ntest_thumbnails_dir = os.path.join(data_dir, \"test_thumbnails\")\ntest_csv_path = os.path.join(data_dir, \"test.csv\")\nsubmission_path = \"submission.csv\"\nmodel_path = \"/kaggle/input/lunit-224-400-tiles-7698/pytorch/default/1/best_model_224_400tiles.pth\"\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nentropy_threshold = 1.5  # Threshold for declaring class 'Other'\n\n# Label mapping\nlabel_mapping = {'HGSC': 0, 'EC': 1, 'CC': 2, 'LGSC': 3, 'MC': 4}\nreverse_label_mapping = {v: k for k, v in label_mapping.items()}\nreverse_label_mapping[5] = 'Other'  # Add \"Other\" for entropy-based predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-05T06:55:03.632202Z","iopub.execute_input":"2025-01-05T06:55:03.632512Z","iopub.status.idle":"2025-01-05T06:55:11.482199Z","shell.execute_reply.started":"2025-01-05T06:55:03.632473Z","shell.execute_reply":"2025-01-05T06:55:11.481208Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Number of tiles and tile size\nn_tiles = 400\ntile_size = 224\nbatch_size = 1  # Adjust batch size as needed\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-05T06:55:11.484202Z","iopub.execute_input":"2025-01-05T06:55:11.484815Z","iopub.status.idle":"2025-01-05T06:55:11.488991Z","shell.execute_reply.started":"2025-01-05T06:55:11.484774Z","shell.execute_reply":"2025-01-05T06:55:11.488166Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"######### MIL Dataset\n\nclass WSI_TMA_MILDataset(Dataset):\n    def __init__(\n        self, \n        image_ids, \n        images_dir, \n        thumbnails_dir, \n        n_tiles, \n        tile_size, \n        transform=None\n    ):\n        \"\"\"\n        image_ids: list of image IDs (strings or integers)\n        images_dir: directory containing the full-resolution images\n        thumbnails_dir: directory containing the thumbnail images\n        n_tiles: number of tiles per image\n        tile_size: size of each tile\n        transform: torchvision transforms\n        \"\"\"\n        self.image_ids = image_ids\n        self.images_dir = images_dir\n        self.thumbnails_dir = thumbnails_dir\n        self.n_tiles = n_tiles\n        self.tile_size = tile_size\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.image_ids)\n\n    def __getitem__(self, idx):\n        \"\"\"\n        Returns:\n            tiles_tensor: A tensor of shape (n_tiles, C, tile_size, tile_size).\n            label: The numeric label (torch.long).\n        \"\"\"\n\n        image_id = self.image_ids[idx]\n        \n        # thumb_path = os.path.join(self.thumbnails_dir, f\"{image_id}_thumbnail.png\")\n        # if os.path.exists(thumb_path):\n        #     img_path = thumb_path\n        # else:\n        #     # Fallback to full image\n        #     img_path = os.path.join(self.images_dir, f\"{image_id}.png\")\n        img_path = os.path.join(self.images_dir,f\"{image_id}.png\")\n        try:\n            # Attempt to open the image\n            # image = Image.open(img_path).convert(\"RGB\")\n            # image = cv2.imread(img_path, cv2.IMREAD_COLOR)\n            image = pyvips.Image.new_from_file(img_path, access='sequential').numpy()\n\n        except (FileNotFoundError, IOError):\n            print(f\"Warning: Unable to load image {img_path}. Skipping.\")\n            return None, image_id  # Return a placeholder for the image with the image_id\n            \n        # img_pil = Image.fromarray(image)\n        # cropped1 = get_cropped_image(img_pil)\n        # cropped2 = crop_wsi_with_otsu(cropped1)\n        # img_rgb = cv2.cvtColor(np.array(img_pil), cv2.COLOR_BGR2RGB)\n        # h, w, c = img_rgb.shape\n\n        tiles = []\n        # Resize the image to 4096x4096\n        # img_resized = Image.fromarray(img_rgb).resize((12288, 12288))\n        # img_resized = np.array(img_resized)\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], image.shape[1])\n        else:\n            resize = A.Resize(image.shape[0]//2, image.shape[1]//2)\n        img_resized = resize(image=image)['image']\n        # Determine the size of each tile (assume 20x20 grid for 400 tiles)\n        tile_size = img_resized.shape[0] // 20\n\n        # Create 400 tiles of equal size\n        for i in range(20):  # Loop over rows\n            for j in range(20):  # Loop over columns\n                # Calculate the coordinates of the current tile\n                x_start = j * tile_size\n                y_start = i * tile_size\n                tile = img_resized[y_start:y_start + tile_size, x_start:x_start + tile_size, :]\n        \n                # Convert the tile to a PIL Image\n                tile_img = Image.fromarray(tile)\n        \n                # Apply transformation if specified\n                if self.transform:\n                    tile_img = self.transform(tile_img)\n        \n                tiles.append(tile_img)\n\n        tiles_tensor = torch.stack(tiles, dim=0)\n        del image, img_resized, resize\n        gc.collect()\n        torch.cuda.empty_cache()\n\n        return tiles_tensor, image_id","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-05T06:55:11.517191Z","iopub.execute_input":"2025-01-05T06:55:11.517479Z","iopub.status.idle":"2025-01-05T06:55:11.528561Z","shell.execute_reply.started":"2025-01-05T06:55:11.517454Z","shell.execute_reply":"2025-01-05T06:55:11.52795Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"########################################\n# MultiPatchViTExtractor\n########################################\nclass MultiPatchViTExtractor:\n    \"\"\"\n    Loads two ViT models (patch8 and patch16) from timm, extracts features from each,\n    and concatenates them into a single vector.\n    \"\"\"\n    def __init__(self, device=\"cpu\"):\n        self.device = device\n        \n        # Example model names\n        self.model_patch8  = timm.create_model(\"vit_small_patch8_224\",  pretrained=False)\n        self.model_patch16 = timm.create_model(\"vit_small_patch16_224\", pretrained=False)\n\n        # load lunit weights\n        self.model_patch8.load_state_dict(\n            torch.load(\"/kaggle/input/lunit-dino-weights/dino_vit_small_patch8_ep200.torch\", \n                       map_location=\"cpu\"), strict=False)\n        self.model_patch16.load_state_dict(\n            torch.load(\"/kaggle/input/lunit-dino-weights/dino_vit_small_patch16_ep200.torch\", \n                       map_location=\"cpu\"), strict=False)\n\n        # Remove classification heads\n        self.model_patch8.head  = nn.Identity()\n        self.model_patch16.head = nn.Identity()\n        \n        self.model_patch8.to(device).eval()\n        self.model_patch16.to(device).eval()\n    \n    @torch.no_grad()\n    def extract_features(self, image_tensor: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Expects image_tensor shape: (1, 3, 224, 224).\n        Returns concatenated features from patch8 & patch16, e.g. shape: (768+768,).\n        \"\"\"\n        feats8  = self.model_patch8(image_tensor).squeeze(0)   # shape (768,)\n        feats16 = self.model_patch16(image_tensor).squeeze(0)  # shape (768,)\n        return torch.cat([feats8, feats16], dim=0)  # shape (1536,)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-05T06:55:11.529521Z","iopub.execute_input":"2025-01-05T06:55:11.529767Z","iopub.status.idle":"2025-01-05T06:55:11.541556Z","shell.execute_reply.started":"2025-01-05T06:55:11.529743Z","shell.execute_reply":"2025-01-05T06:55:11.540856Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#######################################\n# MILAttentionModel (Double ViT Backbone)\n########################################\nclass MILAttentionDoubleDINO(nn.Module):\n    \"\"\"\n    Multi-Instance Learning model that uses the MultiPatchViTExtractor to get features \n    from each tile, applies attention, and then classifies.\n    \"\"\"\n    def __init__(self, device=\"cpu\", num_classes=5, embed_dim=512):\n        super().__init__()\n        # Instead of inception_v3, we use our double-ViT extractor\n        self.extractor = MultiPatchViTExtractor(device=device)\n        \n        # The concatenated output of patch8 and patch16 is 1536 dims\n        in_features = 768\n\n        # Project to an embedding space if needed\n        self.embed = nn.Linear(in_features, embed_dim)\n        \n        # Attention parameters\n        self.attention_A = nn.Linear(embed_dim, 128)\n        self.attention_B = nn.Linear(128, 1)\n        \n        # Classifier\n        self.classifier = nn.Linear(embed_dim, num_classes)\n\n        self.device = device\n\n    def forward(self, x):\n        \"\"\"\n        x shape: [B, N, C, H, W]\n        B: batch size (# of patients/slides)\n        N: # of tiles per slide\n        C, H, W: channels, height, width (e.g., 3, 224, 224)\n        \"\"\"\n        B, N, C, H, W = x.shape\n        \n        # We'll accumulate all tile features in a list, then stack\n        all_features = []\n        for b_idx in range(B):\n            # Extract features for N tiles in the current batch element\n            tile_features = []\n            for n_idx in range(N):\n                # Each tile is shape [C, H, W]\n                tile = x[b_idx, n_idx, ...].unsqueeze(0).to(self.device)  # shape [1, C, H, W]\n                \n                with torch.no_grad():\n                    feats = self.extractor.extract_features(tile)  # shape [1536,]\n                tile_features.append(feats.unsqueeze(0))  # shape [1, 1536]\n            \n            tile_features = torch.cat(tile_features, dim=0)  # shape [N, 1536]\n            all_features.append(tile_features.unsqueeze(0))  # shape [1, N, 1536]\n            del tile_features, tile\n            gc.collect()\n            torch.cuda.empty_cache()\n\n        \n        # Concatenate along batch dimension: [B, N, 1536]\n        all_features = torch.cat(all_features, dim=0).to(self.device)\n        \n        # Project to embedding dimension\n        embeddings = self.embed(all_features)  # [B, N, embed_dim]\n        \n        # Attention\n        A = torch.relu(self.attention_A(embeddings))  # [B, N, 128]\n        A = self.attention_B(A)                       # [B, N, 1]\n        A = torch.softmax(A, dim=1)                   # attention weights over tiles\n        \n        # Weighted sum of embeddings\n        weighted_sum = torch.sum(A * embeddings, dim=1)  # [B, embed_dim]\n\n        del all_features, embeddings\n        gc.collect()\n        torch.cuda.empty_cache()\n\n        # Final classification\n        logits = self.classifier(weighted_sum)  # [B, num_classes]\n        return logits","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-05T06:55:11.542431Z","iopub.execute_input":"2025-01-05T06:55:11.542666Z","iopub.status.idle":"2025-01-05T06:55:11.554851Z","shell.execute_reply.started":"2025-01-05T06:55:11.542643Z","shell.execute_reply":"2025-01-05T06:55:11.554188Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define transforms\ntest_transform = transforms.Compose([\n    transforms.Resize((224,224)),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-05T06:55:11.555774Z","iopub.execute_input":"2025-01-05T06:55:11.556494Z","iopub.status.idle":"2025-01-05T06:55:11.568534Z","shell.execute_reply.started":"2025-01-05T06:55:11.556457Z","shell.execute_reply":"2025-01-05T06:55:11.567784Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Extract all image IDs from the test_images directory\nimage_ids = [os.path.splitext(img)[0] for img in os.listdir(test_images_dir) if img.endswith(\".png\")]\n\n# Create test dataset and loader\ntest_dataset = WSI_TMA_MILDataset(image_ids,\n                                  n_tiles = n_tiles,\n                                  tile_size = tile_size,\n                                  images_dir=test_images_dir, \n                                  thumbnails_dir=test_thumbnails_dir,\n                                  transform=test_transform)\ntest_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=1)\n\n# Load model\nmodel = MILAttentionDoubleDINO(num_classes=5,device=device)\nmodel.load_state_dict(torch.load(model_path, map_location=device))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-05T06:55:11.569397Z","iopub.execute_input":"2025-01-05T06:55:11.569702Z","iopub.status.idle":"2025-01-05T06:55:14.426708Z","shell.execute_reply.started":"2025-01-05T06:55:11.569667Z","shell.execute_reply":"2025-01-05T06:55:14.425856Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with torch.no_grad():\n    model.to(device)\n    model.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-05T06:55:14.427822Z","iopub.execute_input":"2025-01-05T06:55:14.428117Z","iopub.status.idle":"2025-01-05T06:55:14.432724Z","shell.execute_reply.started":"2025-01-05T06:55:14.428091Z","shell.execute_reply":"2025-01-05T06:55:14.431917Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Perform inference with entropy thresholding\npredictions = []\nimage_ids_result = []\n\nwith torch.no_grad():\n    for images, ids in test_loader:\n        valid_images = []\n        valid_ids = []\n        \n        # Filter out None images\n        for img, img_id in zip(images, ids):\n            if img is not None and img_id is not None:\n                valid_images.append(img)\n                valid_ids.append(img_id)\n        \n        if not valid_images:  # Skip batch if no valid images\n            continue\n        \n        valid_images = torch.stack(valid_images).to(device)\n        logits = model(valid_images)\n        \n        # Convert logits to probabilities using softmax\n        probabilities = F.softmax(logits, dim=1).cpu().numpy()\n        print(probabilities)\n        # Compute entropy\n        entropies = -np.sum(probabilities * np.log(probabilities + 1e-9), axis=1)\n        print(entropies)\n        # Predict labels based on entropy\n        preds = np.argmax(probabilities, axis=1)\n        preds[entropies > entropy_threshold] = 5  # Assign \"Other\" class if entropy > threshold\n        \n        predictions.extend(preds)\n        image_ids_result.extend(valid_ids)\n\n        del valid_images, logits\n        gc.collect()\n        torch.cuda.empty_cache()\n\n# Map predictions to labels using reverse_label_mapping\nmapped_predictions = [reverse_label_mapping[pred] for pred in predictions]\n\n# Create submission DataFrame\nsubmission_df = pd.DataFrame({\n    \"image_id\": image_ids_result,\n    \"label\": mapped_predictions\n})\n\n# Save to CSV\nsubmission_df.to_csv(submission_path, index=False)\nprint(f\"Submission saved to {submission_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-05T06:55:14.433766Z","iopub.execute_input":"2025-01-05T06:55:14.434038Z","iopub.status.idle":"2025-01-05T06:55:25.7145Z","shell.execute_reply.started":"2025-01-05T06:55:14.434013Z","shell.execute_reply":"2025-01-05T06:55:25.713497Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}