{"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,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":13253657,"sourceType":"datasetVersion","datasetId":8398519},{"sourceId":262889198,"sourceType":"kernelVersion"},{"sourceId":263000768,"sourceType":"kernelVersion"},{"sourceId":261445297,"sourceType":"kernelVersion"}],"dockerImageVersionId":31090,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":15183.085204,"end_time":"2025-09-15T20:05:06.752817","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2025-09-15T15:52:03.667613","version":"2.6.0"}},"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\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 roc_auc_score\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":89.324528,"end_time":"2025-09-15T15:53:37.09074","exception":false,"start_time":"2025-09-15T15:52:07.766212","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T21:31:06.628325Z","iopub.execute_input":"2025-10-05T21:31:06.629144Z","iopub.status.idle":"2025-10-05T21:32:36.327317Z","shell.execute_reply.started":"2025-10-05T21:31:06.629104Z","shell.execute_reply":"2025-10-05T21:32:36.326665Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"FOLDS = [0]\nSEED = 777\nSIZE = 128\nWSIZE = 256\nRADIUS = 16\nANGLE_AUG = 30\nEPOCHS = 12\nBS = 4\nLR = 1e-5\nWORKERS = 4\nROTATION_PROB_DECAY = .5699 # (1 - p)**2 = p**3 | Equals single and triple rotations during augmentations","metadata":{"papermill":{"duration":0.026471,"end_time":"2025-09-15T15:53:37.138069","exception":false,"start_time":"2025-09-15T15:53:37.111598","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T21:32:36.328448Z","iopub.execute_input":"2025-10-05T21:32:36.32872Z","iopub.status.idle":"2025-10-05T21:32:36.333149Z","shell.execute_reply.started":"2025-10-05T21:32:36.328685Z","shell.execute_reply":"2025-10-05T21:32:36.332474Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"label_columns = [\n    'Other Posterior Circulation',\n    'Basilar Tip',\n    'Right Posterior Communicating Artery',\n    'Left Posterior Communicating Artery',\n    'Right Infraclinoid Internal Carotid Artery',\n    'Left Infraclinoid Internal Carotid Artery',\n    'Right Supraclinoid Internal Carotid Artery',\n    'Left Supraclinoid Internal Carotid Artery',\n    'Right Middle Cerebral Artery',\n    'Left Middle Cerebral Artery',\n    'Right Anterior Cerebral Artery',\n    'Left Anterior Cerebral Artery',\n    'Anterior Communicating Artery'\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T21:32:36.333947Z","iopub.execute_input":"2025-10-05T21:32:36.334291Z","iopub.status.idle":"2025-10-05T21:32:36.351026Z","shell.execute_reply.started":"2025-10-05T21:32:36.334264Z","shell.execute_reply":"2025-10-05T21:32:36.350367Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"source_path = '/kaggle/input/rsna-raw-roi/NPZ/'\ncases = [c[:-4] for c in os.listdir(source_path )]\nprint(len(cases))\ncases[:5]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T21:32:36.352859Z","iopub.execute_input":"2025-10-05T21:32:36.35312Z","iopub.status.idle":"2025-10-05T21:32:36.441763Z","shell.execute_reply.started":"2025-10-05T21:32:36.353103Z","shell.execute_reply":"2025-10-05T21:32:36.441223Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train = pd.read_csv('/kaggle/input/rsna-2d-binary-segmentation-preprocessing/rsna_train_folds.csv')\nroi = pd.read_csv('/kaggle/input/rsna-raw-roi/rsna_roi.csv')\ntrain = roi[roi.case.apply(lambda v: v in cases)].merge(train,left_on='case',right_on='SeriesInstanceUID').reset_index(drop=True)\ntrain.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T21:32:36.442437Z","iopub.execute_input":"2025-10-05T21:32:36.442673Z","iopub.status.idle":"2025-10-05T21:32:36.674334Z","shell.execute_reply.started":"2025-10-05T21:32:36.442651Z","shell.execute_reply":"2025-10-05T21:32:36.673655Z"}},"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)\ntrain.tail()","metadata":{"papermill":{"duration":0.042703,"end_time":"2025-09-15T15:53:37.302135","exception":false,"start_time":"2025-09-15T15:53:37.259432","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T21:32:36.675046Z","iopub.execute_input":"2025-10-05T21:32:36.675305Z","iopub.status.idle":"2025-10-05T21:32:36.696695Z","shell.execute_reply.started":"2025-10-05T21:32:36.675282Z","shell.execute_reply":"2025-10-05T21:32:36.696106Z"}},"outputs":[],"execution_count":null},{"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        self.cases = df['SeriesInstanceUID']\n        self.pmin = df['pmin']\n        self.pmax = df['pmax']\n        self.flip = df['flip']\n        self.axis = df['axis']\n        self.s = df['s']\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        pmin = self.pmin[idx]\n        pmax = self.pmax[idx]\n        axis = self.axis[idx]\n        s = self.s[idx]/4\n        f = self.flip[idx]\n\n        npz =  np.load(source_path + case + '.npz')\n        img = npz['volume']\n        D,H,W = img.shape\n        t = npz['t']\n        l = npz['loc']\n        AP = npz['AP']\n        positive = sum(AP) > 0\n        if positive:\n            t = t[AP > 0]\n            l = l[AP > 0]\n        else:\n            if len(l) > 1:\n                t = t[l > 0]\n                l = l[l > 0]\n\n        if not self.VALID:\n            ul = np.unique(l)\n            k = ul[np.random.randint(len(ul))]\n            kk = np.arange(len(l))[l == k]\n            kk = kk[np.random.randint(len(kk))]\n            d,h,w = np.rint(t[kk]).astype(int)\n#           Random asymmetric zoom\n            sd = s*(.9 + 0.2*np.random.rand())\n            sh = s*(.9 + 0.2*np.random.rand())\n            sw = s*(.9 + 0.2*np.random.rand())\n\n            rd = np.rint(RADIUS*sd/SIZE).astype(int)\n            rh = np.rint(RADIUS*sh/SIZE).astype(int)\n            rw = np.rint(RADIUS*sw/SIZE).astype(int)\n\n            dd = np.rint(sd).astype(int)\n            hh = np.rint(sh).astype(int)\n            ww = np.rint(sw).astype(int)\n\n            ddd = np.rint(sd/2).astype(int)\n            hhh = np.rint(sh/2).astype(int)\n            www = np.rint(sw/2).astype(int)\n            \n            d0 = d - ddd//2 - rd - np.random.randint(dd - 2*rd)\n            h0 = h - hhh//2 - rh - np.random.randint(hh - 2*rh)\n            w0 = w - www//2 - rw - np.random.randint(ww - 2*rw)\n\n            d = d0 + dd + ddd\n            h = h0 + hh + hhh\n            w = w0 + ww + www\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                (torch.from_numpy(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#           Free rotation\n            p = 1\n            t[:,0] += pad_d0 - d0\n            t[:,1] += pad_h0 - h0\n            t[:,2] += pad_w0 - w0\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(t[kk,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#                   Convert angle to radians\n                    angle_rad = math.radians(angle)\n                    cos_angle = math.cos(angle_rad)\n                    sin_angle = math.sin(angle_rad)\n                    py = t[:,2]\n                    px = t[:,1]\n                    cy, cx = c\n#                   Translate point to origin\n                    translated_x = px - cx\n                    translated_y = py - cy    \n#                   Apply rotation\n                    rotated_x = translated_x * cos_angle - translated_y * sin_angle\n                    rotated_y = translated_x * sin_angle + translated_y * cos_angle    \n#                   Translate back\n                    t[:,1] = rotated_x + cx\n                    t[:,2] = rotated_y + cy\n                    p *= ROTATION_PROB_DECAY\n#               Rotate axis\n                img = img.permute(axis_perm)\n                t[:] = t[:,axis_perm]\n\n            img = img[\n                ddd//2:dd+ddd//2,\n                hhh//2:hh+hhh//2,\n                www//2:ww+www//2\n            ]\n            t[:,0] -= ddd//2\n            t[:,1] -= hhh//2\n            t[:,2] -= www//2\n\n        else:\n            d,h,w = np.rint(t[idx%len(t)]).astype(int)\n            dd = hh = ww = np.rint(s).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                (torch.from_numpy(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            t[:,0] += pad_d0 - d0\n            t[:,1] += pad_h0 - h0\n            t[:,2] += pad_w0 - w0\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        t[:,0] *= SIZE/dd\n        t[:,1] *= SIZE/hh\n        t[:,2] *= SIZE/ww\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            for k in range(3,12,2):\n                mr = l == k\n                ml = l == k + 1\n                l[mr] = k + 1\n                l[ml] = k\n#       Mask generation\n        msk = torch.zeros(SIZE,SIZE,SIZE)\n        if positive:\n            r = np.indices((SIZE,SIZE,SIZE)).reshape(1,3,SIZE,SIZE,SIZE) - t.reshape(*t.shape,1,1,1)\n            r = np.sqrt((r*r).sum(1))\n            m = r < RADIUS\n            mm = r.argmin(0)\n            for k in range(len(l)): msk[m[k]*(mm == k)] = l[k]\n#       Reorientation to axial perspective            \n        if axis == 0:\n            img = torch.rot90(torch.rot90(img,1,(-3,-2)),-1,(-2,-1))\n            msk = torch.rot90(torch.rot90(msk,1,(-3,-2)),-1,(-2,-1))\n        if axis == 1:\n            img = torch.rot90(img,1,(-3,-2))\n            msk = torch.rot90(msk,1,(-3,-2))\n#       Flip            \n        if f:\n            img = img.flip(-1)\n            msk = msk.flip(-1)\n        \n        return img, msk.long()","metadata":{"papermill":{"duration":0.045655,"end_time":"2025-09-15T15:53:37.368862","exception":false,"start_time":"2025-09-15T15:53:37.323207","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T21:32:36.697435Z","iopub.execute_input":"2025-10-05T21:32:36.697678Z","iopub.status.idle":"2025-10-05T21:32:36.721455Z","shell.execute_reply.started":"2025-10-05T21:32:36.697652Z","shell.execute_reply":"2025-10-05T21:32:36.720666Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ds = RSNA_Dataset_3D(train)","metadata":{"papermill":{"duration":0.025895,"end_time":"2025-09-15T15:53:37.415731","exception":false,"start_time":"2025-09-15T15:53:37.389836","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T21:32:36.722217Z","iopub.execute_input":"2025-10-05T21:32:36.722469Z","iopub.status.idle":"2025-10-05T21:32:36.738323Z","shell.execute_reply.started":"2025-10-05T21:32:36.72244Z","shell.execute_reply":"2025-10-05T21:32:36.737652Z"}},"outputs":[],"execution_count":null},{"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.782081,"end_time":"2025-09-15T15:53:39.264787","exception":false,"start_time":"2025-09-15T15:53:37.482706","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T21:32:36.738996Z","iopub.execute_input":"2025-10-05T21:32:36.739159Z","iopub.status.idle":"2025-10-05T21:32:37.832527Z","shell.execute_reply.started":"2025-10-05T21:32:36.739146Z","shell.execute_reply":"2025-10-05T21:32:37.831936Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del ds\ngc.collect()","metadata":{"papermill":{"duration":0.264712,"end_time":"2025-09-15T15:53:39.552729","exception":false,"start_time":"2025-09-15T15:53:39.288017","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T21:32:37.834557Z","iopub.execute_input":"2025-10-05T21:32:37.834842Z","iopub.status.idle":"2025-10-05T21:32:38.074664Z","shell.execute_reply.started":"2025-10-05T21:32:37.834827Z","shell.execute_reply":"2025-10-05T21:32:38.073933Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"weights = torch.from_numpy(train[label_columns].values.sum(0)).float()\nweights[[2,3,4,5,6,7,8,9,10,11]] = weights[[2,3,4,5,6,7,8,9,10,11]].view(-1,2).sum(-1).view(5,1).tile(1,2).view(-1)/2\nweights = torch.cat([torch.tensor([len(train)]),weights])\nweights = len(train)/weights\nwmin = weights.min()\nwmax = weights.max()\nweights = 9 * (weights - wmin) / (wmax - wmin) +  1\nprint(f'Class weights: {[np.round(w.item(),3) for w in weights]}')\nds = RSNA_Dataset_3D(train, VALID=True)","metadata":{"papermill":{"duration":0.027617,"end_time":"2025-09-15T15:53:39.603317","exception":false,"start_time":"2025-09-15T15:53:39.5757","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T21:32:38.075527Z","iopub.execute_input":"2025-10-05T21:32:38.075824Z","iopub.status.idle":"2025-10-05T21:32:38.101879Z","shell.execute_reply.started":"2025-10-05T21:32:38.0758Z","shell.execute_reply":"2025-10-05T21:32:38.101342Z"}},"outputs":[],"execution_count":null},{"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":0.887454,"end_time":"2025-09-15T15:53:40.513013","exception":false,"start_time":"2025-09-15T15:53:39.625559","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T21:32:38.102521Z","iopub.execute_input":"2025-10-05T21:32:38.102687Z","iopub.status.idle":"2025-10-05T21:32:39.271399Z","shell.execute_reply.started":"2025-10-05T21:32:38.102673Z","shell.execute_reply":"2025-10-05T21:32:39.270676Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del ds\ngc.collect()","metadata":{"papermill":{"duration":0.275676,"end_time":"2025-09-15T15:53:40.813303","exception":false,"start_time":"2025-09-15T15:53:40.537627","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T21:32:39.272103Z","iopub.execute_input":"2025-10-05T21:32:39.272315Z","iopub.status.idle":"2025-10-05T21:32:39.518636Z","shell.execute_reply.started":"2025-10-05T21:32:39.27229Z","shell.execute_reply":"2025-10-05T21:32:39.517882Z"}},"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":{"papermill":{"duration":0.028998,"end_time":"2025-09-15T15:53:40.866313","exception":false,"start_time":"2025-09-15T15:53:40.837315","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T21:32:39.519556Z","iopub.execute_input":"2025-10-05T21:32:39.519848Z","iopub.status.idle":"2025-10-05T21:32:39.530314Z","shell.execute_reply.started":"2025-10-05T21:32:39.519824Z","shell.execute_reply":"2025-10-05T21:32:39.52979Z"}},"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":{"papermill":{"duration":0.032138,"end_time":"2025-09-15T15:53:40.921821","exception":false,"start_time":"2025-09-15T15:53:40.889683","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T21:32:39.531164Z","iopub.execute_input":"2025-10-05T21:32:39.53139Z","iopub.status.idle":"2025-10-05T21:32:39.546455Z","shell.execute_reply.started":"2025-10-05T21:32:39.531375Z","shell.execute_reply":"2025-10-05T21:32:39.545949Z"}},"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":{"papermill":{"duration":0.037182,"end_time":"2025-09-15T15:53:40.982541","exception":false,"start_time":"2025-09-15T15:53:40.945359","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T21:32:39.547075Z","iopub.execute_input":"2025-10-05T21:32:39.547252Z","iopub.status.idle":"2025-10-05T21:32:39.562526Z","shell.execute_reply.started":"2025-10-05T21:32:39.547238Z","shell.execute_reply":"2025-10-05T21:32:39.561952Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DiceFocalLoss(nn.Module):\n    def __init__(self, weights=None, gamma=1.0, alpha=0.5):\n        super().__init__()\n        self.DL = DiceLoss(weight=weights)\n        self.FL = FocalLoss(alpha=weights,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.029506,"end_time":"2025-09-15T15:53:41.035649","exception":false,"start_time":"2025-09-15T15:53:41.006143","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T21:32:39.563175Z","iopub.execute_input":"2025-10-05T21:32:39.563365Z","iopub.status.idle":"2025-10-05T21:32:39.580668Z","shell.execute_reply.started":"2025-10-05T21:32:39.56335Z","shell.execute_reply":"2025-10-05T21:32:39.580053Z"}},"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            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.03326,"end_time":"2025-09-15T15:53:41.092192","exception":false,"start_time":"2025-09-15T15:53:41.058932","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T21:32:39.58132Z","iopub.execute_input":"2025-10-05T21:32:39.581557Z","iopub.status.idle":"2025-10-05T21:32:39.595299Z","shell.execute_reply.started":"2025-10-05T21:32:39.58154Z","shell.execute_reply":"2025-10-05T21:32:39.594604Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for 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 = convert_2d_to_3d(model)\n    model.load_state_dict(torch.load(f'/kaggle/input/rsna-from-2d-binary-to-3d-full-segmentation-{fold}/best_3d_model_{fold}.pth'))\n    \n    optimizer = optim.Adam(model.parameters(), lr=LR)\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min', patience=0, 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    valid_df = valid_df[~valid_df.flip].reset_index(drop=True)\n\n    weights = torch.from_numpy(train_df[label_columns].values.sum(0)).float()\n    weights[[2,3,4,5,6,7,8,9,10,11]] = weights[[2,3,4,5,6,7,8,9,10,11]].view(-1,2).sum(-1).view(5,1).tile(1,2).view(-1)/2\n    weights = torch.cat([torch.tensor([len(train_df)]),weights])\n    weights = len(train_df)/weights\n    wmin = weights.min()\n    wmax = weights.max()\n    weights = 9 * (weights - wmin) / (wmax - wmin) +  1\n    print(f'Class weights: {[np.round(w.item(),3) for w in weights]}')\n    criterion = DiceFocalLoss(weights=weights.to(device))\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_aneurysm_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#   Validation Inference\n    model.load_state_dict(torch.load(f'best_aneurysm_3d_model_{fold}.pth'))\n    y_true = []\n    y_pred = []\n    with torch.no_grad():\n        for v in tqdm(valid_df[['case','pmin','pmax','s','axis']+label_columns].values):\n            case,pmin,pmax,s,axis = v[:5]\n            y_true.append(v[5:].tolist())\n\n            npz =  np.load(source_path + case + '.npz')\n            img = npz['volume']\n            D,H,W = img.shape\n            t = npz['t']\n            l = npz['loc']\n            d,h,w = np.rint(t[l==0][0]).astype(int)\n            dd = hh = ww = np.rint(s/2).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                (torch.from_numpy(img[d0:d,h0:h,w0:w]) - pmin)/(pmax - pmin),\n                (pad_w0,pad_w,pad_h0,pad_h,pad_d0,pad_d)\n            ).to(device)\n            img = F.interpolate(\n                img.unsqueeze(0).unsqueeze(0),\n                size=(WSIZE,WSIZE,WSIZE),\n                mode='trilinear',\n                align_corners=False\n            )\n#           Reorientation to axial perspective\n            if axis == 0: img = torch.rot90(torch.rot90(img,1,(-3,-2)),-1,(-2,-1))\n            if axis == 1: img = torch.rot90(img,1,(-3,-2))\n#           Prediction\n            out = model(img).softmax(1)\n            out += model(img.flip(-1)).flip(-1).softmax(1)[:,[0,1,2,4,3,6,5,8,7,10,9,12,11,13]]\n            out = out[0].view(14,-1)/2\n            y_pred.append(out[1:].max(-1)[0].tolist())\n\n    y_true = np.array(y_true)\n    y_pred = np.array(y_pred)\n    print(f'Area under the ROC curve mean: {np.mean([roc_auc_score(y_true[:,k],y_pred[:,k]) for k in range(13)])}')\n\n    del model,optimizer,scheduler,criterion,train_dataset,val_dataset,train_loader,val_loader\n    gc.collect()","metadata":{"papermill":{"duration":15081.349325,"end_time":"2025-09-15T20:05:02.465451","exception":false,"start_time":"2025-09-15T15:53:41.116126","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T21:32:39.595981Z","iopub.execute_input":"2025-10-05T21:32:39.596177Z","iopub.status.idle":"2025-10-05T22:53:40.236814Z","shell.execute_reply.started":"2025-10-05T21:32:39.596159Z","shell.execute_reply":"2025-10-05T22:53:40.235978Z"}},"outputs":[],"execution_count":null}]}