{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":"gpu","dataSources":[{"sourceId":99552,"databundleVersionId":13190393,"sourceType":"competition"},{"sourceId":12673868,"sourceType":"datasetVersion","datasetId":8009288},{"sourceId":12758457,"sourceType":"datasetVersion","datasetId":8065380}],"dockerImageVersionId":31090,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"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 nibabel as nib\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\nfrom PIL import Image\nfrom tqdm import tqdm\nfrom scipy.ndimage import gaussian_filter\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":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"FOLDS = [4]\nSEED = 777\nSIZE = 128\nPSIZE = 148\nEPOCHS = 9#10\nBS = 4\nLR = 1e-4\nWORKERS = 4","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"segmentations_path = '/kaggle/input/rsna-intracranial-aneurysm-detection/segmentations/'\ncases = os.listdir(segmentations_path)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Thanks to IAN PAN\n# https://www.kaggle.com/competitions/rsna-intracranial-aneurysm-detection/discussion/593857\nreversed = {}\nfor case in cases:\n    reversed[case] = False\n\nfor case in '''1.2.826.0.1.3680043.8.498.10035643165968342618460849823699311381\n1.2.826.0.1.3680043.8.498.10540586847553109495238524904638776495\n1.2.826.0.1.3680043.8.498.10557880026294057874761753231388788828\n1.2.826.0.1.3680043.8.498.10759842474698331813589731619457567641\n1.2.826.0.1.3680043.8.498.10865391592895615633871689438787039175\n1.2.826.0.1.3680043.8.498.11140496970152788589837488009637704168\n1.2.826.0.1.3680043.8.498.11641438607169452758239778414614826230\n1.2.826.0.1.3680043.8.498.11924949819899884502738782576851659426\n1.2.826.0.1.3680043.8.498.12283701604837916064212605259577798418\n1.2.826.0.1.3680043.8.498.12780116426159918728945213894055885771\n1.2.826.0.1.3680043.8.498.12873050136415197430227722045995986358\n1.2.826.0.1.3680043.8.498.12896910506681881306246412668919668702\n1.2.826.0.1.3680043.8.498.12898332622076283462996059479076432725\n1.2.826.0.1.3680043.8.498.12904246053955178641505906243733756576\n1.2.826.0.1.3680043.8.498.13789305723712362238118274295587312089\n1.2.826.0.1.3680043.8.498.15111820005882064793593034423469604305\n1.2.826.0.1.3680043.8.498.15412988336827906186857260013885503248\n1.2.826.0.1.3680043.8.498.16386250344855221757144432829845114733\n1.2.826.0.1.3680043.8.498.20627322154402566045565159680288078498\n1.2.826.0.1.3680043.8.498.21260453249991608190728327379762807665\n1.2.826.0.1.3680043.8.498.23047023542526806696555440426928375679\n1.2.826.0.1.3680043.8.498.27693546360513068451517048347207987807\n1.2.826.0.1.3680043.8.498.31897325247898403027455884342546675049\n1.2.826.0.1.3680043.8.498.32250259987224176174516959348681094310\n1.2.826.0.1.3680043.8.498.34439485184360273751379923196589017042\n1.2.826.0.1.3680043.8.498.35327124657045713676192746001247576881\n1.2.826.0.1.3680043.8.498.35378146560080702211693278243609271022\n1.2.826.0.1.3680043.8.498.35633450896661854179640200212683653363\n1.2.826.0.1.3680043.8.498.37086262716517957668471635372810376638\n1.2.826.0.1.3680043.8.498.38904475631578710113273863766282479811\n1.2.826.0.1.3680043.8.498.50275403170194436966991630938339966596\n1.2.826.0.1.3680043.8.498.52363954882447190271251269039176558430\n1.2.826.0.1.3680043.8.498.53901203212732811892702239112353256979\n1.2.826.0.1.3680043.8.498.58839417089022860359638460482101293080\n1.2.826.0.1.3680043.8.498.63610643988023140802787347827023957721\n1.2.826.0.1.3680043.8.498.67357468192986203292275214887760889253\n1.2.826.0.1.3680043.8.498.67364033194715249441864636235467322768\n1.2.826.0.1.3680043.8.498.68356160898101066850726244725552676010\n1.2.826.0.1.3680043.8.498.76127804295106714014266113869285421890\n1.2.826.0.1.3680043.8.498.79099213587801933936080747802403048718\n1.2.826.0.1.3680043.8.498.80114244849666367523293067199486077713\n1.2.826.0.1.3680043.8.498.86037975393556827852769300088670915080\n1.2.826.0.1.3680043.8.498.86867376272146805428455638150607288831\n1.2.826.0.1.3680043.8.498.88739296218460643753583291722714541935\n1.2.826.0.1.3680043.8.498.90015157820692758596783999454928886688\n1.2.826.0.1.3680043.8.498.93009153822317083064844213344156801735\n1.2.826.0.1.3680043.8.498.98123758735027035609698227781754927939'''.split('\\n'):\n    reversed[case] = True","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"try:\n    train = pd.read_csv('/kaggle/input/rsna-from-2d-to-3d-segmentation-files/rsna_3d_seg_folds.csv')\n    with open('/kaggle/input/rsna-from-2d-to-3d-segmentation-files/indices.pkl', 'rb') as f:\n        indices = pickle.load(f)\nexcept:\n    train = pd.read_csv('/kaggle/input/rsna-intracranial-aneurysm-detection-seg-slices/rsna_seg_folds.csv').groupby('SeriesInstanceUID').min().reset_index()\n    t = []\n    SIUIDs = []\n    indices = {}\n    for case in tqdm(train['SeriesInstanceUID']):\n        case_dir = segmentations_path+case+'/'\n        nifti_file = os.listdir(case_dir)[0]\n        nifti_path = case_dir + nifti_file\n#       Load segmentation\n        seg = nib.load(nifti_path[:-4] + '_cowseg' + nifti_path[-4:]).get_fdata()\n        if reversed[case]: seg = seg[...,::-1]\n        indices[case] = {}\n        for  k in range(1,14):\n            h,w,d = np.where(seg == k)\n            if len(h > 0):\n                t.append(k)\n                SIUIDs.append(case)\n                indices[case][k] = np.stack([h,w,d])\n    \n    train = pd.DataFrame({\n        'SeriesInstanceUID':SIUIDs,\n        'target':t\n    }).merge(train,on='SeriesInstanceUID')\n    train.to_csv('rsna_3d_seg_folds.csv',index=False)\n    with open('indices.pkl', 'wb') as f:\n        pickle.dump(indices, f)\n\ntrain.tail()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train['flip'] = False\nftrain = train.copy()\nftrain['flip'] = True\ntrain = pd.concat([train,ftrain]).reset_index(drop=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def rotate_3d_tensor(img,msk,device=device):\n    \"\"\"\n    Rotate 3D volume using affine grid transformation.\n    Args:\n        volume: Input tensor of shape (N, C, D, H, W)\n        angles_deg: Rotation angles in degrees for (X, Y, Z) axes\n        mode: 'bilinear' (for images) or 'nearest' (for masks)\n    Returns:\n        Rotated tensor\n    \"\"\"\n#   angles_rad = 2*math.pi*(torch.rand(3,device=device) - .5)\n    angles_rad = math.pi*(torch.rand(3) - .5)/3\n    \n    # Create 3D rotation matrices for each axis\n    def _get_rotation_matrix(angle, axis):\n        cos = torch.cos(angle)\n        sin = torch.sin(angle)\n        if axis == 0:\n            return torch.tensor([\n                [1, 0, 0],\n                [0, cos, -sin],\n                [0, sin, cos]\n            ], device=device)\n        elif axis == 1:\n            return torch.tensor([\n                [cos, 0, sin],\n                [0, 1, 0],\n                [-sin, 0, cos]\n            ], device=device)\n        else:\n            return torch.tensor([\n                [cos, -sin, 0],\n                [sin, cos, 0],\n                [0, 0, 1]\n            ], device=device)\n    \n    R = _get_rotation_matrix(angles_rad[2], 2) @ \\\n        _get_rotation_matrix(angles_rad[1], 1) @ \\\n        _get_rotation_matrix(angles_rad[0], 0)\n    \n    affine_matrix = torch.cat([\n        R,  # 3x3 rotation\n        torch.zeros(3, 1, device=device)\n    ], dim=1).unsqueeze(0)\n    \n    grid = F.affine_grid(affine_matrix, msk.size(), align_corners=False)\n    \n    img = F.grid_sample(img, grid, mode='bilinear', align_corners=False)\n    msk = F.grid_sample(msk, grid, mode='nearest', align_corners=False)\n\n    return img,msk","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RSNA_Dataset_3D(Dataset):\n    def __init__(\n        self,\n        cases,\n        targets,\n        pmin,\n        pmax,\n        flip,\n        VALID=False\n    ):\n        \"\"\"\n        \n        \"\"\"\n        self.cases = cases\n        self.targets = targets\n        self.pmin = pmin\n        self.pmax = pmax\n        self.flip = flip\n        self.VALID = VALID\n\n    def __len__(self):\n        return len(self.cases)\n\n    def __getitem__(self, idx):        \n        case = self.cases[idx]\n        target = self.targets[idx]\n        imin = self.pmin[idx]\n        imax = self.pmax[idx]\n        f = self.flip[idx]\n\n        nifti_path = segmentations_path+case+'/'+case+'.nii'\n        img = nib.load(nifti_path).get_fdata()\n        msk = nib.load(nifti_path[:-4] + '_cowseg' + nifti_path[-4:]).get_fdata()\n\n        img = torch.from_numpy(img).permute(2,0,1).unsqueeze(0).float()\n        msk = torch.from_numpy(msk).permute(2,0,1).unsqueeze(0).float()\n        if reversed[case]: msk = msk.flip(1)\n\n        D,H,W = img.shape[1:]\n        r = min(H,W)\n\n        if not self.VALID:\n            h,w,d = indices[case][target][:,np.random.randint(len(indices[case][target][0]))]\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*PSIZE/512).astype(int)\n            hh = np.rint(rh*PSIZE/512).astype(int)\n            ww = np.rint(rw*PSIZE/512).astype(int)\n\n            ddd = np.rint(rd/4).astype(int)\n            hhh = np.rint(rh/4).astype(int)\n            www = np.rint(rw/4).astype(int)\n\n            dddd = np.rint(rd/8).astype(int)\n            hhhh = np.rint(rh/8).astype(int)\n            wwww = np.rint(rw/8).astype(int)\n\n            d += np.random.randint(dddd) - dddd//2\n            h += np.random.randint(hhhh) - hhhh//2\n            w += np.random.randint(wwww) - wwww//2\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] - imin)/(imax - imin),\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            img,msk = rotate_3d_tensor(\n                img.unsqueeze(0),\n                msk.unsqueeze(0),\n                'cpu'\n            )\n            d = (dd - ddd)//2\n            h = (hh - hhh)//2\n            w = (ww - www)//2\n            img = img[0,:,d:d+ddd,h:h+hhh,w:w+www]\n            msk = msk[0,:,d:d+ddd,h:h+hhh,w:w+www]\n\n        else:\n            h,w,d = np.rint(indices[case][target].mean(1)).astype(int)\n            dd = hh = ww = np.rint(r/4).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] - imin)/(imax - imin),\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),\n            size=(SIZE,SIZE,SIZE),\n            mode='trilinear',\n            align_corners=False\n        )[0]\n        msk = F.interpolate(\n            msk.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                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,SIZE))\n\n#       Flip            \n        if f:\n            img = img.flip(-2)\n            msk = msk.flip(-2)\n            mr = msk == 3\n            ml = msk == 4\n            msk[mr] = 4\n            msk[ml] = 3\n            mr = msk == 5\n            ml = msk == 6\n            msk[mr] = 6\n            msk[ml] = 5\n            mr = msk == 7\n            ml = msk == 8\n            msk[mr] = 8\n            msk[ml] = 7\n            mr = msk == 9\n            ml = msk == 10\n            msk[mr] = 10\n            msk[ml] = 9\n            mr = msk == 11\n            ml = msk == 12\n            msk[mr] = 12\n            msk[ml] = 11\n        \n        return img, msk[0].long()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ds = RSNA_Dataset_3D(\n    cases=train['SeriesInstanceUID'],\n    targets=train['target'],\n    pmin=train['pmin'],\n    pmax=train['pmax'],\n    flip=train['flip']\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img,msk = ds.__getitem__(np.random.randint(len(ds)))\nplt.imshow(img[0].sum(0))\nplt.show()\nplt.imshow((msk > 0).sum(0))\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del ds\ngc.collect()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ds = RSNA_Dataset_3D(\n    cases=train['SeriesInstanceUID'],\n    targets=train['target'],\n    pmin=train['pmin'],\n    pmax=train['pmax'],\n    flip=train['flip'],\n    VALID=True\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img,msk = ds.__getitem__(np.random.randint(len(ds)))\nplt.imshow(img[0].sum(0))\nplt.show()\nplt.imshow((msk > 0).sum(0))\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del ds\ngc.collect()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"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":{"trusted":true},"outputs":[],"execution_count":null},{"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":{"trusted":true},"outputs":[],"execution_count":null},{"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":{"trusted":true},"outputs":[],"execution_count":null},{"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":{"trusted":true},"outputs":[],"execution_count":null},{"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            2,\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":{"trusted":true},"outputs":[],"execution_count":null},{"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=\"imagenet\",\n        in_channels=1,\n        classes=2\n    ).to(device)\n    model.load_state_dict(torch.load(f'/kaggle/input/rsna-from-2d-to-3d-segmentation-files/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=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 = RSNA_Dataset_3D(\n        cases=train_df['SeriesInstanceUID'],\n        targets=train_df['target'],\n        pmin=train_df['pmin'],\n        pmax=train_df['pmax'],\n        flip=train_df['flip']\n    )\n    val_dataset = RSNA_Dataset_3D(\n        cases=valid_df['SeriesInstanceUID'],\n        targets=valid_df['target'],\n        pmin=valid_df['pmin'],\n        pmax=valid_df['pmax'],\n        flip=valid_df['flip'],\n        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\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    y_true = []\n    y_pred = []\n    with torch.no_grad():\n        for images, masks in tqdm(val_loader):\n            images = images.to(device, non_blocking=True)\n            masks = masks.to(device, non_blocking=True)\n            m = masks > 0\n            masks = masks[m]            \n            outputs = model(images).argmax(1)[m]\n            m = outputs > 0\n            y_true = y_true + masks[m].tolist()\n            y_pred = y_pred + outputs[m].tolist()\n\n    cm = confusion_matrix(y_true, y_pred)\n    D = np.sqrt(cm[np.arange(13),np.arange(13)])\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(1, 14), \n        yticklabels=range(1, 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":{"trusted":true},"outputs":[],"execution_count":null}]}