{"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":13762876,"sourceType":"competition"},{"sourceId":261464046,"sourceType":"kernelVersion"},{"sourceId":261515368,"sourceType":"kernelVersion"}],"dockerImageVersionId":31090,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":34044.927269,"end_time":"2025-09-16T08:20:54.826645","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2025-09-15T22:53:29.899376","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"a2b82177","cell_type":"code","source":"!pip install segmentation_models_pytorch==0.3.3\n\nimport os\nimport gc\nimport random\nimport math\nimport pickle\n\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\nfrom tqdm import tqdm\nfrom sklearn.metrics import confusion_matrix\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nimport torchvision.transforms as transforms\n\nfrom torch.utils.data import Dataset, DataLoader\n\nimport segmentation_models_pytorch as smp\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"papermill":{"duration":95.93095,"end_time":"2025-09-15T22:55:09.885783","exception":false,"start_time":"2025-09-15T22:53:33.954833","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"02b2a3e6","cell_type":"code","source":"FOLDS = [0]\nSEED = 777\nSIZE = 128\nANGLE_AUG = 30\nEPOCHS = 20\nBS = 2\nLR = 5e-5\nWORKERS = 4\nROTATION_PROB_DECAY = .5699 # (1 - p)**2 = p**3 | Equals single and triple rotations during augmentations","metadata":{"papermill":{"duration":0.029289,"end_time":"2025-09-15T22:55:09.939272","exception":false,"start_time":"2025-09-15T22:55:09.909983","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"4edfe19d","cell_type":"code","source":"source_path = '/kaggle/input/rsna-3d-full-segmentation-preprocessing/'\ncases = [c[:-4] for c in os.listdir(source_path + 'images')]\ncases[:5]","metadata":{"papermill":{"duration":0.057976,"end_time":"2025-09-15T22:55:10.020334","exception":false,"start_time":"2025-09-15T22:55:09.962358","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"232d1d9d","cell_type":"code","source":"train = pd.read_csv(source_path + 'rsna_3d_seg_folds.csv')\nwith open(source_path + 'centroids.pkl', 'rb') as f:\n    centroids = pickle.load(f)\n\ntrain.tail()","metadata":{"papermill":{"duration":0.088963,"end_time":"2025-09-15T22:55:10.13328","exception":false,"start_time":"2025-09-15T22:55:10.044317","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"5b864239","cell_type":"code","source":"train['flip'] = False\nftrain = train.copy()\nftrain['flip'] = True\ntrain = pd.concat([train,ftrain]).reset_index(drop=True)\ntrain.tail()","metadata":{"papermill":{"duration":0.047531,"end_time":"2025-09-15T22:55:10.205151","exception":false,"start_time":"2025-09-15T22:55:10.15762","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"9688ec3a","cell_type":"code","source":"class RSNA_Dataset_3D(Dataset):\n    def __init__(\n        self,\n        df,\n        VALID=False\n    ):\n        \"\"\"\n        A PyTorch Dataset for 3D semantic segmentation of intracranial arteries from angiographic scans.\n\n        This dataset handles preprocessed 3D volumes from the RSNA challenge. The task involves multi-class\n        segmentation to identify and label 13 individual intracranial artery segments (classes 1-13) against\n        the background (class 0).\n\n        The data has been preprocessed from original NIfTI files into 3D numpy arrays:\n        - Images: Angiographic scans (e.g., CTA or MRA) representing vascular intensity\n        - Labels: Multi-class masks where each voxel is assigned a class (0-13)\n\n        The dataset performs extensive on-the-fly augmentation during training, including:\n        - Random asymmetric zoom and cropping around each class centroid (for the specified target)\n        - 3D rotation and flipping (with symmetric label remapping for paired vessels)\n        - Intensity adjustments (inversion, contrast, brightness, noise)\n\n        For validation, a deterministic center-crop around the target centroid is used for consistent evaluation.\n\n        Args:\n            cases: List of case identifiers (SeriesInstanceUID)\n            targets: The specific target class label around which to center the crop\n            pmin: Minimum intensity value for normalization per case\n            pmax: Maximum intensity value for normalization per case\n            flip: Boolean flag indicating whether to apply left-right flipping\n            d0, h0, w0: Original segmentation mask offset coordinates in the full volume\n            d, h, w: Original segmentation mask dimensions in the full volume\n            VALID: If True, disables augmentation for validation/inference mode\n        \"\"\"\n        if VALID:\n            self.df = df[df['target'].apply(lambda v:v in [\n                1, # Other Posterior Circulation\n                2, # Basilar Tip\n                5, # Right Infraclinoid Internal Carotid Artery\n                6, # Left Infraclinoid Internal Carotid Artery\n                9, # Right Middle Cerebral Artery\n                10,# Left Middle Cerebral Artery\n                13 # Anterior Communicating Artery\n            ])].reset_index(drop=True)\n        else:\n            self.df = df\n        self.VALID = VALID\n\n    def __len__(self):\n        return len(self.df)\n        \n    def __getitem__(self, idx):        \n        case = self.df['SeriesInstanceUID'][idx]\n        target = self.df['target'][idx]\n        pmin = self.df['pmin'][idx]\n        pmax = self.df['pmax'][idx]\n        f = self.df['flip'][idx]\n\n        img = torch.from_numpy(np.load(source_path + 'images/' + case + '.npy')).float()\n        msk = torch.zeros_like(img)\n        d = self.df['d'][idx]\n        h = self.df['h'][idx]\n        w = self.df['w'][idx]\n        r = (d*h*w)**(1./3)\n        d0 = self.df['d0'][idx]\n        h0 = self.df['h0'][idx]\n        w0 = self.df['w0'][idx]\n        msk[d0:d0+d,h0:h0+h,w0:w0+w] = torch.from_numpy(np.load(source_path + 'labels/' + case + '.npy'))\n\n        D,H,W = img.shape\n\n        if not self.VALID:\n            z,y,x = centroids[case][target]\n            d,h,w = np.rint(centroids[case][target]).astype(int)\n#           Random asymmetric zoom\n            rd = r*(.9 + 0.2*np.random.rand())\n            rh = r*(.9 + 0.2*np.random.rand())\n            rw = r*(.9 + 0.2*np.random.rand())\n\n            dd = np.rint(rd).astype(int)\n            hh = np.rint(rh).astype(int)\n            ww = np.rint(rw).astype(int)\n\n            ddd = np.rint(rd/2).astype(int)\n            hhh = np.rint(rh/2).astype(int)\n            www = np.rint(rw/2).astype(int)\n\n            dd += ddd\n            hh += hhh\n            ww += www\n\n            d += np.random.randint(ddd//2) - ddd//4\n            h += np.random.randint(hhh//2) - hhh//4\n            w += np.random.randint(www//2) - www//4\n\n            d0 = d - dd//2\n            h0 = h - hh//2\n            w0 = w - ww//2\n\n            d = d0 + dd\n            h = h0 + hh\n            w = w0 + ww\n\n            pad_d0 = max(0, -d0)\n            pad_h0 = max(0, -h0)\n            pad_w0 = max(0, -w0)\n        \n            pad_d = max(0, d - D)\n            pad_h = max(0, h - H)\n            pad_w = max(0, w - W)\n        \n            d0 = max(0, d0)\n            h0 = max(0, h0)\n            w0 = max(0, w0)\n        \n            d = min(D, d)\n            h = min(H, h)\n            w = min(W, w)\n                \n            img = torch.nn.functional.pad(\n                (img[d0:d,h0:h,w0:w] - pmin)/(pmax - pmin),\n                (pad_w0,pad_w,pad_h0,pad_h,pad_d0,pad_d)\n            )\n            msk = torch.nn.functional.pad(\n                msk[d0:d,h0:h,w0:w],\n                (pad_w0,pad_w,pad_h0,pad_h,pad_d0,pad_d)\n            )\n#           Free rotation\n            p = 1\n            center = np.array([\n                z - d0 + pad_d0,\n                y - h0 + pad_h0,\n                x - w0 + pad_w0\n            ])\n            axis_perm = [\n                [2,0,1],\n                [1,2,0]\n            ][np.random.randint(2)]\n            for _ in range(3):\n                if np.random.rand() < p:\n                    c = list(center[1:][::-1])\n                    angle = 2*ANGLE_AUG*(np.random.rand() - .5)\n                    img = transforms.functional.rotate(\n                        img,\n                        angle,\n                        transforms.InterpolationMode.BILINEAR,\n                        center=c\n                    )\n                    msk = transforms.functional.rotate(\n                        msk,\n                        angle,\n                        transforms.InterpolationMode.NEAREST,\n                        center=c\n                    )\n                    p *= ROTATION_PROB_DECAY\n#               Rotate axis\n                img = img.permute(axis_perm)\n                msk = msk.permute(axis_perm)\n                center = center[axis_perm]\n                \n            img = img[ddd//2:dd-ddd//2,hhh//2:hh-hhh//2,www//2:ww-www//2]\n            msk = msk[ddd//2:dd-ddd//2,hhh//2:hh-hhh//2,www//2:ww-www//2]\n\n        else:\n            d,h,w = np.rint(centroids[case][target]).astype(int)\n            dd = hh = ww = np.rint(r).astype(int)\n\n            d0 = d - dd//2\n            h0 = h - hh//2\n            w0 = w - ww//2\n\n            d = d0 + dd\n            h = h0 + hh\n            w = w0 + ww\n\n            pad_d0 = max(0, -d0)\n            pad_h0 = max(0, -h0)\n            pad_w0 = max(0, -w0)\n        \n            pad_d = max(0, d - D)\n            pad_h = max(0, h - H)\n            pad_w = max(0, w - W)\n        \n            d0 = max(0, d0)\n            h0 = max(0, h0)\n            w0 = max(0, w0)\n        \n            d = min(D, d)\n            h = min(H, h)\n            w = min(W, w)\n                \n            img = torch.nn.functional.pad(\n                (img[d0:d,h0:h,w0:w] - pmin)/(pmax - pmin),\n                (pad_w0,pad_w,pad_h0,pad_h,pad_d0,pad_d)\n            )\n            msk = torch.nn.functional.pad(\n                msk[d0:d,h0:h,w0:w],\n                (pad_w0,pad_w,pad_h0,pad_h,pad_d0,pad_d)\n            )\n        img = F.interpolate(\n            img.unsqueeze(0).unsqueeze(0),\n            size=(SIZE,SIZE,SIZE),\n            mode='trilinear',\n            align_corners=False\n        )[0]\n        msk = F.interpolate(\n            msk.unsqueeze(0).unsqueeze(0),\n            size=(SIZE,SIZE,SIZE),\n            mode='nearest'\n        )[0]\n\n        if not self.VALID:\n#           Intensity Inversion\n            if np.random.rand() < .1:\n                img = 1 - img\n#           Contrast\n            if np.random.rand() < .5:\n                x = .9 + .2*torch.rand(1)\n                img *= x - (x - 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,SIZE))\n#       Flip            \n        if f:\n            img = img.flip(-1)\n            msk = msk.flip(-1)\n            for t in range(3,12,2):\n                mr = msk == t\n                ml = msk == t + 1\n                msk[mr] = t + 1\n                msk[ml] = t\n        \n        return img, msk[0].long()","metadata":{"papermill":{"duration":0.046602,"end_time":"2025-09-15T22:55:10.275971","exception":false,"start_time":"2025-09-15T22:55:10.229369","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"cd1bda6d","cell_type":"code","source":"ds = RSNA_Dataset_3D(train)","metadata":{"papermill":{"duration":0.030694,"end_time":"2025-09-15T22:55:10.378219","exception":false,"start_time":"2025-09-15T22:55:10.347525","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"55ddfda6","cell_type":"code","source":"img,msk = ds.__getitem__(np.random.randint(len(ds)))\nYX = msk.max(0)[0]\nZX = msk.max(1)[0]\nZY = msk.max(2)[0]\nYX[0,:14] = ZX[0,:14] = ZY[0,:14] = torch.arange(14)\n_, axs = plt.subplots(2, 3)\naxs[0,0].imshow(img[0].max(0)[0])\naxs[0,1].imshow(img[0].max(1)[0])\naxs[0,2].imshow(img[0].max(2)[0])\naxs[1,0].imshow(YX,cmap='turbo')\naxs[1,1].imshow(ZX,cmap='turbo')\naxs[1,2].imshow(ZY,cmap='turbo')\nplt.show()\nplt.imshow(np.arange(14).reshape(1,14),cmap='turbo')\nplt.gca().get_yaxis().set_visible(False)\nplt.xticks(ticks=np.arange(14), labels=np.arange(14))\nplt.show()","metadata":{"papermill":{"duration":2.231914,"end_time":"2025-09-15T22:55:12.634306","exception":false,"start_time":"2025-09-15T22:55:10.402392","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"88127487","cell_type":"code","source":"del ds\ngc.collect()","metadata":{"papermill":{"duration":0.27162,"end_time":"2025-09-15T22:55:12.933169","exception":false,"start_time":"2025-09-15T22:55:12.661549","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"e6c00e82","cell_type":"code","source":"ds = RSNA_Dataset_3D(train, VALID=True)","metadata":{"papermill":{"duration":0.037621,"end_time":"2025-09-15T22:55:12.998066","exception":false,"start_time":"2025-09-15T22:55:12.960445","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"4e6f87da","cell_type":"code","source":"img,msk = ds.__getitem__(np.random.randint(len(ds)))\nYX = msk.max(0)[0]\nZX = msk.max(1)[0]\nZY = msk.max(2)[0]\nYX[0,:14] = ZX[0,:14] = ZY[0,:14] = torch.arange(14)\n_, axs = plt.subplots(2, 3)\naxs[0,0].imshow(img[0].max(0)[0])\naxs[0,1].imshow(img[0].max(1)[0])\naxs[0,2].imshow(img[0].max(2)[0])\naxs[1,0].imshow(YX,cmap='turbo')\naxs[1,1].imshow(ZX,cmap='turbo')\naxs[1,2].imshow(ZY,cmap='turbo')\nplt.show()\nplt.imshow(np.arange(14).reshape(1,14),cmap='turbo')\nplt.gca().get_yaxis().set_visible(False)\nplt.xticks(ticks=np.arange(14), labels=np.arange(14))\nplt.show()","metadata":{"papermill":{"duration":1.152977,"end_time":"2025-09-15T22:55:14.178123","exception":false,"start_time":"2025-09-15T22:55:13.025146","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"a1c4f432","cell_type":"code","source":"del ds\ngc.collect()","metadata":{"papermill":{"duration":0.278721,"end_time":"2025-09-15T22:55:14.487394","exception":false,"start_time":"2025-09-15T22:55:14.208673","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"851f62b2","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.034392,"end_time":"2025-09-15T22:55:14.550718","exception":false,"start_time":"2025-09-15T22:55:14.516326","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"d82fdb4d","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.037102,"end_time":"2025-09-15T22:55:14.616554","exception":false,"start_time":"2025-09-15T22:55:14.579452","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"136f59bb","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.040304,"end_time":"2025-09-15T22:55:14.684769","exception":false,"start_time":"2025-09-15T22:55:14.644465","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"0fb4d163","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.034305,"end_time":"2025-09-15T22:55:14.74749","exception":false,"start_time":"2025-09-15T22:55:14.713185","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"bcaea337","cell_type":"code","source":"def convert_2d_to_3d(model, is_top_level=True):\n    for name, module in model.named_children():\n        # Recursively convert child modules first\n        convert_2d_to_3d(module, False)\n\n        # Replace Conv2d with Conv3d\n        if isinstance(module, nn.Conv2d):\n            # Handle cases where kernel_size/stride/padding are ints (not tuples)\n            kernel_size = module.kernel_size[0]\n            stride = module.stride[0]\n            padding = module.padding[0]\n\n            # New Conv3d layer with expanded kernel\n            new_conv = nn.Conv3d(\n                in_channels=module.in_channels,\n                out_channels=module.out_channels,\n                kernel_size=kernel_size,\n                stride=stride,\n                padding=padding,\n                bias=False\n            )\n            \n            # Initialize weights: tile 2D weights along depth and average\n            weight_2d = module.weight.data\n            weight_3d = weight_2d.unsqueeze(2).repeat(1, 1, kernel_size, 1, 1) / kernel_size\n            new_conv.weight.data = weight_3d\n            \n            setattr(model, name, new_conv)\n\n        # Replace BatchNorm2d with BatchNorm3d\n        elif isinstance(module, nn.BatchNorm2d):\n            new_bn = nn.BatchNorm3d(\n                num_features=module.num_features,\n                eps=module.eps,\n                momentum=module.momentum,\n                affine=module.affine,\n                track_running_stats=module.track_running_stats\n            ).to(device)\n            # Copy existing parameters\n            new_bn.load_state_dict(module.state_dict())\n            setattr(model, name, new_bn)\n\n        # Replace MaxPool2d with MaxPool3d (anisotropic)\n        elif isinstance(module, nn.MaxPool2d):\n            # Handle int vs. tuple for kernel_size, stride, padding\n            kernel_size = module.kernel_size\n            stride = module.stride\n            padding = module.padding\n\n            new_pool = nn.MaxPool3d(\n                kernel_size=kernel_size,\n                stride=stride,\n                padding=padding,\n                dilation=1,\n                ceil_mode=False\n            )\n            setattr(model, name, new_pool)\n\n    if is_top_level and hasattr(model, 'segmentation_head'):\n        old_weight = model.segmentation_head[0].weight.data\n        new_weight = torch.cat([\n            old_weight[0:1],  # Class 0 (foreground, unchanged)\n            old_weight[1:2].repeat(13, 1, *[1] * (old_weight.dim() - 2))  # Repeat Class 1 for 13 new positives\n        ], dim=0)\n        model.segmentation_head[0] = nn.Conv3d(\n            16,\n            14,\n            kernel_size=3,\n            stride=1,\n            padding=1,\n            bias=False\n        )\n        model.segmentation_head[0].weight.data = new_weight\n\n    return model","metadata":{"papermill":{"duration":0.039446,"end_time":"2025-09-15T22:55:14.815372","exception":false,"start_time":"2025-09-15T22:55:14.775926","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"f112dd30","cell_type":"code","source":"criterion = DiceFocalLoss()\nfor fold in FOLDS:\n    seed_everything(SEED)\n    model = smp.Unet(\n        encoder_name=\"resnet18\",\n        encoder_weights=None,\n        in_channels=1,\n        classes=2\n    ).to(device)\n    model.load_state_dict(torch.load(f'/kaggle/input/rsna-2d-binary-segmentation-training-{fold}/best_model_{fold}.pth'))\n    model = convert_2d_to_3d(model)\n    \n    optimizer = optim.Adam(model.parameters(), lr=LR)\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min', patience=1, factor=.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 = RSNA_Dataset_3D(train_df)\n    val_dataset = RSNA_Dataset_3D(valid_df, VALID=True)\n    \n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=BS,\n        shuffle=True,\n        num_workers=WORKERS,\n        pin_memory=True\n    )\n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=BS,\n        shuffle=False,\n        num_workers=WORKERS,\n        pin_memory=True\n    )\n#   Metrics tracking\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_3d_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#   Confusion Matrix of the best model\n    model.load_state_dict(torch.load(f'best_3d_model_{fold}.pth'))\n    cm = np.zeros((14, 14))\n    labels = np.arange(14)\n    with torch.no_grad():\n        for images, masks in tqdm(val_loader):\n            images = images.to(device, non_blocking=True) \n            outputs = model(images).argmax(1)\n            cm += confusion_matrix(\n                y_true=masks.flatten().tolist(),\n                y_pred=outputs.flatten().tolist(),\n                labels=labels\n            )\n    D = np.sqrt(cm[np.arange(14),np.arange(14)])\n    cm = cm/(D.reshape(-1,1))\n    cm = cm/(D.reshape(1,-1))\n\n    plt.figure(figsize=(12, 10))\n    ax = sns.heatmap(\n        cm, \n        annot=True, \n        fmt=\".2f\",\n        cmap='Blues', \n        vmin=0, \n        vmax=1,\n        xticklabels=range(14), \n        yticklabels=range(14),\n        cbar_kws={'label': 'Normalized Value'}\n    )\n    plt.title('Class-Normalized Confusion Matrix', pad=20, fontsize=14)\n    plt.xlabel('Predicted Label', fontsize=12)\n    plt.ylabel('True Label', fontsize=12)\n    plt.xticks(rotation=45, ha='right')\n    plt.yticks(rotation=0)\n\n    plt.tight_layout()\n    plt.savefig(f' confusion_matrix_{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":33936.672625,"end_time":"2025-09-16T08:20:51.517142","exception":false,"start_time":"2025-09-15T22:55:14.844517","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null}]}