{"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":15091289,"datasetId":9662057,"databundleVersionId":15975653},{"sourceType":"datasetVersion","sourceId":15091146,"datasetId":9661939,"databundleVersionId":15975493},{"sourceType":"datasetVersion","sourceId":14514638,"datasetId":9270464,"databundleVersionId":15341772},{"sourceType":"datasetVersion","sourceId":15054271,"datasetId":9637746,"databundleVersionId":15934772}],"dockerImageVersionId":31287,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ================================\n# IMPORT LIBRARIES\n# ================================\n\nimport os\nimport gc\nimport cv2\nimport math\nimport random\nimport numpy as np\nimport pandas as pd\n\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom sklearn.model_selection import train_test_split\n\nprint(\"PyTorch version:\", torch.__version__)\nprint(\"GPU:\", torch.cuda.is_available())\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:12.069573Z","iopub.execute_input":"2026-03-08T14:49:12.070322Z","iopub.status.idle":"2026-03-08T14:49:12.077031Z","shell.execute_reply.started":"2026-03-08T14:49:12.070286Z","shell.execute_reply":"2026-03-08T14:49:12.076089Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_DIR = \"/kaggle/input/competitions/sartorius-cell-instance-segmentation\"\n\ntrain_df = pd.read_csv(f\"{DATA_DIR}/train.csv\")\n\nprint(\"Train rows:\", len(train_df))\ndisplay(train_df.head())\ntrain_df.cell_type.value_counts()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:12.078430Z","iopub.execute_input":"2026-03-08T14:49:12.078670Z","iopub.status.idle":"2026-03-08T14:49:12.463043Z","shell.execute_reply.started":"2026-03-08T14:49:12.078650Z","shell.execute_reply":"2026-03-08T14:49:12.462376Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def 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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:12.463834Z","iopub.execute_input":"2026-03-08T14:49:12.464054Z","iopub.status.idle":"2026-03-08T14:49:12.468780Z","shell.execute_reply.started":"2026-03-08T14:49:12.464034Z","shell.execute_reply":"2026-03-08T14:49:12.468141Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"row = train_df.iloc[0]\n\nmask = rle_decode(\n    row.annotation,\n    (int(row.height),\n    int(row.width))\n)\n\nprint(mask.shape)\nprint(mask.sum())\n\nplt.imshow(mask)\nplt.title(\"single cell\")\nplt.show()\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:12.470335Z","iopub.execute_input":"2026-03-08T14:49:12.470564Z","iopub.status.idle":"2026-03-08T14:49:12.646444Z","shell.execute_reply.started":"2026-03-08T14:49:12.470544Z","shell.execute_reply":"2026-03-08T14:49:12.645856Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_mask(image_id):\n\n    df_img = train_df[train_df[\"id\"] == image_id]\n\n    height = int(df_img.iloc[0].height)\n    width = int(df_img.iloc[0].width)\n\n    mask = np.zeros((height, width), dtype=np.uint8)\n\n    for rle in df_img[\"annotation\"]:\n\n        if pd.isna(rle):\n            continue\n\n        mask += rle_decode(rle, (height, width))\n\n    mask = np.clip(mask, 0, 1)\n\n    return mask","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:12.647259Z","iopub.execute_input":"2026-03-08T14:49:12.647594Z","iopub.status.idle":"2026-03-08T14:49:12.652645Z","shell.execute_reply.started":"2026-03-08T14:49:12.647530Z","shell.execute_reply":"2026-03-08T14:49:12.651891Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img_id = train_df.iloc[0].id\ndf_img = train_df[train_df.id == img_id]\n\nimage = cv2.imread(f\"{DATA_DIR}/train/{img_id}.png\")\nimage = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\nmask = build_mask(img_id)\n\nplt.figure(figsize=(15,5))\n\nplt.subplot(1,3,1)\nplt.imshow(image)\nplt.title(\"image\")\n\nplt.subplot(1,3,2)\nplt.imshow(mask)\nplt.title(\"mask\")\n\nplt.subplot(1,3,3)\nplt.imshow(image)\nplt.imshow(mask, alpha=0.4)\nplt.title(\"overlay\")\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:12.653596Z","iopub.execute_input":"2026-03-08T14:49:12.654390Z","iopub.status.idle":"2026-03-08T14:49:13.195903Z","shell.execute_reply.started":"2026-03-08T14:49:12.654358Z","shell.execute_reply":"2026-03-08T14:49:13.195082Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_ids = train_df[\"id\"].unique()\n\ntrain_ids, valid_ids = train_test_split(\n    image_ids,\n    test_size=0.1,\n    random_state=42\n)\n\nprint(len(train_ids), len(valid_ids))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:13.196846Z","iopub.execute_input":"2026-03-08T14:49:13.197331Z","iopub.status.idle":"2026-03-08T14:49:13.206002Z","shell.execute_reply.started":"2026-03-08T14:49:13.197305Z","shell.execute_reply":"2026-03-08T14:49:13.205254Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\ntrain_transform = A.Compose([\n    A.Resize(512,512),\n\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.RandomRotate90(p=0.5),\n\n    A.Normalize(mean=(0,0,0), std=(1,1,1)),\n    ToTensorV2()\n])\n\nvalid_transform = A.Compose([\n    A.Resize(512,512),\n    A.Normalize(mean=(0,0,0), std=(1,1,1)),\n    ToTensorV2()\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:13.206889Z","iopub.execute_input":"2026-03-08T14:49:13.207132Z","iopub.status.idle":"2026-03-08T14:49:13.223582Z","shell.execute_reply.started":"2026-03-08T14:49:13.207105Z","shell.execute_reply":"2026-03-08T14:49:13.222835Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"IMG_SIZE = 512\n\nclass CellDataset(Dataset):\n\n    def __init__(self, image_ids, transform=None):\n\n        self.image_ids = image_ids\n        self.transform = transform\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        image = cv2.imread(f\"{DATA_DIR}/train/{image_id}.png\")\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\n        mask = build_mask(image_id)\n\n        if self.transform:\n\n            aug = self.transform(image=image, mask=mask)\n\n            image = aug[\"image\"]\n            mask = aug[\"mask\"].unsqueeze(0)\n\n        else:\n\n            image = cv2.resize(image,(IMG_SIZE,IMG_SIZE))\n            mask = cv2.resize(mask,(IMG_SIZE,IMG_SIZE),interpolation=cv2.INTER_NEAREST)\n\n            image = image / 255.0\n\n            image = torch.tensor(image).permute(2,0,1).float()\n            mask = torch.tensor(mask).unsqueeze(0).float()\n\n        return image, mask","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:13.226320Z","iopub.execute_input":"2026-03-08T14:49:13.227046Z","iopub.status.idle":"2026-03-08T14:49:13.237689Z","shell.execute_reply.started":"2026-03-08T14:49:13.227024Z","shell.execute_reply":"2026-03-08T14:49:13.236990Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = CellDataset(\n    train_ids,\n    transform=train_transform\n)\n\nvalid_dataset = CellDataset(\n    valid_ids,\n    transform=valid_transform\n)\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=4,\n    shuffle=True,\n    num_workers=2\n)\n\nvalid_loader = DataLoader(\n    valid_dataset,\n    batch_size=4,\n    shuffle=False\n)\n\nprint(\"Train batches:\",len(train_loader))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:13.238555Z","iopub.execute_input":"2026-03-08T14:49:13.238788Z","iopub.status.idle":"2026-03-08T14:49:13.257200Z","shell.execute_reply.started":"2026-03-08T14:49:13.238756Z","shell.execute_reply":"2026-03-08T14:49:13.256615Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchvision.models import convnext_tiny\n\nconvnext_tiny = convnext_tiny(weights=None)\nconvnext_tiny.load_state_dict(torch.load(\"/kaggle/input/datasets/mrdeptrai/convnext-tiny/convnext_tiny.pth\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:13.258073Z","iopub.execute_input":"2026-03-08T14:49:13.258334Z","iopub.status.idle":"2026-03-08T14:49:13.750621Z","shell.execute_reply.started":"2026-03-08T14:49:13.258306Z","shell.execute_reply":"2026-03-08T14:49:13.749906Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\nclass ConvBlock(nn.Module):\n\n    def __init__(self, in_c, out_c):\n        super().__init__()\n\n        self.conv = nn.Sequential(\n            nn.Conv2d(in_c, out_c, 3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True),\n\n            nn.Conv2d(out_c, out_c, 3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        return self.conv(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:13.751574Z","iopub.execute_input":"2026-03-08T14:49:13.751899Z","iopub.status.idle":"2026-03-08T14:49:13.756966Z","shell.execute_reply.started":"2026-03-08T14:49:13.751866Z","shell.execute_reply":"2026-03-08T14:49:13.756305Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Up(nn.Module):\n\n    def __init__(self, in_c, out_c):\n        super().__init__()\n\n        self.up = nn.ConvTranspose2d(in_c, out_c, 2, stride=2)\n\n    def forward(self, x):\n        return self.up(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:13.758001Z","iopub.execute_input":"2026-03-08T14:49:13.758195Z","iopub.status.idle":"2026-03-08T14:49:13.773955Z","shell.execute_reply.started":"2026-03-08T14:49:13.758176Z","shell.execute_reply":"2026-03-08T14:49:13.773456Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class NestedBlock(nn.Module):\n\n    def __init__(self, in_c, mid_c, out_c):\n        super().__init__()\n\n        self.conv = nn.Sequential(\n            nn.Conv2d(in_c, mid_c, 3, padding=1),\n            nn.BatchNorm2d(mid_c),\n            nn.ReLU(inplace=True),\n\n            nn.Conv2d(mid_c, out_c, 3, padding=1),\n            nn.BatchNorm2d(out_c),\n            nn.ReLU(inplace=True),\n        )\n\n    def forward(self, x):\n        return self.conv(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:13.774719Z","iopub.execute_input":"2026-03-08T14:49:13.774907Z","iopub.status.idle":"2026-03-08T14:49:13.786778Z","shell.execute_reply.started":"2026-03-08T14:49:13.774878Z","shell.execute_reply":"2026-03-08T14:49:13.786246Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport timm\n\n\nclass EfficientNetB3Encoder(nn.Module):\n\n    def __init__(self, weight_path=None):\n        super().__init__()\n\n        # tạo model architecture\n        self.backbone = timm.create_model(\n            \"efficientnet_b3\",\n            pretrained=False,\n            features_only=True\n        )\n\n        # load weight nếu có\n        if weight_path is not None:\n\n            state_dict = torch.load(weight_path, map_location=\"cpu\")\n\n            self.backbone.load_state_dict(state_dict, strict=False)\n\n            print(\"Loaded encoder weights from:\", weight_path)\n\n    def forward(self, x):\n\n        features = self.backbone(x)\n\n        return features","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:13.787549Z","iopub.execute_input":"2026-03-08T14:49:13.787800Z","iopub.status.idle":"2026-03-08T14:49:13.800135Z","shell.execute_reply.started":"2026-03-08T14:49:13.787769Z","shell.execute_reply":"2026-03-08T14:49:13.799455Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class UNetPlusPlus(nn.Module):\n\n    def __init__(self, num_classes=1):\n        super().__init__()\n\n        self.encoder = EfficientNetB3Encoder(\n            weight_path=\"/kaggle/input/datasets/mrdeptrai/efficientnet-b3/efficientnet_b3.pth\"\n        )\n\n        filters = [24, 32, 48, 136, 384]  # đúng với EfficientNetB3\n\n        self.up4 = Up(filters[4], filters[3])\n        self.up3 = Up(filters[3], filters[2])\n        self.up2 = Up(filters[2], filters[1])\n        self.up1 = Up(filters[1], filters[0])\n\n        self.conv4 = ConvBlock(filters[3] + filters[3], filters[3])\n        self.conv3 = ConvBlock(filters[2] + filters[2], filters[2])\n        self.conv2 = ConvBlock(filters[1] + filters[1], filters[1])\n        self.conv1 = ConvBlock(filters[0] + filters[0], filters[0])\n\n        self.final = nn.Conv2d(filters[0], num_classes, 1)\n\n    def forward(self, x):\n\n        feats = self.encoder(x)\n\n        f0, f1, f2, f3, f4 = feats  # 256,128,64,32,16\n\n        x4 = self.up4(f4)\n        x4 = torch.cat([x4, f3], dim=1)\n        x4 = self.conv4(x4)\n\n        x3 = self.up3(x4)\n        x3 = torch.cat([x3, f2], dim=1)\n        x3 = self.conv3(x3)\n\n        x2 = self.up2(x3)\n        x2 = torch.cat([x2, f1], dim=1)\n        x2 = self.conv2(x2)\n\n        x1 = self.up1(x2)\n        x1 = torch.cat([x1, f0], dim=1)\n        x1 = self.conv1(x1)\n\n        out = self.final(x1)\n\n        # upsample cuối về đúng size input\n        out = F.interpolate(\n            out,\n            size=x.shape[-2:],\n            mode=\"bilinear\",\n            align_corners=False\n        )\n\n        return out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:13.801146Z","iopub.execute_input":"2026-03-08T14:49:13.801341Z","iopub.status.idle":"2026-03-08T14:49:13.815348Z","shell.execute_reply.started":"2026-03-08T14:49:13.801322Z","shell.execute_reply":"2026-03-08T14:49:13.814844Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = UNetPlusPlus().cuda()\n\nx = torch.randn(2,3,512,512).cuda()\n\ny = model(x)\n\nprint(y.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:13.816117Z","iopub.execute_input":"2026-03-08T14:49:13.816330Z","iopub.status.idle":"2026-03-08T14:49:14.144127Z","shell.execute_reply.started":"2026-03-08T14:49:13.816307Z","shell.execute_reply":"2026-03-08T14:49:14.143552Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\n# =========================\n# Dice Loss (improved)\n# =========================\n\nclass DiceLoss(nn.Module):\n\n    def __init__(self, smooth=1e-6):\n        super().__init__()\n        self.smooth = smooth\n\n    def forward(self, preds, targets):\n\n        preds = torch.sigmoid(preds)\n\n        # reshape per image\n        preds = preds.contiguous().view(preds.size(0), -1)\n        targets = targets.contiguous().view(targets.size(0), -1)\n\n        intersection = (preds * targets).sum(dim=1)\n\n        dice = (2 * intersection + self.smooth) / (\n            preds.sum(dim=1) + targets.sum(dim=1) + self.smooth\n        )\n\n        dice_loss = 1 - dice\n\n        return dice_loss.mean()\n\n\n# =========================\n# Loss Functions\n# =========================\n\nbce_loss_fn = nn.BCEWithLogitsLoss(\n    pos_weight=torch.tensor([3.0]).to(device)  # giảm weight cho ổn định\n)\n\ndice_loss_fn = DiceLoss()\n\n\n# =========================\n# Combined Loss\n# =========================\n\ndef loss_fn(preds, masks):\n\n    bce = bce_loss_fn(preds, masks)\n    dice = dice_loss_fn(preds, masks)\n\n    # weighted sum (thường tốt hơn)\n    return 0.5 * bce + 0.5 * dice","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:14.144897Z","iopub.execute_input":"2026-03-08T14:49:14.145099Z","iopub.status.idle":"2026-03-08T14:49:14.159892Z","shell.execute_reply.started":"2026-03-08T14:49:14.145080Z","shell.execute_reply":"2026-03-08T14:49:14.159344Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EPOCHS = 80\n\ndevice = \"cuda\"\n\nmodel = UNetPlusPlus().to(device)\n\noptimizer = torch.optim.AdamW(model.parameters(),lr=1e-3)\n\nscheduler = CosineAnnealingLR(\n    optimizer,\n    T_max=EPOCHS\n)\n\nscaler = torch.cuda.amp.GradScaler()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:14.160714Z","iopub.execute_input":"2026-03-08T14:49:14.160959Z","iopub.status.idle":"2026-03-08T14:49:14.423296Z","shell.execute_reply.started":"2026-03-08T14:49:14.160926Z","shell.execute_reply":"2026-03-08T14:49:14.422472Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def dice_score(preds, targets, smooth=1e-6):\n\n    preds = torch.sigmoid(preds)\n    preds = (preds > 0.5).float()\n\n    preds = preds.view(-1)\n    targets = targets.view(-1)\n\n    intersection = (preds * targets).sum()\n\n    dice = (2. * intersection + smooth) / (\n        preds.sum() + targets.sum() + smooth\n    )\n\n    return dice.item()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:14.424186Z","iopub.execute_input":"2026-03-08T14:49:14.424479Z","iopub.status.idle":"2026-03-08T14:49:14.428926Z","shell.execute_reply.started":"2026-03-08T14:49:14.424454Z","shell.execute_reply":"2026-03-08T14:49:14.428296Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_dice = 0\n\npatience = 10\nearly_stop_counter = 0\n\nfor epoch in range(EPOCHS):\n\n    # ================= TRAIN =================\n    model.train()\n\n    train_loss = 0\n    train_dice = 0\n\n    loop = tqdm(train_loader)\n\n    for images, masks in loop:\n\n        images = images.to(device).float()\n        masks = masks.to(device).float()\n\n        optimizer.zero_grad()\n\n        with torch.cuda.amp.autocast():\n\n            preds = model(images)\n            loss = loss_fn(preds, masks)\n\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        train_loss += loss.item()\n\n        with torch.no_grad():\n            dice = dice_score(preds, masks)\n\n        train_dice += dice\n\n        loop.set_description(f\"Epoch {epoch+1}/{EPOCHS}\")\n        loop.set_postfix(loss=loss.item(), dice=dice)\n\n    train_loss /= len(train_loader)\n    train_dice /= len(train_loader)\n\n    print(f\"Train Loss {train_loss:.4f} Dice {train_dice:.4f}\")\n\n    # ================= VALID =================\n    model.eval()\n\n    valid_loss = 0\n    valid_dice = 0\n\n    with torch.no_grad():\n\n        for images, masks in valid_loader:\n\n            images = images.to(device).float()\n            masks = masks.to(device).float()\n\n            preds = model(images)\n            loss = loss_fn(preds, masks)\n\n            valid_loss += loss.item()\n\n            dice = dice_score(preds, masks)\n            valid_dice += dice\n\n    valid_loss /= len(valid_loader)\n    valid_dice /= len(valid_loader)\n\n    print(f\"Valid Loss {valid_loss:.4f} Dice {valid_dice:.4f}\")\n\n    # ================= SAVE BEST MODEL =================\n    if valid_dice > best_dice:\n    \n        best_dice = valid_dice\n        early_stop_counter = 0\n    \n        torch.save(\n            model.state_dict(),\n            \"best_model.pth\"\n        )\n    \n        print(f\"✅ Saved Best Model | Dice {best_dice:.4f}\")\n    \n    else:\n    \n        early_stop_counter += 1\n        print(f\"EarlyStopping Counter: {early_stop_counter}/{patience}\")\n    \n        if early_stop_counter >= patience:\n    \n            print(\"🛑 Early stopping triggered\")\n            break\n    # ================= SCHEDULER =================\n    scheduler.step()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:14.429879Z","iopub.execute_input":"2026-03-08T14:49:14.430156Z","iopub.status.idle":"2026-03-08T14:49:29.619365Z","shell.execute_reply.started":"2026-03-08T14:49:14.430129Z","shell.execute_reply":"2026-03-08T14:49:29.617914Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.load_state_dict(torch.load(\"best_model.pth\"))\nmodel.eval()\n\nimage, mask = valid_dataset[0]\n\nwith torch.no_grad():\n    pred = model(image.unsqueeze(0).to(device))\n\npred = torch.sigmoid(pred)\n\nprint(mask.unique())\nprint(\"min:\", pred.min().item())\nprint(\"max:\", pred.max().item())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:29.620106Z","iopub.status.idle":"2026-03-08T14:49:29.620393Z","shell.execute_reply.started":"2026-03-08T14:49:29.620253Z","shell.execute_reply":"2026-03-08T14:49:29.620269Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\n\ndef evaluate_threshold(model, loader, device, thresholds):\n\n    model.eval()\n\n    best_t = 0\n    best_dice = 0\n\n    with torch.no_grad():\n\n        for t in thresholds:\n\n            dices = []\n\n            for images, masks in loader:\n\n                images = images.to(device)\n                masks = masks.to(device)\n\n                pred1 = model(images)\n                pred2 = torch.flip(model(torch.flip(images,[3])),[3])\n                pred3 = torch.flip(model(torch.flip(images,[2])),[2])\n                \n                preds = (pred1 + pred2 + pred3) / 3\n                preds = torch.sigmoid(preds)\n\n                pred_mask = (preds > t).float()\n\n                dice = dice_score(pred_mask, masks)\n\n                dices.append(dice)\n\n            mean_dice = np.mean(dices)\n\n            print(f\"Threshold {t:.2f} Dice {mean_dice:.4f}\")\n\n            if mean_dice > best_dice:\n\n                best_dice = mean_dice\n                best_t = t\n\n    print(\"Best threshold:\", best_t)\n    print(\"Best dice:\", best_dice)\n\n    return best_t\n\n\nthresholds = np.arange(0.2,0.8,0.05)\n\nbest_threshold = evaluate_threshold(\n    model,\n    valid_loader,\n    device,\n    thresholds\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:29.621925Z","iopub.status.idle":"2026-03-08T14:49:29.622239Z","shell.execute_reply.started":"2026-03-08T14:49:29.622104Z","shell.execute_reply":"2026-03-08T14:49:29.622118Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_prediction():\n\n    model.eval()\n\n    image, mask = valid_dataset[0]\n\n    with torch.no_grad():\n        pred = model(image.unsqueeze(0).to(device))\n\n    prob = torch.sigmoid(pred).squeeze().cpu().numpy()\n\n    print(\"prob min:\", prob.min())\n    print(\"prob max:\", prob.max())\n\n    pred_mask = (prob > best_threshold).astype(\"uint8\")\n\n    plt.figure(figsize=(16,4))\n\n    plt.subplot(1,4,1)\n    plt.title(\"Image\")\n    plt.imshow(image.permute(1,2,0))\n    plt.axis(\"off\")\n\n    plt.subplot(1,4,2)\n    plt.title(\"GT\")\n    plt.imshow(mask.squeeze())\n    plt.axis(\"off\")\n\n    plt.subplot(1,4,3)\n    plt.title(\"Probability\")\n    plt.imshow(prob)\n    plt.axis(\"off\")\n\n    plt.subplot(1,4,4)\n    plt.title(\"Pred Mask\")\n    plt.imshow(pred_mask, cmap=\"gray\")\n    plt.axis(\"off\")\n\n    plt.show()\n\nvisualize_prediction()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:29.623546Z","iopub.status.idle":"2026-03-08T14:49:29.623783Z","shell.execute_reply.started":"2026-03-08T14:49:29.623674Z","shell.execute_reply":"2026-03-08T14:49:29.623687Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def mask_to_instances(mask):\n\n    mask = (mask>best_threshold).astype(np.uint8)\n\n    num_labels,labels = cv2.connectedComponents(mask)\n\n    masks = []\n\n    for i in range(1,num_labels):\n\n        m = (labels==i).astype(np.uint8)\n\n        masks.append(m)\n\n    return masks","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:29.625774Z","iopub.status.idle":"2026-03-08T14:49:29.626124Z","shell.execute_reply.started":"2026-03-08T14:49:29.625946Z","shell.execute_reply":"2026-03-08T14:49:29.625968Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def 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-08T14:49:29.627627Z","iopub.status.idle":"2026-03-08T14:49:29.627988Z","shell.execute_reply.started":"2026-03-08T14:49:29.627808Z","shell.execute_reply":"2026-03-08T14:49:29.627830Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nTEST_DIR = \"/kaggle/input/competitions/sartorius-cell-instance-segmentation/test\"\n\ntest_ids = []\n\nfor file in os.listdir(TEST_DIR):\n    if file.endswith(\".png\"):\n        test_ids.append(file.replace(\".png\",\"\"))\n\nprint(\"Test images:\",len(test_ids))\nprint(test_ids[:5])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:29.628929Z","iopub.status.idle":"2026-03-08T14:49:29.629175Z","shell.execute_reply.started":"2026-03-08T14:49:29.629062Z","shell.execute_reply":"2026-03-08T14:49:29.629077Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport random\nimport cv2\nimport torch\nimport numpy as np\n\nmodel.eval()\n\n# chọn ngẫu nhiên vài ảnh test\nsample_ids = random.sample(test_ids, 3)\n\nfor image_id in sample_ids:\n\n    image = cv2.imread(f\"{TEST_DIR}/{image_id}.png\")\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\n    image_resized = cv2.resize(image,(IMG_SIZE,IMG_SIZE))\n    image_norm = image_resized / 255.0\n\n    tensor = torch.tensor(image_norm).permute(2,0,1).unsqueeze(0).float().to(device)\n\n    with torch.no_grad():\n        pred = model(tensor)\n\n    prob = torch.sigmoid(pred).cpu().numpy()[0,0]\n\n    mask = prob > best_threshold\n\n    # plot\n    fig, ax = plt.subplots(1,3,figsize=(15,5))\n\n    ax[0].imshow(image_resized)\n    ax[0].set_title(\"Image\")\n\n    ax[1].imshow(prob, cmap=\"viridis\")\n    ax[1].set_title(\"Probability\")\n\n    ax[2].imshow(mask, cmap=\"gray\")\n    ax[2].set_title(\"Mask\")\n\n    for a in ax:\n        a.axis(\"off\")\n\n    plt.suptitle(image_id)\n    plt.show()\n\n    print(\"prob min:\", prob.min(), \"prob max:\", prob.max())\n    print(\"mask pixels:\", mask.sum())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:29.630042Z","iopub.status.idle":"2026-03-08T14:49:29.630269Z","shell.execute_reply.started":"2026-03-08T14:49:29.630163Z","shell.execute_reply":"2026-03-08T14:49:29.630176Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from skimage import measure\n\nsubmission = []\n\nfor image_id in test_ids:\n\n    image = cv2.imread(f\"{TEST_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,(IMG_SIZE,IMG_SIZE))\n    image = image / 255.0\n\n    tensor = torch.tensor(image).permute(2,0,1).unsqueeze(0).float().to(device)\n\n    with torch.no_grad():\n    \n        # original\n        pred1 = model(tensor)\n    \n        # horizontal flip\n        tensor_h = torch.flip(tensor,[3])\n        pred2 = model(tensor_h)\n        pred2 = torch.flip(pred2,[3])\n    \n        # vertical flip\n        tensor_v = torch.flip(tensor,[2])\n        pred3 = model(tensor_v)\n        pred3 = torch.flip(pred3,[2])\n    \n        # average\n        pred = (pred1 + pred2 + pred3) / 3\n        \n    prob = torch.sigmoid(pred).cpu().numpy()[0,0]\n    \n    prob = cv2.GaussianBlur(prob,(5,5),0)\n    mask = prob > best_threshold\n\n    mask = cv2.resize(\n        mask.astype(np.uint8),\n        (w,h),\n        interpolation=cv2.INTER_NEAREST\n    )\n\n    labels = measure.label(mask)\n\n    count = 0\n\n    for i in range(1, labels.max()+1):\n\n        cell = (labels == i).astype(np.uint8)\n\n        if cell.sum() < 30:\n            continue\n\n        rle = rle_encode(cell)\n\n        submission.append([image_id, rle])\n\n        count += 1\n\n    if count == 0:\n        submission.append([image_id, \"\"])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:29.631879Z","iopub.status.idle":"2026-03-08T14:49:29.632106Z","shell.execute_reply.started":"2026-03-08T14:49:29.631998Z","shell.execute_reply":"2026-03-08T14:49:29.632011Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub = pd.DataFrame(submission, columns=[\"id\",\"predicted\"])\n\nsub.to_csv(\"submission.csv\", index=False)\n\nprint(sub.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:29.632951Z","iopub.status.idle":"2026-03-08T14:49:29.633188Z","shell.execute_reply.started":"2026-03-08T14:49:29.633077Z","shell.execute_reply":"2026-03-08T14:49:29.633090Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"submission rows:\", len(submission))\nprint(\"unique images:\", len(test_ids))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:49:29.635260Z","iopub.status.idle":"2026-03-08T14:49:29.635804Z","shell.execute_reply.started":"2026-03-08T14:49:29.635523Z","shell.execute_reply":"2026-03-08T14:49:29.635547Z"}},"outputs":[],"execution_count":null}]}