{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":18647,"databundleVersionId":1126921}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\n# # It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# # For example, here's several helpful packages to load\n# import numpy as np # linear algebra\n# import pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# # Input data files are available in the read-only \"../input/\" directory\n# # For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\n# import os\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n\n# # You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# # You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-08-08T04:53:00.037463Z","iopub.execute_input":"2026-08-08T04:53:00.03837Z","iopub.status.idle":"2026-08-08T04:53:34.842967Z","shell.execute_reply.started":"2026-08-08T04:53:00.038333Z","shell.execute_reply":"2026-08-08T04:53:34.841598Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # ── CELL 1: Imports ──────────────────────────────────────────────────────────\n# import os, numpy as np, pandas as pd\n# from PIL import Image\n# import cv2, torch, torch.nn as nn, torch.nn.functional as F\n# from torch.utils.data import Dataset, DataLoader\n# from sklearn.model_selection import train_test_split\n# from sklearn.metrics import cohen_kappa_score\n# import timm\n# from tqdm import tqdm\n\n# Image.MAX_IMAGE_PIXELS = None\n# device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n# print(\"Device:\", device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T16:57:36.850154Z","iopub.execute_input":"2026-08-08T16:57:36.850548Z","iopub.status.idle":"2026-08-08T16:57:58.133811Z","shell.execute_reply.started":"2026-08-08T16:57:36.850503Z","shell.execute_reply":"2026-08-08T16:57:58.133047Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # ── CELL 2: Config (easy to tune) ────────────────────────────────────────────\n# PATCH_SIZE   = 224   # no resize needed in forward pass\n# NUM_PATCHES  = 32    # reduced from 64 to save memory\n# NUM_SLIDES   = 1000  # keep 1000 for speed\n# BATCH_SIZE   = 16    # safer than 32 on T4/P100\n# EPOCHS       = 10\n# LR           = 3e-5\n# SAVE_DIR     = \"/kaggle/working/patches\"\n\n# MEAN = [0.485, 0.456, 0.406]\n# STD  = [0.229, 0.224, 0.225]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T13:05:26.024423Z","iopub.execute_input":"2026-08-08T13:05:26.025263Z","iopub.status.idle":"2026-08-08T13:05:26.029832Z","shell.execute_reply.started":"2026-08-08T13:05:26.025224Z","shell.execute_reply":"2026-08-08T13:05:26.029054Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # ── CELL 3: Load & split data ─────────────────────────────────────────────────\n# train_csv   = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train.csv\"\n# images_path = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train_images\"\n\n# df = pd.read_csv(train_csv)\n# train_df, val_df = train_test_split(\n#     df, test_size=0.2, stratify=df[\"isup_grade\"], random_state=42\n# )\n# train_df = train_df.sample(NUM_SLIDES, random_state=42).reset_index(drop=True)\n# val_df   = val_df.sample(200, random_state=42).reset_index(drop=True)\n# print(\"Train slides:\", len(train_df), \"| Val slides:\", len(val_df))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T13:05:30.532537Z","iopub.execute_input":"2026-08-08T13:05:30.532952Z","iopub.status.idle":"2026-08-08T13:05:30.579374Z","shell.execute_reply.started":"2026-08-08T13:05:30.532919Z","shell.execute_reply":"2026-08-08T13:05:30.578712Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # ── CELL 4: Tissue-aware patch extraction ────────────────────────────────────\n# def is_tissue(patch, threshold=0.7):\n#     \"\"\"Return True if patch has enough tissue (not white background).\"\"\"\n#     gray = cv2.cvtColor(patch, cv2.COLOR_RGB2GRAY)\n#     return (gray < 230).mean() > threshold   # >70% non-white pixels\n\n# def extract_patches(df, save_dir, num_patches=NUM_PATCHES):\n#     os.makedirs(save_dir, exist_ok=True)\n#     for _, row in tqdm(df.iterrows(), total=len(df)):\n#         img_id = row[\"image_id\"]\n#         label  = row[\"isup_grade\"]\n#         path   = os.path.join(images_path, img_id + \".tiff\")\n#         try:\n#             slide = np.array(Image.open(path).reduce(8))\n#         except Exception:\n#             continue\n#         h, w = slide.shape[:2]  \n#         if h <= PATCH_SIZE or w <= PATCH_SIZE:     \n#             continue                \n#         saved = 0 \n#         attempts = 0\n#         while saved < num_patches and attempts < num_patches * 10:\n#             y = np.random.randint(0, max(1, h - PATCH_SIZE))\n#             x = np.random.randint(0, max(1, w - PATCH_SIZE))\n#             patch = slide[y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n#             if is_tissue(patch):\n#                 cv2.imwrite(f\"{save_dir}/{img_id}_{saved}_{label}.png\",\n#                             cv2.cvtColor(patch, cv2.COLOR_RGB2BGR))\n#                 saved += 1\n#             attempts += 1\n\n# extract_patches(train_df, SAVE_DIR + \"/train\")\n# extract_patches(val_df,   SAVE_DIR + \"/val\", num_patches=16)\n# print(\"Train patches:\", len(os.listdir(SAVE_DIR+\"/train\")))\n# print(\"Val patches  :\", len(os.listdir(SAVE_DIR+\"/val\")))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T13:05:42.919062Z","iopub.execute_input":"2026-08-08T13:05:42.919448Z","iopub.status.idle":"2026-08-08T14:21:36.260841Z","shell.execute_reply.started":"2026-08-08T13:05:42.919419Z","shell.execute_reply":"2026-08-08T14:21:36.260137Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # ── CELL 5: Dataset with proper normalization ─────────────────────────────────\n# import torchvision.transforms as T\n\n# train_transforms = T.Compose([\n#     T.RandomHorizontalFlip(),\n#     T.RandomVerticalFlip(),\n#     T.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.1),\n#     T.ToTensor(),\n#     T.Normalize(mean=MEAN, std=STD),\n# ])\n\n# val_transforms = T.Compose([\n#     T.ToTensor(),\n#     T.Normalize(mean=MEAN, std=STD),\n# ])\n\n# class PatchDataset(Dataset):\n#     def __init__(self, patch_dir, transform=None):\n#         self.files     = [f for f in os.listdir(patch_dir) if f.endswith(\".png\")]\n#         self.patch_dir = patch_dir\n#         self.transform = transform\n\n#     def __len__(self): return len(self.files)\n\n#     def __getitem__(self, idx):\n#         f     = self.files[idx]\n#         img   = Image.open(os.path.join(self.patch_dir, f)).convert(\"RGB\")\n#         label = int(f.split(\"_\")[-1].replace(\".png\", \"\"))\n#         if self.transform:\n#             img = self.transform(img)\n#         return img, label\n\n# train_ds = PatchDataset(SAVE_DIR+\"/train\", train_transforms)\n# val_ds   = PatchDataset(SAVE_DIR+\"/val\",   val_transforms)\n\n# train_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True,\n#                           num_workers=2, pin_memory=True)\n# val_loader   = DataLoader(val_ds,   batch_size=BATCH_SIZE, shuffle=False,\n#                           num_workers=2, pin_memory=True)\n# print(\"Train batches:\", len(train_loader), \"| Val batches:\", len(val_loader))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T14:21:55.289857Z","iopub.execute_input":"2026-08-08T14:21:55.29027Z","iopub.status.idle":"2026-08-08T14:21:55.305558Z","shell.execute_reply.started":"2026-08-08T14:21:55.290238Z","shell.execute_reply":"2026-08-08T14:21:55.304707Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # ── Attention Pooling Module (the hierarchy magic) ─────────────────────────\n# class AttentionPool(nn.Module):\n#     \"\"\"\n#     Takes N patch features from one slide,\n#     learns which patches matter most, and\n#     produces a single slide-level vector.\n#     Inspired by HIPT's region-level aggregation.\n#     \"\"\"\n#     def __init__(self, feat_dim=768, hidden_dim=256):\n#         super().__init__()\n#         self.attention = nn.Sequential(\n#             nn.Linear(feat_dim, hidden_dim),\n#             nn.Tanh(),\n#             nn.Linear(hidden_dim, 1)\n#         )\n\n#     def forward(self, x):\n#         # x: (N_patches, feat_dim)\n#         attn_weights = self.attention(x)          # (N, 1)\n#         attn_weights = torch.softmax(attn_weights, dim=0)  # normalize\n#         slide_feat   = (attn_weights * x).sum(dim=0)       # weighted sum\n#         return slide_feat, attn_weights","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T14:21:58.982466Z","iopub.execute_input":"2026-08-08T14:21:58.983204Z","iopub.status.idle":"2026-08-08T14:21:58.988208Z","shell.execute_reply.started":"2026-08-08T14:21:58.983171Z","shell.execute_reply":"2026-08-08T14:21:58.987563Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # ── Hierarchical Swin Model ────────────────────────────────────────────────\n# class HierarchicalSwinPANDA(nn.Module):\n#     \"\"\"\n#     Level 1: Swin-Tiny encodes each patch → patch feature (768-d)\n#     Level 2: Attention pooling aggregates patches → slide feature\n#     Level 3: Classifier predicts ISUP grade\n#     \"\"\"\n#     def __init__(self, num_classes=6):\n#         super().__init__()\n#         self.patch_encoder = timm.create_model(\n#             \"swin_tiny_patch4_window7_224\",\n#             pretrained=True,\n#             num_classes=0\n#         )\n#         self.attn_pool  = AttentionPool(feat_dim=768, hidden_dim=256)\n#         self.drop       = nn.Dropout(0.3)\n#         self.classifier = nn.Linear(768, num_classes)\n\n#     def forward(self, patches):\n#         # patches: (N_patches, 3, 224, 224) — all patches from ONE slide\n#         patch_feats = self.patch_encoder(patches)     # (N, 768)\n#         slide_feat, attn_w = self.attn_pool(patch_feats)  # (768,)\n#         logits = self.classifier(self.drop(slide_feat))    # (6,)\n#         return logits, attn_w","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T14:22:04.097195Z","iopub.execute_input":"2026-08-08T14:22:04.097602Z","iopub.status.idle":"2026-08-08T14:22:04.103561Z","shell.execute_reply.started":"2026-08-08T14:22:04.097571Z","shell.execute_reply":"2026-08-08T14:22:04.102586Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # ── Slide-level Dataset (key difference from before) ──────────────────────\n# class SlideDataset(Dataset):\n#     \"\"\"\n#     Each item = ALL patches from ONE slide (not individual patches).\n#     This is what makes it hierarchical.\n#     \"\"\"\n#     def __init__(self, patch_dir, transform=None):\n#         self.transform  = transform\n#         self.slide_data = {}  # {image_id: {\"patches\": [...], \"label\": int}}\n\n#         for f in os.listdir(patch_dir):\n#             if not f.endswith(\".png\"):\n#                 continue\n#             parts    = f.replace(\".png\",\"\").split(\"_\")\n#             image_id = \"_\".join(parts[:-2])   # handle underscores in ID\n#             label    = int(parts[-1])\n\n#             if image_id not in self.slide_data:\n#                 self.slide_data[image_id] = {\"patches\": [], \"label\": label}\n#             self.slide_data[image_id][\"patches\"].append(\n#                 os.path.join(patch_dir, f)\n#             )\n\n#         self.slides = list(self.slide_data.keys())\n\n#     def __len__(self):\n#         return len(self.slides)\n\n#     def __getitem__(self, idx):\n#         sid   = self.slides[idx]\n#         info  = self.slide_data[sid]\n#         label = info[\"label\"]\n\n#         patch_tensors = []\n#         for p in info[\"patches\"]:\n#             img = Image.open(p).convert(\"RGB\")\n#             if self.transform:\n#                 img = self.transform(img)\n#             patch_tensors.append(img)\n\n#         # Stack all patches: (N_patches, 3, 224, 224)\n#         patches = torch.stack(patch_tensors, dim=0)\n#         return patches, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T14:22:08.477167Z","iopub.execute_input":"2026-08-08T14:22:08.477673Z","iopub.status.idle":"2026-08-08T14:22:08.485255Z","shell.execute_reply.started":"2026-08-08T14:22:08.47763Z","shell.execute_reply":"2026-08-08T14:22:08.484666Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # ── Slide-level DataLoader ─────────────────────────────────────────────────\n# # batch_size=1 because each slide has variable number of patches\n# train_slide_ds = SlideDataset(SAVE_DIR+\"/train\", train_transforms)\n# val_slide_ds   = SlideDataset(SAVE_DIR+\"/val\",   val_transforms)\n\n# train_slide_loader = DataLoader(train_slide_ds, batch_size=1,\n#                                 shuffle=True, num_workers=2)\n# val_slide_loader   = DataLoader(val_slide_ds,   batch_size=1,\n#                                 shuffle=False, num_workers=2)\n\n# print(\"Train slides:\", len(train_slide_ds))\n# print(\"Val slides  :\", len(val_slide_ds))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T14:22:14.338237Z","iopub.execute_input":"2026-08-08T14:22:14.33898Z","iopub.status.idle":"2026-08-08T14:22:14.364288Z","shell.execute_reply.started":"2026-08-08T14:22:14.338946Z","shell.execute_reply":"2026-08-08T14:22:14.363693Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # ── Training loop (slide-level) ───────────────────────────────────────────\n# model     = HierarchicalSwinPANDA().to(device)\n# criterion = nn.CrossEntropyLoss()\n# optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=1e-4)\n# scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)\n\n# def run_epoch(loader, train=True):\n#     model.train() if train else model.eval()\n#     total_loss, all_preds, all_labels = 0, [], []\n#     ctx = torch.enable_grad() if train else torch.no_grad()\n#     with ctx:\n#         for patches, labels in loader:\n#             # patches: (1, N_patches, 3, 224, 224)\n#             patches = patches.squeeze(0).to(device)  # (N, 3, 224, 224)\n#             labels  = labels.to(device)\n\n#             if train:\n#                 optimizer.zero_grad()\n\n#             logits, _ = model(patches)\n#             logits    = logits.unsqueeze(0)          # (1, 6) for loss\n#             loss      = criterion(logits, labels)\n\n#             if train:\n#                 loss.backward()\n#                 optimizer.step()\n\n#             total_loss += loss.item()\n#             all_preds  += torch.argmax(logits, 1).cpu().tolist()\n#             all_labels += labels.cpu().tolist()\n\n#     kappa = cohen_kappa_score(all_labels, all_preds,\n#                               weights=\"quadratic\", labels=list(range(6)))\n#     return total_loss / len(loader), kappa       \n\n# best_kappa = -1\n# for epoch in range(EPOCHS):\n#     tr_loss, tr_kappa = run_epoch(train_slide_loader, train=True)\n#     vl_loss, vl_kappa = run_epoch(val_slide_loader,   train=False)\n#     scheduler.step()\n#     print(f\"Epoch {epoch+1}/{EPOCHS} | \"\n#           f\"Train Loss {tr_loss:.4f} QWK {tr_kappa:.4f} | \"\n#           f\"Val Loss {vl_loss:.4f} QWK {vl_kappa:.4f}\")\n#     if vl_kappa > best_kappa:\n#         best_kappa = vl_kappa\n#         torch.save(model.state_dict(), \"/kaggle/working/hier_swin_best.pth\")\n#         print(\"  ✅ Best model saved\")\n\n# print(\"\\nBest Val QWK:\", best_kappa)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T14:22:18.035536Z","iopub.execute_input":"2026-08-08T14:22:18.0361Z","iopub.status.idle":"2026-08-08T14:41:59.130494Z","shell.execute_reply.started":"2026-08-08T14:22:18.036066Z","shell.execute_reply":"2026-08-08T14:41:59.129505Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # ── Training loop (slide-level) ───────────────────────────────────────────\n# model     = HierarchicalSwinPANDA().to(device)\n# criterion = nn.CrossEntropyLoss()\n# optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=1e-4)\n# scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)\n\n# def run_epoch(loader, train=True):\n#     model.train() if train else model.eval()\n#     total_loss, all_preds, all_labels = 0, [], []\n#     ctx = torch.enable_grad() if train else torch.no_grad()\n#     with ctx:\n#         for patches, labels in loader:\n#             # patches: (1, N_patches, 3, 224, 224)\n#             patches = patches.squeeze(0).to(device)  # (N, 3, 224, 224)\n#             labels  = labels.to(device)\n\n#             if train:\n#                 optimizer.zero_grad()\n\n#             logits, _ = model(patches)\n#             logits    = logits.unsqueeze(0)          # (1, 6) for loss\n#             loss      = criterion(logits, labels)\n\n#             if train:\n#                 loss.backward()\n#                 optimizer.step()\n\n#             total_loss += loss.item()\n#             all_preds  += torch.argmax(logits, 1).cpu().tolist()\n#             all_labels += labels.cpu().tolist()\n\n#     kappa = cohen_kappa_score(all_labels, all_preds,\n#                               weights=\"quadratic\", labels=list(range(6)))\n#     return total_loss / len(loader), kappa\n\n# best_kappa = -1\n# for epoch in range(EPOCHS):\n#     tr_loss, tr_kappa = run_epoch(train_slide_loader, train=True)\n#     vl_loss, vl_kappa = run_epoch(val_slide_loader,   train=False)\n#     scheduler.step()\n#     print(f\"Epoch {epoch+1}/{EPOCHS} | \"\n#           f\"Train Loss {tr_loss:.4f} QWK {tr_kappa:.4f} | \"\n#           f\"Val Loss {vl_loss:.4f} QWK {vl_kappa:.4f}\")\n#     if vl_kappa > best_kappa:\n#         best_kappa = vl_kappa\n#         torch.save(model.state_dict(), \"/kaggle/working/hier_swin_best.pth\")\n#         print(\"  ✅ Best model saved\")\n\n# print(\"\\nBest Val QWK:\", best_kappa)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-22T20:28:03.629937Z","iopub.execute_input":"2026-04-22T20:28:03.630681Z","iopub.status.idle":"2026-04-22T20:49:55.051091Z","shell.execute_reply.started":"2026-04-22T20:28:03.63065Z","shell.execute_reply":"2026-04-22T20:49:55.049929Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── CELL 8: Save & wrap up ────────────────────────────────────────────────────\nprint(\"Best Validation QWK:\", best_kappa)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-22T20:49:55.059723Z","iopub.execute_input":"2026-04-22T20:49:55.060135Z","iopub.status.idle":"2026-04-22T20:49:55.0875Z","shell.execute_reply.started":"2026-04-22T20:49:55.060071Z","shell.execute_reply":"2026-04-22T20:49:55.086824Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import matplotlib.pyplot as plt\n\n# train_qwk = [0.2305, 0.5275, 0.6528, 0.7010, 0.7358, \n#              0.7921, 0.8530, 0.8922, 0.9126, 0.9290]\n# val_qwk   = [0.5454, 0.6098, 0.6802, 0.6842, 0.7025,\n#              0.7409, 0.7562, 0.7319, 0.7385, 0.7334]\n\n# plt.figure(figsize=(8,5))\n# plt.plot(range(1,11), train_qwk, 'b-o', label='Train QWK')\n# plt.plot(range(1,11), val_qwk,   'r-o', label='Val QWK')\n# plt.axhline(y=0.756, color='green', linestyle='--', label='Best Val QWK = 0.756')\n# plt.xlabel('Epoch')\n# plt.ylabel('Quadratic Weighted Kappa')\n# plt.title('Hierarchical Swin-Tiny on PANDA (1000 slides)')\n# plt.legend()\n# plt.grid(True)\n# plt.tight_layout()\n# plt.savefig('/kaggle/working/training_curve.png', dpi=150)\n# plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T14:42:45.679405Z","iopub.execute_input":"2026-08-08T14:42:45.680222Z","iopub.status.idle":"2026-08-08T14:42:46.068141Z","shell.execute_reply.started":"2026-08-08T14:42:45.680178Z","shell.execute_reply":"2026-08-08T14:42:46.067321Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import matplotlib.pyplot as plt\n# import seaborn as sns\n# from sklearn.metrics import confusion_matrix\n# import numpy as np\n\n# # ── Re-run validation to collect predictions ──────────────────────────────\n# model.load_state_dict(torch.load(\"/kaggle/working/hier_swin_best.pth\"))\n# model.eval()\n\n# all_preds, all_labels = [], []\n\n# with torch.no_grad():\n#     for patches, labels in val_slide_loader:f\n#         patches = patches.squeeze(0).to(device)\n#         logits, _ = model(patches)\n#         logits = logits.unsqueeze(0)\n#         all_preds  += torch.argmax(logits, 1).cpu().tolist()\n#         all_labels += labels.tolist()\n\n# # ── Plot confusion matrix ─────────────────────────────────────────────────\n# cm = confusion_matrix(all_labels, all_preds, labels=list(range(6)))\n\n# plt.figure(figsize=(8, 6))\n# sns.heatmap(\n#     cm,\n#     annot=True,\n#     fmt='d',\n#     cmap='Blues',\n#     xticklabels=[f'Pred {i}' for i in range(6)],\n#     yticklabels=[f'True {i}' for i in range(6)]\n# )\n# plt.title('Confusion Matrix — Hierarchical Swin-Tiny on PANDA\\n(Val QWK = 0.756)')\n# plt.ylabel('True ISUP Grade')\n# plt.xlabel('Predicted ISUP Grade')\n# plt.tight_layout()\n# plt.savefig('/kaggle/working/confusion_matrix.png', dpi=150)\n# plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T14:42:53.304202Z","iopub.execute_input":"2026-08-08T14:42:53.304764Z","iopub.status.idle":"2026-08-08T14:42:59.082162Z","shell.execute_reply.started":"2026-08-08T14:42:53.304732Z","shell.execute_reply":"2026-08-08T14:42:59.081331Z"}},"outputs":[],"execution_count":null},{"cell_type":"raw","source":"","metadata":{}},{"cell_type":"code","source":"# # ==============================================================================\n# # TWO-LEVEL HIERARCHICAL SWIN MODEL — PANDA ISUP GRADING\n# # Patch -> Region (spatial) -> Slide\n# # Built on top of your baseline-swin-model notebook. Run cell-by-cell on Kaggle.\n# # ==============================================================================\n\n# import os, numpy as np, pandas as pd  \n# from PIL import Image\n# import cv2, torch, torch.nn as nn, torch.nn.functional as F\n# from torch.utils.data import Dataset, DataLoader\n# from sklearn.model_selection import train_test_split\n# from sklearn.metrics import cohen_kappa_score\n# import timm\n# from tqdm import tqdm\n# import torchvision.transforms as T\n\n# Image.MAX_IMAGE_PIXELS = None\n# device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n# print(\"Device:\", device)\n\n# # ── CELL 1: Config ───────────────────────────────────────────────────────────\n# PATCH_SIZE     = 224\n# GRID_SIZE      = 4          # slide divided into GRID_SIZE x GRID_SIZE regions (e.g. 4x4 = 16 regions)\n# PATCHES_PER_REGION = 4      # patches sampled per region (tissue-permitting)\n# NUM_SLIDES     = 1000       # keep consistent with your baseline for comparison\n# VAL_SLIDES     = 200\n# TEST_SLIDES    = 200\n# BATCH_SIZE     = 1          # 1 slide per batch (variable region/patch counts)\n# EPOCHS         = 15\n# LR             = 3e-5\n# WEIGHT_DECAY   = 5e-4       # increased from 1e-4 to fight overfitting\n# PATIENCE       = 5          # early stopping patience (epochs without val improvement)\n# SAVE_DIR       = \"/kaggle/working/regions\"\n# MEAN = [0.485, 0.456, 0.406]\n# STD  = [0.229, 0.224, 0.225]\n\n# train_csv   = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train.csv\"\n# images_path = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train_images\"\n\n# # ── CELL 2: Three-way split (train/val/test — your baseline only had train/val) ──\n# df = pd.read_csv(train_csv)\n\n# train_df, temp_df = train_test_split(\n#     df, test_size=0.3, stratify=df[\"isup_grade\"], random_state=42\n# )\n# val_df, test_df = train_test_split(\n#     temp_df, test_size=0.5, stratify=temp_df[\"isup_grade\"], random_state=42\n# )\n\n# train_df = train_df.sample(min(NUM_SLIDES, len(train_df)), random_state=42).reset_index(drop=True)\n# val_df   = val_df.sample(min(VAL_SLIDES, len(val_df)), random_state=42).reset_index(drop=True)\n# test_df  = test_df.sample(min(TEST_SLIDES, len(test_df)), random_state=42).reset_index(drop=True)\n\n# print(f\"Train slides: {len(train_df)} | Val slides: {len(val_df)} | Test slides: {len(test_df)}\")\n\n# # ── CELL 3: Tissue-aware SPATIAL region + patch extraction ──────────────────\n# def is_tissue(patch, threshold=0.6):\n#     gray = cv2.cvtColor(patch, cv2.COLOR_RGB2GRAY)\n#     return (gray < 230).mean() > threshold\n\n# def extract_spatial_patches(df, save_dir, grid_size=GRID_SIZE, patches_per_region=PATCHES_PER_REGION):\n#     \"\"\"\n#     Divides each slide into a grid_size x grid_size grid of REGIONS.\n#     Within each region, samples up to patches_per_region tissue patches.\n#     Filename encodes: imgid_region{r}_patch{p}_label{l}.png\n#     Regions with no tissue are simply skipped (not padded) — the model\n#     handles variable region/patch counts.\n#     \"\"\"\n#     os.makedirs(save_dir, exist_ok=True)\n#     manifest = []\n#     for _, row in tqdm(df.iterrows(), total=len(df)):\n#         img_id = row[\"image_id\"]\n#         label  = row[\"isup_grade\"]\n#         path   = os.path.join(images_path, img_id + \".tiff\")\n#         try:\n#             slide = np.array(Image.open(path).reduce(8))\n#         except Exception:\n#             continue\n#         h, w = slide.shape[:2]\n#         if h <= PATCH_SIZE * 2 or w <= PATCH_SIZE * 2:\n#             continue\n\n#         region_h = h // grid_size\n#         region_w = w // grid_size\n\n#         for r_idx in range(grid_size * grid_size):\n#             ry, rx = divmod(r_idx, grid_size)\n#             y0, y1 = ry * region_h, min((ry + 1) * region_h, h)\n#             x0, x1 = rx * region_w, min((rx + 1) * region_w, w)\n#             if (y1 - y0) <= PATCH_SIZE or (x1 - x0) <= PATCH_SIZE:\n#                 continue\n\n#             region_slide = slide[y0:y1, x0:x1]\n#             saved, attempts = 0, 0\n#             while saved < patches_per_region and attempts < patches_per_region * 8:\n#                 yy = np.random.randint(0, max(1, (y1 - y0) - PATCH_SIZE))\n#                 xx = np.random.randint(0, max(1, (x1 - x0) - PATCH_SIZE))\n#                 patch = region_slide[yy:yy+PATCH_SIZE, xx:xx+PATCH_SIZE]\n#                 if patch.shape[0] == PATCH_SIZE and patch.shape[1] == PATCH_SIZE and is_tissue(patch):\n#                     fname = f\"{img_id}_r{r_idx}_p{saved}_{label}.png\"\n#                     cv2.imwrite(os.path.join(save_dir, fname),\n#                                 cv2.cvtColor(patch, cv2.COLOR_RGB2BGR))\n#                     saved += 1\n#                 attempts += 1\n\n# extract_spatial_patches(train_df, SAVE_DIR + \"/train\")\n# extract_spatial_patches(val_df,   SAVE_DIR + \"/val\")\n# extract_spatial_patches(test_df,  SAVE_DIR + \"/test\")\n\n# print(\"Train patches:\", len(os.listdir(SAVE_DIR + \"/train\")))\n# print(\"Val patches  :\", len(os.listdir(SAVE_DIR + \"/val\")))\n# print(\"Test patches :\", len(os.listdir(SAVE_DIR + \"/test\")))\n\n# # ── CELL 4: Transforms ───────────────────────────────────────────────────────\n# train_transforms = T.Compose([\n#     T.RandomHorizontalFlip(),\n#     T.RandomVerticalFlip(),\n#     T.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.1),\n#     T.ToTensor(),\n#     T.Normalize(mean=MEAN, std=STD),\n# ])\n# eval_transforms = T.Compose([\n#     T.ToTensor(),\n#     T.Normalize(mean=MEAN, std=STD),\n# ])\n\n# # ── CELL 5: Slide dataset with REGION structure ──────────────────────────────\n# class HierSlideDataset(Dataset):\n#     \"\"\"\n#     Each item = one slide, structured as a list of regions,\n#     each region being a list of patch file paths.\n#     Returns: (list[Tensor(N_r, 3, 224, 224)] per region, label)\n#     \"\"\"\n#     def __init__(self, patch_dir, transform=None):\n#         self.transform = transform\n#         # slide_id -> region_id -> [patch paths]\n#         self.slide_data = {}\n#         for f in os.listdir(patch_dir):\n#             if not f.endswith(\".png\"):\n#                 continue\n#             stem = f.replace(\".png\", \"\")\n#             # format: imgid_r{r}_p{p}_{label}\n#             parts = stem.split(\"_\")\n#             label = int(parts[-1])\n#             p_tag = parts[-2]      # p{p}\n#             r_tag = parts[-3]      # r{r}\n#             img_id = \"_\".join(parts[:-3])\n#             region_id = int(r_tag[1:])\n\n#             self.slide_data.setdefault(img_id, {\"label\": label, \"regions\": {}})\n#             self.slide_data[img_id][\"regions\"].setdefault(region_id, [])\n#             self.slide_data[img_id][\"regions\"][region_id].append(os.path.join(patch_dir, f))\n\n#         self.slides = list(self.slide_data.keys())\n\n#     def __len__(self):\n#         return len(self.slides)\n\n#     def __getitem__(self, idx):\n#         sid = self.slides[idx]\n#         info = self.slide_data[sid]\n#         label = info[\"label\"]\n\n#         region_tensors = []\n#         for region_id, paths in info[\"regions\"].items():\n#             patch_tensors = []\n#             for p in paths:\n#                 img = Image.open(p).convert(\"RGB\")\n#                 if self.transform:\n#                     img = self.transform(img)\n#                 patch_tensors.append(img)\n#             region_tensors.append(torch.stack(patch_tensors, dim=0))  # (N_r, 3, 224, 224)\n\n#         return region_tensors, label\n\n# def collate_single(batch):\n#     # batch_size=1 always, just unwrap\n#     region_tensors, label = batch[0]\n#     return region_tensors, label\n\n# train_ds = HierSlideDataset(SAVE_DIR + \"/train\", train_transforms)\n# val_ds   = HierSlideDataset(SAVE_DIR + \"/val\",   eval_transforms)\n# test_ds  = HierSlideDataset(SAVE_DIR + \"/test\",  eval_transforms)\n\n# train_loader = DataLoader(train_ds, batch_size=1, shuffle=True,  num_workers=2, collate_fn=collate_single)\n# val_loader   = DataLoader(val_ds,   batch_size=1, shuffle=False, num_workers=2, collate_fn=collate_single)\n# test_loader  = DataLoader(test_ds,  batch_size=1, shuffle=False, num_workers=2, collate_fn=collate_single)\n\n# print(\"Train slides:\", len(train_ds), \"| Val slides:\", len(val_ds), \"| Test slides:\", len(test_ds))\n\n# # ── CELL 6: Attention pooling (reused at both levels) ────────────────────────\n# class AttentionPool(nn.Module):\n#     def __init__(self, feat_dim=768, hidden_dim=256):\n#         super().__init__()\n#         self.attention = nn.Sequential(\n#             nn.Linear(feat_dim, hidden_dim),\n#             nn.Tanh(),\n#             nn.Linear(hidden_dim, 1)\n#         )\n#     def forward(self, x):\n#         # x: (N, feat_dim)\n#         attn_weights = self.attention(x)\n#         attn_weights = torch.softmax(attn_weights, dim=0)\n#         pooled = (attn_weights * x).sum(dim=0)\n#         return pooled, attn_weights\n\n# # ── CELL 7: Two-level hierarchical model ─────────────────────────────────────\n# class HierarchicalSwin2Level(nn.Module):\n#     \"\"\"\n#     Level 1 (patch):  Swin-Tiny encodes patches -> patch features (768-d)\n#     Level 2 (region): AttentionPool over patches within each region -> region feature\n#     Level 3 (slide):  AttentionPool over region features -> slide feature\n#     Classifier: slide feature -> ISUP logits\n#     \"\"\"\n#     def __init__(self, num_classes=6, freeze_encoder=False):\n#         super().__init__()\n#         self.patch_encoder = timm.create_model(\n#             \"swin_tiny_patch4_window7_224\",\n#             pretrained=True,\n#             num_classes=0\n#         )\n#         if freeze_encoder:\n#             for p in self.patch_encoder.parameters():\n#                 p.requires_grad = False\n\n#         self.region_pool = AttentionPool(feat_dim=768, hidden_dim=256)\n#         self.slide_pool  = AttentionPool(feat_dim=768, hidden_dim=256)\n#         self.drop        = nn.Dropout(0.4)          # increased from 0.3\n#         self.classifier  = nn.Linear(768, num_classes)\n\n#     def forward(self, region_tensors):\n#         \"\"\"\n#         region_tensors: list of Tensors, each (N_r, 3, 224, 224), one per region\n#         \"\"\"\n#         region_feats = []\n#         region_attn  = []\n#         for patches in region_tensors:\n#             patches = patches.to(next(self.parameters()).device)\n#             patch_feats = self.patch_encoder(patches)         # (N_r, 768)\n#             region_feat, attn_w = self.region_pool(patch_feats)\n#             region_feats.append(region_feat)\n#             region_attn.append(attn_w)\n\n#         region_feats = torch.stack(region_feats, dim=0)       # (R, 768)\n#         slide_feat, slide_attn = self.slide_pool(region_feats)\n\n#         logits = self.classifier(self.drop(slide_feat))\n#         return logits, slide_attn, region_attn\n\n# # ── CELL 8: Training loop with early stopping ────────────────────────────────\n# model     = HierarchicalSwin2Level(freeze_encoder=False).to(device)\n# criterion = nn.CrossEntropyLoss()\n# optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\n# scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)\n\n# def run_epoch(loader, train=True):\n#     model.train() if train else model.eval()\n#     total_loss, all_preds, all_labels = 0, [], []\n#     ctx = torch.enable_grad() if train else torch.no_grad()\n#     with ctx:\n#         for region_tensors, label in loader:\n#             label_t = torch.tensor([label], device=device)\n#             if train:\n#                 optimizer.zero_grad()\n#             logits, _, _ = model(region_tensors)\n#             logits = logits.unsqueeze(0)\n#             loss = criterion(logits, label_t)\n#             if train:\n#                 loss.backward()\n#                 optimizer.step()\n#             total_loss += loss.item()\n#             all_preds  += torch.argmax(logits, 1).cpu().tolist()\n#             all_labels += [label]\n#     kappa = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\", labels=list(range(6)))\n#     return total_loss / max(1, len(loader)), kappa\n\n# history = {\"train_loss\": [], \"train_qwk\": [], \"val_loss\": [], \"val_qwk\": []}\n# best_kappa = -1\n# epochs_no_improve = 0\n\n# for epoch in range(EPOCHS):\n#     tr_loss, tr_kappa = run_epoch(train_loader, train=True)\n#     vl_loss, vl_kappa = run_epoch(val_loader,   train=False)\n#     scheduler.step()\n\n#     history[\"train_loss\"].append(tr_loss)\n#     history[\"train_qwk\"].append(tr_kappa)\n#     history[\"val_loss\"].append(vl_loss)\n#     history[\"val_qwk\"].append(vl_kappa)\n\n#     print(f\"Epoch {epoch+1}/{EPOCHS} | \"\n#           f\"Train Loss {tr_loss:.4f} QWK {tr_kappa:.4f} | \"\n#           f\"Val Loss {vl_loss:.4f} QWK {vl_kappa:.4f}\")\n\n#     if vl_kappa > best_kappa:\n#         best_kappa = vl_kappa\n#         epochs_no_improve = 0\n#         torch.save(model.state_dict(), \"/kaggle/working/hier_swin_2level_best.pth\")\n#         print(\"  Best model saved\")\n#     else:\n#         epochs_no_improve += 1\n#         if epochs_no_improve >= PATIENCE:\n#             print(f\"  Early stopping triggered at epoch {epoch+1}\")\n#             break\n\n# print(\"\\nBest Val QWK:\", best_kappa)\n\n# # ── CELL 9: Final TEST evaluation (held out, never touched during training) ──\n# model.load_state_dict(torch.load(\"/kaggle/working/hier_swin_2level_best.pth\"))\n# test_loss, test_kappa = run_epoch(test_loader, train=False)\n# print(f\"\\nFINAL TEST RESULTS -> Loss: {test_loss:.4f} | QWK: {test_kappa:.4f}\")\n\n# # ── CELL 10: Plot train/val curves for your professor ────────────────────────\n# import matplotlib.pyplot as plt\n# epochs_ran = range(1, len(history[\"train_qwk\"]) + 1)\n# plt.figure(figsize=(7,5))\n# plt.plot(epochs_ran, history[\"train_qwk\"], \"b-o\", label=\"Train QWK\")\n# plt.plot(epochs_ran, history[\"val_qwk\"], \"r-o\", label=\"Val QWK\")\n# plt.axhline(best_kappa, color=\"g\", linestyle=\"--\", label=f\"Best Val QWK = {best_kappa:.3f}\")\n# plt.axhline(test_kappa, color=\"purple\", linestyle=\":\", label=f\"Test QWK = {test_kappa:.3f}\")\n# plt.xlabel(\"Epoch\")\n# plt.ylabel(\"Quadratic Weighted Kappa\")\n# plt.title(f\"2-Level Hierarchical Swin-Tiny on PANDA ({NUM_SLIDES} slides, {GRID_SIZE}x{GRID_SIZE} regions)\")\n# plt.legend()                                     \n# plt.grid(True)            \n# plt.savefig(\"/kaggle/working/hier_swin_2level_curves.png\", dpi=150, bbox_inches=\"tight\")\n# plt.show()                          ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T19:55:20.203237Z","iopub.execute_input":"2026-08-08T19:55:20.204065Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # ==============================================================================\n# # TWO-LEVEL HIERARCHICAL SWIN MODEL — PANDA ISUP GRADING\n# # Patch -> Region (spatial) -> Slide\n# # Built on top of your baseline-swin-model notebook. Run cell-by-cell on Kaggle.\n# # ==============================================================================\n\n# y stopping patience (epochs without val improvement)\n# SAVE_DIR       = \"/kaggle/working/regions\"\n# MEAN = [0.485, 0.456, 0.406]\n# STD  = [0.229, 0.224, 0.225]\n# import os, numpy as np, pandas as pd\n# from PIL import Image\n# import cv2, torch, torch.nn as nn, torch.nn.functional as F\n# from torch.utils.data import Dataset, DataLoader\n# from sklearn.model_selection import train_test_split\n# from sklearn.metrics import cohen_kappa_score\n# import timm\n# from tqdm import tqdm\n# import torchvision.transforms as T\n\n# Image.MAX_IMAGE_PIXELS = None\n# device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n# print(\"Device:\", device)\n\n# # ── CELL 1: Config ───────────────────────────────────────────────────────────\n# PATCH_SIZE     = 224\n# DOWNSAMPLE     = 4          # was 8 (reduce(8)) - halving this keeps regions big enough to hold patches\n# GRID_SIZE      = 3          # was 4 -> 3x3 = 9 regions, each region stays well above PATCH_SIZE\n# PATCHES_PER_REGION = 5      # patches sampled per region (tissue-permitting)\n# NUM_SLIDES     = 50       # keep consistent with your baseline for comparison\n# VAL_SLIDES     = 20\n# TEST_SLIDES    = 20\n# BATCH_SIZE     = 1          # 1 slide per batch (variable region/patch counts)\n# EPOCHS         = 15\n# LR             = 3e-5\n# WEIGHT_DECAY   = 5e-4       # increased from 1e-4 to fight overfitting\n# PATIENCE       = 5          # earl\n# train_csv   = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train.csv\"\n# images_path = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train_images\"\n\n# # ── CELL 2: Three-way split (train/val/test — your baseline only had train/val) ──\n# df = pd.read_csv(train_csv)\n\n# train_df, temp_df = train_test_split(\n#     df, test_size=0.3, stratify=df[\"isup_grade\"], random_state=42\n# )\n# val_df, test_df = train_test_split(\n#     temp_df, test_size=0.5, stratify=temp_df[\"isup_grade\"], random_state=42\n# )\n\n# train_df = train_df.sample(min(NUM_SLIDES, len(train_df)), random_state=42).reset_index(drop=True)\n# val_df   = val_df.sample(min(VAL_SLIDES, len(val_df)), random_state=42).reset_index(drop=True)\n# test_df  = test_df.sample(min(TEST_SLIDES, len(test_df)), random_state=42).reset_index(drop=True)\n\n# print(f\"Train slides: {len(train_df)} | Val slides: {len(val_df)} | Test slides: {len(test_df)}\")\n\n# # ── CELL 3: Tissue-aware SPATIAL region + patch extraction (OpenSlide, fast) ──\n# # Uses OpenSlide to read directly from a lower pyramid level instead of decoding\n# # the full-resolution TIFF and downsampling in memory (which is what PIL.reduce()\n# # was doing — and why extraction took ~58 min for 1000 slides).\n# try:\n#     import openslide\n# except ImportError:\n#     import subprocess, sys\n#     subprocess.run([sys.executable, \"-m\", \"pip\", \"install\", \"openslide-python\", \"-q\"])\n#     import openslide\n\n# PYRAMID_LEVEL = 1   # 0=full res, 1=~4x downsampled, 2=~16x downsampled (varies per slide)\n\n# def is_tissue(patch, threshold=0.6):\n#     gray = cv2.cvtColor(patch, cv2.COLOR_RGB2GRAY)\n#     return (gray < 230).mean() > threshold\n\n# def extract_spatial_patches(df, save_dir, grid_size=GRID_SIZE, patches_per_region=PATCHES_PER_REGION):\n#     \"\"\"\n#     Same region/patch logic as before, but reads slides via OpenSlide at a\n#     pre-downsampled pyramid level instead of decoding full resolution with PIL.\n#     This is typically 5-10x faster on large WSIs.\n#     \"\"\"\n#     os.makedirs(save_dir, exist_ok=True)\n#     for _, row in tqdm(df.iterrows(), total=len(df)):\n#         img_id = row[\"image_id\"]\n#         label  = row[\"isup_grade\"]\n#         path   = os.path.join(images_path, img_id + \".tiff\")\n#         try:\n#             slide_obj = openslide.OpenSlide(path)\n#             level = min(PYRAMID_LEVEL, slide_obj.level_count - 1)\n#             w, h = slide_obj.level_dimensions[level]\n#             slide = np.array(slide_obj.read_region((0, 0), level, (w, h)).convert(\"RGB\"))\n#             slide_obj.close()\n#         except Exception:\n#             continue\n\n#         min_dim = int(PATCH_SIZE * 1.5) * grid_size\n#         if h <= min_dim or w <= min_dim:\n#             continue\n\n#         region_h = h // grid_size\n#         region_w = w // grid_size\n\n#         for r_idx in range(grid_size * grid_size):\n#             ry, rx = divmod(r_idx, grid_size)\n#             y0, y1 = ry * region_h, min((ry + 1) * region_h, h)\n#             x0, x1 = rx * region_w, min((rx + 1) * region_w, w)\n#             if (y1 - y0) <= int(PATCH_SIZE * 1.2) or (x1 - x0) <= int(PATCH_SIZE * 1.2):\n#                 continue\n\n#             region_slide = slide[y0:y1, x0:x1]\n#             saved, attempts = 0, 0\n#             while saved < patches_per_region and attempts < patches_per_region * 8:\n#                 yy = np.random.randint(0, max(1, (y1 - y0) - PATCH_SIZE))\n#                 xx = np.random.randint(0, max(1, (x1 - x0) - PATCH_SIZE))\n#                 patch = region_slide[yy:yy+PATCH_SIZE, xx:xx+PATCH_SIZE]\n#                 if patch.shape[0] == PATCH_SIZE and patch.shape[1] == PATCH_SIZE and is_tissue(patch):\n#                     fname = f\"{img_id}_r{r_idx}_p{saved}_{label}.png\"\n#                     cv2.imwrite(os.path.join(save_dir, fname),\n#                                 cv2.cvtColor(patch, cv2.COLOR_RGB2BGR))\n#                     saved += 1\n#                 attempts += 1\n\n# extract_spatial_patches(train_df, SAVE_DIR + \"/train\")\n# extract_spatial_patches(val_df,   SAVE_DIR + \"/val\")\n# extract_spatial_patches(test_df,  SAVE_DIR + \"/test\")\n\n# n_train_patches = len(os.listdir(SAVE_DIR + \"/train\"))\n# n_val_patches   = len(os.listdir(SAVE_DIR + \"/val\"))\n# n_test_patches  = len(os.listdir(SAVE_DIR + \"/test\"))\n# print(\"Train patches:\", n_train_patches, f\"(~{n_train_patches/max(1,len(train_df)):.1f} per slide)\")\n# print(\"Val patches  :\", n_val_patches,   f\"(~{n_val_patches/max(1,len(val_df)):.1f} per slide)\")\n# print(\"Test patches :\", n_test_patches,  f\"(~{n_test_patches/max(1,len(test_df)):.1f} per slide)\")\n# print(\"NOTE: if per-slide patch count is well below GRID_SIZE^2 * PATCHES_PER_REGION,\")\n# print(\"regions are still being skipped as too small -> check slide dimension distribution.\")\n\n# # ── CELL 4: Transforms ───────────────────────────────────────────────────────\n# train_transforms = T.Compose([\n#     T.RandomHorizontalFlip(),\n#     T.RandomVerticalFlip(),\n#     T.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.1),\n#     T.ToTensor(),\n#     T.Normalize(mean=MEAN, std=STD),\n# ])\n# eval_transforms = T.Compose([\n#     T.ToTensor(),\n#     T.Normalize(mean=MEAN, std=STD),\n# ])\n\n# # ── CELL 5: Slide dataset with REGION structure ──────────────────────────────\n# class HierSlideDataset(Dataset):\n#     \"\"\"\n#     Each item = one slide, structured as a list of regions,\n#     each region being a list of patch file paths.\n#     Returns: (list[Tensor(N_r, 3, 224, 224)] per region, label)\n#     \"\"\"\n#     def __init__(self, patch_dir, transform=None):\n#         self.transform = transform\n#         # slide_id -> region_id -> [patch paths]\n#         self.slide_data = {}\n#         for f in os.listdir(patch_dir):\n#             if not f.endswith(\".png\"):\n#                 continue\n#             stem = f.replace(\".png\", \"\")\n#             # format: imgid_r{r}_p{p}_{label}\n#             parts = stem.split(\"_\")\n#             label = int(parts[-1])\n#             p_tag = parts[-2]      # p{p}\n#             r_tag = parts[-3]      # r{r}\n#             img_id = \"_\".join(parts[:-3])\n#             region_id = int(r_tag[1:])\n\n#             self.slide_data.setdefault(img_id, {\"label\": label, \"regions\": {}})\n#             self.slide_data[img_id][\"regions\"].setdefault(region_id, [])\n#             self.slide_data[img_id][\"regions\"][region_id].append(os.path.join(patch_dir, f))\n\n#         self.slides = list(self.slide_data.keys())\n\n#     def __len__(self):\n#         return len(self.slides)\n\n#     def __getitem__(self, idx):\n#         sid = self.slides[idx]\n#         info = self.slide_data[sid]\n#         label = info[\"label\"]\n\n#         region_tensors = []\n#         for region_id, paths in info[\"regions\"].items():\n#             patch_tensors = []\n#             for p in paths:\n#                 img = Image.open(p).convert(\"RGB\")\n#                 if self.transform:\n#                     img = self.transform(img)\n#                 patch_tensors.append(img)\n#             region_tensors.append(torch.stack(patch_tensors, dim=0))  # (N_r, 3, 224, 224)\n\n#         return region_tensors, label\n\n# def collate_single(batch):\n#     # batch_size=1 always, just unwrap\n#     region_tensors, label = batch[0]\n#     return region_tensors, label\n\n# train_ds = HierSlideDataset(SAVE_DIR + \"/train\", train_transforms)\n# val_ds   = HierSlideDataset(SAVE_DIR + \"/val\",   eval_transforms)\n# test_ds  = HierSlideDataset(SAVE_DIR + \"/test\",  eval_transforms)\n\n# train_loader = DataLoader(train_ds, batch_size=1, shuffle=True,  num_workers=2, collate_fn=collate_single)\n# val_loader   = DataLoader(val_ds,   batch_size=1, shuffle=False, num_workers=2, collate_fn=collate_single)\n# test_loader  = DataLoader(test_ds,  batch_size=1, shuffle=False, num_workers=2, collate_fn=collate_single)\n\n# print(\"Train slides:\", len(train_ds), \"| Val slides:\", len(val_ds), \"| Test slides:\", len(test_ds))\n\n# # ── CELL 6: Attention pooling (reused at both levels) ────────────────────────\n# class AttentionPool(nn.Module):\n#     def __init__(self, feat_dim=768, hidden_dim=256):\n#         super().__init__()\n#         self.attention = nn.Sequential(\n#             nn.Linear(feat_dim, hidden_dim),\n#             nn.Tanh(),\n#             nn.Linear(hidden_dim, 1)\n#         )\n#     def forward(self, x):\n#         # x: (N, feat_dim)\n#         attn_weights = self.attention(x)\n#         attn_weights = torch.softmax(attn_weights, dim=0)\n#         pooled = (attn_weights * x).sum(dim=0)\n#         return pooled, attn_weights\n\n# # ── CELL 7: Two-level hierarchical model ─────────────────────────────────────\n# class HierarchicalSwin2Level(nn.Module):\n#     \"\"\"\n#     Level 1 (patch):  Swin-Tiny encodes patches -> patch features (768-d)\n#     Level 2 (region): AttentionPool over patches within each region -> region feature\n#     Level 3 (slide):  AttentionPool over region features -> slide feature\n#     Classifier: slide feature -> ISUP logits\n#     \"\"\"\n#     def __init__(self, num_classes=6, freeze_encoder=False):\n#         super().__init__()\n#         self.patch_encoder = timm.create_model(\n#             \"swin_tiny_patch4_window7_224\",\n#             pretrained=True,\n#             num_classes=0\n#         )\n#         if freeze_encoder:\n#             for p in self.patch_encoder.parameters():\n#                 p.requires_grad = False\n\n#         self.region_pool = AttentionPool(feat_dim=768, hidden_dim=256)\n#         self.slide_pool  = AttentionPool(feat_dim=768, hidden_dim=256)\n#         self.drop        = nn.Dropout(0.4)          # increased from 0.3\n#         self.classifier  = nn.Linear(768, num_classes)\n\n#     def forward(self, region_tensors):\n#         \"\"\"\n#         region_tensors: list of Tensors, each (N_r, 3, 224, 224), one per region\n#         \"\"\"\n#         region_feats = []\n#         region_attn  = []\n#         for patches in region_tensors:\n#             patches = patches.to(next(self.parameters()).device)\n#             patch_feats = self.patch_encoder(patches)         # (N_r, 768)\n#             region_feat, attn_w = self.region_pool(patch_feats)\n#             region_feats.append(region_feat)\n#             region_attn.append(attn_w)\n\n#         region_feats = torch.stack(region_feats, dim=0)       # (R, 768)\n#         slide_feat, slide_attn = self.slide_pool(region_feats)\n\n#         logits = self.classifier(self.drop(slide_feat))\n#         return logits, slide_attn, region_attn\n\n# # ── CELL 8: Training loop with early stopping + checkpoint-resume ────────────\n# CKPT_PATH = \"/kaggle/working/hier_swin_2level_checkpoint.pth\"   # full resumable state\n# BEST_PATH = \"/kaggle/working/hier_swin_2level_best.pth\"          # best weights only\n\n# model     = HierarchicalSwin2Level(freeze_encoder=False).to(device)\n# criterion = nn.CrossEntropyLoss()\n# optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\n# scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)\n\n# start_epoch = 0\n# history = {\"train_loss\": [], \"train_qwk\": [], \"val_loss\": [], \"val_qwk\": []}\n# best_kappa = -1\n# epochs_no_improve = 0\n\n# # If a checkpoint exists from a previous (crashed/interrupted) session, resume from it.\n# # On Kaggle, /kaggle/working persists within a session but NOT automatically across\n# # a hard crash unless you've saved this file to a Dataset — see notes below the loop.\n# if os.path.exists(CKPT_PATH):\n#     ckpt = torch.load(CKPT_PATH, map_location=device)\n#     model.load_state_dict(ckpt[\"model_state\"])\n#     optimizer.load_state_dict(ckpt[\"optimizer_state\"])\n#     scheduler.load_state_dict(ckpt[\"scheduler_state\"])\n#     start_epoch = ckpt[\"epoch\"] + 1\n#     history = ckpt[\"history\"]\n#     best_kappa = ckpt[\"best_kappa\"]\n#     epochs_no_improve = ckpt[\"epochs_no_improve\"]\n#     print(f\"Resumed from checkpoint at epoch {start_epoch} | best_kappa so far: {best_kappa:.4f}\")\n# else:\n#     print(\"No checkpoint found — starting fresh.\")\n\n# def run_epoch(loader, train=True):\n#     model.train() if train else model.eval()\n#     total_loss, all_preds, all_labels = 0, [], []\n#     ctx = torch.enable_grad() if train else torch.no_grad()\n#     with ctx:\n#         for region_tensors, label in loader:\n#             label_t = torch.tensor([label], device=device)\n#             if train:\n#                 optimizer.zero_grad()\n#             logits, _, _ = model(region_tensors)\n#             logits = logits.unsqueeze(0)\n#             loss = criterion(logits, label_t)\n#             if train:\n#                 loss.backward()\n#                 optimizer.step()\n#             total_loss += loss.item()\n#             all_preds  += torch.argmax(logits, 1).cpu().tolist()\n#             all_labels += [label]\n#     kappa = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\", labels=list(range(6)))\n#     return total_loss / max(1, len(loader)), kappa\n\n# for epoch in range(start_epoch, EPOCHS):\n#     tr_loss, tr_kappa = run_epoch(train_loader, train=True)\n#     vl_loss, vl_kappa = run_epoch(val_loader,   train=False)\n#     scheduler.step()\n\n#     history[\"train_loss\"].append(tr_loss)\n#     history[\"train_qwk\"].append(tr_kappa)\n#     history[\"val_loss\"].append(vl_loss)\n#     history[\"val_qwk\"].append(vl_kappa)\n\n#     print(f\"Epoch {epoch+1}/{EPOCHS} | \"\n#           f\"Train Loss {tr_loss:.4f} QWK {tr_kappa:.4f} | \"\n#           f\"Val Loss {vl_loss:.4f} QWK {vl_kappa:.4f}\")\n\n#     if vl_kappa > best_kappa:\n#         best_kappa = vl_kappa\n#         epochs_no_improve = 0\n#         torch.save(model.state_dict(), BEST_PATH)\n#         print(\"  Best model saved\")\n#     else:\n#         epochs_no_improve += 1\n\n#     # Save full resumable checkpoint EVERY epoch, regardless of improvement,\n#     # so a crash never costs you more than one epoch of work.\n#     torch.save({\n#         \"epoch\": epoch,\n#         \"model_state\": model.state_dict(),\n#         \"optimizer_state\": optimizer.state_dict(),\n#         \"scheduler_state\": scheduler.state_dict(),\n#         \"history\": history,\n#         \"best_kappa\": best_kappa,\n#         \"epochs_no_improve\": epochs_no_improve,\n#     }, CKPT_PATH)\n\n#     if epochs_no_improve >= PATIENCE:\n#         print(f\"  Early stopping triggered at epoch {epoch+1}\")\n#         break\n\n# print(\"\\nBest Val QWK:\", best_kappa)\n\n# # IMPORTANT: after this cell finishes (or if you're about to lose your session),\n# # go to your notebook's Output pane and click \"New Dataset\" on /kaggle/working\n# # to persist hier_swin_2level_checkpoint.pth and hier_swin_2level_best.pth.\n# # Next session: add that dataset as input, copy the .pth files back into\n# # /kaggle/working before re-running this cell, and it will auto-resume.\n\n# # ── CELL 9: Final TEST evaluation (held out, never touched during training) ──\n# model.load_state_dict(torch.load(BEST_PATH))\n# test_loss, test_kappa = run_epoch(test_loader, train=False)\n# print(f\"\\nFINAL TEST RESULTS -> Loss: {test_loss:.4f} | QWK: {test_kappa:.4f}\")\n\n# # ── CELL 10: Plot train/val curves for your professor ────────────────────────\n# import matplotlib.pyplot as plt\n# epochs_ran = range(1, len(history[\"train_qwk\"]) + 1)\n# plt.figure(figsize=(7,5))\n# plt.plot(epochs_ran, history[\"train_qwk\"], \"b-o\", label=\"Train QWK\")\n# plt.plot(epochs_ran, history[\"val_qwk\"], \"r-o\", label=\"Val QWK\")\n# plt.axhline(best_kappa, color=\"g\", linestyle=\"--\", label=f\"Best Val QWK = {best_kappa:.3f}\")\n# plt.axhline(test_kappa, color=\"purple\", linestyle=\":\", label=f\"Test QWK = {test_kappa:.3f}\")\n# plt.xlabel(\"Epoch\")\n# plt.ylabel(\"Quadratic Weighted Kappa\")\n# plt.title(f\"2-Level Hierarchical Swin-Tiny on PANDA ({NUM_SLIDES} slides, {GRID_SIZE}x{GRID_SIZE} regions)\")\n# plt.legend()\n# plt.grid(True)\n# plt.savefig(\"/kaggle/working/hier_swin_2level_curves.png\", dpi=150, bbox_inches=\"tight\")\n# plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import os, numpy as np, pandas as pd\n# from PIL import Image\n# import cv2, torch, torch.nn as nn, torch.nn.functional as F\n# from torch.utils.data import Dataset, DataLoader\n# from sklearn.model_selection import train_test_split\n# from sklearn.metrics import cohen_kappa_score\n# import timm\n# from tqdm import tqdm\n# import torchvision.transforms as T\n\n# Image.MAX_IMAGE_PIXELS = None\n# device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n# print(\"Device:\", device)\n\n# # ── CELL 1: Config ───────────────────────────────────────────────────────────\n# PATCH_SIZE     = 224\n# DOWNSAMPLE     = 4          # was 8 (reduce(8)) - halving this keeps regions big enough to hold patches\n# GRID_SIZE      = 3          # was 4 -> 3x3 = 9 regions, each region stays well above PATCH_SIZE\n# PATCHES_PER_REGION = 5      # patches sampled per region (tissue-permitting)\n# NUM_SLIDES     = 50       # keep consistent with your baseline for comparison\n# VAL_SLIDES     = 20\n# TEST_SLIDES    = 20\n# BATCH_SIZE     = 1          # 1 slide per batch (variable region/patch counts)\n# EPOCHS         = 15\n# LR             = 3e-5\n# WEIGHT_DECAY   = 5e-4       # increased from 1e-4 to fight overfitting\n# PATIENCE       = 5          # earl\n# train_csv   = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train.csv\"\n# images_path = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train_images\"\n\n# # ── CELL 2: Three-way split (train/val/test — your baseline only had train/val) ──\n# df = pd.read_csv(train_csv)\n\n# train_df, temp_df = train_test_split(\n#     df, test_size=0.3, stratify=df[\"isup_grade\"], random_state=42\n# )\n# val_df, test_df = train_test_split(\n#     temp_df, test_size=0.5, stratify=temp_df[\"isup_grade\"], random_state=42\n# )\n\n# train_df = train_df.sample(min(NUM_SLIDES, len(train_df)), random_state=42).reset_index(drop=True)\n# val_df   = val_df.sample(min(VAL_SLIDES, len(val_df)), random_state=42).reset_index(drop=True)\n# test_df  = test_df.sample(min(TEST_SLIDES, len(test_df)), random_state=42).reset_index(drop=True)\n\n# print(f\"Train slides: {len(train_df)} | Val slides: {len(val_df)} | Test slides: {len(test_df)}\")\n\n# # ── CELL 3: Tissue-aware SPATIAL region + patch extraction (OpenSlide, fast) ──\n# # Uses OpenSlide to read directly from a lower pyramid level instead of decoding\n# # the full-resolution TIFF and downsampling in memory (which is what PIL.reduce()\n# # was doing — and why extraction took ~58 min for 1000 slides).\n# try:\n#     import openslide\n# except ImportError:\n#     import subprocess, sys\n#     subprocess.run([sys.executable, \"-m\", \"pip\", \"install\", \"openslide-python\", \"-q\"])\n#     import openslide\n\n# PYRAMID_LEVEL = 1   # 0=full res, 1=~4x downsampled, 2=~16x downsampled (varies per slide)\n\n# def is_tissue(patch, threshold=0.6):\n#     gray = cv2.cvtColor(patch, cv2.COLOR_RGB2GRAY)\n#     return (gray < 230).mean() > threshold\n\n# def extract_spatial_patches(df, save_dir, grid_size=GRID_SIZE, patches_per_region=PATCHES_PER_REGION):\n#     \"\"\"\n#     Same region/patch logic as before, but reads slides via OpenSlide at a\n#     pre-downsampled pyramid level instead of decoding full resolution with PIL.\n#     This is typically 5-10x faster on large WSIs.\n#     \"\"\"\n#     os.makedirs(save_dir, exist_ok=True)\n#     for _, row in tqdm(df.iterrows(), total=len(df)):\n#         img_id = row[\"image_id\"]\n#         label  = row[\"isup_grade\"]\n#         path   = os.path.join(images_path, img_id + \".tiff\")\n#         try:\n#             slide_obj = openslide.OpenSlide(path)\n#             level = min(PYRAMID_LEVEL, slide_obj.level_count - 1)\n#             w, h = slide_obj.level_dimensions[level]\n#             slide = np.array(slide_obj.read_region((0, 0), level, (w, h)).convert(\"RGB\"))\n#             slide_obj.close()\n#         except Exception:\n#             continue\n\n#         min_dim = int(PATCH_SIZE * 1.5) * grid_size\n#         if h <= min_dim or w <= min_dim:\n#             continue\n\n#         region_h = h // grid_size\n#         region_w = w // grid_size\n\n#         for r_idx in range(grid_size * grid_size):\n#             ry, rx = divmod(r_idx, grid_size)\n#             y0, y1 = ry * region_h, min((ry + 1) * region_h, h)\n#             x0, x1 = rx * region_w, min((rx + 1) * region_w, w)\n#             if (y1 - y0) <= int(PATCH_SIZE * 1.2) or (x1 - x0) <= int(PATCH_SIZE * 1.2):\n#                 continue\n\n#             region_slide = slide[y0:y1, x0:x1]\n#             saved, attempts = 0, 0\n#             while saved < patches_per_region and attempts < patches_per_region * 8:\n#                 yy = np.random.randint(0, max(1, (y1 - y0) - PATCH_SIZE))\n#                 xx = np.random.randint(0, max(1, (x1 - x0) - PATCH_SIZE))\n#                 patch = region_slide[yy:yy+PATCH_SIZE, xx:xx+PATCH_SIZE]\n#                 if patch.shape[0] == PATCH_SIZE and patch.shape[1] == PATCH_SIZE and is_tissue(patch):\n#                     fname = f\"{img_id}_r{r_idx}_p{saved}_{label}.png\"\n#                     cv2.imwrite(os.path.join(save_dir, fname),\n#                                 cv2.cvtColor(patch, cv2.COLOR_RGB2BGR))\n#                     saved += 1\n#                 attempts += 1\n\n# extract_spatial_patches(train_df, SAVE_DIR + \"/train\")\n# extract_spatial_patches(val_df,   SAVE_DIR + \"/val\")\n# extract_spatial_patches(test_df,  SAVE_DIR + \"/test\")\n\n# n_train_patches = len(os.listdir(SAVE_DIR + \"/train\"))\n# n_val_patches   = len(os.listdir(SAVE_DIR + \"/val\"))\n# n_test_patches  = len(os.listdir(SAVE_DIR + \"/test\"))\n# print(\"Train patches:\", n_train_patches, f\"(~{n_train_patches/max(1,len(train_df)):.1f} per slide)\")\n# print(\"Val patches  :\", n_val_patches,   f\"(~{n_val_patches/max(1,len(val_df)):.1f} per slide)\")\n# print(\"Test patches :\", n_test_patches,  f\"(~{n_test_patches/max(1,len(test_df)):.1f} per slide)\")\n# print(\"NOTE: if per-slide patch count is well below GRID_SIZE^2 * PATCHES_PER_REGION,\")\n# print(\"regions are still being skipped as too small -> check slide dimension distribution.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T23:57:28.272326Z","iopub.execute_input":"2026-08-08T23:57:28.272726Z","iopub.status.idle":"2026-08-08T23:57:38.957829Z","shell.execute_reply.started":"2026-08-08T23:57:28.272694Z","shell.execute_reply":"2026-08-08T23:57:38.956875Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import os, numpy as np, pandas as pd\n# from PIL import Image\n# import cv2, torch, torch.nn as nn, torch.nn.functional as F\n# from torch.utils.data import Dataset, DataLoader\n# from sklearn.model_selection import train_test_split\n# from sklearn.metrics import cohen_kappa_score\n# import timm\n# from tqdm import tqdm\n# import torchvision.transforms as T\n \n# Image.MAX_IMAGE_PIXELS = None\n# device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n# print(\"Device:\", device)\n \n# # ── CELL 1: Config ───────────────────────────────────────────────────────────\n# PATCH_SIZE     = 224\n# DOWNSAMPLE     = 4          # was 8 (reduce(8)) - halving this keeps regions big enough to hold patches\n# GRID_SIZE      = 3          # was 4 -> 3x3 = 9 regions, each region stays well above PATCH_SIZE\n# PATCHES_PER_REGION = 5      # patches sampled per region (tissue-permitting)\n# NUM_SLIDES     = 50         # TEMP: small test run to verify OpenSlide works & check timing\n# VAL_SLIDES     = 20         # scaled down to match — revert to 1000/200/200 once verified\n# TEST_SLIDES    = 20\n# BATCH_SIZE     = 1          # 1 slide per batch (variable region/patch counts)\n# EPOCHS         = 15\n# LR             = 3e-5\n# WEIGHT_DECAY   = 5e-4       # increased from 1e-4 to fight overfitting\n# PATIENCE       = 5          # early stopping patience (epochs without val improvement)\n# SAVE_DIR       = \"/kaggle/working/regions\"\n# MEAN = [0.485, 0.456, 0.406]\n# STD  = [0.229, 0.224, 0.225]\n \n# train_csv   = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train.csv\"\n# images_path = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train_images\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T23:58:46.244443Z","iopub.execute_input":"2026-08-08T23:58:46.244903Z","iopub.status.idle":"2026-08-08T23:58:46.25283Z","shell.execute_reply.started":"2026-08-08T23:58:46.244871Z","shell.execute_reply":"2026-08-08T23:58:46.251936Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# df = pd.read_csv(train_csv)\n \n# train_df, temp_df = train_test_split(\n#     df, test_size=0.3, stratify=df[\"isup_grade\"], random_state=42\n# )\n# val_df, test_df = train_test_split(\n#     temp_df, test_size=0.5, stratify=temp_df[\"isup_grade\"], random_state=42\n# )\n \n# train_df = train_df.sample(min(NUM_SLIDES, len(train_df)), random_state=42).reset_index(drop=True)\n# val_df   = val_df.sample(min(VAL_SLIDES, len(val_df)), random_state=42).reset_index(drop=True)\n# test_df  = test_df.sample(min(TEST_SLIDES, len(test_df)), random_state=42).reset_index(drop=True)\n \n# print(f\"Train slides: {len(train_df)} | Val slides: {len(val_df)} | Test slides: {len(test_df)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T23:59:06.959958Z","iopub.execute_input":"2026-08-08T23:59:06.960218Z","iopub.status.idle":"2026-08-08T23:59:06.990689Z","shell.execute_reply.started":"2026-08-08T23:59:06.960194Z","shell.execute_reply":"2026-08-08T23:59:06.99008Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # ── CELL 3: Tissue-aware SPATIAL region + patch extraction (OpenSlide, fast) ──\n# # Uses OpenSlide to read directly from a lower pyramid level instead of decoding\n# # the full-resolution TIFF and downsampling in memory (which is what PIL.reduce()\n# # was doing — and why extraction took ~58 min for 1000 slides).\n# try:\n#     import openslide\n# except ImportError:\n#     import subprocess, sys\n#     subprocess.run([sys.executable, \"-m\", \"pip\", \"install\", \"openslide-python\", \"-q\"])\n#     import openslide\n \n# PYRAMID_LEVEL = 1   # 0=full res, 1=~4x downsampled, 2=~16x downsampled (varies per slide)\n \n# def is_tissue(patch, threshold=0.6):\n#     gray = cv2.cvtColor(patch, cv2.COLOR_RGB2GRAY)\n#     return (gray < 230).mean() > threshold\n \n# def extract_spatial_patches(df, save_dir, grid_size=GRID_SIZE, patches_per_region=PATCHES_PER_REGION):\n#     \"\"\"\n#     Same region/patch logic as before, but reads slides via OpenSlide at a\n#     pre-downsampled pyramid level instead of decoding full resolution with PIL.\n#     This is typically 5-10x faster on large WSIs.\n#     \"\"\"\n#     os.makedirs(save_dir, exist_ok=True)\n#     for _, row in tqdm(df.iterrows(), total=len(df)):\n#         img_id = row[\"image_id\"]\n#         label  = row[\"isup_grade\"]\n#         path   = os.path.join(images_path, img_id + \".tiff\")\n#         try:\n#             slide_obj = openslide.OpenSlide(path)\n#             level = min(PYRAMID_LEVEL, slide_obj.level_count - 1)\n#             w, h = slide_obj.level_dimensions[level]\n#             slide = np.array(slide_obj.read_region((0, 0), level, (w, h)).convert(\"RGB\"))\n#             slide_obj.close()\n#         except Exception:\n#             continue\n \n#         min_dim = int(PATCH_SIZE * 1.5) * grid_size\n#         if h <= min_dim or w <= min_dim:\n#             continue\n \n#         region_h = h // grid_size\n#         region_w = w // grid_size\n \n#         for r_idx in range(grid_size * grid_size):\n#             ry, rx = divmod(r_idx, grid_size)\n#             y0, y1 = ry * region_h, min((ry + 1) * region_h, h)\n#             x0, x1 = rx * region_w, min((rx + 1) * region_w, w)\n#             if (y1 - y0) <= int(PATCH_SIZE * 1.2) or (x1 - x0) <= int(PATCH_SIZE * 1.2):\n#                 continue\n \n#             region_slide = slide[y0:y1, x0:x1]\n#             saved, attempts = 0, 0\n#             while saved < patches_per_region and attempts < patches_per_region * 8:\n#                 yy = np.random.randint(0, max(1, (y1 - y0) - PATCH_SIZE))\n#                 xx = np.random.randint(0, max(1, (x1 - x0) - PATCH_SIZE))\n#                 patch = region_slide[yy:yy+PATCH_SIZE, xx:xx+PATCH_SIZE]\n#                 if patch.shape[0] == PATCH_SIZE and patch.shape[1] == PATCH_SIZE and is_tissue(patch):\n#                     fname = f\"{img_id}_r{r_idx}_p{saved}_{label}.png\"\n#                     cv2.imwrite(os.path.join(save_dir, fname),\n#                                 cv2.cvtColor(patch, cv2.COLOR_RGB2BGR))\n#                     saved += 1\n#                 attempts += 1\n \n# extract_spatial_patches(train_df, SAVE_DIR + \"/train\")\n# extract_spatial_patches(val_df,   SAVE_DIR + \"/val\")\n# extract_spatial_patches(test_df,  SAVE_DIR + \"/test\")    \n \n# n_train_patches = len(os.listdir(SAVE_DIR + \"/train\"))\n# n_val_patches   = len(os.listdir(SAVE_DIR + \"/val\")) \n# n_test_patches  = len(os.listdir(SAVE_DIR + \"/test\"))\n# print(\"Train patches:\", n_train_patches, f\"(~{n_train_patches/max(1,len(train_df)):.1f} per slide)\")\n# print(\"Val patches  :\", n_val_patches,   f\"(~{n_val_patches/max(1,len(val_df)):.1f} per slide)\")\n# print(\"Test patches :\", n_test_patches,  f\"(~{n_test_patches/max(1,len(test_df)):.1f} per slide)\")\n# print(\"NOTE: if per-slide patch count is well below GRID_SIZE^2 * PATCHES_PER_REGION,\")\n# print(\"regions are still being skipped as too small -> check slide dimension distribution.\")\n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T23:59:31.243433Z","iopub.execute_input":"2026-08-08T23:59:31.243702Z","iopub.status.idle":"2026-08-09T00:02:32.180808Z","shell.execute_reply.started":"2026-08-08T23:59:31.243678Z","shell.execute_reply":"2026-08-09T00:02:32.180128Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import os, numpy as np, pandas as pd\n# from PIL import Image\n# import cv2, torch, torch.nn as nn, torch.nn.functional as F\n# from torch.utils.data import Dataset, DataLoader\n# from sklearn.model_selection import train_test_split\n# from sklearn.metrics import cohen_kappa_score\n# import timm\n# from tqdm import tqdm\n# import torchvision.transforms as T\n \n# Image.MAX_IMAGE_PIXELS = None\n# device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n# print(\"Device:\", device)\n \n# # ── CELL 1: Config ───────────────────────────────────────────────────────────\n# PATCH_SIZE     = 224\n# DOWNSAMPLE     = 4          # was 8 (reduce(8)) - halving this keeps regions big enough to hold patches\n# GRID_SIZE      = 3          # was 4 -> 3x3 = 9 regions, each region stays well above PATCH_SIZE\n# PATCHES_PER_REGION = 5      # patches sampled per region (tissue-permitting)\n# NUM_SLIDES     = 1000       # back to full run — extraction speed & density confirmed on 50-slide test\n# VAL_SLIDES     = 200\n# TEST_SLIDES    = 200\n# BATCH_SIZE     = 1          # 1 slide per batch (variable region/patch counts)\n# EPOCHS         = 15\n# LR             = 3e-5\n# WEIGHT_DECAY   = 5e-4       # increased from 1e-4 to fight overfitting\n# PATIENCE       = 5          # early stopping patience (epochs without val improvement)\n# SAVE_DIR       = \"/kaggle/working/regions\"\n# MEAN = [0.485, 0.456, 0.406]\n# STD  = [0.229, 0.224, 0.225]\n \n# train_csv   = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train.csv\"\n# images_path = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train_images\"\n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T00:04:57.83358Z","iopub.execute_input":"2026-08-09T00:04:57.83418Z","iopub.status.idle":"2026-08-09T00:04:57.841115Z","shell.execute_reply.started":"2026-08-09T00:04:57.834152Z","shell.execute_reply":"2026-08-09T00:04:57.840262Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# df = pd.read_csv(train_csv)\n \n# train_df, temp_df = train_test_split(\n#     df, test_size=0.3, stratify=df[\"isup_grade\"], random_state=42\n# )\n# val_df, test_df = train_test_split(\n#     temp_df, test_size=0.5, stratify=temp_df[\"isup_grade\"], random_state=42\n# )\n \n# train_df = train_df.sample(min(NUM_SLIDES, len(train_df)), random_state=42).reset_index(drop=True)\n# val_df   = val_df.sample(min(VAL_SLIDES, len(val_df)), random_state=42).reset_index(drop=True)\n# test_df  = test_df.sample(min(TEST_SLIDES, len(test_df)), random_state=42).reset_index(drop=True)\n \n# print(f\"Train slides: {len(train_df)} | Val slides: {len(val_df)} | Test slides: {len(test_df)}\")\n \n# # ── CELL 3: Tissue-aware SPATIAL region + patch extraction (OpenSlide, fast) ──\n# # Uses OpenSlide to read directly from a lower pyramid level instead of decoding\n# # the full-resolution TIFF and downsampling in memory (which is what PIL.reduce()\n# # was doing — and why extraction took ~58 min for 1000 slides).\n# try:\n#     import openslide\n# except ImportError:\n#     import subprocess, sys\n#     subprocess.run([sys.executable, \"-m\", \"pip\", \"install\", \"openslide-python\", \"-q\"])\n#     import openslide\n \n# PYRAMID_LEVEL = 1   # 0=full res, 1=~4x downsampled, 2=~16x downsampled (varies per slide)\n \n# def is_tissue(patch, threshold=0.6):\n#     gray = cv2.cvtColor(patch, cv2.COLOR_RGB2GRAY)\n#     return (gray < 230).mean() > threshold\n \n# def extract_spatial_patches(df, save_dir, grid_size=GRID_SIZE, patches_per_region=PATCHES_PER_REGION):\n#     \"\"\"\n#     Same region/patch logic as before, but reads slides via OpenSlide at a\n#     pre-downsampled pyramid level instead of decoding full resolution with PIL.\n#     This is typically 5-10x faster on large WSIs.\n#     \"\"\"\n#     os.makedirs(save_dir, exist_ok=True)\n#     for _, row in tqdm(df.iterrows(), total=len(df)):\n#         img_id = row[\"image_id\"]\n#         label  = row[\"isup_grade\"]\n#         path   = os.path.join(images_path, img_id + \".tiff\")\n#         try:\n#             slide_obj = openslide.OpenSlide(path)\n#             level = min(PYRAMID_LEVEL, slide_obj.level_count - 1)\n#             w, h = slide_obj.level_dimensions[level]\n#             slide = np.array(slide_obj.read_region((0, 0), level, (w, h)).convert(\"RGB\"))\n#             slide_obj.close()\n#         except Exception:\n#             continue\n \n#         min_dim = int(PATCH_SIZE * 1.5) * grid_size\n#         if h <= min_dim or w <= min_dim:\n#             continue\n \n#         region_h = h // grid_size\n#         region_w = w // grid_size\n \n#         for r_idx in range(grid_size * grid_size):\n#             ry, rx = divmod(r_idx, grid_size)\n#             y0, y1 = ry * region_h, min((ry + 1) * region_h, h)\n#             x0, x1 = rx * region_w, min((rx + 1) * region_w, w)\n#             if (y1 - y0) <= int(PATCH_SIZE * 1.2) or (x1 - x0) <= int(PATCH_SIZE * 1.2):\n#                 continue\n \n#             region_slide = slide[y0:y1, x0:x1]\n#             saved, attempts = 0, 0\n#             while saved < patches_per_region and attempts < patches_per_region * 8:\n#                 yy = np.random.randint(0, max(1, (y1 - y0) - PATCH_SIZE))\n#                 xx = np.random.randint(0, max(1, (x1 - x0) - PATCH_SIZE))\n#                 patch = region_slide[yy:yy+PATCH_SIZE, xx:xx+PATCH_SIZE]\n#                 if patch.shape[0] == PATCH_SIZE and patch.shape[1] == PATCH_SIZE and is_tissue(patch):\n#                     fname = f\"{img_id}_r{r_idx}_p{saved}_{label}.png\"\n#                     cv2.imwrite(os.path.join(save_dir, fname),\n#                                 cv2.cvtColor(patch, cv2.COLOR_RGB2BGR))\n#                     saved += 1\n#                 attempts += 1\n \n# extract_spatial_patches(train_df, SAVE_DIR + \"/train\")\n# extract_spatial_patches(val_df,   SAVE_DIR + \"/val\")\n# extract_spatial_patches(test_df,  SAVE_DIR + \"/test\")\n \n# n_train_patches = len(os.listdir(SAVE_DIR + \"/train\"))\n# n_val_patches   = len(os.listdir(SAVE_DIR + \"/val\"))                    \n# n_test_patches  = len(os.listdir(SAVE_DIR + \"/test\"))\n# print(\"Train patches:\", n_train_patches, f\"(~{n_train_patches/max(1,len(train_df)):.1f} per slide)\")\n# print(\"Val patches  :\", n_val_patches,   f\"(~{n_val_patches/max(1,len(val_df)):.1f} per slide)\")\n# print(\"Test patches :\", n_test_patches,  f\"(~{n_test_patches/max(1,len(test_df)):.1f} per slide)\")\n# print(\"NOTE: if per-slide patch count is well below GRID_SIZE^2 * PATCHES_PER_REGION,\")\n# print(\"regions are still being skipped as too small -> check slide dimension distribution.\")\n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T00:05:27.200253Z","iopub.execute_input":"2026-08-09T00:05:27.200792Z","iopub.status.idle":"2026-08-09T00:56:39.808365Z","shell.execute_reply.started":"2026-08-09T00:05:27.200725Z","shell.execute_reply":"2026-08-09T00:56:39.807685Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# train_transforms = T.Compose([\n#     T.RandomHorizontalFlip(),\n#     T.RandomVerticalFlip(),\n#     T.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.1),\n#     T.ToTensor(),\n#     T.Normalize(mean=MEAN, std=STD),\n# ])\n# eval_transforms = T.Compose([\n#     T.ToTensor(),\n#     T.Normalize(mean=MEAN, std=STD),\n# ])\n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T01:03:13.811949Z","iopub.execute_input":"2026-08-09T01:03:13.81236Z","iopub.status.idle":"2026-08-09T01:03:13.817405Z","shell.execute_reply.started":"2026-08-09T01:03:13.812331Z","shell.execute_reply":"2026-08-09T01:03:13.816497Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# class HierSlideDataset(Dataset):\n#     \"\"\"\n#     Each item = one slide, structured as a list of regions,\n#     each region being a list of patch file paths.\n#     Returns: (list[Tensor(N_r, 3, 224, 224)] per region, label)\n#     \"\"\"\n#     def __init__(self, patch_dir, transform=None):\n#         self.transform = transform\n#         # slide_id -> region_id -> [patch paths]\n#         self.slide_data = {}\n#         for f in os.listdir(patch_dir):\n#             if not f.endswith(\".png\"):\n#                 continue\n#             stem = f.replace(\".png\", \"\")\n#             # format: imgid_r{r}_p{p}_{label}\n#             parts = stem.split(\"_\")\n#             label = int(parts[-1])\n#             p_tag = parts[-2]      # p{p}\n#             r_tag = parts[-3]      # r{r}\n#             img_id = \"_\".join(parts[:-3])\n#             region_id = int(r_tag[1:])\n \n#             self.slide_data.setdefault(img_id, {\"label\": label, \"regions\": {}})\n#             self.slide_data[img_id][\"regions\"].setdefault(region_id, [])\n#             self.slide_data[img_id][\"regions\"][region_id].append(os.path.join(patch_dir, f))\n \n#         self.slides = list(self.slide_data.keys())\n \n#     def __len__(self):\n#         return len(self.slides)\n \n#     def __getitem__(self, idx):\n#         sid = self.slides[idx]\n#         info = self.slide_data[sid]\n#         label = info[\"label\"]\n \n#         region_tensors = []\n#         for region_id, paths in info[\"regions\"].items():\n#             patch_tensors = []\n#             for p in paths:\n#                 img = Image.open(p).convert(\"RGB\")\n#                 if self.transform:\n#                     img = self.transform(img)\n#                 patch_tensors.append(img)\n#             region_tensors.append(torch.stack(patch_tensors, dim=0))  # (N_r, 3, 224, 224)\n \n#         return region_tensors, label\n \n# def collate_single(batch):\n#     # batch_size=1 always, just unwrap\n#     region_tensors, label = batch[0]\n#     return region_tensors, label\n \n# train_ds = HierSlideDataset(SAVE_DIR + \"/train\", train_transforms)\n# val_ds   = HierSlideDataset(SAVE_DIR + \"/val\",   eval_transforms)\n# test_ds  = HierSlideDataset(SAVE_DIR + \"/test\",  eval_transforms)\n \n# train_loader = DataLoader(train_ds, batch_size=1, shuffle=True,  num_workers=2, collate_fn=collate_single)\n# val_loader   = DataLoader(val_ds,   batch_size=1, shuffle=False, num_workers=2, collate_fn=collate_single)\n# test_loader  = DataLoader(test_ds,  batch_size=1, shuffle=False, num_workers=2, collate_fn=collate_single)\n \n# print(\"Train slides:\", len(train_ds), \"| Val slides:\", len(val_ds), \"| Test slides:\", len(test_ds))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T01:03:37.108514Z","iopub.execute_input":"2026-08-09T01:03:37.109132Z","iopub.status.idle":"2026-08-09T01:03:37.17519Z","shell.execute_reply.started":"2026-08-09T01:03:37.109102Z","shell.execute_reply":"2026-08-09T01:03:37.174403Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# class AttentionPool(nn.Module):\n#     def __init__(self, feat_dim=768, hidden_dim=256):\n#         super().__init__()\n#         self.attention = nn.Sequential(\n#             nn.Linear(feat_dim, hidden_dim),\n#             nn.Tanh(),\n#             nn.Linear(hidden_dim, 1)\n#         )\n#     def forward(self, x):\n#         # x: (N, feat_dim)\n#         attn_weights = self.attention(x)\n#         attn_weights = torch.softmax(attn_weights, dim=0)\n#         pooled = (attn_weights * x).sum(dim=0)\n#         return pooled, attn_weights","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T01:04:47.208485Z","iopub.execute_input":"2026-08-09T01:04:47.208911Z","iopub.status.idle":"2026-08-09T01:04:47.21386Z","shell.execute_reply.started":"2026-08-09T01:04:47.208883Z","shell.execute_reply":"2026-08-09T01:04:47.21319Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# class HierarchicalSwin2Level(nn.Module):\n#     \"\"\"\n#     Level 1 (patch):  Swin-Tiny encodes patches -> patch features (768-d)\n#     Level 2 (region): AttentionPool over patches within each region -> region feature\n#     Level 3 (slide):  AttentionPool over region features -> slide feature\n#     Classifier: slide feature -> ISUP logits\n#     \"\"\"\n#     def __init__(self, num_classes=6, freeze_encoder=False):\n#         super().__init__()\n#         self.patch_encoder = timm.create_model(\n#             \"swin_tiny_patch4_window7_224\",\n#             pretrained=True,\n#             num_classes=0\n#         )\n#         if freeze_encoder:\n#             for p in self.patch_encoder.parameters():\n#                 p.requires_grad = False\n \n#         self.region_pool = AttentionPool(feat_dim=768, hidden_dim=256)\n#         self.slide_pool  = AttentionPool(feat_dim=768, hidden_dim=256)\n#         self.drop        = nn.Dropout(0.4)          # increased from 0.3\n#         self.classifier  = nn.Linear(768, num_classes)\n \n#     def forward(self, region_tensors):\n#         \"\"\"\n#         region_tensors: list of Tensors, each (N_r, 3, 224, 224), one per region\n#         \"\"\"\n#         region_feats = []\n#         region_attn  = []\n#         for patches in region_tensors:\n#             patches = patches.to(next(self.parameters()).device)\n#             patch_feats = self.patch_encoder(patches)         # (N_r, 768)\n#             region_feat, attn_w = self.region_pool(patch_feats)\n#             region_feats.append(region_feat)\n#             region_attn.append(attn_w)\n \n#         region_feats = torch.stack(region_feats, dim=0)       # (R, 768)\n#         slide_feat, slide_attn = self.slide_pool(region_feats)\n \n#         logits = self.classifier(self.drop(slide_feat))\n#         return logits, slide_attn, region_attn","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T01:05:15.072746Z","iopub.execute_input":"2026-08-09T01:05:15.073543Z","iopub.status.idle":"2026-08-09T01:05:15.081293Z","shell.execute_reply.started":"2026-08-09T01:05:15.073507Z","shell.execute_reply":"2026-08-09T01:05:15.080331Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CKPT_PATH = \"/kaggle/working/hier_swin_2level_checkpoint.pth\"   # full resumable state\n# BEST_PATH = \"/kaggle/working/hier_swin_2level_best.pth\"          # best weights only\n \n# model     = HierarchicalSwin2Level(freeze_encoder=False).to(device)\n# criterion = nn.CrossEntropyLoss()\n# optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\n# scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)\n \n# start_epoch = 0\n# history = {\"train_loss\": [], \"train_qwk\": [], \"val_loss\": [], \"val_qwk\": []}\n# best_kappa = -1\n# epochs_no_improve = 0\n \n# # If a checkpoint exists from a previous (crashed/interrupted) session, resume from it.\n# # On Kaggle, /kaggle/working persists within a session but NOT automatically across\n# # a hard crash unless you've saved this file to a Dataset — see notes below the loop.\n# if os.path.exists(CKPT_PATH):\n#     ckpt = torch.load(CKPT_PATH, map_location=device)\n#     model.load_state_dict(ckpt[\"model_state\"])\n#     optimizer.load_state_dict(ckpt[\"optimizer_state\"])\n#     scheduler.load_state_dict(ckpt[\"scheduler_state\"])\n#     start_epoch = ckpt[\"epoch\"] + 1\n#     history = ckpt[\"history\"]\n#     best_kappa = ckpt[\"best_kappa\"]\n#     epochs_no_improve = ckpt[\"epochs_no_improve\"]\n#     print(f\"Resumed from checkpoint at epoch {start_epoch} | best_kappa so far: {best_kappa:.4f}\")\n# else:\n#     print(\"No checkpoint found — starting fresh.\")\n \n# def run_epoch(loader, train=True):\n#     model.train() if train else model.eval()\n#     total_loss, all_preds, all_labels = 0, [], []\n#     ctx = torch.enable_grad() if train else torch.no_grad()\n#     with ctx:\n#         for region_tensors, label in loader:\n#             label_t = torch.tensor([label], device=device)\n#             if train:\n#                 optimizer.zero_grad()\n#             logits, _, _ = model(region_tensors)\n#             logits = logits.unsqueeze(0)\n#             loss = criterion(logits, label_t)\n#             if train:\n#                 loss.backward()\n#                 optimizer.step()\n#             total_loss += loss.item()\n#             all_preds  += torch.argmax(logits, 1).cpu().tolist()\n#             all_labels += [label]\n#     kappa = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\", labels=list(range(6)))\n#     return total_loss / max(1, len(loader)), kappa\n \n# for epoch in range(start_epoch, EPOCHS):\n#     tr_loss, tr_kappa = run_epoch(train_loader, train=True)\n#     vl_loss, vl_kappa = run_epoch(val_loader,   train=False)\n#     scheduler.step()\n \n#     history[\"train_loss\"].append(tr_loss)\n#     history[\"train_qwk\"].append(tr_kappa)\n#     history[\"val_loss\"].append(vl_loss)\n#     history[\"val_qwk\"].append(vl_kappa)\n \n#     print(f\"Epoch {epoch+1}/{EPOCHS} | \"\n#           f\"Train Loss {tr_loss:.4f} QWK {tr_kappa:.4f} | \"\n#           f\"Val Loss {vl_loss:.4f} QWK {vl_kappa:.4f}\")\n \n#     if vl_kappa > best_kappa:\n#         best_kappa = vl_kappa\n#         epochs_no_improve = 0\n#         torch.save(model.state_dict(), BEST_PATH)\n#         print(\"  Best model saved\")\n#     else:\n#         epochs_no_improve += 1\n \n#     # Save full resumable checkpoint EVERY epoch, regardless of improvement,\n#     # so a crash never costs you more than one epoch of work.\n#     torch.save({\n#         \"epoch\": epoch,\n#         \"model_state\": model.state_dict(),\n#         \"optimizer_state\": optimizer.state_dict(),\n#         \"scheduler_state\": scheduler.state_dict(),\n#         \"history\": history,\n#         \"best_kappa\": best_kappa,\n#         \"epochs_no_improve\": epochs_no_improve,\n#     }, CKPT_PATH)\n \n#     if epochs_no_improve >= PATIENCE:\n#         print(f\"  Early stopping triggered at epoch {epoch+1}\")\n#         break\n \n# print(\"\\nBest Val QWK:\", best_kappa)\n \n# # IMPORTANT: after this cell finishes (or if you're about to lose your session),\n# # go to your notebook's Output pane and click \"New Dataset\" on /kaggle/working\n# # to persist hier_swin_2level_checkpoint.pth and hier_swin_2level_best.pth.\n# # Next session: add that dataset as input, copy the .pth files back into\n# # /kaggle/working before re-running this cell, and it will auto-resume.","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T01:06:21.095744Z","iopub.execute_input":"2026-08-09T01:06:21.096083Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.load_state_dict(torch.load(BEST_PATH))\ntest_loss, test_kappa = run_epoch(test_loader, train=False)\nprint(f\"\\nFINAL TEST RESULTS -> Loss: {test_loss:.4f} | QWK: {test_kappa:.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # ==============================================================================\n# # TWO-LEVEL HIERARCHICAL SWIN MODEL — PANDA ISUP GRADING\n# # Patch -> Region (spatial) -> Slide\n# # Built on top of your baseline-swin-model notebook. Run cell-by-cell on Kaggle.\n# # ==============================================================================\n\n# import os, numpy as np, pandas as pd\n# from PIL import Image\n# import cv2, torch, torch.nn as nn, torch.nn.functional as F\n# from torch.utils.data import Dataset, DataLoader\n# from sklearn.model_selection import train_test_split\n# from sklearn.metrics import cohen_kappa_score\n# import timm\n# from tqdm import tqdm\n# import torchvision.transforms as T\n\n# Image.MAX_IMAGE_PIXELS = None\n# device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n# print(\"Device:\", device)\n\n# # ── CELL 1: Config ───────────────────────────────────────────────────────────\n# PATCH_SIZE     = 224\n# DOWNSAMPLE     = 4          # was 8 (reduce(8)) - halving this keeps regions big enough to hold patches\n# GRID_SIZE      = 3          # was 4 -> 3x3 = 9 regions, each region stays well above PATCH_SIZE\n# PATCHES_PER_REGION = 5      # patches sampled per region (tissue-permitting)\n# NUM_SLIDES     = 1000       # back to full run — extraction speed & density confirmed on 50-slide test\n# VAL_SLIDES     = 200\n# TEST_SLIDES    = 200\n# BATCH_SIZE     = 1          # 1 slide per batch (variable region/patch counts)\n# EPOCHS         = 15\n# LR             = 3e-5\n# WEIGHT_DECAY   = 5e-4       # increased from 1e-4 to fight overfitting\n# PATIENCE       = 5          # early stopping patience (epochs without val improvement)\n# SAVE_DIR       = \"/kaggle/working/regions\"\n# MEAN = [0.485, 0.456, 0.406]\n# STD  = [0.229, 0.224, 0.225]\n\n# train_csv   = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train.csv\"\n# images_path = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train_images\"\n\n# # ── CELL 2: Three-way split (train/val/test — your baseline only had train/val) ──\n# df = pd.read_csv(train_csv)\n\n# train_df, temp_df = train_test_split(\n#     df, test_size=0.3, stratify=df[\"isup_grade\"], random_state=42\n# )\n# val_df, test_df = train_test_split(\n#     temp_df, test_size=0.5, stratify=temp_df[\"isup_grade\"], random_state=42\n# )\n\n# train_df = train_df.sample(min(NUM_SLIDES, len(train_df)), random_state=42).reset_index(drop=True)\n# val_df   = val_df.sample(min(VAL_SLIDES, len(val_df)), random_state=42).reset_index(drop=True)\n# test_df  = test_df.sample(min(TEST_SLIDES, len(test_df)), random_state=42).reset_index(drop=True)\n\n# print(f\"Train slides: {len(train_df)} | Val slides: {len(val_df)} | Test slides: {len(test_df)}\")\n\n# # ── CELL 3: Tissue-aware SPATIAL region + patch extraction (OpenSlide, fast) ──\n# # Uses OpenSlide to read directly from a lower pyramid level instead of decoding\n# # the full-resolution TIFF and downsampling in memory (which is what PIL.reduce()\n# # was doing — and why extraction took ~58 min for 1000 slides).\n# try:\n#     import openslide\n# except ImportError:\n#     import subprocess, sys\n#     subprocess.run([sys.executable, \"-m\", \"pip\", \"install\", \"openslide-python\", \"-q\"])\n#     import openslide\n\n# PYRAMID_LEVEL = 1   # 0=full res, 1=~4x downsampled, 2=~16x downsampled (varies per slide)\n\n# def is_tissue(patch, threshold=0.6):\n#     gray = cv2.cvtColor(patch, cv2.COLOR_RGB2GRAY)\n#     return (gray < 230).mean() > threshold\n\n# def extract_spatial_patches(df, save_dir, grid_size=GRID_SIZE, patches_per_region=PATCHES_PER_REGION):\n#     \"\"\"\n#     Same region/patch logic as before, but reads slides via OpenSlide at a\n#     pre-downsampled pyramid level instead of decoding full resolution with PIL.\n#     This is typically 5-10x faster on large WSIs.\n#     \"\"\"\n#     os.makedirs(save_dir, exist_ok=True)\n#     for _, row in tqdm(df.iterrows(), total=len(df)):\n#         img_id = row[\"image_id\"]\n#         label  = row[\"isup_grade\"]\n#         path   = os.path.join(images_path, img_id + \".tiff\")\n#         try:\n#             slide_obj = openslide.OpenSlide(path)\n#             level = min(PYRAMID_LEVEL, slide_obj.level_count - 1)\n#             w, h = slide_obj.level_dimensions[level]\n#             slide = np.array(slide_obj.read_region((0, 0), level, (w, h)).convert(\"RGB\"))\n#             slide_obj.close()\n#         except Exception:\n#             continue\n\n#         min_dim = int(PATCH_SIZE * 1.5) * grid_size\n#         if h <= min_dim or w <= min_dim:\n#             continue\n\n#         region_h = h // grid_size\n#         region_w = w // grid_size\n\n#         for r_idx in range(grid_size * grid_size):\n#             ry, rx = divmod(r_idx, grid_size)\n#             y0, y1 = ry * region_h, min((ry + 1) * region_h, h)\n#             x0, x1 = rx * region_w, min((rx + 1) * region_w, w)\n#             if (y1 - y0) <= int(PATCH_SIZE * 1.2) or (x1 - x0) <= int(PATCH_SIZE * 1.2):\n#                 continue\n\n#             region_slide = slide[y0:y1, x0:x1]\n#             saved, attempts = 0, 0\n#             while saved < patches_per_region and attempts < patches_per_region * 8:\n#                 yy = np.random.randint(0, max(1, (y1 - y0) - PATCH_SIZE))\n#                 xx = np.random.randint(0, max(1, (x1 - x0) - PATCH_SIZE))\n#                 patch = region_slide[yy:yy+PATCH_SIZE, xx:xx+PATCH_SIZE]\n#                 if patch.shape[0] == PATCH_SIZE and patch.shape[1] == PATCH_SIZE and is_tissue(patch):\n#                     fname = f\"{img_id}_r{r_idx}_p{saved}_{label}.png\"\n#                     cv2.imwrite(os.path.join(save_dir, fname),\n#                                 cv2.cvtColor(patch, cv2.COLOR_RGB2BGR))\n#                     saved += 1\n#                 attempts += 1\n\n# extract_spatial_patches(train_df, SAVE_DIR + \"/train\")\n# extract_spatial_patches(val_df,   SAVE_DIR + \"/val\")\n# extract_spatial_patches(test_df,  SAVE_DIR + \"/test\")\n\n# n_train_patches = len(os.listdir(SAVE_DIR + \"/train\"))\n# n_val_patches   = len(os.listdir(SAVE_DIR + \"/val\"))\n# n_test_patches  = len(os.listdir(SAVE_DIR + \"/test\"))\n# print(\"Train patches:\", n_train_patches, f\"(~{n_train_patches/max(1,len(train_df)):.1f} per slide)\")\n# print(\"Val patches  :\", n_val_patches,   f\"(~{n_val_patches/max(1,len(val_df)):.1f} per slide)\")\n# print(\"Test patches :\", n_test_patches,  f\"(~{n_test_patches/max(1,len(test_df)):.1f} per slide)\")\n# print(\"NOTE: if per-slide patch count is well below GRID_SIZE^2 * PATCHES_PER_REGION,\")\n# print(\"regions are still being skipped as too small -> check slide dimension distribution.\")\n\n# # ── CELL 4: Transforms ───────────────────────────────────────────────────────\n# train_transforms = T.Compose([\n#     T.RandomHorizontalFlip(),\n#     T.RandomVerticalFlip(),\n#     T.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.1),\n#     T.ToTensor(),\n#     T.Normalize(mean=MEAN, std=STD),\n# ])\n# eval_transforms = T.Compose([\n#     T.ToTensor(),\n#     T.Normalize(mean=MEAN, std=STD),\n# ])\n\n# # ── CELL 5: Slide dataset with REGION structure ──────────────────────────────\n# class HierSlideDataset(Dataset):\n#     \"\"\"\n#     Each item = one slide, structured as a list of regions,\n#     each region being a list of patch file paths.\n#     Returns: (list[Tensor(N_r, 3, 224, 224)] per region, label)\n#     \"\"\"\n#     def __init__(self, patch_dir, transform=None):\n#         self.transform = transform\n#         # slide_id -> region_id -> [patch paths]\n#         self.slide_data = {}\n#         for f in os.listdir(patch_dir):\n#             if not f.endswith(\".png\"):\n#                 continue\n#             stem = f.replace(\".png\", \"\")\n#             # format: imgid_r{r}_p{p}_{label}\n#             parts = stem.split(\"_\")\n#             label = int(parts[-1])\n#             p_tag = parts[-2]      # p{p}\n#             r_tag = parts[-3]      # r{r}\n#             img_id = \"_\".join(parts[:-3])\n#             region_id = int(r_tag[1:])\n\n#             self.slide_data.setdefault(img_id, {\"label\": label, \"regions\": {}})\n#             self.slide_data[img_id][\"regions\"].setdefault(region_id, [])\n#             self.slide_data[img_id][\"regions\"][region_id].append(os.path.join(patch_dir, f))\n\n#         self.slides = list(self.slide_data.keys())\n\n#     def __len__(self):\n#         return len(self.slides)\n\n#     def __getitem__(self, idx):\n#         sid = self.slides[idx]\n#         info = self.slide_data[sid]\n#         label = info[\"label\"]\n\n#         region_tensors = []\n#         for region_id, paths in info[\"regions\"].items():\n#             patch_tensors = []\n#             for p in paths:\n#                 img = Image.open(p).convert(\"RGB\")\n#                 if self.transform:\n#                     img = self.transform(img)\n#                 patch_tensors.append(img)\n#             region_tensors.append(torch.stack(patch_tensors, dim=0))  # (N_r, 3, 224, 224)\n\n#         return region_tensors, label\n\n# def collate_single(batch):\n#     # batch_size=1 always, just unwrap\n#     region_tensors, label = batch[0]\n#     return region_tensors, label\n\n# train_ds = HierSlideDataset(SAVE_DIR + \"/train\", train_transforms)\n# val_ds   = HierSlideDataset(SAVE_DIR + \"/val\",   eval_transforms)\n# test_ds  = HierSlideDataset(SAVE_DIR + \"/test\",  eval_transforms)\n\n# train_loader = DataLoader(train_ds, batch_size=1, shuffle=True,  num_workers=2, collate_fn=collate_single)\n# val_loader   = DataLoader(val_ds,   batch_size=1, shuffle=False, num_workers=2, collate_fn=collate_single)\n# test_loader  = DataLoader(test_ds,  batch_size=1, shuffle=False, num_workers=2, collate_fn=collate_single)\n\n# print(\"Train slides:\", len(train_ds), \"| Val slides:\", len(val_ds), \"| Test slides:\", len(test_ds))\n\n# # ── CELL 6: Attention pooling (reused at both levels) ────────────────────────\n# class AttentionPool(nn.Module):\n#     def __init__(self, feat_dim=768, hidden_dim=256):\n#         super().__init__()\n#         self.attention = nn.Sequential(\n#             nn.Linear(feat_dim, hidden_dim),\n#             nn.Tanh(),\n#             nn.Linear(hidden_dim, 1)\n#         )\n#     def forward(self, x):\n#         # x: (N, feat_dim)\n#         attn_weights = self.attention(x)\n#         attn_weights = torch.softmax(attn_weights, dim=0)\n#         pooled = (attn_weights * x).sum(dim=0)\n#         return pooled, attn_weights\n\n# # ── CELL 7: Two-level hierarchical model ─────────────────────────────────────\n# class HierarchicalSwin2Level(nn.Module):\n#     \"\"\"\n#     Level 1 (patch):  Swin-Tiny encodes patches -> patch features (768-d)\n#     Level 2 (region): AttentionPool over patches within each region -> region feature\n#     Level 3 (slide):  AttentionPool over region features -> slide feature\n#     Classifier: slide feature -> ISUP logits\n#     \"\"\"\n#     def __init__(self, num_classes=6, freeze_encoder=False):\n#         super().__init__()\n#         self.patch_encoder = timm.create_model(\n#             \"swin_tiny_patch4_window7_224\",\n#             pretrained=True,\n#             num_classes=0\n#         )\n#         if freeze_encoder:\n#             for p in self.patch_encoder.parameters():\n#                 p.requires_grad = False\n\n#         self.region_pool = AttentionPool(feat_dim=768, hidden_dim=256)\n#         self.slide_pool  = AttentionPool(feat_dim=768, hidden_dim=256)\n#         self.drop        = nn.Dropout(0.4)          # increased from 0.3\n#         self.classifier  = nn.Linear(768, num_classes)\n\n#     def forward(self, region_tensors):\n#         \"\"\"\n#         region_tensors: list of Tensors, each (N_r, 3, 224, 224), one per region\n#         \"\"\"\n#         region_feats = []\n#         region_attn  = []\n#         for patches in region_tensors:\n#             patches = patches.to(next(self.parameters()).device)\n#             patch_feats = self.patch_encoder(patches)         # (N_r, 768)\n#             region_feat, attn_w = self.region_pool(patch_feats)\n#             region_feats.append(region_feat)\n#             region_attn.append(attn_w)\n\n#         region_feats = torch.stack(region_feats, dim=0)       # (R, 768)\n#         slide_feat, slide_attn = self.slide_pool(region_feats)\n\n#         logits = self.classifier(self.drop(slide_feat))\n#         return logits, slide_attn, region_attn\n\n# # ── CELL 8: Training loop with early stopping + checkpoint-resume ────────────\n# CKPT_PATH = \"/kaggle/working/hier_swin_2level_checkpoint.pth\"   # full resumable state\n# BEST_PATH = \"/kaggle/working/hier_swin_2level_best.pth\"          # best weights only\n\n# model     = HierarchicalSwin2Level(freeze_encoder=False).to(device)\n# criterion = nn.CrossEntropyLoss()\n# optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\n# scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)\n\n# start_epoch = 0\n# history = {\"train_loss\": [], \"train_qwk\": [], \"val_loss\": [], \"val_qwk\": []}\n# best_kappa = -1\n# epochs_no_improve = 0\n\n# # If a checkpoint exists from a previous (crashed/interrupted) session, resume from it.\n# # On Kaggle, /kaggle/working persists within a session but NOT automatically across\n# # a hard crash unless you've saved this file to a Dataset — see notes below the loop.\n# if os.path.exists(CKPT_PATH):\n#     ckpt = torch.load(CKPT_PATH, map_location=device)\n#     model.load_state_dict(ckpt[\"model_state\"])\n#     optimizer.load_state_dict(ckpt[\"optimizer_state\"])\n#     scheduler.load_state_dict(ckpt[\"scheduler_state\"])\n#     start_epoch = ckpt[\"epoch\"] + 1\n#     history = ckpt[\"history\"]\n#     best_kappa = ckpt[\"best_kappa\"]\n#     epochs_no_improve = ckpt[\"epochs_no_improve\"]\n#     print(f\"Resumed from checkpoint at epoch {start_epoch} | best_kappa so far: {best_kappa:.4f}\")\n# else:\n#     print(\"No checkpoint found — starting fresh.\")\n\n# def run_epoch(loader, train=True):\n#     model.train() if train else model.eval()\n#     total_loss, all_preds, all_labels = 0, [], []\n#     ctx = torch.enable_grad() if train else torch.no_grad()\n#     with ctx:\n#         for region_tensors, label in loader:\n#             label_t = torch.tensor([label], device=device)\n#             if train:\n#                 optimizer.zero_grad()\n#             logits, _, _ = model(region_tensors)\n#             logits = logits.unsqueeze(0)\n#             loss = criterion(logits, label_t)\n#             if train:\n#                 loss.backward()\n#                 optimizer.step()\n#             total_loss += loss.item()\n#             all_preds  += torch.argmax(logits, 1).cpu().tolist()\n#             all_labels += [label]\n#     kappa = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\", labels=list(range(6)))\n#     return total_loss / max(1, len(loader)), kappa\n\n# for epoch in range(start_epoch, EPOCHS):\n#     tr_loss, tr_kappa = run_epoch(train_loader, train=True)\n#     vl_loss, vl_kappa = run_epoch(val_loader,   train=False)\n#     scheduler.step()\n\n#     history[\"train_loss\"].append(tr_loss)\n#     history[\"train_qwk\"].append(tr_kappa)\n#     history[\"val_loss\"].append(vl_loss)\n#     history[\"val_qwk\"].append(vl_kappa)\n\n#     print(f\"Epoch {epoch+1}/{EPOCHS} | \"\n#           f\"Train Loss {tr_loss:.4f} QWK {tr_kappa:.4f} | \"\n#           f\"Val Loss {vl_loss:.4f} QWK {vl_kappa:.4f}\")\n\n#     if vl_kappa > best_kappa:\n#         best_kappa = vl_kappa\n#         epochs_no_improve = 0\n#         torch.save(model.state_dict(), BEST_PATH)\n#         print(\"  Best model saved\")\n#     else:\n#         epochs_no_improve += 1\n\n#     # Save full resumable checkpoint EVERY epoch, regardless of improvement,\n#     # so a crash never costs you more than one epoch of work.\n#     torch.save({\n#         \"epoch\": epoch,\n#         \"model_state\": model.state_dict(),\n#         \"optimizer_state\": optimizer.state_dict(),\n#         \"scheduler_state\": scheduler.state_dict(),\n#         \"history\": history,\n#         \"best_kappa\": best_kappa,\n#         \"epochs_no_improve\": epochs_no_improve,\n#     }, CKPT_PATH)\n\n#     if epochs_no_improve >= PATIENCE:\n#         print(f\"  Early stopping triggered at epoch {epoch+1}\")\n#         break\n\n# print(\"\\nBest Val QWK:\", best_kappa)\n\n# # IMPORTANT: after this cell finishes (or if you're about to lose your session),\n# # go to your notebook's Output pane and click \"New Dataset\" on /kaggle/working\n# # to persist hier_swin_2level_checkpoint.pth and hier_swin_2level_best.pth.\n# # Next session: add that dataset as input, copy the .pth files back into\n# # /kaggle/working before re-running this cell, and it will auto-resume.\n\n# # ── CELL 9: Final TEST evaluation (held out, never touched during training) ──\n# model.load_state_dict(torch.load(BEST_PATH))\n# test_loss, test_kappa = run_epoch(test_loader, train=False)\n# print(f\"\\nFINAL TEST RESULTS -> Loss: {test_loss:.4f} | QWK: {test_kappa:.4f}\")\n\n# # ── CELL 10: Plot train/val curves for your professor ────────────────────────\n# import matplotlib.pyplot as plt\n# epochs_ran = range(1, len(history[\"train_qwk\"]) + 1)\n# plt.figure(figsize=(7,5))\n# plt.plot(epochs_ran, history[\"train_qwk\"], \"b-o\", label=\"Train QWK\")\n# plt.plot(epochs_ran, history[\"val_qwk\"], \"r-o\", label=\"Val QWK\")\n# plt.axhline(best_kappa, color=\"g\", linestyle=\"--\", label=f\"Best Val QWK = {best_kappa:.3f}\")\n# plt.axhline(test_kappa, color=\"purple\", linestyle=\":\", label=f\"Test QWK = {test_kappa:.3f}\")\n# plt.xlabel(\"Epoch\")\n# plt.ylabel(\"Quadratic Weighted Kappa\")\n# plt.title(f\"2-Level Hierarchical Swin-Tiny on PANDA ({NUM_SLIDES} slides, {GRID_SIZE}x{GRID_SIZE} regions)\")\n# plt.legend()\n# plt.grid(True)\n# plt.savefig(\"/kaggle/working/hier_swin_2level_curves.png\", dpi=150, bbox_inches=\"tight\")\n# plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T01:16:33.882974Z","iopub.execute_input":"2026-08-09T01:16:33.883383Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#NEW FINE TUNED MODEL ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T10:18:22.023762Z","iopub.execute_input":"2026-08-09T10:18:22.024041Z","iopub.status.idle":"2026-08-09T10:18:22.028166Z","shell.execute_reply.started":"2026-08-09T10:18:22.023994Z","shell.execute_reply":"2026-08-09T10:18:22.027516Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, numpy as np, pandas as pd\nfrom PIL import Image\nimport cv2, torch, torch.nn as nn, torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import cohen_kappa_score\nimport timm\nfrom tqdm import tqdm\nimport torchvision.transforms as T\n \nImage.MAX_IMAGE_PIXELS = None\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Device:\", device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T10:18:43.146033Z","iopub.execute_input":"2026-08-09T10:18:43.146787Z","iopub.status.idle":"2026-08-09T10:19:03.302156Z","shell.execute_reply.started":"2026-08-09T10:18:43.146756Z","shell.execute_reply":"2026-08-09T10:19:03.301426Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PATCH_SIZE     = 224\nDOWNSAMPLE     = 4          # was 8 (reduce(8)) - halving this keeps regions big enough to hold patches\nGRID_SIZE      = 3          # was 4 -> 3x3 = 9 regions, each region stays well above PATCH_SIZE\nPATCHES_PER_REGION = 5      # patches sampled per region (tissue-permitting)\nNUM_SLIDES     = 1000       # back to full run — extraction speed & density confirmed on 50-slide test\nVAL_SLIDES     = 200\nTEST_SLIDES    = 200\nBATCH_SIZE     = 1          # 1 slide per batch (variable region/patch counts)\nEPOCHS         = 15\nLR             = 3e-5\nWEIGHT_DECAY   = 5e-4       # increased from 1e-4 to fight overfitting\nPATIENCE       = 5          # early stopping patience (epochs without val improvement)\nACCUM_STEPS    = 8          # gradient accumulation — smooths noisy single-slide updates\nUNFREEZE_LAST_N_LAYERS = 2  # only fine-tune last N Swin-Tiny stages; rest stay frozen\nSAVE_DIR       = \"/kaggle/working/regions\"\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n \ntrain_csv   = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train.csv\"\nimages_path = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train_images\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T10:21:32.774026Z","iopub.execute_input":"2026-08-09T10:21:32.775096Z","iopub.status.idle":"2026-08-09T10:21:32.780284Z","shell.execute_reply.started":"2026-08-09T10:21:32.775061Z","shell.execute_reply":"2026-08-09T10:21:32.779433Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(train_csv)\n \ntrain_df, temp_df = train_test_split(\n    df, test_size=0.3, stratify=df[\"isup_grade\"], random_state=42\n)\nval_df, test_df = train_test_split(\n    temp_df, test_size=0.5, stratify=temp_df[\"isup_grade\"], random_state=42\n)\n \ntrain_df = train_df.sample(min(NUM_SLIDES, len(train_df)), random_state=42).reset_index(drop=True)\nval_df   = val_df.sample(min(VAL_SLIDES, len(val_df)), random_state=42).reset_index(drop=True)\ntest_df  = test_df.sample(min(TEST_SLIDES, len(test_df)), random_state=42).reset_index(drop=True)\n \nprint(f\"Train slides: {len(train_df)} | Val slides: {len(val_df)} | Test slides: {len(test_df)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T10:21:52.114803Z","iopub.execute_input":"2026-08-09T10:21:52.115211Z","iopub.status.idle":"2026-08-09T10:21:52.191499Z","shell.execute_reply.started":"2026-08-09T10:21:52.115184Z","shell.execute_reply":"2026-08-09T10:21:52.190866Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"try:\n    import openslide\nexcept ImportError:\n    import subprocess, sys\n    subprocess.run([sys.executable, \"-m\", \"pip\", \"install\", \"openslide-python\", \"-q\"])\n    import openslide\n \nPYRAMID_LEVEL = 1   # 0=full res, 1=~4x downsampled, 2=~16x downsampled (varies per slide)\n \ndef is_tissue(patch, threshold=0.6):\n    gray = cv2.cvtColor(patch, cv2.COLOR_RGB2GRAY)\n    return (gray < 230).mean() > threshold\n \ndef extract_spatial_patches(df, save_dir, grid_size=GRID_SIZE, patches_per_region=PATCHES_PER_REGION):\n    \"\"\"\n    Same region/patch logic as before, but reads slides via OpenSlide at a\n    pre-downsampled pyramid level instead of decoding full resolution with PIL.\n    This is typically 5-10x faster on large WSIs.\n    \"\"\"\n    os.makedirs(save_dir, exist_ok=True)\n    for _, row in tqdm(df.iterrows(), total=len(df)):\n        img_id = row[\"image_id\"]\n        label  = row[\"isup_grade\"]\n        path   = os.path.join(images_path, img_id + \".tiff\")\n        try:\n            slide_obj = openslide.OpenSlide(path)\n            level = min(PYRAMID_LEVEL, slide_obj.level_count - 1)\n            w, h = slide_obj.level_dimensions[level]\n            slide = np.array(slide_obj.read_region((0, 0), level, (w, h)).convert(\"RGB\"))\n            slide_obj.close()\n        except Exception:\n            continue\n \n        min_dim = int(PATCH_SIZE * 1.5) * grid_size\n        if h <= min_dim or w <= min_dim:\n            continue\n \n        region_h = h // grid_size\n        region_w = w // grid_size\n \n        for r_idx in range(grid_size * grid_size):\n            ry, rx = divmod(r_idx, grid_size)\n            y0, y1 = ry * region_h, min((ry + 1) * region_h, h)\n            x0, x1 = rx * region_w, min((rx + 1) * region_w, w)\n            if (y1 - y0) <= int(PATCH_SIZE * 1.2) or (x1 - x0) <= int(PATCH_SIZE * 1.2):\n                continue\n \n            region_slide = slide[y0:y1, x0:x1]\n            saved, attempts = 0, 0\n            while saved < patches_per_region and attempts < patches_per_region * 8:\n                yy = np.random.randint(0, max(1, (y1 - y0) - PATCH_SIZE))\n                xx = np.random.randint(0, max(1, (x1 - x0) - PATCH_SIZE))\n                patch = region_slide[yy:yy+PATCH_SIZE, xx:xx+PATCH_SIZE]\n                if patch.shape[0] == PATCH_SIZE and patch.shape[1] == PATCH_SIZE and is_tissue(patch):\n                    fname = f\"{img_id}_r{r_idx}_p{saved}_{label}.png\"\n                    cv2.imwrite(os.path.join(save_dir, fname),\n                                cv2.cvtColor(patch, cv2.COLOR_RGB2BGR))\n                    saved += 1\n                attempts += 1\n \nextract_spatial_patches(train_df, SAVE_DIR + \"/train\")    \nextract_spatial_patches(val_df,   SAVE_DIR + \"/val\")  \nextract_spatial_patches(test_df,  SAVE_DIR + \"/test\")   \n        \nn_train_patches = len(os.listdir(SAVE_DIR + \"/train\"))      \nn_val_patches   = len(os.listdir(SAVE_DIR + \"/val\"))       \nn_test_patches  = len(os.listdir(SAVE_DIR + \"/test\"))     \nprint(\"Train patches:\", n_train_patches, f\"(~{n_train_patches/max(1,len(train_df)):.1f} per slide)\")\nprint(\"Val patches  :\", n_val_patches,   f\"(~{n_val_patches/max(1,len(val_df)):.1f} per slide)\")\nprint(\"Test patches :\", n_test_patches,  f\"(~{n_test_patches/max(1,len(test_df)):.1f} per slide)\")\nprint(\"NOTE: if per-slide patch count is well below GRID_SIZE^2 * PATCHES_PER_REGION,\")\nprint(\"regions are still being skipped as too small -> check slide dimension distribution.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T10:22:14.001882Z","iopub.execute_input":"2026-08-09T10:22:14.002306Z","iopub.status.idle":"2026-08-09T11:27:13.273669Z","shell.execute_reply.started":"2026-08-09T10:22:14.002278Z","shell.execute_reply":"2026-08-09T11:27:13.272983Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}