{"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":"gpu","dataSources":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"},{"sourceId":7120688,"sourceType":"datasetVersion","datasetId":4107067},{"sourceId":226547,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":193210,"modelId":215152}],"dockerImageVersionId":30805,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import random\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom PIL import Image\nimport cv2\nimport timm\nfrom torchvision import transforms\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import balanced_accuracy_score\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch.optim.lr_scheduler as lr_scheduler\nfrom tqdm import tqdm\nfrom sklearn.metrics import balanced_accuracy_score, average_precision_score\nfrom sklearn.metrics import roc_curve, auc, precision_recall_curve\nimport matplotlib.pyplot as plt\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tile_size= 224\nbatch_size = 3\nn_tiles = 400\nnum_workers = 4","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def crop_wsi_with_otsu(img_pil, output_path=None):\n    \"\"\"\n    Load a large image (e.g., a WSI thumbnail or a TMA),\n    generate an Otsu mask in grayscale, find the bounding box\n    of the tissue region, and crop the entire image to that box.\n\n    input_path: full path to the .png (or .tif, .jpg, etc.) image\n    output_path: if provided, save the cropped image to this path\n                 if None, just return the cropped PIL Image object\n    \"\"\"\n    # # 1) Load image in RGB (PIL -> NumPy)\n    # img_pil = Image.open(input_path).convert(\"RGB\")\n    img_rgb = np.array(img_pil)  # shape: (H, W, 3)\n\n    # 2) Convert to grayscale for Otsu\n    gray = cv2.cvtColor(img_rgb, cv2.COLOR_RGB2GRAY)\n\n    # 3) Otsu threshold - background vs. tissue\n    \n    _, raw_mask = cv2.threshold(gray, 0, 255, cv2.THRESH_BINARY + cv2.THRESH_OTSU)\n    fraction_white = (raw_mask > 0).mean()  # fraction of pixels that are 255\n\n    if fraction_white > 0.5:\n        mask = 255 - raw_mask  # invert\n    else:\n        mask = raw_mask\n    \n    # Convert to binary {0,1}\n    mask_binary = (mask > 0).astype(np.uint8)\n    # 4) Find bounding box\n    #    We want the min/max row/col where mask is non-zero\n    coords = np.where(mask > 0)  # returns (row_indices, col_indices)\n    if len(coords[0]) == 0:\n        # No tissue found, fallback or skip\n        print(\"Warning: Otsu found no tissue; returning original image\")\n        cropped_img = img_pil\n    else:\n        y_min, y_max = coords[0].min(), coords[0].max()\n        x_min, x_max = coords[1].min(), coords[1].max()\n\n        # 5) Crop the original RGB\n        cropped_img = img_pil.crop((x_min, y_min, x_max+1, y_max+1))\n        # Note: +1 to include that pixel.\n\n    # 6) Save or return\n    # if output_path:\n    #     cropped_img.save(output_path)\n    #     print(f\"Cropped image saved to: {output_path}\")\n    return cropped_img\n\n# Example usage:\n# cropped_image = crop_wsi_with_otsu(\"slide1.png\", \"slide1_cropped.png\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_cropped_image(image, th_area=1000):\n    # Calculate the aspect ratio\n    as_ratio = image.size[0] / image.size[1]\n    \n    if as_ratio >= 1.5:\n        # Create a mask using maximum value condition\n        mask = np.max(np.array(image) > 0, axis=-1).astype(np.uint8)\n        \n        # Find connected components in the mask\n        retval, labels = cv2.connectedComponents(mask)\n        \n        for label in range(1, retval):\n            # Skip small components\n            area = np.sum(labels == label)\n            if area < th_area:\n                continue\n            \n            # Get coordinates of the first valid connected component\n            x, y = np.meshgrid(np.arange(image.size[0]), np.arange(image.size[1]))\n            xs, ys = x[labels == label], y[labels == label]\n            \n            # Calculate cropping boundaries\n            sx, ex = np.min(xs), np.max(xs)\n            cx = (sx + ex) // 2\n            crop_size = image.size[1]\n            sx = max(0, cx - crop_size // 2)\n            ex = min(sx + crop_size - 1, image.size[0] - 1)\n            sx = ex - crop_size + 1\n            sy, ey = 0, image.size[1] - 1\n            \n            # Crop the image and return\n            cropped_image = image.crop((sx, sy, ex + 1, ey + 1))\n            return cropped_image\n    else:\n        # If aspect ratio is less than 1.5, use the entire image\n        return image","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n########################################\n# Dataset\n########################################\nclass WSI_TMA_MILDataset(Dataset):\n    def __init__(self, df, data_dir, n_tiles, tile_size, transform=None):\n        \"\"\"\n        df: DataFrame with columns ['image_id', 'label', 'is_tma']\n        data_dir: directory containing 'train_images' and 'train_thumbnails'\n        n_tiles: number of tiles per image\n        tile_size: size of each tile\n        transform: torchvision transforms\n        \"\"\"\n        self.df = df.reset_index(drop=True)\n        self.data_dir = data_dir\n        self.n_tiles = n_tiles\n        self.tile_size = tile_size\n        self.transform = transform\n        # Adjust mapping as per your classes\n        self.label_mapping = { 'HGSC':0, 'EC':1, 'CC':2, 'LGSC':3, 'MC':4 }\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        image_id = row['image_id']\n        label_str = row['label']\n        label = self.label_mapping[label_str]\n        is_tma = row['is_tma']\n\n        if is_tma:\n            # TMA image\n            img_path = os.path.join(self.data_dir, 'train_images', f\"{image_id}.png\")\n        else:\n            # WSI thumbnail\n            img_path = os.path.join(self.data_dir, 'train_thumbnails', f\"{image_id}_thumbnail.png\")\n\n        img = cv2.imread(img_path, cv2.IMREAD_COLOR)\n        if img is None:\n            raise FileNotFoundError(f\"Image not found: {img_path}\")\n        img_pil = Image.fromarray(img)\n        cropped1 = get_cropped_image(img_pil)\n        cropped2 = crop_wsi_with_otsu(cropped1)\n        img_rgb = cv2.cvtColor(np.array(cropped2), 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((8192, 8192))\n        img_resized = np.array(img_resized)\n\n        # Determine the size of each tile (assume 20x20 grid for 400 tiles)\n        tile_size = 8192 // 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        return tiles_tensor, torch.tensor(label, dtype=torch.long)","metadata":{"trusted":true},"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,)\n\n\n","metadata":{"trusted":true},"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, 160)\n        self.attention_B = nn.Linear(160, 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        \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        # Final classification\n        logits = self.classifier(weighted_sum)  # [B, num_classes]\n        return logits","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"########################################\n# Setup\n########################################\ndata_dir = \"/kaggle/input/UBC-OCEAN/\"\ndf = pd.read_csv(os.path.join(data_dir, \"train.csv\"))\n\n# Example label mapping already defined in dataset class\ny = df['label']\n\ntrain_transform = transforms.Compose([\n    transforms.ColorJitter(brightness=.2,contrast=.2,saturation=.2,hue=.2),\n    transforms.Resize((224,224)),\n    # transforms.RandomHorizontalFlip(),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225])\n])\n\nval_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])\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(device)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Split WSI data\nwsi_df = df[df['is_tma'] == 0]\ntrain_wsi, val_wsi = train_test_split(wsi_df, test_size=0.3, stratify=wsi_df['label'], random_state=42)\n\n# Split TMA data\ntma_df = df[df['is_tma'] == 1]\ntrain_tma, val_tma = train_test_split(tma_df, test_size=0.6, stratify=tma_df['label'], random_state=42)\n\n# Combine splits\ntrain_df = pd.concat([train_wsi, train_tma]).reset_index(drop=True)\nval_df = pd.concat([val_wsi, val_tma]).reset_index(drop=True)\n\n# Load data\ntrain_dataset = WSI_TMA_MILDataset(train_df, data_dir=data_dir, n_tiles=n_tiles, tile_size=tile_size, transform=train_transform)\nval_dataset = WSI_TMA_MILDataset(val_df, data_dir=data_dir, n_tiles=n_tiles, tile_size=tile_size, transform=val_transform)\n\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers)\nval_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=num_workers)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"###################################\n# Train and validate\n###################################\nmodel = MILAttentionDoubleDINO(num_classes=5, device=device).to(device)\ncriterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(model.parameters(), 1e-4)#, weight_decay=5e-4)\n# scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=20, eta_min=5e-5)\nresume_best_model_path = \"/kaggle/input/resume-after-epoch-5/pytorch/default/1/best_model_224_400tiles_tuned (2).pth\"\n\nmodel.load_state_dict(torch.load(resume_best_model_path))\n\nbest_balanced_acc = 0.0\nbest_auprc = 0.0\npatience = 4\nepochs_no_improve = 0\nbest_model_path = \"best_model_224_400tiles_tuned.pth\"\n\nfor epoch in range(6):  # Up to 20 epochs, adjust as needed\n    print(f\"\\nEpoch {epoch+1}/{6}\")\n    # print(f\"Current Learning Rate: {scheduler.get_last_lr()[0]:.6f}\")\n    \n    # Train\n    model.train()\n    running_loss = 0.0\n    for tiles, labels in train_loader:\n        tiles, labels = tiles.to(device), labels.to(device)\n        optimizer.zero_grad()\n        logits = model(tiles)\n        loss = criterion(logits, labels)\n        loss.backward()\n        optimizer.step()\n        running_loss += loss.item() * tiles.size(0)\n\n    train_loss = running_loss / len(train_loader.dataset)\n\n    # Validation\n    model.eval()\n    running_loss_val = 0.0\n    all_preds = []\n    all_labels = []\n    all_probs = []  # To store probabilities for AUPRC calculation\n    with torch.no_grad():\n        for tiles, labels in val_loader:\n            tiles, labels = tiles.to(device), labels.to(device)\n            logits = model(tiles)\n            loss = criterion(logits, labels)\n            running_loss_val += loss.item() * tiles.size(0)\n            \n            probs = torch.softmax(logits, dim=1).cpu().numpy()  # Get probabilities\n            preds = torch.argmax(logits, dim=1).cpu().numpy()\n            \n            all_probs.append(probs)\n            all_preds.append(preds)\n            all_labels.append(labels.cpu().numpy())\n            \n\n    val_loss = running_loss_val / len(val_loader.dataset)\n    all_preds = np.concatenate(all_preds)\n    all_labels = np.concatenate(all_labels)\n    all_probs = np.concatenate(all_probs)  # Shape: (num_samples, num_classes)\n\n    # Calculate metrics\n    balanced_acc = balanced_accuracy_score(all_labels, all_preds)\n    # Calculate AUPRC (macro-average across all classes)\n    auprc = average_precision_score(\n        np.eye(len(all_probs[0]))[all_labels],  # One-hot encode labels\n        all_probs,\n        average=\"macro\"\n    )\n\n    print(f\"Epoch {epoch+1} - Train Loss: {train_loss:.4f}, Val Loss: {val_loss:.4f}, \"\n          f\"Val Balanced Acc: {balanced_acc:.4f}, Val AUPRC: {auprc:.4f}\")\n\n    # Check for improvement\n    if balanced_acc > best_balanced_acc or auprc > best_auprc:\n        if balanced_acc > best_balanced_acc:\n            best_balanced_acc = balanced_acc\n        if auprc > best_auprc:\n            best_auprc = auprc\n        epochs_no_improve = 0\n        torch.save(model.state_dict(), best_model_path)\n        print(\"  * Best model saved.\")\n    else:\n        epochs_no_improve += 1\n        if epochs_no_improve >= patience:\n            print(\"Early stopping.\")\n            break\n\n    # Step the scheduler\n    # scheduler.step()\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load best model and evaluate\nmodel.load_state_dict(torch.load(best_model_path))\n\n# Evaluate final performance on validation set\nmodel.eval()\nall_preds = []\nall_labels = []\nall_probs = []\nwith torch.no_grad():\n    for tiles, labels in val_loader:\n        tiles, labels = tiles.to(device), labels.to(device)\n        logits = model(tiles)\n        probs = torch.softmax(logits, dim=1).cpu().numpy()\n        preds = torch.argmax(logits, dim=1).cpu().numpy()\n        all_probs.append(probs)\n        all_preds.append(preds)\n        all_labels.append(labels.cpu().numpy())\nall_preds = np.concatenate(all_preds)\nall_labels = np.concatenate(all_labels)\nall_probs = np.concatenate(all_probs)\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Final metrics\nfinal_balanced_acc = balanced_accuracy_score(all_labels, all_preds)\nfinal_auprc = average_precision_score(\n    np.eye(len(all_probs[0]))[all_labels],\n    all_probs,\n    average=\"macro\"\n)\nprint(f\"Final Balanced Accuracy: {final_balanced_acc:.4f}\")\nprint(f\"Final AUPRC: {final_auprc:.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"np.save('/kaggle/working/all_labels.npy', all_labels)\nnp.save('/kaggle/working/all_probs.npy', all_probs)\n\nprint(\"Files saved successfully: all_labels.npy, all_probs.npy\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}