{"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":30201,"databundleVersionId":2750748},{"sourceType":"datasetVersion","sourceId":15102478,"datasetId":9669428,"databundleVersionId":15988088}],"dockerImageVersionId":31287,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# CELL 1 — IMPORT LIBRARIES","metadata":{}},{"cell_type":"code","source":"# ============================================\n# IMPORTS\n# ============================================\n\nimport os\nimport gc\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nfrom sklearn.model_selection import train_test_split\n\nfrom tqdm import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\n\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom torchvision.models.detection import maskrcnn_resnet50_fpn\nfrom torchvision.models.detection.rpn import AnchorGenerator\n\nfrom torchvision.transforms import functional as TF\n\nfrom torch.cuda.amp import autocast, GradScaler","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T13:53:12.241856Z","iopub.execute_input":"2026-03-10T13:53:12.242279Z","iopub.status.idle":"2026-03-10T13:53:12.247687Z","shell.execute_reply.started":"2026-03-10T13:53:12.242252Z","shell.execute_reply":"2026-03-10T13:53:12.246794Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CELL 2 — CONFIG","metadata":{}},{"cell_type":"code","source":"# ============================================\n# CONFIG\n# ============================================\n\nimport torch\n\nclass CFG:\n\n    # paths\n    DATA_DIR = \"/kaggle/input/competitions/sartorius-cell-instance-segmentation\"\n    TRAIN_IMG_DIR = f\"{DATA_DIR}/train\"\n    TEST_IMG_DIR = f\"{DATA_DIR}/test\"\n    TRAIN_CSV = f\"{DATA_DIR}/train.csv\"\n\n    # pretrained weights (local dataset)\n    BACKBONE_WEIGHTS = \"/kaggle/input/datasets/mrdeptrai/mask-r-cnn-resnet50/maskrcnn_resnet50_fpn.pth\"\n\n    # training\n    IMG_SIZE = 512\n    PATCH_SIZE = 256\n    STRIDE = 128\n\n    BATCH_SIZE = 2\n    NUM_WORKERS = 2\n\n    EPOCHS = 20\n    LR = 2e-4\n    WEIGHT_DECAY = 1e-4\n\n    NUM_CLASSES = 2\n\n    AMP = True\n\n    EARLY_STOPPING = 5\n\n    GRAD_CLIP = 1.0\n\n    DEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n    THRESHOLD = 0.5\n\n    THRESHOLD_MIN = 0.30\n\n    THRESHOLD_MAX = 0.75\n\n    MIN_MASK_AREA = 30\n\n    N_ROUNDS = 5          # số lần random subset\n    \n    SUBSET_SIZE = 50     # số ảnh mỗi subset\n    \n\nprint(\"Device:\", CFG.DEVICE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T13:53:12.249276Z","iopub.execute_input":"2026-03-10T13:53:12.249567Z","iopub.status.idle":"2026-03-10T13:53:12.264778Z","shell.execute_reply.started":"2026-03-10T13:53:12.249545Z","shell.execute_reply":"2026-03-10T13:53:12.264098Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CELL 3 — LOAD DATASET","metadata":{}},{"cell_type":"code","source":"# ============================================\n# LOAD DATA\n# ============================================\n\ntrain_df = pd.read_csv(CFG.TRAIN_CSV)\n\nprint(\"Train rows:\", len(train_df))\n\n\nimage_ids = train_df[\"id\"].unique()\n\ntrain_ids, val_ids = train_test_split(\n    image_ids,\n    test_size=0.2,\n    random_state=42\n)\n\nprint(\"Train images:\", len(train_ids))\nprint(\"Val images:\", len(val_ids))\nprint(\"Number of images:\", len(image_ids))\n\ntrain_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T13:53:12.265555Z","iopub.execute_input":"2026-03-10T13:53:12.265791Z","iopub.status.idle":"2026-03-10T13:53:12.607335Z","shell.execute_reply.started":"2026-03-10T13:53:12.265759Z","shell.execute_reply":"2026-03-10T13:53:12.606676Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CELL 4 — RLE UTILS","metadata":{}},{"cell_type":"code","source":"# ============================================\n# RLE FUNCTIONS\n# ============================================\n\ndef rle_decode(mask_rle, shape):\n\n    s = mask_rle.split()\n\n    starts = np.asarray(s[0::2], dtype=int)\n    lengths = np.asarray(s[1::2], dtype=int)\n\n    starts -= 1\n    ends = starts + lengths\n\n    img = np.zeros(shape[0] * shape[1], dtype=np.uint8)\n\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n\n    # FIX: remove transpose\n    return img.reshape(shape)\n\n\ndef rle_encode(img):\n\n    pixels = img.flatten()\n\n    pixels = np.concatenate([[0], pixels, [0]])\n\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n\n    runs[1::2] -= runs[::2]\n\n    return \" \".join(str(x) for x in runs)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T13:53:12.608693Z","iopub.execute_input":"2026-03-10T13:53:12.608897Z","iopub.status.idle":"2026-03-10T13:53:12.614880Z","shell.execute_reply.started":"2026-03-10T13:53:12.608878Z","shell.execute_reply":"2026-03-10T13:53:12.614228Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CELL 5 — VISUALIZATION","metadata":{}},{"cell_type":"code","source":"# ============================================\n# VISUALIZE DATA\n# ============================================\n\ndef visualize_sample(image_id):\n\n    image_path = f\"{CFG.TRAIN_IMG_DIR}/{image_id}.png\"\n\n    image = cv2.imread(image_path)\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\n    records = train_df[train_df[\"id\"] == image_id]\n\n    mask = np.zeros(image.shape[:2], dtype=np.uint8)\n\n    for rle in records[\"annotation\"]:\n        m = rle_decode(rle, image.shape[:2])\n        mask = np.maximum(mask, m)\n\n    overlay = image.copy()\n    overlay[mask == 1] = [255,0,0]\n\n    fig, ax = plt.subplots(1,3, figsize=(18,6))\n\n    ax[0].imshow(image)\n    ax[0].set_title(\"Original\")\n\n    ax[1].imshow(mask, cmap=\"gray\")\n    ax[1].set_title(\"Mask\")\n\n    ax[2].imshow(overlay)\n    ax[2].set_title(\"Overlay\")\n\n    plt.show()\n\n\nvisualize_sample(image_ids[0])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T13:53:12.615720Z","iopub.execute_input":"2026-03-10T13:53:12.616078Z","iopub.status.idle":"2026-03-10T13:53:13.363985Z","shell.execute_reply.started":"2026-03-10T13:53:12.616034Z","shell.execute_reply":"2026-03-10T13:53:13.363101Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CELL 6 — PATCH GENERATION","metadata":{}},{"cell_type":"code","source":"# ============================================\n# PATCH GENERATOR\n# ============================================\n\ndef generate_patches(img, mask):\n\n    patches = []\n\n    h, w = img.shape[:2]\n\n    for y in range(0, h-CFG.PATCH_SIZE+1, CFG.STRIDE):\n        for x in range(0, w-CFG.PATCH_SIZE+1, CFG.STRIDE):\n\n            img_patch = img[y:y+CFG.PATCH_SIZE, x:x+CFG.PATCH_SIZE]\n            mask_patch = mask[y:y+CFG.PATCH_SIZE, x:x+CFG.PATCH_SIZE]\n\n            patches.append((img_patch, mask_patch))\n\n    return patches","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T13:53:13.364819Z","iopub.execute_input":"2026-03-10T13:53:13.365031Z","iopub.status.idle":"2026-03-10T13:53:13.370017Z","shell.execute_reply.started":"2026-03-10T13:53:13.365010Z","shell.execute_reply":"2026-03-10T13:53:13.369313Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CELL 7 — AUGMENTATION","metadata":{}},{"cell_type":"code","source":"# ============================================\n# SIMPLE AUGMENTATION\n# ============================================\n\ndef augment(image, mask):\n\n    if np.random.rand() < 0.5:\n        image = np.fliplr(image)\n        mask = np.fliplr(mask)\n\n    if np.random.rand() < 0.5:\n        image = np.flipud(image)\n        mask = np.flipud(mask)\n\n    k = np.random.randint(4)\n    image = np.rot90(image, k)\n    mask = np.rot90(mask, k)\n\n    return image.copy(), mask.copy()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T13:53:13.370861Z","iopub.execute_input":"2026-03-10T13:53:13.371501Z","iopub.status.idle":"2026-03-10T13:53:13.384580Z","shell.execute_reply.started":"2026-03-10T13:53:13.371477Z","shell.execute_reply":"2026-03-10T13:53:13.384084Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CELL 8 — DATASET CLASS","metadata":{}},{"cell_type":"code","source":"# ============================================\n# DATASET\n# ============================================\n\nclass CellDataset(Dataset):\n\n    def __init__(self, image_ids, train=True):\n\n        self.image_ids = image_ids\n        self.train = train\n\n    def __len__(self):\n        return len(self.image_ids)\n\n    def __getitem__(self, idx):\n\n        image_id = self.image_ids[idx]\n\n        path = f\"{CFG.TRAIN_IMG_DIR}/{image_id}.png\"\n\n        image = cv2.imread(path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\n        records = train_df[train_df[\"id\"]==image_id]\n\n        mask = np.zeros(image.shape[:2], dtype=np.uint8)\n\n        masks = []\n\n        for rle in records[\"annotation\"]:\n\n            m = rle_decode(rle, image.shape[:2])\n            masks.append(m)\n\n        if len(masks)==0:\n            masks.append(np.zeros_like(mask))\n\n        masks = np.stack(masks)\n\n        boxes = []\n\n        for m in masks:\n\n            pos = np.where(m)\n\n            xmin = np.min(pos[1])\n            xmax = np.max(pos[1])\n\n            ymin = np.min(pos[0])\n            ymax = np.max(pos[0])\n\n            boxes.append([xmin,ymin,xmax,ymax])\n\n        boxes = torch.as_tensor(boxes, dtype=torch.float32)\n\n        labels = torch.ones((len(boxes),), dtype=torch.int64)\n\n        masks = torch.as_tensor(masks, dtype=torch.uint8)\n\n        image = TF.to_tensor(image)\n\n        target = {}\n\n        target[\"boxes\"] = boxes\n        target[\"labels\"] = labels\n        target[\"masks\"] = masks\n\n        return image, target","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T13:53:13.385515Z","iopub.execute_input":"2026-03-10T13:53:13.385724Z","iopub.status.idle":"2026-03-10T13:53:13.395645Z","shell.execute_reply.started":"2026-03-10T13:53:13.385705Z","shell.execute_reply":"2026-03-10T13:53:13.394998Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CELL 9 — MODEL","metadata":{}},{"cell_type":"code","source":"# ============================================\n# MODEL\n# ============================================\n\nfrom torchvision.models.detection import maskrcnn_resnet50_fpn\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nfrom torchvision.models.detection.mask_rcnn import MaskRCNNPredictor\n\nprint(\"Initializing Mask R-CNN...\")\n\n# 1️⃣ Build architecture (no internet weights)\nmodel = maskrcnn_resnet50_fpn(weights=None, weights_backbone=None)\n\n# 2️⃣ Load local pretrained weights\nprint(\"Loading pretrained MaskRCNN weights from:\", CFG.BACKBONE_WEIGHTS)\n\nstate_dict = torch.load(CFG.BACKBONE_WEIGHTS, map_location=\"cpu\")\n\nmodel.load_state_dict(state_dict)\n\nprint(\"Pretrained weights loaded!\")\n\n# --------------------------------------------\n# 3️⃣ Replace classification head\n# --------------------------------------------\n\nin_features = model.roi_heads.box_predictor.cls_score.in_features\n\nmodel.roi_heads.box_predictor = FastRCNNPredictor(\n    in_features,\n    CFG.NUM_CLASSES\n)\n\n# --------------------------------------------\n# 4️⃣ Replace mask head\n# --------------------------------------------\n\nin_features_mask = model.roi_heads.mask_predictor.conv5_mask.in_channels\n\nmodel.roi_heads.mask_predictor = MaskRCNNPredictor(\n    in_features_mask,\n    256,\n    CFG.NUM_CLASSES\n)\n\n# --------------------------------------------\n# 5️⃣ Move to device\n# --------------------------------------------\n\nmodel.to(CFG.DEVICE)\n\nprint(\"Model ready on device:\", CFG.DEVICE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T13:53:13.397430Z","iopub.execute_input":"2026-03-10T13:53:13.397643Z","iopub.status.idle":"2026-03-10T13:53:14.161760Z","shell.execute_reply.started":"2026-03-10T13:53:13.397623Z","shell.execute_reply":"2026-03-10T13:53:14.161172Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CELL 10 — DATALOADER","metadata":{}},{"cell_type":"code","source":"# ============================================\n# DATALOADER\n# ============================================\n\ntrain_dataset = CellDataset(train_ids)\nval_dataset = CellDataset(val_ids)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=CFG.BATCH_SIZE,\n    shuffle=True,\n    num_workers=CFG.NUM_WORKERS,\n    collate_fn=lambda x: tuple(zip(*x))\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=CFG.BATCH_SIZE,\n    shuffle=False,\n    num_workers=CFG.NUM_WORKERS,\n    collate_fn=lambda x: tuple(zip(*x))\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T13:53:14.162842Z","iopub.execute_input":"2026-03-10T13:53:14.163146Z","iopub.status.idle":"2026-03-10T13:53:14.167886Z","shell.execute_reply.started":"2026-03-10T13:53:14.163112Z","shell.execute_reply":"2026-03-10T13:53:14.167202Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CELL 11 — OPTIMIZER","metadata":{}},{"cell_type":"code","source":"# ============================================\n# OPTIMIZER\n# ============================================\n\noptimizer = optim.AdamW(\n    model.parameters(),\n    lr=CFG.LR,\n    weight_decay=CFG.WEIGHT_DECAY\n)\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer,\n    T_max=CFG.EPOCHS\n)\n\nscaler = GradScaler()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T13:53:14.168677Z","iopub.execute_input":"2026-03-10T13:53:14.168936Z","iopub.status.idle":"2026-03-10T13:53:14.181278Z","shell.execute_reply.started":"2026-03-10T13:53:14.168914Z","shell.execute_reply":"2026-03-10T13:53:14.180513Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CELL 12 — TRAIN LOOP","metadata":{}},{"cell_type":"code","source":"# ============================================\n# TRAIN\n# ============================================\n\nbest_loss = 1e9\npatience = 0\n\nfor epoch in range(CFG.EPOCHS):\n\n    # =========================\n    # TRAIN\n    # =========================\n    model.train()\n\n    running_loss = 0\n\n    pbar = tqdm(train_loader)\n\n    for images, targets in pbar:\n\n        images = [img.to(CFG.DEVICE) for img in images]\n        targets = [{k:v.to(CFG.DEVICE) for k,v in t.items()} for t in targets]\n\n        optimizer.zero_grad()\n\n        with autocast(enabled=CFG.AMP):\n\n            loss_dict = model(images, targets)\n            loss = sum(loss for loss in loss_dict.values())\n\n        scaler.scale(loss).backward()\n\n        torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.GRAD_CLIP)\n\n        scaler.step(optimizer)\n        scaler.update()\n\n        running_loss += loss.item()\n\n        pbar.set_description(f\"train loss {loss.item():.4f}\")\n\n    epoch_loss = running_loss / len(train_loader)\n\n    # =========================\n    # VALIDATION\n    # =========================\n    \n    model.train()   # MaskRCNN cần train mode để trả loss\n    \n    val_loss = 0\n    \n    with torch.no_grad():\n    \n        for images, targets in val_loader:\n    \n            images = [img.to(CFG.DEVICE) for img in images]\n            targets = [{k:v.to(CFG.DEVICE) for k,v in t.items()} for t in targets]\n    \n            with autocast(enabled=CFG.AMP):\n    \n                loss_dict = model(images, targets)\n                loss = sum(loss for loss in loss_dict.values())\n    \n            val_loss += loss.item()\n    \n    val_loss = val_loss / len(val_loader)\n    \n    # =========================\n    # SAVE BEST MODEL\n    # =========================\n    if val_loss < best_loss:\n\n        best_loss = val_loss\n        patience = 0\n\n        torch.save(model.state_dict(), \"best_model.pth\")\n\n        print(\"Best model saved\")\n\n    else:\n\n        patience += 1\n\n        if patience >= CFG.EARLY_STOPPING:\n            print(\"Early stopping triggered\")\n            break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T13:53:14.182253Z","iopub.execute_input":"2026-03-10T13:53:14.182626Z","iopub.status.idle":"2026-03-10T13:53:32.308287Z","shell.execute_reply.started":"2026-03-10T13:53:14.182593Z","shell.execute_reply":"2026-03-10T13:53:32.307329Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CELL 13 — INFERENCE","metadata":{}},{"cell_type":"code","source":"# ============================================\n# INFERENCE\n# ============================================\n\nmodel.load_state_dict(torch.load(\"best_model.pth\"))\n\nmodel.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T13:53:32.308966Z","iopub.status.idle":"2026-03-10T13:53:32.309236Z","shell.execute_reply.started":"2026-03-10T13:53:32.309116Z","shell.execute_reply":"2026-03-10T13:53:32.309131Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CELL 14 — THRESHOLD TUNING","metadata":{}},{"cell_type":"code","source":"# ============================================\n# DICE SCORE\n# ============================================\n\ndef dice_score(pred, target):\n\n    pred = pred.astype(np.bool_)\n    target = target.astype(np.bool_)\n\n    intersection = (pred & target).sum()\n\n    return (2. * intersection) / (pred.sum() + target.sum() + 1e-6)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T13:53:32.310112Z","iopub.status.idle":"2026-03-10T13:53:32.310343Z","shell.execute_reply.started":"2026-03-10T13:53:32.310235Z","shell.execute_reply":"2026-03-10T13:53:32.310249Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================\n# COMBINE INSTANCE MASKS\n# ============================================\n\ndef combine_masks(masks, threshold):\n\n    combined = np.zeros(masks.shape[-2:], dtype=np.uint8)\n\n    for m in masks:\n        m = m[0]\n        m = (m > threshold).astype(np.uint8)\n\n        combined = np.maximum(combined, m)\n\n    return combined","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T13:53:32.311966Z","iopub.status.idle":"2026-03-10T13:53:32.312447Z","shell.execute_reply.started":"2026-03-10T13:53:32.312255Z","shell.execute_reply":"2026-03-10T13:53:32.312272Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================\n# THRESHOLD TUNING (Random Subsets)\n# ============================================\n\nimport random\n\nthresholds = np.arange(CFG.THRESHOLD_MIN, CFG.THRESHOLD_MAX, 0.05)\n\n\nresults = {}\n\nmodel.eval()\n\nfor thr in thresholds:\n\n    round_scores = []\n\n    for r in range(CFG.N_ROUNDS):\n\n        subset_ids = random.sample(list(val_ids), CFG.SUBSET_SIZE)\n\n        scores = []\n\n        for image_id in subset_ids:\n\n            path = f\"{CFG.TRAIN_IMG_DIR}/{image_id}.png\"\n\n            image = cv2.imread(path)\n            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\n            image_tensor = TF.to_tensor(image).to(CFG.DEVICE)\n\n            with torch.no_grad():\n                pred = model([image_tensor])[0]\n\n            masks = pred[\"masks\"].cpu().numpy()\n\n            pred_mask = combine_masks(masks, thr)\n\n            # GT mask\n            records = train_df[train_df[\"id\"] == image_id]\n\n            gt_mask = np.zeros(image.shape[:2], dtype=np.uint8)\n\n            for rle in records[\"annotation\"]:\n                gt_mask = np.maximum(\n                    gt_mask,\n                    rle_decode(rle, image.shape[:2])\n                )\n\n            score = dice_score(pred_mask, gt_mask)\n\n            scores.append(score)\n\n        subset_score = np.mean(scores)\n\n        round_scores.append(subset_score)\n\n    final_score = np.mean(round_scores)\n\n    results[thr] = final_score\n\n    print(f\"Threshold {thr:.2f} -> Dice {final_score:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T13:53:32.313609Z","iopub.status.idle":"2026-03-10T13:53:32.314087Z","shell.execute_reply.started":"2026-03-10T13:53:32.313890Z","shell.execute_reply":"2026-03-10T13:53:32.313917Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_thr = max(results, key=results.get)\n\nprint(\"Best threshold:\", best_thr)\nprint(\"Best Dice:\", results[best_thr])\n\nCFG.THRESHOLD = best_thr","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T13:53:32.315021Z","iopub.status.idle":"2026-03-10T13:53:32.315537Z","shell.execute_reply.started":"2026-03-10T13:53:32.315368Z","shell.execute_reply":"2026-03-10T13:53:32.315394Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CELL 15 — TTA PREDICTION","metadata":{}},{"cell_type":"code","source":"# ============================================\n# TTA\n# ============================================\n\ndef predict_image(image):\n\n    image = TF.to_tensor(image).to(CFG.DEVICE)\n\n    with torch.no_grad():\n\n        pred = model([image])[0]\n\n    masks = pred[\"masks\"].cpu().numpy()\n\n    masks = masks > CFG.THRESHOLD\n\n    return masks","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T13:53:32.316969Z","iopub.status.idle":"2026-03-10T13:53:32.317325Z","shell.execute_reply.started":"2026-03-10T13:53:32.317147Z","shell.execute_reply":"2026-03-10T13:53:32.317168Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CELL 16 — VISUALIZE PREDICTION","metadata":{}},{"cell_type":"code","source":"# ============================================\n# VISUALIZE PREDICTION\n# ============================================\n\ndef visualize_prediction(image_id):\n\n    path = f\"{CFG.TRAIN_IMG_DIR}/{image_id}.png\"\n\n    image = cv2.imread(path)\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\n    masks = predict_image(image)\n\n    combined = np.zeros(image.shape[:2])\n\n    for m in masks:\n        combined = np.maximum(combined, m[0])\n\n    overlay = image.copy()\n    overlay[combined>0] = [0,255,0]\n\n    fig,ax = plt.subplots(1,3,figsize=(18,6))\n\n    ax[0].imshow(image)\n    ax[1].imshow(combined)\n    ax[2].imshow(overlay)\n\n    plt.show()\n\n\nvisualize_prediction(image_ids[0])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T13:53:32.318343Z","iopub.status.idle":"2026-03-10T13:53:32.318585Z","shell.execute_reply.started":"2026-03-10T13:53:32.318470Z","shell.execute_reply":"2026-03-10T13:53:32.318485Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CELL 17 — SUBMISSION","metadata":{}},{"cell_type":"code","source":"import os\n\ntest_ids = []\n\nfor f in os.listdir(CFG.TEST_IMG_DIR):\n    if f.endswith(\".png\"):\n        test_ids.append(f.replace(\".png\",\"\"))\n\nprint(\"Test images:\", len(test_ids))\nprint(test_ids[:5])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T13:53:32.319923Z","iopub.status.idle":"2026-03-10T13:53:32.320196Z","shell.execute_reply.started":"2026-03-10T13:53:32.320052Z","shell.execute_reply":"2026-03-10T13:53:32.320073Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval()\n\nsubmission = []\n\nfor image_id in test_ids:\n\n    image = cv2.imread(f\"{CFG.TEST_IMG_DIR}/{image_id}.png\")\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\n    h, w = image.shape[:2]\n\n    image = cv2.resize(image, (CFG.IMG_SIZE, CFG.IMG_SIZE))\n    image = image / 255.0\n\n    tensor = torch.tensor(image).permute(2,0,1).unsqueeze(0).float().to(CFG.DEVICE)\n\n    with torch.no_grad():\n        output = model(tensor)[0]\n\n    masks = output[\"masks\"].cpu().numpy()\n    scores = output[\"scores\"].cpu().numpy()\n\n    used = np.zeros((h,w), dtype=np.uint8)\n\n    count = 0\n\n    for mask, score in zip(masks, scores):\n\n        if score < 0.5:\n            continue\n\n        mask = mask[0]\n\n        mask = cv2.resize(\n            mask,\n            (w,h),\n            interpolation=cv2.INTER_LINEAR\n        )\n\n        mask = mask > CFG.THRESHOLD\n\n        if mask.sum() < CFG.MIN_MASK_AREA:\n            continue\n\n        # remove overlap (IMPORTANT)\n        mask = mask & (~used)\n\n        if mask.sum() == 0:\n            continue\n\n        used = used | mask\n\n        rle = rle_encode(mask.astype(np.uint8))\n\n        submission.append([image_id, rle])\n\n        count += 1\n\n    if count == 0:\n        submission.append([image_id, \"1 1\"])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T13:53:32.321380Z","iopub.status.idle":"2026-03-10T13:53:32.321747Z","shell.execute_reply.started":"2026-03-10T13:53:32.321534Z","shell.execute_reply":"2026-03-10T13:53:32.321562Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --------------------------------\n# build dataframe\n# --------------------------------\n\nsub = pd.DataFrame(submission, columns=[\"id\",\"predicted\"])\n\nsub[\"predicted\"] = sub[\"predicted\"].astype(str)\n\nsub = sub.sort_values(\"id\").reset_index(drop=True)\n\nsub.to_csv(\"submission.csv\", index=False)\n\nprint(sub.head())\nprint(\"rows:\", len(sub))\nprint(\"unique images:\", sub[\"id\"].nunique())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T13:53:32.323333Z","iopub.status.idle":"2026-03-10T13:53:32.323582Z","shell.execute_reply.started":"2026-03-10T13:53:32.323469Z","shell.execute_reply":"2026-03-10T13:53:32.323484Z"}},"outputs":[],"execution_count":null}]}