{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.11.13"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":99552,"databundleVersionId":13694723,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":261445297,"sourceType":"kernelVersion"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":28999.646669,"end_time":"2025-09-06T08:17:08.189823","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2025-09-06T00:13:48.543154","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"c4b066b4","cell_type":"code","source":"!pip install segmentation_models_pytorch==0.3.3\n\nimport os\nimport gc\nimport random\n\nimport numpy as np\nimport pandas as pd\nimport nibabel as nib\nfrom scipy.ndimage import gaussian_filter\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\n\nimport segmentation_models_pytorch as smp\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"papermill":{"duration":95.539163,"end_time":"2025-09-06T00:15:28.210613","exception":false,"start_time":"2025-09-06T00:13:52.67145","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"bc28cd6c","cell_type":"code","source":"FOLDS = [0]\nSEED = 777\nSIZE = 512\nEPOCHS = 50\nBS = 32\nLR = 1e-4","metadata":{"papermill":{"duration":0.028791,"end_time":"2025-09-06T00:15:28.262938","exception":false,"start_time":"2025-09-06T00:15:28.234147","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"a2c5a9a4","cell_type":"code","source":"source_path = '/kaggle/input/rsna-2d-binary-segmentation-preprocessing/'\noutput_images_dir = source_path + 'images/'\noutput_labels_dir = source_path + 'labels/'\ntrain = pd.read_csv(source_path + 'rsna_2d_seg_folds.csv')\ntrain.tail()","metadata":{"papermill":{"duration":0.2048,"end_time":"2025-09-06T00:15:28.490395","exception":false,"start_time":"2025-09-06T00:15:28.285595","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"3c148d2d","cell_type":"code","source":"def elastic_deform(\n    img: torch.Tensor,  # [1, H, W]\n    msk: torch.Tensor,  # [1, H, W]\n    alpha: float = 400.0,\n    sigma: float = 10.0,\n    grid_step: int = 16,\n):\n    _, H, W = img.shape\n    \n    # 1. Create coarse displacement grid (shape: [H//grid_step, W//grid_step])\n    grid_h, grid_w = H // grid_step, W // grid_step\n    dx = torch.randn(grid_h, grid_w) * alpha\n    dy = torch.randn(grid_h, grid_w) * alpha\n    \n    # 2. Smooth coarse displacements\n    dx = torch.from_numpy(gaussian_filter(dx.numpy(), sigma=sigma))\n    dy = torch.from_numpy(gaussian_filter(dy.numpy(), sigma=sigma))\n    \n    # 3. Upsample to full resolution using bilinear interpolation\n    dx_full = F.interpolate(\n        dx.unsqueeze(0).unsqueeze(0),  # [1, 1, grid_h, grid_w]\n        size=(H, W),\n        mode='bilinear',\n        align_corners=False\n    ).squeeze()  # [H, W]\n    \n    dy_full = F.interpolate(\n        dy.unsqueeze(0).unsqueeze(0),\n        size=(H, W),\n        mode='bilinear',\n        align_corners=False\n    ).squeeze()\n    \n    # 4. Create normalized grid (same as before)\n    x, y = torch.meshgrid(torch.arange(H), torch.arange(W), indexing='ij')\n    grid = torch.stack([y + dy_full, x + dx_full], dim=-1).float()  # [H, W, 2]\n    \n    # Normalize and apply deformation (unchanged)\n    grid[..., 0] = 2.0 * grid[..., 0] / (W - 1) - 1.0\n    grid[..., 1] = 2.0 * grid[..., 1] / (H - 1) - 1.0\n    grid = grid.unsqueeze(0)  # [1, H, W, 2]\n    \n    img_deformed = F.grid_sample(\n        img.unsqueeze(0),\n        grid,\n        mode='bilinear',\n        align_corners=False\n    ).squeeze(0)\n    \n    msk_deformed = F.grid_sample(\n        msk.unsqueeze(0).float(),\n        grid,\n        mode='nearest',\n        align_corners=False\n    ).squeeze(0).long()\n    \n    return img_deformed, msk_deformed","metadata":{"papermill":{"duration":0.033644,"end_time":"2025-09-06T00:15:28.548478","exception":false,"start_time":"2025-09-06T00:15:28.514834","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"c98cbcec","cell_type":"code","source":"class SliceDatasetNPY(Dataset):\n    def __init__(\n        self,\n        fname,\n        pmin,\n        pmax,\n        h,\n        w,\n        h0,\n        w0,\n        VALID=False\n    ):\n        \"\"\"\n        PyTorch Dataset for loading 2D slices from 3D medical volumes (numpy format) with \n        on-the-fly augmentation for intracranial aneurysm segmentation.\n    \n        Features:\n        - Loads preprocessed .npy slices and corresponding masks\n        - Normalizes intensities using case-specific percentile values (pmin/pmax)\n        - Training mode includes:\n          * Random asymmetric zoom/crop\n          * Full rotation (0-360°)\n          * Elastic deformations\n          * Intensity variations\n        - Validation mode uses center-cropping around mask centroid\n        - Automatic resizing to 512x512 resolution\n        - Supports all orthogonal views (YX, YZ, XZ)\n    \n        Args:\n            fname: List of numpy filenames\n            pmin: List of 1st percentile values for normalization\n            pmax: List of 99th percentile values for normalization\n            VALID: If True, disables augmentations for validation\n        \"\"\"\n        self.fname = fname\n        self.pmin = pmin\n        self.pmax = pmax\n        self.h = h\n        self.w = w\n        self.h0 = h0\n        self.w0 = w0\n        self.VALID = VALID\n        self.img_resize = transforms.Resize((SIZE, SIZE), interpolation=transforms.InterpolationMode.BILINEAR)\n        self.msk_resize = transforms.Resize((SIZE, SIZE), interpolation=transforms.InterpolationMode.NEAREST)\n\n    def __len__(self):\n        return len(self.fname)\n\n    def __getitem__(self, idx):\n        fname = self.fname[idx]\n        img = torch.from_numpy(np.load(output_images_dir+fname)).unsqueeze(0).float()\n        H,W = img.shape[1:]\n        msk = torch.zeros((1,H,W), dtype=torch.long)\n        h = self.h[idx]\n        w = self.w[idx]\n        h0 = self.h0[idx]\n        w0 = self.w0[idx]\n        msk[0,h0:h0+h,w0:w0+w] = torch.from_numpy(np.load(output_labels_dir+fname)).long()\n\n        imin = self.pmin[idx]\n        imax = self.pmax[idx]\n        img = (img - imin) / (imax - imin + 1e-6)\n\n        if H > W:\n            h = W\n            w = W\n        else:\n            h = H\n            w = H\n\n        if not self.VALID:\n#           Random asymmetric zoom\n            h = np.rint(h*(.8 + 0.4*np.random.rand())).astype(int)\n            w = np.rint(w*(.8 + 0.4*np.random.rand())).astype(int)\n            \n            if H > h:\n                h0 = np.random.randint(H - h)\n                pad_h = 0\n            else:\n                h0 = 0\n                pad_h = h - H\n            if W > w:\n                w0 = np.random.randint(W - w)\n                pad_w = 0\n            else:\n                w0 = 0\n                pad_w = w - H\n\n            img = torch.nn.functional.pad(\n                img,\n                (pad_w//2,pad_w - pad_w//2,pad_h//2,pad_h - pad_h//2)\n            )\n            msk = torch.nn.functional.pad(\n                msk,\n                (pad_w//2,pad_w - pad_w//2,pad_h//2,pad_h - pad_h//2)\n            )\n#           Free rotation\n            angle = 360*np.random.rand() - 180\n            center = [w0 + w//2,h0 + h//2]\n            img = transforms.functional.rotate(\n                img,\n                angle,\n                transforms.InterpolationMode.BILINEAR,\n                center=center\n            )\n            msk = transforms.functional.rotate(\n                msk,\n                angle,\n                transforms.InterpolationMode.NEAREST,\n                center=center\n            )\n\n        else:\n            hh,ww = torch.where(msk[0])\n            h0 = np.rint(hh.float().mean()).long() - h//2\n            w0 = np.rint(ww.float().mean()).long() - w//2\n            if h0 < 0: h0 = 0\n            if w0 < 0: w0 = 0\n            if h0 > H - h: h0 = H - h\n            if w0 > W - w: w0 = W - w\n\n        img = self.img_resize(img[:,h0:h0 + h,w0:w0 + w])\n        msk = self.msk_resize(msk[:,h0:h0 + h,w0:w0 + w])\n\n        if not self.VALID:\n#           Intensity Inversion\n            if np.random.rand() < .1:\n                img = 1 - img\n#           Elastic deformation\n            if np.random.rand() < .5:\n                img,msk = elastic_deform(img,msk)\n#           Random flip            \n            if np.random.rand() < .5:\n                img = img.flip(-1)\n                msk = msk.flip(-1)\n#           Contrast\n            if np.random.rand() < .5:\n                f = .9 + .2*torch.rand(1)\n                img *= f - (f - 1)/2\n#           Brightness\n            if np.random.rand() < .5:\n                img += .2*torch.rand(1) - .1\n#           Gaussian Noise\n            if np.random.rand() < .5:\n                img += torch.normal(torch.tensor(0.),torch.tensor(.05),(1,SIZE,SIZE))\n        \n        return img, msk[0].long()","metadata":{"papermill":{"duration":0.039563,"end_time":"2025-09-06T00:15:28.612277","exception":false,"start_time":"2025-09-06T00:15:28.572714","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"b8cf2ea5","cell_type":"code","source":"ds = SliceDatasetNPY(\n    fname=train['fname'],\n    pmin=train['pmin'],\n    pmax=train['pmax'],\n    h=train['h'],\n    w=train['w'],\n    h0=train['h0'],\n    w0=train['w0']\n)","metadata":{"papermill":{"duration":0.031505,"end_time":"2025-09-06T00:15:28.667787","exception":false,"start_time":"2025-09-06T00:15:28.636282","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"9e219b0d","cell_type":"code","source":"img,msk = ds.__getitem__(np.random.randint(len(ds)))\nplt.imshow(img[0] + msk)","metadata":{"papermill":{"duration":0.417485,"end_time":"2025-09-06T00:15:29.16542","exception":false,"start_time":"2025-09-06T00:15:28.747935","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"24d305e9","cell_type":"code","source":"del ds\ngc.collect()","metadata":{"papermill":{"duration":0.267719,"end_time":"2025-09-06T00:15:29.460629","exception":false,"start_time":"2025-09-06T00:15:29.19291","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"4bfeb018","cell_type":"code","source":"ds = SliceDatasetNPY(\n    fname=train['fname'],\n    pmin=train['pmin'],\n    pmax=train['pmax'],\n    h=train['h'],\n    w=train['w'],\n    h0=train['h0'],\n    w0=train['w0'],\n    VALID=True\n)","metadata":{"papermill":{"duration":0.032394,"end_time":"2025-09-06T00:15:29.519594","exception":false,"start_time":"2025-09-06T00:15:29.4872","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"39571356","cell_type":"code","source":"img,msk = ds.__getitem__(np.random.randint(len(ds)))\nplt.imshow(img[0] + msk)","metadata":{"papermill":{"duration":0.246569,"end_time":"2025-09-06T00:15:29.792802","exception":false,"start_time":"2025-09-06T00:15:29.546233","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"41c5dd64","cell_type":"code","source":"del ds\ngc.collect()","metadata":{"papermill":{"duration":0.265498,"end_time":"2025-09-06T00:15:30.087889","exception":false,"start_time":"2025-09-06T00:15:29.822391","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"bd590b23","cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False","metadata":{"papermill":{"duration":0.033404,"end_time":"2025-09-06T00:15:30.149761","exception":false,"start_time":"2025-09-06T00:15:30.116357","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"34bfc291","cell_type":"code","source":"# https://github.com/shuaizzZ/Dice-Loss-PyTorch/blob/master/dice_loss.py\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\n\n\nclass DiceLoss(nn.Module):\n    \"\"\"Dice Loss PyTorch\n        Created by: Zhang Shuai\n        Email: shuaizzz666@gmail.com\n        dice_loss = 1 - 2*p*t / (p^2 + t^2). p and t represent predict and target.\n    Args:\n        weight: An array of shape [C,]\n        predict: A float32 tensor of shape [N, C, *], for Semantic segmentation task is [N, C, H, W]\n        target: A int64 tensor of shape [N, *], for Semantic segmentation task is [N, H, W]\n    Return:\n        diceloss\n    \"\"\"\n    def __init__(self, weight=None):\n        super(DiceLoss, self).__init__()\n        if weight is not None:\n            weight = torch.Tensor(weight)\n            self.weight = weight / torch.sum(weight) # Normalized weight\n        self.smooth = 1e-5\n\n    def forward(self, predict, target):\n        N, C = predict.size()[:2]\n        predict = predict.view(N, C, -1) # (N, C, *)\n        target = target.view(N, 1, -1) # (N, 1, *)\n\n        predict = F.softmax(predict, dim=1) # (N, C, *) ==> (N, C, *)\n        ## convert target(N, 1, *) into one hot vector (N, C, *)\n        target_onehot = torch.zeros(predict.size()).cuda()  # (N, 1, *) ==> (N, C, *)\n        target_onehot.scatter_(1, target, 1)  # (N, C, *)\n\n        intersection = torch.sum(predict * target_onehot, dim=2)  # (N, C)\n        union = torch.sum(predict.pow(2), dim=2) + torch.sum(target_onehot, dim=2)  # (N, C)\n        ## p^2 + t^2 >= 2*p*t, target_onehot^2 == target_onehot\n        dice_coef = (2 * intersection + self.smooth) / (union + self.smooth)  # (N, C)\n\n        if hasattr(self, 'weight'):\n            if self.weight.type() != predict.type():\n                self.weight = self.weight.type_as(predict)\n            dice_coef = dice_coef * self.weight * C  # (N, C)\n        dice_loss = 1 - torch.mean(dice_coef)  # 1\n\n        return dice_loss","metadata":{"papermill":{"duration":0.053574,"end_time":"2025-09-06T00:15:30.237106","exception":false,"start_time":"2025-09-06T00:15:30.183532","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"adb83012","cell_type":"code","source":"# https://github.com/AdeelH/pytorch-multi-class-focal-loss\nfrom typing import Optional, Sequence\n\nimport torch\nfrom torch import Tensor\nfrom torch import nn\nfrom torch.nn import functional as F\n\n\nclass FocalLoss(nn.Module):\n    \"\"\" Focal Loss, as described in https://arxiv.org/abs/1708.02002.\n\n    It is essentially an enhancement to cross entropy loss and is\n    useful for classification tasks when there is a large class imbalance.\n    x is expected to contain raw, unnormalized scores for each class.\n    y is expected to contain class labels.\n\n    Shape:\n        - x: (batch_size, C) or (batch_size, C, d1, d2, ..., dK), K > 0.\n        - y: (batch_size,) or (batch_size, d1, d2, ..., dK), K > 0.\n    \"\"\"\n\n    def __init__(self,\n                 alpha: Optional[Tensor] = None,\n                 gamma: float = 0.,\n                 reduction: str = 'mean',\n                 ignore_index: int = -100):\n        \"\"\"Constructor.\n\n        Args:\n            alpha (Tensor, optional): Weights for each class. Defaults to None.\n            gamma (float, optional): A constant, as described in the paper.\n                Defaults to 0.\n            reduction (str, optional): 'mean', 'sum' or 'none'.\n                Defaults to 'mean'.\n            ignore_index (int, optional): class label to ignore.\n                Defaults to -100.\n        \"\"\"\n        if reduction not in ('mean', 'sum', 'none'):\n            raise ValueError(\n                'Reduction must be one of: \"mean\", \"sum\", \"none\".')\n\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.ignore_index = ignore_index\n        self.reduction = reduction\n\n        self.nll_loss = nn.NLLLoss(\n            weight=alpha, reduction='none', ignore_index=ignore_index)\n\n    def __repr__(self):\n        arg_keys = ['alpha', 'gamma', 'ignore_index', 'reduction']\n        arg_vals = [self.__dict__[k] for k in arg_keys]\n        arg_strs = [f'{k}={v!r}' for k, v in zip(arg_keys, arg_vals)]\n        arg_str = ', '.join(arg_strs)\n        return f'{type(self).__name__}({arg_str})'\n\n    def forward(self, x: Tensor, y: Tensor) -> Tensor:\n        if x.ndim > 2:\n            # (N, C, d1, d2, ..., dK) --> (N * d1 * ... * dK, C)\n            c = x.shape[1]\n            x = x.permute(0, *range(2, x.ndim), 1).reshape(-1, c)\n            # (N, d1, d2, ..., dK) --> (N * d1 * ... * dK,)\n            y = y.view(-1)\n\n        unignored_mask = y != self.ignore_index\n        y = y[unignored_mask]\n        if len(y) == 0:\n            return torch.tensor(0.)\n        x = x[unignored_mask]\n\n        # compute weighted cross entropy term: -alpha * log(pt)\n        # (alpha is already part of self.nll_loss)\n        log_p = F.log_softmax(x, dim=-1)\n        ce = self.nll_loss(log_p, y)\n\n        # get true class column from each row\n        all_rows = torch.arange(len(x))\n        log_pt = log_p[all_rows, y]\n\n        # compute focal term: (1 - pt)^gamma\n        pt = log_pt.exp()\n        focal_term = (1 - pt)**self.gamma\n\n        # the full loss: -alpha * ((1 - pt)^gamma) * log(pt)\n        loss = focal_term * ce\n\n        if self.reduction == 'mean':\n            loss = loss.mean()\n        elif self.reduction == 'sum':\n            loss = loss.sum()\n\n        return loss\n\n\ndef focal_loss(alpha: Optional[Sequence] = None,\n               gamma: float = 0.,\n               reduction: str = 'mean',\n               ignore_index: int = -100,\n               device='cpu',\n               dtype=torch.float32) -> FocalLoss:\n    \"\"\"Factory function for FocalLoss.\n\n    Args:\n        alpha (Sequence, optional): Weights for each class. Will be converted\n            to a Tensor if not None. Defaults to None.\n        gamma (float, optional): A constant, as described in the paper.\n            Defaults to 0.\n        reduction (str, optional): 'mean', 'sum' or 'none'.\n            Defaults to 'mean'.\n        ignore_index (int, optional): class label to ignore.\n            Defaults to -100.\n        device (str, optional): Device to move alpha to. Defaults to 'cpu'.\n        dtype (torch.dtype, optional): dtype to cast alpha to.\n            Defaults to torch.float32.\n\n    Returns:\n        A FocalLoss object\n    \"\"\"\n    if alpha is not None:\n        if not isinstance(alpha, Tensor):\n            alpha = torch.tensor(alpha)\n        alpha = alpha.to(device=device, dtype=dtype)\n\n    fl = FocalLoss(\n        alpha=alpha,\n        gamma=gamma,\n        reduction=reduction,\n        ignore_index=ignore_index)\n    return fl","metadata":{"papermill":{"duration":0.041529,"end_time":"2025-09-06T00:15:30.322557","exception":false,"start_time":"2025-09-06T00:15:30.281028","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"e4857896","cell_type":"code","source":"class DiceFocalLoss(nn.Module):\n    def __init__(self, gamma=1.0, alpha=0.5):\n        super().__init__()\n        self.DL = DiceLoss()\n        self.FL = FocalLoss(gamma=gamma)\n        self.alpha = alpha\n\n    def forward(self, inputs, targets):\n        DL = self.DL(inputs,targets)\n        FL = self.FL(inputs,targets)\n\n        return self.alpha * DL + (1 - self.alpha) * FL","metadata":{"papermill":{"duration":0.034208,"end_time":"2025-09-06T00:15:30.384865","exception":false,"start_time":"2025-09-06T00:15:30.350657","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"81fc6160","cell_type":"code","source":"criterion = DiceFocalLoss()\n\nfor fold in FOLDS:\n    seed_everything(SEED)\n    model = smp.Unet(\n        encoder_name=\"resnet18\",\n        encoder_weights=\"imagenet\",\n        in_channels=1,\n        classes=2\n    )\n    model = model.to(device)\n    optimizer = optim.Adam(model.parameters(), lr=LR)\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min', patience=3, factor=0.5, min_lr=1e-6)\n\n    train_df = train[train['fold'] != fold].reset_index(drop=True)\n    valid_df = train[train['fold'] == fold].reset_index(drop=True)\n\n    train_dataset = SliceDatasetNPY(\n        train_df['fname'],\n        train_df['pmin'],\n        train_df['pmax'],\n        h=train_df['h'],\n        w=train_df['w'],\n        h0=train_df['h0'],\n        w0=train_df['w0']\n    )\n    val_dataset = SliceDatasetNPY(\n        valid_df['fname'],\n        valid_df['pmin'],\n        valid_df['pmax'],\n        h=valid_df['h'],\n        w=valid_df['w'],\n        h0=valid_df['h0'],\n        w0=valid_df['w0'],\n        VALID=True\n    )\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=BS,\n        shuffle=True,\n        num_workers=4,\n        pin_memory=True\n    )\n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=BS,\n        shuffle=False,\n        num_workers=4,\n        pin_memory=True\n    )\n\n    train_loss_history = []\n    val_loss_history = []\n    best_val_loss = float('inf')\n\n#   Training loop\n    for epoch in range(EPOCHS):\n        model.train()\n        epoch_train_loss = 0.0\n    \n#       Training phase\n        for images, masks in tqdm(train_loader, desc=f'Epoch {epoch+1}/{EPOCHS}'):\n            images = images.to(device, non_blocking=True)\n            masks = masks.to(device, non_blocking=True)\n        \n            optimizer.zero_grad()\n        \n            outputs = model(images)\n            loss = criterion(outputs, masks)\n        \n            loss.backward()\n            optimizer.step()\n        \n            epoch_train_loss += loss.item() * images.size(0)\n    \n#       Validation phase\n        model.eval()\n        epoch_val_loss = 0.0\n        with torch.no_grad():\n            for images, masks in tqdm(val_loader, desc=f'Epoch {epoch+1}/{EPOCHS}'):\n                images = images.to(device, non_blocking=True)\n                masks = masks.to(device, non_blocking=True)\n            \n                outputs = model(images)\n                loss = criterion(outputs, masks)\n            \n                epoch_val_loss += loss.item() * images.size(0)\n    \n#       Calculate epoch metrics\n        epoch_train_loss /= len(train_loader.dataset)\n        epoch_val_loss /= len(val_loader.dataset)\n    \n        train_loss_history.append(epoch_train_loss)\n        val_loss_history.append(epoch_val_loss)\n    \n#       Update learning rate\n        scheduler.step(epoch_val_loss)\n    \n#       Save best model\n        if epoch_val_loss < best_val_loss:\n            best_val_loss = epoch_val_loss\n            torch.save(model.state_dict(), f'best_model_{fold}.pth')\n    \n        print(f'Epoch {epoch+1}/{EPOCHS} - '\n              f'Train Loss: {epoch_train_loss:.4f} - '\n              f'Val Loss: {epoch_val_loss:.4f} - '\n              f'LR: {optimizer.param_groups[0][\"lr\"]:.2e}')\n\n#   Plot training history\n    plt.figure(figsize=(10, 5))\n    plt.plot(train_loss_history, label='Train Loss')\n    plt.plot(val_loss_history, label='Validation Loss')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.legend()\n    plt.title('Training History')\n    plt.savefig(f'training_history_{fold}.png')\n    plt.show()\n\n    del model,optimizer,scheduler,train_dataset,val_dataset,train_loader,val_loader\n    gc.collect()","metadata":{"papermill":{"duration":28893.454034,"end_time":"2025-09-06T08:17:03.86634","exception":false,"start_time":"2025-09-06T00:15:30.412306","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null}]}