{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"}],"dockerImageVersionId":31089,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\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\nimport numpy as np # linear algebra\nimport 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\nfor dirname, dirnames, _ in os.walk('/kaggle/input'):\n    for d in dirnames:\n        print(os.path.join(dirname, d))\n\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":"2025-08-20T16:09:57.705875Z","iopub.execute_input":"2025-08-20T16:09:57.706259Z","iopub.status.idle":"2025-08-20T16:10:05.579328Z","shell.execute_reply.started":"2025-08-20T16:09:57.706221Z","shell.execute_reply":"2025-08-20T16:10:05.577857Z"},"_kg_hide-output":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install torch torchvision albumentations tifffile opencv-python tqdm","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- imports\nimport os, re, glob, math, csv, gc\nimport numpy as np\nimport pandas as pd\nimport tifffile as tiff\nimport cv2\nfrom tqdm import tqdm\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\n\n# =========================\n# Utils: RLE encode/decode\n# =========================\ndef rle_encode(mask):\n    # mask: 2D uint8 array with {0,1}\n    pixels = mask.flatten(order='F')  # Fortran order per common competitions\n    pads = np.pad(pixels, (1,1), mode='constant')\n    runs = np.where(pads[1:] != pads[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\n\ndef rle_decode(rle, shape):\n    if (rle is None) or (rle == ''):\n        return np.zeros(shape, dtype=np.uint8)\n    s = list(map(int, rle.split()))\n    starts, lengths = s[0::2], s[1::2]\n    starts = np.asarray(starts) - 1\n    ends = starts + lengths\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n    return img.reshape(shape, order='F')\n\n# =========================\n# Paths\n# =========================\nDATA_DIR = \"/kaggle/input/blood-vessel-segmentation\"  # change if local\nTRAIN_DIR = os.path.join(DATA_DIR, \"train\")\nTEST_DIR  = os.path.join(DATA_DIR, \"test\")\n\n# Get all datasets under train\ntrain_sets = [\n    \"kidney_1_dense\",\n    \"kidney_1_voi\",\n    \"kidney_2\",\n    \"kidney_3_dense\",\n    \"kidney_3_sparse\",\n]\n\n# =========================\n# Data indexing\n# =========================\ndef list_train_items():\n    items = []\n    for ds in train_sets:\n        img_dir = os.path.join(TRAIN_DIR, ds, \"images\")\n        lbl_dir = os.path.join(TRAIN_DIR, ds, \"labels\")\n        # kidney_3_dense only has labels (images are in kidney_3_sparse/images)\n        if ds == \"kidney_3_dense\" and not os.path.exists(img_dir):\n            img_dir = os.path.join(TRAIN_DIR, \"kidney_3_sparse\", \"images\")\n\n        if not os.path.exists(img_dir) or not os.path.exists(lbl_dir):\n            # some subsets may be missing (rare), skip safely\n            continue\n\n        for p in sorted(glob.glob(os.path.join(img_dir, \"*.tif\"))):\n            fname = os.path.basename(p)  # e.g. 0001.tif\n            lbl_p = os.path.join(lbl_dir, fname)\n            if os.path.exists(lbl_p):\n                items.append((ds, p, lbl_p))\n    return items\n\ndef list_test_items():\n    items = []\n    for ds in [\"kidney_5\",\"kidney_6\"]:\n        img_dir = os.path.join(TEST_DIR, ds, \"images\")\n        if not os.path.exists(img_dir): \n            continue\n        for p in sorted(glob.glob(os.path.join(img_dir, \"*.tif\"))):\n            fname = os.path.basename(p)\n            slice_id = f\"{ds}_{os.path.splitext(fname)[0]}\"\n            items.append((slice_id, p))\n    return items\n\ntrain_items = list_train_items()\ntest_items  = list_test_items()\nprint(f\"Train slices: {len(train_items)}, Test slices: {len(test_items)}\")\n\n# =========================\n# Dataset\n# =========================\nimport albumentations as A\n\nclass VesselDataset(Dataset):\n    \"\"\"\n    items: list of tuples -> (dataset_name, image_path, label_path)\n    aug:   bool, if True apply training augs; if False, deterministic preprocessing only\n    size:  int, the max side after LongestMaxSize (square canvas after PadIfNeeded)\n    \"\"\"\n\n    def __init__(self, items, aug=False, size=768):\n        self.items = items\n        self.aug = aug\n        self.size = int(size)\n\n        if aug:\n            # Stable on Kaggle across Albumentations versions\n            self.transform = A.Compose([\n                A.LongestMaxSize(max_size=self.size),\n                A.PadIfNeeded(min_height=self.size, min_width=self.size,\n                              border_mode=cv2.BORDER_REFLECT),\n                A.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.10, rotate_limit=15,\n                                   border_mode=cv2.BORDER_REFLECT, p=0.5),\n                A.HorizontalFlip(p=0.5),\n                A.VerticalFlip(p=0.5),\n                A.RandomBrightnessContrast(p=0.3),\n            ])\n        else:\n            self.transform = A.Compose([\n                A.LongestMaxSize(max_size=self.size),\n                A.PadIfNeeded(min_height=self.size, min_width=self.size,\n                              border_mode=cv2.BORDER_REFLECT),\n            ])\n\n    def __len__(self):\n        return len(self.items)\n\n    def _read_pair(self, img_path, lbl_path):\n        # Read grayscale TIFFs\n        img = tiff.imread(img_path)  # H x W\n        msk = tiff.imread(lbl_path)  # H x W\n\n        # Normalize image to [0, 1] safely\n        img = img.astype(np.float32)\n        mn, mx = float(img.min()), float(img.max())\n        if mx > mn:\n            img = (img - mn) / (mx - mn)\n        else:\n            img = np.zeros_like(img, dtype=np.float32)\n\n        # Ensure binary mask {0,1}\n        msk = (msk > 0).astype(np.uint8)\n        return img, msk\n\n    def __getitem__(self, idx):\n        ds_name, img_path, lbl_path = self.items[idx]\n        img, msk = self._read_pair(img_path, lbl_path)\n\n        # Albumentations expects HWC; replicate channels\n        img3 = np.stack([img, img, img], axis=-1)\n\n        data = self.transform(image=img3, mask=msk)\n        img3 = data[\"image\"]            # H x W x 3, float32 in [0,1]\n        msk1 = data[\"mask\"].astype(np.uint8)  # H x W, {0,1}\n\n        # To tensors (CHW)\n        img_chw = np.transpose(img3, (2, 0, 1))  # 3 x H x W\n        img_t = torch.from_numpy(img_chw).float()\n        msk_t = torch.from_numpy(msk1[None, ...]).float()  # 1 x H x W\n\n        return img_t, msk_t\n\n# Simple split\nnp.random.seed(42)\nperm = np.random.permutation(len(train_items))\ncut = int(0.9*len(perm))\ntr_idx, va_idx = perm[:cut], perm[cut:]\ntrain_ds = VesselDataset([train_items[i] for i in tr_idx], aug=True)\nvalid_ds = VesselDataset([train_items[i] for i in va_idx], aug=False)\n\ntrain_dl = DataLoader(train_ds, batch_size=4, shuffle=True, num_workers=2, pin_memory=True)\nvalid_dl = DataLoader(valid_ds, batch_size=2, shuffle=False, num_workers=2, pin_memory=True)\n\n# =========================\n# U-Net (small)\n# =========================\nclass DoubleConv(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Conv2d(in_ch, out_ch, 3, padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_ch, out_ch, 3, padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True),\n        )\n    def forward(self, x): return self.net(x)\n\nclass UNet(nn.Module):\n    def __init__(self, in_ch=3, out_ch=1, chs=(32,64,128,256)):\n        super().__init__()\n        self.down1 = DoubleConv(in_ch, chs[0])\n        self.pool1 = nn.MaxPool2d(2)\n        self.down2 = DoubleConv(chs[0], chs[1])\n        self.pool2 = nn.MaxPool2d(2)\n        self.down3 = DoubleConv(chs[1], chs[2])\n        self.pool3 = nn.MaxPool2d(2)\n\n        self.bottleneck = DoubleConv(chs[2], chs[3])\n\n        self.up3 = nn.ConvTranspose2d(chs[3], chs[2], 2, 2)\n        self.conv3 = DoubleConv(chs[3], chs[2])\n        self.up2 = nn.ConvTranspose2d(chs[2], chs[1], 2, 2)\n        self.conv2 = DoubleConv(chs[2], chs[1])\n        self.up1 = nn.ConvTranspose2d(chs[1], chs[0], 2, 2)\n        self.conv1 = DoubleConv(chs[1], chs[0])\n\n        self.out = nn.Conv2d(chs[0], out_ch, 1)\n\n    def forward(self, x):\n        d1 = self.down1(x)\n        d2 = self.down2(self.pool1(d1))\n        d3 = self.down3(self.pool2(d2))\n        bn = self.bottleneck(self.pool3(d3))\n\n        u3 = self.up3(bn); u3 = torch.cat([u3, d3], dim=1); u3 = self.conv3(u3)\n        u2 = self.up2(u3); u2 = torch.cat([u2, d2], dim=1); u2 = self.conv2(u2)\n        u1 = self.up1(u2); u1 = torch.cat([u1, d1], dim=1); u1 = self.conv1(u1)\n        return self.out(u1)\n\n# =========================\n# Loss (Dice + BCE)\n# =========================\nclass DiceLoss(nn.Module):\n    def __init__(self, eps=1e-6): super().__init__(); self.eps=eps\n    def forward(self, logits, targets):\n        probs = torch.sigmoid(logits)\n        num = 2*(probs*targets).sum(dim=(2,3)) + self.eps\n        den = (probs+targets).sum(dim=(2,3)) + self.eps\n        dice = 1 - (num/den).mean()\n        return dice\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = UNet().to(device)\nbce   = nn.BCEWithLogitsLoss()\ndice  = DiceLoss()\nopt   = torch.optim.AdamW(model.parameters(), lr=1e-3)\n\n# =========================\n# Training loop\n# =========================\ndef train_epoch():\n    model.train()\n    loss_all=0\n    for x,y in tqdm(train_dl, desc=\"train\"):\n        x=x.to(device); y=y.to(device)\n        opt.zero_grad()\n        logit = model(x)\n        loss = 0.5*bce(logit,y) + 0.5*dice(logit,y)\n        loss.backward()\n        opt.step()\n        loss_all += loss.item()*x.size(0)\n    return loss_all/len(train_dl.dataset)\n\n@torch.no_grad()\ndef valid_epoch():\n    model.eval()\n    loss_all=0\n    for x,y in tqdm(valid_dl, desc=\"valid\"):\n        x=x.to(device); y=y.to(device)\n        logit = model(x)\n        loss = 0.5*bce(logit,y) + 0.5*dice(logit,y)\n        loss_all += loss.item()*x.size(0)\n    return loss_all/len(valid_dl.dataset)\n\nbest = 1e9\nfor ep in range(8):  # start with ~8 epochs; extend if you have time/GPU\n    tr = train_epoch()\n    va = valid_epoch()\n    print(f\"Epoch {ep}: train {tr:.4f} | valid {va:.4f}\")\n    if va < best:\n        best = va\n        torch.save(model.state_dict(), \"unet_best.pth\")\nprint(\"Best valid:\", best)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-20T16:15:29.145759Z","iopub.execute_input":"2025-08-20T16:15:29.146194Z","execution_failed":"2025-08-20T16:19:36.885Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- imports\nimport os, re, glob, math, csv, gc\nimport numpy as np\nimport pandas as pd\nimport tifffile as tiff\nimport cv2\nfrom tqdm import tqdm\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\n\n# =========================\n# Utils: RLE encode/decode\n# =========================\ndef rle_encode(mask):\n    # mask: 2D uint8 array with {0,1}\n    pixels = mask.flatten(order='F')  # Fortran order per common competitions\n    pads = np.pad(pixels, (1,1), mode='constant')\n    runs = np.where(pads[1:] != pads[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\n\ndef rle_decode(rle, shape):\n    if (rle is None) or (rle == ''):\n        return np.zeros(shape, dtype=np.uint8)\n    s = list(map(int, rle.split()))\n    starts, lengths = s[0::2], s[1::2]\n    starts = np.asarray(starts) - 1\n    ends = starts + lengths\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n    return img.reshape(shape, order='F')\n\n# =========================\n# Paths\n# =========================\nDATA_DIR = \"/kaggle/input/blood-vessel-segmentation\"  # change if local\nTRAIN_DIR = os.path.join(DATA_DIR, \"train\")\nTEST_DIR  = os.path.join(DATA_DIR, \"test\")\n\n# Get all datasets under train\ntrain_sets = [\n    \"kidney_1_dense\",\n    \"kidney_1_voi\",\n    \"kidney_2\",\n    \"kidney_3_dense\",\n    \"kidney_3_sparse\",\n]\n\n# =========================\n# Data indexing\n# =========================\ndef list_train_items():\n    items = []\n    for ds in train_sets:\n        img_dir = os.path.join(TRAIN_DIR, ds, \"images\")\n        lbl_dir = os.path.join(TRAIN_DIR, ds, \"labels\")\n        # kidney_3_dense only has labels (images are in kidney_3_sparse/images)\n        if ds == \"kidney_3_dense\" and not os.path.exists(img_dir):\n            img_dir = os.path.join(TRAIN_DIR, \"kidney_3_sparse\", \"images\")\n\n        if not os.path.exists(img_dir) or not os.path.exists(lbl_dir):\n            # some subsets may be missing (rare), skip safely\n            continue\n\n        for p in sorted(glob.glob(os.path.join(img_dir, \"*.tif\"))):\n            fname = os.path.basename(p)  # e.g. 0001.tif\n            lbl_p = os.path.join(lbl_dir, fname)\n            if os.path.exists(lbl_p):\n                items.append((ds, p, lbl_p))\n    return items\n\ndef list_test_items():\n    items = []\n    for ds in [\"kidney_5\",\"kidney_6\"]:\n        img_dir = os.path.join(TEST_DIR, ds, \"images\")\n        if not os.path.exists(img_dir): \n            continue\n        for p in sorted(glob.glob(os.path.join(img_dir, \"*.tif\"))):\n            fname = os.path.basename(p)\n            slice_id = f\"{ds}_{os.path.splitext(fname)[0]}\"\n            items.append((slice_id, p))\n    return items\n\ntrain_items = list_train_items()\ntest_items  = list_test_items()\nprint(f\"Train slices (all): {len(train_items)}, Test slices (all): {len(test_items)}\")\n\n# -------------------------\n# CHANGED: limit training pool to exactly 500 TIFF pairs (random, reproducible)\n# -------------------------\nnp.random.seed(42)\nif len(train_items) > 500:\n    idx = np.random.permutation(len(train_items))[:500]\n    train_items = [train_items[i] for i in idx]\nprint(f\"Train slices (limited to 500): {len(train_items)}\")\n\n# =========================\n# Dataset\n# =========================\nimport albumentations as A\n\nclass VesselDataset(Dataset):\n    \"\"\"\n    items: list of tuples -> (dataset_name, image_path, label_path)\n    aug:   bool, if True apply training augs; if False, deterministic preprocessing only\n    size:  int, the max side after LongestMaxSize (square canvas after PadIfNeeded)\n    \"\"\"\n\n    def __init__(self, items, aug=False, size=768):\n        self.items = items\n        self.aug = aug\n        self.size = int(size)\n\n        if aug:\n            # Stable on Kaggle across Albumentations versions\n            self.transform = A.Compose([\n                A.LongestMaxSize(max_size=self.size),\n                A.PadIfNeeded(min_height=self.size, min_width=self.size,\n                              border_mode=cv2.BORDER_REFLECT),\n                A.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.10, rotate_limit=15,\n                                   border_mode=cv2.BORDER_REFLECT, p=0.5),\n                A.HorizontalFlip(p=0.5),\n                A.VerticalFlip(p=0.5),\n                A.RandomBrightnessContrast(p=0.3),\n            ])\n        else:\n            self.transform = A.Compose([\n                A.LongestMaxSize(max_size=self.size),\n                A.PadIfNeeded(min_height=self.size, min_width=self.size,\n                              border_mode=cv2.BORDER_REFLECT),\n            ])\n\n    def __len__(self):\n        return len(self.items)\n\n    def _read_pair(self, img_path, lbl_path):\n        # Read grayscale TIFFs\n        img = tiff.imread(img_path)  # H x W\n        msk = tiff.imread(lbl_path)  # H x W\n\n        # Normalize image to [0, 1] safely\n        img = img.astype(np.float32)\n        mn, mx = float(img.min()), float(img.max())\n        if mx > mn:\n            img = (img - mn) / (mx - mn)\n        else:\n            img = np.zeros_like(img, dtype=np.float32)\n\n        # Ensure binary mask {0,1}\n        msk = (msk > 0).astype(np.uint8)\n        return img, msk\n\n    def __getitem__(self, idx):\n        ds_name, img_path, lbl_path = self.items[idx]\n        img, msk = self._read_pair(img_path, lbl_path)\n\n        # Albumentations expects HWC; replicate channels\n        img3 = np.stack([img, img, img], axis=-1)\n\n        data = self.transform(image=img3, mask=msk)\n        img3 = data[\"image\"]            # H x W x 3, float32 in [0,1]\n        msk1 = data[\"mask\"].astype(np.uint8)  # H x W, {0,1}\n\n        # To tensors (CHW)\n        img_chw = np.transpose(img3, (2, 0, 1))  # 3 x H x W\n        img_t = torch.from_numpy(img_chw).float()\n        msk_t = torch.from_numpy(msk1[None, ...]).float()  # 1 x H x W\n\n        return img_t, msk_t\n\n# =========================\n# Dataloaders\n# =========================\n# CHANGED: batch sizes tuned for 500-slice training (stability vs. speed)\nTRAIN_BATCH = 4   # was 2/4 before; 4 is good for 500 samples\nVALID_BATCH = 2\n\n# Simple split (90/10 on the limited 500)\nnp.random.seed(42)\nperm = np.random.permutation(len(train_items))\ncut = int(0.9*len(perm))\ntr_idx, va_idx = perm[:cut], perm[cut:]\ntrain_ds = VesselDataset([train_items[i] for i in tr_idx], aug=True,  size=768)\nvalid_ds = VesselDataset([train_items[i] for i in va_idx], aug=False, size=768)\n\ntrain_dl = DataLoader(train_ds, batch_size=TRAIN_BATCH, shuffle=True,  num_workers=2, pin_memory=True)\nvalid_dl = DataLoader(valid_ds, batch_size=VALID_BATCH, shuffle=False, num_workers=2, pin_memory=True)\n\nprint(f\"Train batches/epoch: {len(train_dl)} (batch_size={TRAIN_BATCH})\")\nprint(f\"Valid batches/epoch: {len(valid_dl)} (batch_size={VALID_BATCH})\")\n\n# =========================\n# U-Net (small)\n# =========================\nclass DoubleConv(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Conv2d(in_ch, out_ch, 3, padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_ch, out_ch, 3, padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True),\n        )\n    def forward(self, x): return self.net(x)\n\nclass UNet(nn.Module):\n    def __init__(self, in_ch=3, out_ch=1, chs=(32,64,128,256)):\n        super().__init__()\n        self.down1 = DoubleConv(in_ch, chs[0])\n        self.pool1 = nn.MaxPool2d(2)\n        self.down2 = DoubleConv(chs[0], chs[1])\n        self.pool2 = nn.MaxPool2d(2)\n        self.down3 = DoubleConv(chs[1], chs[2])\n        self.pool3 = nn.MaxPool2d(2)\n\n        self.bottleneck = DoubleConv(chs[2], chs[3])\n\n        self.up3 = nn.ConvTranspose2d(chs[3], chs[2], 2, 2)\n        self.conv3 = DoubleConv(chs[3], chs[2])\n        self.up2 = nn.ConvTranspose2d(chs[2], chs[1], 2, 2)\n        self.conv2 = DoubleConv(chs[2], chs[1])\n        self.up1 = nn.ConvTranspose2d(chs[1], chs[0], 2, 2)\n        self.conv1 = DoubleConv(chs[1], chs[0])\n\n        self.out = nn.Conv2d(chs[0], out_ch, 1)\n\n    def forward(self, x):\n        d1 = self.down1(x)\n        d2 = self.down2(self.pool1(d1))\n        d3 = self.down3(self.pool2(d2))\n        bn = self.bottleneck(self.pool3(d3))\n\n        u3 = self.up3(bn); u3 = torch.cat([u3, d3], dim=1); u3 = self.conv3(u3)\n        u2 = self.up2(u3); u2 = torch.cat([u2, d2], dim=1); u2 = self.conv2(u2)\n        u1 = self.up1(u2); u1 = torch.cat([u1, d1], dim=1); u1 = self.conv1(u1)\n        return self.out(u1)\n\n# =========================\n# Loss (Dice + BCE)\n# =========================\nclass DiceLoss(nn.Module):\n    def __init__(self, eps=1e-6): super().__init__(); self.eps=eps\n    def forward(self, logits, targets):\n        probs = torch.sigmoid(logits)\n        num = 2*(probs*targets).sum(dim=(2,3)) + self.eps\n        den = (probs+targets).sum(dim=(2,3)) + self.eps\n        dice = 1 - (num/den).mean()\n        return dice\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = UNet().to(device)\nbce   = nn.BCEWithLogitsLoss()\ndice  = DiceLoss()\nopt   = torch.optim.AdamW(model.parameters(), lr=1e-3)\n\n# =========================\n# Training loop\n# =========================\ndef train_epoch():\n    model.train()\n    loss_all=0\n    for x,y in tqdm(train_dl, desc=\"train\"):\n        x=x.to(device); y=y.to(device)\n        opt.zero_grad()\n        logit = model(x)\n        loss = 0.5*bce(logit,y) + 0.5*dice(logit,y)\n        loss.backward()\n        opt.step()\n        loss_all += loss.item()*x.size(0)\n    return loss_all/len(train_dl.dataset)\n\n@torch.no_grad()\ndef valid_epoch():\n    model.eval()\n    loss_all=0\n    for x,y in tqdm(valid_dl, desc=\"valid\"):\n        x=x.to(device); y=y.to(device)\n        logit = model(x)\n        loss = 0.5*bce(logit,y) + 0.5*dice(logit,y)\n        loss_all += loss.item()*x.size(0)\n    return loss_all/len(valid_dl.dataset)\n\n# =========================\n# Epochs\n# =========================\n# CHANGED: more epochs help when training on a small subset (500)\nEPOCHS = 10\n\nbest = 1e9\nfor ep in range(EPOCHS):\n    tr = train_epoch()\n    va = valid_epoch()\n    print(f\"Epoch {ep}: train {tr:.4f} | valid {va:.4f}\")\n    if va < best:\n        best = va\n        torch.save(model.state_dict(), \"unet_best.pth\")\nprint(\"Best valid:\", best)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-20T16:41:27.361731Z","iopub.execute_input":"2025-08-20T16:41:27.362041Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}