{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"}],"dockerImageVersionId":30786,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q git+https://github.com/qubvel/segmentation_models.pytorch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T08:10:51.070552Z","iopub.execute_input":"2026-01-18T08:10:51.070927Z","iopub.status.idle":"2026-01-18T08:11:10.668217Z","shell.execute_reply.started":"2026-01-18T08:10:51.070879Z","shell.execute_reply":"2026-01-18T08:11:10.667263Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\nimport pydicom\nimport numpy as np\nimport os\nimport glob\nfrom tqdm import tqdm\nimport gc\nimport sklearn\n\nimport torchvision\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset\nfrom fastai.vision.all import *\nimport segmentation_models_pytorch as smp\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"_uuid":"d28a3186-897b-4159-9d44-67c70bb80d7a","_cell_guid":"cec8538e-3ac2-4af4-8978-9307d733bc98","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-01-18T08:11:10.669911Z","iopub.execute_input":"2026-01-18T08:11:10.670203Z","iopub.status.idle":"2026-01-18T08:11:17.706822Z","shell.execute_reply.started":"2026-01-18T08:11:10.670173Z","shell.execute_reply":"2026-01-18T08:11:17.705917Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SEED = 1337\nFOLDS = [1,2,3,4,5]\nPATH = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/'\nTRAIN_PATH = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/'\nENCODER_NAME = \"resnet18\"\nPATCH_H = 512\nPATCH_W = 512\nANGLE = 30\npatch_size = 64\nBS = 24\nEPOCHS = 2","metadata":{"_uuid":"180893b2-c54e-45a0-9b86-3ec965f041c8","_cell_guid":"9a11e263-8200-4794-8e27-57bf82c9218b","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-01-18T08:11:48.518548Z","iopub.execute_input":"2026-01-18T08:11:48.519240Z","iopub.status.idle":"2026-01-18T08:11:48.523670Z","shell.execute_reply.started":"2026-01-18T08:11:48.519207Z","shell.execute_reply":"2026-01-18T08:11:48.522757Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"spinal = [\n    'spinal_canal_stenosis_l1_l2',\n    'spinal_canal_stenosis_l2_l3',\n    'spinal_canal_stenosis_l3_l4',\n    'spinal_canal_stenosis_l4_l5',\n    'spinal_canal_stenosis_l5_s1'\n]\ncoor = [\n    'x_L1L2',\n    'y_L1L2',\n    'x_L2L3',\n    'y_L2L3',\n    'x_L3L4',\n    'y_L3L4',\n    'x_L4L5',\n    'y_L4L5',\n    'x_L5S1',\n    'y_L5S1'\n]","metadata":{"_uuid":"be7633ba-a031-4749-9ed1-3dcc0aad0c56","_cell_guid":"5db99b5a-2010-450c-9955-3ae34783bba2","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-01-18T08:11:49.045600Z","iopub.execute_input":"2026-01-18T08:11:49.045933Z","iopub.status.idle":"2026-01-18T08:11:49.050160Z","shell.execute_reply.started":"2026-01-18T08:11:49.045905Z","shell.execute_reply":"2026-01-18T08:11:49.049368Z"}},"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\n    \ndef augment_image(image):\n#   Randomly rotate the image.\n    angle = torch.as_tensor(random.uniform(-ANGLE, ANGLE))\n    image = torchvision.transforms.functional.rotate(\n        image,angle.item(),\n        interpolation=torchvision.transforms.InterpolationMode.BILINEAR\n    )\n    return image\n\ndef my_collate_fn(data):\n    collation = [torch.cat(s) for s in zip(*data)]\n    return collation\n\ndef display_images(images, title, max_images_per_row=4):\n    # Calculate the number of rows needed\n    num_images = len(images)\n    num_rows = (num_images + max_images_per_row - 1) // max_images_per_row  # Ceiling division\n\n    # Create a subplot grid\n    fig, axes = plt.subplots(num_rows, max_images_per_row, figsize=(5, 1.5 * num_rows))\n    \n    # Flatten axes array for easier looping if there are multiple rows\n    if num_rows > 1:\n        axes = axes.flatten()\n    else:\n        axes = [axes]  # Make it iterable for consistency\n\n    # Plot each image\n    for idx, image in enumerate(images):\n        ax = axes[idx]\n        ax.imshow(image, cmap='gray')  # Assuming grayscale for simplicity, change cmap as needed\n        ax.axis('off')  # Hide axes\n\n    # Turn off unused subplots\n    for idx in range(num_images, len(axes)):\n        axes[idx].axis('off')\n    fig.suptitle(title, fontsize=16)\n\n    plt.tight_layout()","metadata":{"_uuid":"461ba614-e9ed-46a5-9919-eb1f2020cce4","_cell_guid":"6310a5e2-a75f-40c8-acaa-8b674d296bc8","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-01-18T08:11:49.217076Z","iopub.execute_input":"2026-01-18T08:11:49.217870Z","iopub.status.idle":"2026-01-18T08:11:49.226553Z","shell.execute_reply.started":"2026-01-18T08:11:49.217835Z","shell.execute_reply":"2026-01-18T08:11:49.225514Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\ntrain = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv\")\nn_folds = 5\ntrain[\"fold\"] = (np.arange(len(train)) % n_folds) + 1\ntrain.tail()","metadata":{"_uuid":"ad8194f5-ddb8-4a94-9246-e3f063999655","_cell_guid":"dde26a15-6bac-4146-8a66-7f1d304cbc93","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-01-18T08:11:49.369689Z","iopub.execute_input":"2026-01-18T08:11:49.370653Z","iopub.status.idle":"2026-01-18T08:11:49.402070Z","shell.execute_reply.started":"2026-01-18T08:11:49.370608Z","shell.execute_reply":"2026-01-18T08:11:49.401309Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_coor = pd.read_csv(PATH + 'train_label_coordinates.csv')\ndf_coor.tail()","metadata":{"_uuid":"f022d908-1896-4daf-9aab-a7120e46286c","_cell_guid":"2df5872e-79b9-4760-a2c9-f24cf307d046","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-01-18T08:11:49.517719Z","iopub.execute_input":"2026-01-18T08:11:49.518014Z","iopub.status.idle":"2026-01-18T08:11:49.581524Z","shell.execute_reply.started":"2026-01-18T08:11:49.517988Z","shell.execute_reply":"2026-01-18T08:11:49.580692Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Instance number here tells us the slice where maximum information or positive indicator for the condition is present\nS = df_coor[\n    df_coor['condition'] == 'Spinal Canal Stenosis'\n].sort_values([\n    'study_id',\n    'series_id',\n    'level'\n]).reset_index(drop=True)\nS.tail()","metadata":{"_uuid":"5e4139a1-9309-4844-a091-8bea6e0244a1","_cell_guid":"67a2acbb-9cd6-49bd-8624-39bae7762495","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-01-18T08:11:49.672871Z","iopub.execute_input":"2026-01-18T08:11:49.673177Z","iopub.status.idle":"2026-01-18T08:11:49.692292Z","shell.execute_reply.started":"2026-01-18T08:11:49.673148Z","shell.execute_reply":"2026-01-18T08:11:49.691631Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S['x_mean_fraction'] = S['x']/S.groupby(['study_id','series_id'])['x'].mean().loc[[(study_id,series_id) for study_id,series_id in S[['study_id','series_id']].values]].values\nS.tail()","metadata":{"_uuid":"6ebb52c3-d448-4bd9-ac25-0613d83de3f8","_cell_guid":"2ab3cc3c-d59b-4912-8c29-47d301678407","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-01-18T08:11:49.810267Z","iopub.execute_input":"2026-01-18T08:11:49.810930Z","iopub.status.idle":"2026-01-18T08:11:49.866142Z","shell.execute_reply.started":"2026-01-18T08:11:49.810899Z","shell.execute_reply":"2026-01-18T08:11:49.865390Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S = S[S['x_mean_fraction'] > .8]\nS.tail()","metadata":{"_uuid":"9708d077-5b53-4248-a24b-cde8a7cb7c6f","_cell_guid":"a36fa186-3317-4c96-b545-668b06f19271","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-01-18T08:11:49.963080Z","iopub.execute_input":"2026-01-18T08:11:49.963391Z","iopub.status.idle":"2026-01-18T08:11:49.975758Z","shell.execute_reply.started":"2026-01-18T08:11:49.963363Z","shell.execute_reply":"2026-01-18T08:11:49.975059Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S['instance_number'] = S['instance_number'] - 1\nS.tail()","metadata":{"_uuid":"280797d6-d2c5-47a1-ac3c-1667605d1816","_cell_guid":"bafd7484-0a39-4de3-801e-851784aa13ad","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-01-18T08:11:50.219922Z","iopub.execute_input":"2026-01-18T08:11:50.220264Z","iopub.status.idle":"2026-01-18T08:11:50.231574Z","shell.execute_reply.started":"2026-01-18T08:11:50.220230Z","shell.execute_reply":"2026-01-18T08:11:50.230702Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"coordinates = {}\nfor study_id,df in S.groupby('study_id'):\n    coordinates[study_id] = {}\nfor (study_id,series_id),df in tqdm(S.groupby(['study_id','series_id'])):\n    coordinates[study_id][series_id] = {\n                'L1/L2':{\n                    'x':torch.nan,\n                    'y':torch.nan,\n                    'instance_number':torch.nan\n                },\n                    'L2/L3':{\n                    'x':torch.nan,\n                    'y':torch.nan,\n                    'instance_number':torch.nan\n                },\n                'L3/L4':{\n                    'x':torch.nan,\n                    'y':torch.nan,\n                    'instance_number':torch.nan\n                },\n                'L4/L5':{\n                    'x':torch.nan,\n                    'y':torch.nan,\n                    'instance_number':torch.nan\n                },\n                'L5/S1':{\n                    'x':torch.nan,\n                    'y':torch.nan,\n                    'instance_number':torch.nan\n                }\n    }\n    \n    for i in range(len(df)):\n        row = df.iloc[i]\n        coordinates[row['study_id']][row['series_id']][row['level']]['x'] = row['x']\n        coordinates[row['study_id']][row['series_id']][row['level']]['y'] = row['y']\n        coordinates[row['study_id']][row['series_id']][row['level']]['instance_number'] = row['instance_number']","metadata":{"_uuid":"d1bfee32-528c-4149-92b2-bcf07f22d90f","_cell_guid":"128c5155-e9fc-48fc-969b-f7b7a6d1a785","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-01-18T08:11:52.809462Z","iopub.execute_input":"2026-01-18T08:11:52.809806Z","iopub.status.idle":"2026-01-18T08:11:53.642281Z","shell.execute_reply.started":"2026-01-18T08:11:52.809777Z","shell.execute_reply":"2026-01-18T08:11:53.641433Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S =  S[[\n    'study_id',\n    'series_id'\n]].groupby([\n    'study_id',\n    'series_id'\n]).count().reset_index()\nS.tail()","metadata":{"_uuid":"5e1c19f5-4f06-413b-a2db-3c91f65c4e98","_cell_guid":"9194cf59-e045-4822-9242-895f3abd7460","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-01-18T08:11:53.643624Z","iopub.execute_input":"2026-01-18T08:11:53.643892Z","iopub.status.idle":"2026-01-18T08:11:53.655743Z","shell.execute_reply.started":"2026-01-18T08:11:53.643866Z","shell.execute_reply":"2026-01-18T08:11:53.654857Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"v = np.zeros((len(S),15))\nfor i in tqdm(range(len(S))):\n    row = S.iloc[i]\n    k = 0\n    for level in coordinates[row['study_id']][row['series_id']]:\n        v[i,k:k+3] = list(coordinates[row['study_id']][row['series_id']][level].values())\n        k += 3","metadata":{"_uuid":"ce9a0015-fad4-4d6a-b614-8a672400a997","_cell_guid":"dfcbd937-00a3-4016-be7c-73d7f9a7f342","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-01-18T08:11:53.656785Z","iopub.execute_input":"2026-01-18T08:11:53.657068Z","iopub.status.idle":"2026-01-18T08:11:53.792759Z","shell.execute_reply.started":"2026-01-18T08:11:53.657029Z","shell.execute_reply":"2026-01-18T08:11:53.791944Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S[[\n    'x_L1L2',\n    'y_L1L2',\n    'i_L1L2',\n    'x_L2L3',\n    'y_L2L3',\n    'i_L2L3',\n    'x_L3L4',\n    'y_L3L4',\n    'i_L3L4',\n    'x_L4L5',\n    'y_L4L5',\n    'i_L4L5',\n    'x_L5S1',\n    'y_L5S1',\n    'i_L5S1'\n]] = v\nS.tail()","metadata":{"_uuid":"c200ff17-f121-4273-b1b4-af83fa8d8763","_cell_guid":"7d9d9eee-cf1e-4ed0-ad21-700e4092c136","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-01-18T08:11:53.794335Z","iopub.execute_input":"2026-01-18T08:11:53.794625Z","iopub.status.idle":"2026-01-18T08:11:53.817365Z","shell.execute_reply.started":"2026-01-18T08:11:53.794599Z","shell.execute_reply":"2026-01-18T08:11:53.816340Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S = S.merge(train,left_on='study_id',right_on='study_id')\nS.tail()","metadata":{"_uuid":"63ae0b6c-2aa9-4073-b0ef-1dd947bf8b45","_cell_guid":"a760c3ef-f86f-4a22-add5-9ea6ac92995e","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-01-18T08:11:53.818609Z","iopub.execute_input":"2026-01-18T08:11:53.819381Z","iopub.status.idle":"2026-01-18T08:11:53.842100Z","shell.execute_reply.started":"2026-01-18T08:11:53.819335Z","shell.execute_reply":"2026-01-18T08:11:53.841306Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mask = S[[\n    'x_L1L2',\n    'x_L2L3',\n    'x_L3L4',\n    'x_L4L5',\n    'x_L5S1'\n]].isna().values\nmask += S[[\n    'y_L1L2',\n    'y_L2L3',\n    'y_L3L4',\n    'y_L4L5',\n    'y_L5S1'\n]].isna().values\nmask += S[[\n    'i_L1L2',\n    'i_L2L3',\n    'i_L3L4',\n    'i_L4L5',\n    'i_L5S1'\n]].isna().values\nmask = mask > 0\nmask[-5:]\nv = S[spinal].values\nv[mask] = 'UNK'\nS[spinal] = v","metadata":{"_uuid":"35b558c5-5ae0-46a1-8c43-02fa6dcf338e","_cell_guid":"dfe868f4-b03b-475c-9daa-aebd7c8279d7","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-01-18T08:11:54.892687Z","iopub.execute_input":"2026-01-18T08:11:54.893529Z","iopub.status.idle":"2026-01-18T08:11:54.903201Z","shell.execute_reply.started":"2026-01-18T08:11:54.893480Z","shell.execute_reply":"2026-01-18T08:11:54.902247Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S.groupby('fold').count()","metadata":{"_uuid":"07f9494d-c64d-4d31-ae1c-63d8a7596510","_cell_guid":"3dd160a4-7618-44bb-8bc4-e626c50b2b48","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-01-18T08:11:55.049154Z","iopub.execute_input":"2026-01-18T08:11:55.049485Z","iopub.status.idle":"2026-01-18T08:11:55.071527Z","shell.execute_reply.started":"2026-01-18T08:11:55.049429Z","shell.execute_reply":"2026-01-18T08:11:55.070748Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Sagittal_T2_spine_discriminator_Dataset(Dataset):\n    def __init__(self, df, VALID=False, P=patch_size):\n        self.data = df\n        self.VALID = VALID\n        self.P = P\n        self.resize = torchvision.transforms.Resize((PATCH_H,PATCH_W),antialias=True)\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, index):\n        row = self.data.iloc[index]\n        \n        \n        sample = TRAIN_PATH + str(row['study_id']) + '/'+str(row['series_id'])\n\n        images = [x.replace('\\\\','/') for x in glob.glob(sample+'/*.dcm')]\n        images.sort(reverse=False, key=lambda x: int(x.split('/')[-1].replace('.dcm', '')))\n\n        instance_numbers = row[[\n            'i_L1L2',\n            'i_L2L3',\n            'i_L3L4',\n            'i_L4L5',\n            'i_L5S1'\n        ]].values\n\n        image = torch.stack([\n            torch.as_tensor(pydicom.dcmread(x).pixel_array.astype(np.float32)) for x in images\n        ]).float().to(device)\n        image = image/image.max()\n        D,H,W = image.shape\n\n        c = torch.as_tensor([x for x in row[coor]]).view(5,2).float()\n        missing = c.isnan().sum(1) > 0\n        c[missing] = torch.as_tensor([H/2,W/2])\n\n        if H > W:\n            d = W\n            h = (H - d)//2\n            image = image[:,h:h+d]\n            c[:,1] -= h\n            H = W\n        elif H < W:\n            d = H\n            w = (W - d)//2\n            image = image[:,:,w:w+d]\n            c[:,0] -= w\n            W = H\n\n        image = self.resize(image)\n        image = nn.functional.pad(image,[self.P]*4,'reflect')\n        c[:,1] = c[:,1]*PATCH_H/H + self.P\n        c[:,0] = c[:,0]*PATCH_W/W + self.P\n        c = c.long()\n\n        crops = torch.stack([\n            image[\n                :,\n                xy[1]-self.P:xy[1]+self.P,\n                xy[0]-self.P:xy[0]+self.P\n            ] for xy in c\n        ])\n\n        image = torch.zeros(5,D,2*self.P,2*self.P).to(device)\n        label = torch.zeros(5,D).long().to(device) - 100\n        for i in range(5):\n            if ~missing[i]:\n                instance_number = instance_numbers[i].astype(int)\n                pickeable = torch.ones(D).bool()\n                pickeable[instance_number] = False\n                image[i,0] = crops[i,instance_number]\n                k = 1\n                if instance_number > 0:\n                    pickeable[instance_number-1] = False\n                    image[i,1] = crops[i,instance_number-1]\n                    k += 1\n                if instance_number < D - 1:\n                    pickeable[instance_number+1] = False\n                    image[i,2] = crops[i,instance_number+1]\n                    k += 1\n                if instance_number > 1: pickeable[instance_number-2] = False\n                if instance_number < D - 2: pickeable[instance_number+2] = False\n                \n                pickeable = torch.arange(D)[pickeable]\n                picked = pickeable\n                image[i,k:len(picked)+k] = crops[i,picked]\n\n                label[i,:k] = 1\n                label[i,k:len(picked)+k] = 0\n\n                if not self.VALID:\n                    image[i] = augment_image(image[i].reshape(-1,2*self.P,2*self.P)).reshape(D,2*self.P,2*self.P)\n            \n        image = image[:,:,self.P//2:self.P//2+self.P,self.P//2:self.P//2+self.P]\n\n        mask = label != -100\n        image = image[mask]\n        label = label[mask]\n\n        return image,label","metadata":{"_uuid":"852859ef-f0fc-4ef7-ba34-a44b3d589999","_cell_guid":"e3dd34e4-c837-4946-8444-3020ba284219","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-01-18T08:11:55.221034Z","iopub.execute_input":"2026-01-18T08:11:55.221336Z","iopub.status.idle":"2026-01-18T08:11:55.236569Z","shell.execute_reply.started":"2026-01-18T08:11:55.221310Z","shell.execute_reply":"2026-01-18T08:11:55.235623Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ds = Sagittal_T2_spine_discriminator_Dataset(S)","metadata":{"_uuid":"b3aec249-1e43-4893-8d28-a6bc00c75c74","_cell_guid":"1f041335-9da3-41e2-8e94-6a1b2d834711","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-01-18T08:11:57.104360Z","iopub.execute_input":"2026-01-18T08:11:57.105568Z","iopub.status.idle":"2026-01-18T08:11:57.110361Z","shell.execute_reply.started":"2026-01-18T08:11:57.105511Z","shell.execute_reply":"2026-01-18T08:11:57.109398Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample,label = ds.__getitem__(np.random.randint(len(ds)))\nprint(label)\ndisplay_images(sample[label == 1].cpu(),'Positives',max_images_per_row=5)\ndisplay_images(sample[label == 0].cpu(),'Negatives',max_images_per_row=5)","metadata":{"_uuid":"a9c50b2e-74a9-4092-b8c6-ae2bb7ef1370","_cell_guid":"91e7680a-c571-4926-8581-cc711d03457e","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-01-18T08:11:57.270181Z","iopub.execute_input":"2026-01-18T08:11:57.270510Z","iopub.status.idle":"2026-01-18T08:12:00.485160Z","shell.execute_reply.started":"2026-01-18T08:11:57.270479Z","shell.execute_reply":"2026-01-18T08:12:00.484290Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del ds\ngc.collect()","metadata":{"_uuid":"e5aca9a2-ce29-42d2-98e2-0e43a63e0077","_cell_guid":"e93db9c6-0c68-40ff-b07a-540e4b4f29cd","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-01-18T08:12:00.486951Z","iopub.execute_input":"2026-01-18T08:12:00.487302Z","iopub.status.idle":"2026-01-18T08:12:00.749251Z","shell.execute_reply.started":"2026-01-18T08:12:00.487266Z","shell.execute_reply":"2026-01-18T08:12:00.748359Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Sagittal_T2_spine_Discriminator(nn.Module):\n    def __init__(self, dim=512):\n        super().__init__()\n        CNN = torchvision.models.resnet18(weights='DEFAULT')\n        W = nn.Parameter(CNN.conv1.weight.sum(1, keepdim=True))\n        CNN.conv1 = nn.Conv2d(1, patch_size, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)\n        CNN.conv1.weight = W\n        CNN.fc = nn.Identity()\n        self.emb = CNN.to(device)\n        self.proj_out = nn.Linear(dim,2).to(device)\n    \n    def forward(self, x):        \n        x = self.emb(x.view(-1,1,patch_size,patch_size))\n        x = self.proj_out(x.view(-1,512))\n        return x","metadata":{"_uuid":"27d6e2f8-70bb-44e1-a9ff-13ddaf1ed108","_cell_guid":"b2646bf9-9179-43b2-9ef1-e888d39b4a4e","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-01-18T08:12:00.750857Z","iopub.execute_input":"2026-01-18T08:12:00.751215Z","iopub.status.idle":"2026-01-18T08:12:00.759214Z","shell.execute_reply.started":"2026-01-18T08:12:00.751176Z","shell.execute_reply":"2026-01-18T08:12:00.758571Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"negatives = 116965\npositives = 29207\ntotal = negatives + positives\n\ndef myLoss(preds,target):\n    target = target.view(-1)\n    return nn.CrossEntropyLoss(weight=torch.as_tensor([total/(negatives*2),total/(positives*2)]).to(device))(preds,target)","metadata":{"_uuid":"fe995948-0c57-46b5-a039-8f3313979fc1","_cell_guid":"ee96601c-ba43-46ca-85da-78fb654ad829","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-01-18T08:12:00.760480Z","iopub.execute_input":"2026-01-18T08:12:00.760719Z","iopub.status.idle":"2026-01-18T08:12:00.771993Z","shell.execute_reply.started":"2026-01-18T08:12:00.760695Z","shell.execute_reply":"2026-01-18T08:12:00.771112Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for f in FOLDS:\n    seed_everything(SEED)\n    model = Sagittal_T2_spine_Discriminator()\n    \n    df = S\n    tdf = df[df.fold != f]\n    vdf = df[df.fold == f]\n    tds = Sagittal_T2_spine_discriminator_Dataset(tdf)\n    vds = Sagittal_T2_spine_discriminator_Dataset(vdf,VALID=True)\n    tdl = torch.utils.data.DataLoader(\n        tds,\n        batch_size=BS,\n        shuffle=True,\n        drop_last=True,\n        collate_fn=my_collate_fn\n    )\n    vdl = torch.utils.data.DataLoader(\n        vds,\n        batch_size=BS,\n        shuffle=False,\n        collate_fn=my_collate_fn\n    )\n    dls = DataLoaders(tdl,vdl)\n\n    n_iter = len(tds)//BS\n\n    learn = Learner(\n        dls,\n        model,\n        loss_func=myLoss,\n        cbs=[\n            ShowGraphCallback(),\n            GradientClip(3.0)\n        ]\n    )\n    learn.fit_one_cycle(EPOCHS, lr_max=1e-3, wd=0.05, pct_start=0.02)\n    torch.save(model,'Sagittal_T2_spine_discriminator_'+str(f))\n    del model,df,tdf,vdf,tds,vds,tdl,vdl,dls,learn\n    gc.collect()","metadata":{"_uuid":"d2a07a6f-2eaf-46af-bdac-3be85d3c0b5c","_cell_guid":"26735e30-ba33-4d77-b954-b0e532457b3d","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-01-18T08:12:04.958304Z","iopub.execute_input":"2026-01-18T08:12:04.958710Z","iopub.status.idle":"2026-01-18T09:12:59.128875Z","shell.execute_reply.started":"2026-01-18T08:12:04.958679Z","shell.execute_reply":"2026-01-18T09:12:59.127867Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"y_pred = []\ny_true = []\nfor f in FOLDS:\n    model = torch.load('Sagittal_T2_spine_discriminator_'+str(f)) \n    df = S\n    vdf = df[df.fold == f]\n    vds = Sagittal_T2_spine_discriminator_Dataset(vdf,VALID=True)\n    vdl = torch.utils.data.DataLoader(\n        vds,\n        batch_size=BS,\n        shuffle=False,\n        collate_fn=my_collate_fn\n    )\n    with torch.no_grad():\n        for images,target in tqdm(vdl):\n            target = target.view(-1).tolist()\n            preds = model(images).argmax(-1).tolist()\n            y_pred = y_pred + preds\n            y_true = y_true + target\n\n    del model,df,vdf,vds,vdl\n    gc.collect()\n\nsklearn.metrics.confusion_matrix(y_true, y_pred)","metadata":{"_uuid":"c0c25747-88a0-4b32-b554-8d4387569c50","_cell_guid":"90467392-0092-4445-9ac0-d18abebc0c39","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"negatives = sum([y == 0 for y in y_true])\npositives = sum([y == 1 for y in y_true])\ntotal = negatives + positives\nprint(negatives,positives,total)","metadata":{"_uuid":"6540a352-a97d-4852-a4c6-7a4d438028eb","_cell_guid":"61ecb263-76f4-4e54-ac07-57ebadb6fa77","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null}]}