{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":14535717,"sourceType":"datasetVersion","datasetId":9009659}],"dockerImageVersionId":31260,"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":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:08:35.550290Z","iopub.execute_input":"2026-01-23T00:08:35.550595Z","iopub.status.idle":"2026-01-23T00:08:46.867371Z","shell.execute_reply.started":"2026-01-23T00:08:35.550569Z","shell.execute_reply":"2026-01-23T00:08:46.866445Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:08:46.868927Z","iopub.execute_input":"2026-01-23T00:08:46.869237Z","iopub.status.idle":"2026-01-23T00:09:02.640763Z","shell.execute_reply.started":"2026-01-23T00:08:46.869207Z","shell.execute_reply":"2026-01-23T00:09:02.639888Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SEED = 1337\nFOLDS = [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\nLmax = 15\nLR_MAX = 5e-6\nBS = 40\nEPOCHS = 10","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:02.641689Z","iopub.execute_input":"2026-01-23T00:09:02.641952Z","iopub.status.idle":"2026-01-23T00:09:02.646595Z","shell.execute_reply.started":"2026-01-23T00:09:02.641925Z","shell.execute_reply":"2026-01-23T00:09:02.645741Z"}},"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]\nlabels = {\n    'Normal/Mild':0,\n    'Moderate':1,\n    'Severe':2,\n    'UNK':-100\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:02.648614Z","iopub.execute_input":"2026-01-23T00:09:02.648852Z","iopub.status.idle":"2026-01-23T00:09:02.669214Z","shell.execute_reply.started":"2026-01-23T00:09:02.648825Z","shell.execute_reply":"2026-01-23T00:09:02.668471Z"}},"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    \n# CosineAnnealingAlpha\ndef nt(nmin,nmax,tcur,tmax):\n    return (nmax - .5*(nmax-nmin)*(1+np.cos(tcur*np.pi/tmax))).astype(np.float32)\n\nplt.plot(nt(.25,1,np.arange(EPOCHS),EPOCHS))\nplt.show()\n\n# callback to update alpha during training\ndef cb(self):\n    alpha = torch.as_tensor(nt(.25,1,learn.train_iter,EPOCHS*n_iter))\n    learn.dls.train_ds.alpha = alpha\nalpha_cb = Callback(before_batch=cb)\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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:02.670159Z","iopub.execute_input":"2026-01-23T00:09:02.670564Z","iopub.status.idle":"2026-01-23T00:09:03.092404Z","shell.execute_reply.started":"2026-01-23T00:09:02.670521Z","shell.execute_reply":"2026-01-23T00:09:03.091540Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train = 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 = train[['study_id','fold']+spinal][train[spinal].isna().sum(1) < len(spinal)].reset_index(drop=True)\ntrain.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:03.093456Z","iopub.execute_input":"2026-01-23T00:09:03.093857Z","iopub.status.idle":"2026-01-23T00:09:03.162376Z","shell.execute_reply.started":"2026-01-23T00:09:03.093815Z","shell.execute_reply":"2026-01-23T00:09:03.161488Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train = train.fillna('UNK')\ntrain[(train[spinal] == 'UNK').sum(1)>0].reset_index(drop=True).tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:03.163534Z","iopub.execute_input":"2026-01-23T00:09:03.163894Z","iopub.status.idle":"2026-01-23T00:09:03.176131Z","shell.execute_reply.started":"2026-01-23T00:09:03.163848Z","shell.execute_reply":"2026-01-23T00:09:03.175423Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_coor = pd.read_csv(PATH + 'train_label_coordinates.csv')\ndf_coor.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:03.177143Z","iopub.execute_input":"2026-01-23T00:09:03.177485Z","iopub.status.idle":"2026-01-23T00:09:03.304496Z","shell.execute_reply.started":"2026-01-23T00:09:03.177436Z","shell.execute_reply":"2026-01-23T00:09:03.303660Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S = 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":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:03.305567Z","iopub.execute_input":"2026-01-23T00:09:03.305810Z","iopub.status.idle":"2026-01-23T00:09:03.329322Z","shell.execute_reply.started":"2026-01-23T00:09:03.305782Z","shell.execute_reply":"2026-01-23T00:09:03.328601Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:03.332446Z","iopub.execute_input":"2026-01-23T00:09:03.332812Z","iopub.status.idle":"2026-01-23T00:09:03.391876Z","shell.execute_reply.started":"2026-01-23T00:09:03.332772Z","shell.execute_reply":"2026-01-23T00:09:03.391163Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S = S[S['x_mean_fraction'] > .8]\nS.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:03.392933Z","iopub.execute_input":"2026-01-23T00:09:03.393717Z","iopub.status.idle":"2026-01-23T00:09:03.406040Z","shell.execute_reply.started":"2026-01-23T00:09:03.393669Z","shell.execute_reply":"2026-01-23T00:09:03.405237Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S['instance_number'] = S['instance_number'] - 1\nS.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:03.407087Z","iopub.execute_input":"2026-01-23T00:09:03.407415Z","iopub.status.idle":"2026-01-23T00:09:03.422732Z","shell.execute_reply.started":"2026-01-23T00:09:03.407370Z","shell.execute_reply":"2026-01-23T00:09:03.422020Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:03.423665Z","iopub.execute_input":"2026-01-23T00:09:03.423861Z","iopub.status.idle":"2026-01-23T00:09:04.197460Z","shell.execute_reply.started":"2026-01-23T00:09:03.423837Z","shell.execute_reply":"2026-01-23T00:09:04.196635Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:04.198601Z","iopub.execute_input":"2026-01-23T00:09:04.198961Z","iopub.status.idle":"2026-01-23T00:09:04.211914Z","shell.execute_reply.started":"2026-01-23T00:09:04.198909Z","shell.execute_reply":"2026-01-23T00:09:04.211078Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:04.212950Z","iopub.execute_input":"2026-01-23T00:09:04.213293Z","iopub.status.idle":"2026-01-23T00:09:04.345125Z","shell.execute_reply.started":"2026-01-23T00:09:04.213261Z","shell.execute_reply":"2026-01-23T00:09:04.344267Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:04.346052Z","iopub.execute_input":"2026-01-23T00:09:04.346418Z","iopub.status.idle":"2026-01-23T00:09:04.371148Z","shell.execute_reply.started":"2026-01-23T00:09:04.346373Z","shell.execute_reply":"2026-01-23T00:09:04.370419Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S = S.merge(train,left_on='study_id',right_on='study_id')\nS.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:04.372177Z","iopub.execute_input":"2026-01-23T00:09:04.372416Z","iopub.status.idle":"2026-01-23T00:09:04.411754Z","shell.execute_reply.started":"2026-01-23T00:09:04.372378Z","shell.execute_reply":"2026-01-23T00:09:04.411048Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:04.412912Z","iopub.execute_input":"2026-01-23T00:09:04.413650Z","iopub.status.idle":"2026-01-23T00:09:04.424061Z","shell.execute_reply.started":"2026-01-23T00:09:04.413602Z","shell.execute_reply":"2026-01-23T00:09:04.423288Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S['flip'] = False\nfS = S.copy()\nfS['flip'] = True\nS = pd.concat([S,fS]).reset_index(drop=True)\nS.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:04.425159Z","iopub.execute_input":"2026-01-23T00:09:04.425493Z","iopub.status.idle":"2026-01-23T00:09:04.453341Z","shell.execute_reply.started":"2026-01-23T00:09:04.425451Z","shell.execute_reply":"2026-01-23T00:09:04.452589Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S.groupby('fold').count()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:04.454377Z","iopub.execute_input":"2026-01-23T00:09:04.454739Z","iopub.status.idle":"2026-01-23T00:09:04.472448Z","shell.execute_reply.started":"2026-01-23T00:09:04.454692Z","shell.execute_reply":"2026-01-23T00:09:04.471697Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Sagittal_T2_Spinal_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        label = torch.as_tensor([labels[x] for x in row[spinal]]).to(device)\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        label[missing] = - 100\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        if not self.VALID: c += torch.normal(0,5,(5,2)).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,Lmax,2*self.P,2*self.P).to(device)\n        slices_mask = torch.ones(5,Lmax).bool().to(device)\n        for i in range(5):\n            if ~missing[i]:\n                instance_number = instance_numbers[i].astype(int)\n                start = max([0,instance_number - Lmax//2])\n                end = min([D,start + Lmax])\n                crop = crops[i,start:end]\n\n                if row.flip:\n                    image[i,:len(crop)] = crop.flip(0)\n                else:\n                    image[i,:len(crop)] = crop\n                slices_mask[i,:len(crop)] = False\n\n                if not self.VALID:\n                    image[i] = augment_image(image[i].reshape(-1,2*self.P,2*self.P)).reshape(Lmax,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        return [image,slices_mask],label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:04.473456Z","iopub.execute_input":"2026-01-23T00:09:04.474052Z","iopub.status.idle":"2026-01-23T00:09:04.488710Z","shell.execute_reply.started":"2026-01-23T00:09:04.474012Z","shell.execute_reply":"2026-01-23T00:09:04.487821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ds = Sagittal_T2_Spinal_Dataset(S)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:04.489809Z","iopub.execute_input":"2026-01-23T00:09:04.490159Z","iopub.status.idle":"2026-01-23T00:09:04.506512Z","shell.execute_reply.started":"2026-01-23T00:09:04.490114Z","shell.execute_reply":"2026-01-23T00:09:04.505649Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"[sample,mask],label = ds.__getitem__(np.random.randint(len(ds)))\nprint(label)\nfor k in range(5):\n    fig, axes = plt.subplots(1,Lmax, figsize=(10,2))\n    for i in range(2):\n        for j in range(Lmax):\n            axes[j].imshow(sample.cpu()[k,j])\n    plt.show()\n\nplt.imshow(mask.cpu().view(-1,Lmax)) ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:04.507464Z","iopub.execute_input":"2026-01-23T00:09:04.507817Z","iopub.status.idle":"2026-01-23T00:09:09.366977Z","shell.execute_reply.started":"2026-01-23T00:09:04.507790Z","shell.execute_reply":"2026-01-23T00:09:09.366190Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del ds\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:09.368209Z","iopub.execute_input":"2026-01-23T00:09:09.368595Z","iopub.status.idle":"2026-01-23T00:09:09.668299Z","shell.execute_reply.started":"2026-01-23T00:09:09.368554Z","shell.execute_reply":"2026-01-23T00:09:09.667433Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:09.669262Z","iopub.execute_input":"2026-01-23T00:09:09.669615Z","iopub.status.idle":"2026-01-23T00:09:09.686798Z","shell.execute_reply.started":"2026-01-23T00:09:09.669586Z","shell.execute_reply":"2026-01-23T00:09:09.686028Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SinusoidalPosEmb(nn.Module):\n    def __init__(self, dim=16, M=10000):\n        super().__init__()\n        self.dim = dim\n        self.M = M\n\n    def forward(self, x):\n        device = x.device\n        half_dim = self.dim // 2\n        emb = math.log(self.M) / half_dim\n        emb = torch.exp(torch.arange(half_dim, device=device) * (-emb))\n        emb = x[...,None] * emb[None,...]\n        emb = torch.cat((emb.sin(), emb.cos()), dim=-1)\n        return emb\n\nclass Sagittal_T2_Spinal_ViT(nn.Module):\n    def __init__(\n            self,\n            ENCODER,\n            dim=512,\n            depth=24,\n            head_size=64\n        ):\n        super().__init__()\n        self.ENCODER = ENCODER\n        self.slices_enc = SinusoidalPosEmb(dim)(torch.arange(Lmax, device=device).unsqueeze(0))\n        self.slices_enc = nn.Parameter(self.slices_enc)\n        pos_enc = SinusoidalPosEmb(dim)(torch.arange(5, device=device).unsqueeze(0))\n        self.pos_enc = nn.Parameter(pos_enc)\n        self.slices_transformer = nn.TransformerEncoder(\n                nn.TransformerEncoderLayer(d_model=dim, nhead=dim//head_size, dim_feedforward=4*dim,\n                dropout=0.1, activation=nn.GELU(), batch_first=True, norm_first=True, device=device), depth)\n        self.transformer = nn.TransformerEncoder(\n                nn.TransformerEncoderLayer(d_model=dim, nhead=dim//head_size, dim_feedforward=4*dim,\n                dropout=0.1, activation=nn.GELU(), batch_first=True, norm_first=True, device=device), depth)\n        self.proj_out = nn.Linear(dim,3).to(device)\n    \n    def forward(self, x):\n        x,slices_mask = x\n        slices_mask = slices_mask.view(-1,Lmax)\n        mask = slices_mask.sum(-1) < Lmax\n        \n        x = self.ENCODER(x.view(-1,1,patch_size,patch_size))\n\n        x = x.view(-1,Lmax,512)\n        x = x + self.slices_enc\n        x[mask] = self.slices_transformer(x[mask],src_key_padding_mask=slices_mask[mask])\n\n        x[slices_mask] = 0\n        d = (~slices_mask).sum(1).unsqueeze(-1).tile(1,512)\n        x = x.sum(1)\n        x[d > 0] = x[d > 0]/d[d > 0]\n\n        level_mask = (slices_mask.sum(1) == Lmax).view(-1,5)\n        x = x.view(-1,5,512) + self.pos_enc\n        x = self.transformer(x,src_key_padding_mask=level_mask)\n        x = self.proj_out(x.view(-1,512)).view(-1,5,3)\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:09.687694Z","iopub.execute_input":"2026-01-23T00:09:09.688305Z","iopub.status.idle":"2026-01-23T00:09:09.713845Z","shell.execute_reply.started":"2026-01-23T00:09:09.688273Z","shell.execute_reply":"2026-01-23T00:09:09.712909Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def myLoss(preds,target):\n    target = target.view(-1)\n    preds = preds.view(-1,3)\n    \n    return nn.CrossEntropyLoss(weight=torch.as_tensor([1.,2.,4.]).to(device))(preds,target)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:09.714788Z","iopub.execute_input":"2026-01-23T00:09:09.715494Z","iopub.status.idle":"2026-01-23T00:09:09.732000Z","shell.execute_reply.started":"2026-01-23T00:09:09.715464Z","shell.execute_reply":"2026-01-23T00:09:09.731040Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm.auto import tqdm\nimport time\nfrom sklearn.metrics import (\n    confusion_matrix, classification_report,\n    roc_auc_score, roc_curve, auc,\n    cohen_kappa_score, matthews_corrcoef,\n    precision_recall_fscore_support, accuracy_score\n)\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport torch\nimport seaborn as sns\nfrom sklearn.metrics import roc_curve, auc\nimport pandas as pd\n%matplotlib inline\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:09.736098Z","iopub.execute_input":"2026-01-23T00:09:09.736388Z","iopub.status.idle":"2026-01-23T00:09:09.996903Z","shell.execute_reply.started":"2026-01-23T00:09:09.736359Z","shell.execute_reply":"2026-01-23T00:09:09.996132Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================\n#              PAPER-READY ENHANCEMENTS\n# ================================================\nimport seaborn as sns\nfrom sklearn.metrics import roc_curve, auc\nfrom sklearn.preprocessing import label_binarize\nimport pandas as pd\nfrom scipy.stats import ttest_rel\n\n# For combined confusion matrix later\nall_fold_labels = []\nall_fold_preds = []\n\n# Publication style\nplt.rcParams.update({\n    \"font.size\": 11,\n    \"axes.labelsize\": 11,\n    \"axes.titlesize\": 12,\n    \"legend.fontsize\": 9,\n    \"figure.dpi\": 300\n})\nsns.set_style(\"whitegrid\")\n\n\n# ======================= FOLD LOOP =======================\n\nfold_results = []\ntotal_start_time = time.time()\n\nfor fold_idx, f in enumerate(FOLDS):\n    print(f\"\\n================== Fold {f} ({fold_idx+1}/{len(FOLDS)}) ==================\")\n    fold_start = time.time()\n\n    seed_everything(SEED)\n\n    # ----- Load discriminator -----\n    discriminator = torch.load(\n        \"/kaggle/input/lumbar-spine-keypoint-detection-models/Sagittal_T2_spine_discriminator_3\",\n        weights_only=False\n    )\n    model = Sagittal_T2_Spinal_ViT(discriminator.emb)\n\n    # ----- Split data -----\n    df = S\n    tdf = df[df.fold != f]\n    vdf = df[df.fold == f]\n\n    tds = Sagittal_T2_Spinal_Dataset(tdf)\n    vds = Sagittal_T2_Spinal_Dataset(vdf, VALID=True)\n\n    tdl = torch.utils.data.DataLoader(tds, batch_size=BS, shuffle=True, drop_last=True)\n    vdl = torch.utils.data.DataLoader(vds, batch_size=BS, shuffle=False)\n\n    dls = DataLoaders(tdl, vdl)\n\n    n_iter = len(tds)//BS\n\n    # ----- Learner -----\n    learn = Learner(\n        dls,\n        model,\n        loss_func=myLoss,\n        cbs=[ShowGraphCallback(), GradientClip(3.0), alpha_cb]\n    )\n\n    # ----- Train -----\n    print(\"Training...\")\n    train_start = time.time()\n    learn.fit_one_cycle(EPOCHS, lr_max=LR_MAX, wd=0.05, pct_start=0.02)\n    train_end = time.time()\n    print(f\"Training time: {(train_end - train_start):.2f} sec\")\n\n    # ----- Save loss curves -----\n    train_losses = learn.recorder.values\n    plt.figure()\n    plt.plot([x[0] for x in train_losses], label=\"train_loss\")\n    plt.plot([x[1] for x in train_losses], label=\"valid_loss\")\n    plt.title(f\"Loss Curve Fold {f}\")\n    plt.legend()\n    plt.savefig(f\"loss_curve_fold_{f}.png\")\n    plt.close()\n\n    # =================== VALIDATION ===================\n    \n    print(\"Validating...\")\n    model.eval()\n    all_preds = []\n    all_labels = []\n    all_probs_batches = []\n\n    val_start = time.time()\n\n    with torch.no_grad():\n        for (data, mask), labels in tqdm(vdl, desc=f\"Fold {f} Validation\", leave=True):\n            probs = model([data, mask]).softmax(-1)  # (B,5,3)\n            preds = probs.argmax(-1).cpu().numpy().reshape(-1)\n            true = labels.cpu().numpy().reshape(-1)\n\n            all_preds.append(preds)\n            all_labels.append(true)\n            all_probs_batches.append(probs.cpu().numpy().reshape(-1, 3))\n\n    val_end = time.time()\n    print(f\"Validation time: {(val_end - val_start):.2f} sec\")\n\n    # Merge batches\n    all_preds = np.concatenate(all_preds)\n    all_labels = np.concatenate(all_labels)\n    all_probs = np.concatenate(all_probs_batches)\n\n    # Remove ignored labels (-100)\n    valid_mask = all_labels != -100\n    all_preds = all_preds[valid_mask]\n    all_labels = all_labels[valid_mask]\n    all_probs = all_probs[valid_mask]\n\n    # =================== METRICS ===================\n\n    # Accuracy\n    acc = accuracy_score(all_labels, all_preds)\n\n    # Precision / Recall / F1\n    prec_macro, rec_macro, f1_macro, _ = precision_recall_fscore_support(\n        all_labels, all_preds, average='macro'\n    )\n    f1_weighted = precision_recall_fscore_support(\n        all_labels, all_preds, average='weighted'\n    )[2]\n\n    # Cohen Kappa & MCC\n    kappa = cohen_kappa_score(all_labels, all_preds)\n    mcc = matthews_corrcoef(all_labels, all_preds)\n\n    # AUC (OvR)\n    num_classes = 3\n    y_true_oh = np.eye(num_classes)[all_labels]  # one-hot\n\n    auc_macro = roc_auc_score(y_true_oh, all_probs, average='macro', multi_class='ovr')\n    auc_micro = roc_auc_score(y_true_oh, all_probs, average='micro', multi_class='ovr')\n\n    # Confusion Matrix\n    cm = confusion_matrix(all_labels, all_preds)\n    print(\"\\nConfusion Matrix:\")\n    print(cm)\n\n    # Classification report\n    print(\"\\nClassification Report:\")\n    print(classification_report(all_labels, all_preds, digits=3))\n\n    # =================== ADDED FOR PAPER ===================\n    # Confusion heatmap\n    plt.figure(figsize=(5,4))\n    sns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\",\n                xticklabels=range(num_classes), yticklabels=range(num_classes))\n    plt.xlabel(\"Predicted\")\n    plt.ylabel(\"True\")\n    plt.title(f\"Confusion Matrix Fold {f}\")\n    plt.tight_layout()\n    plt.savefig(f\"confusion_matrix_heatmap_fold_{f}.png\")\n    plt.close()\n\n    # ROC Curve\n    y_bin = label_binarize(all_labels, classes=list(range(num_classes)))\n    plt.figure(figsize=(6,5))\n    for i in range(num_classes):\n        fpr, tpr, _ = roc_curve(y_bin[:, i], all_probs[:, i])\n        roc_auc = auc(fpr, tpr)\n        plt.plot(fpr, tpr, lw=2, label=f\"Class {i} (AUC={roc_auc:.3f})\")\n    plt.plot([0,1],[0,1],'k--', lw=1)\n    plt.xlabel('False Positive Rate')\n    plt.ylabel('True Positive Rate')\n    plt.title(f'ROC Curve Fold {f}')\n    plt.legend()\n    plt.tight_layout()\n    plt.savefig(f\"roc_curve_fold_{f}.png\")\n    plt.close()\n\n    # Save for combined later\n    all_fold_labels.append(all_labels)\n    all_fold_preds.append(all_preds)\n    # =======================================================\n\n    # Save fold metrics\n    fold_results.append({\n        \"fold\": f,\n        \"accuracy\": acc,\n        \"macro_f1\": f1_macro,\n        \"weighted_f1\": f1_weighted,\n        \"macro_auc\": auc_macro,\n        \"micro_auc\": auc_micro,\n        \"kappa\": kappa,\n        \"mcc\": mcc\n    })\n\n    # ----- Fold Timing -----\n    fold_end = time.time()\n    fold_duration = fold_end - fold_start\n    print(f\"Fold {f} duration: {fold_duration/60:.2f} min\")\n\n    # ----- ETA -----\n    folds_done = fold_idx + 1\n    folds_left = len(FOLDS) - folds_done\n    avg_fold_time = (fold_end - total_start_time) / folds_done\n\n    eta_min = (avg_fold_time * folds_left) / 60\n    eta_hr = (avg_fold_time * folds_left) / 3600\n    print(f\"ETA remaining: {eta_min:.1f} minutes ({eta_hr:.2f} hours)\")\n\n    # ----- Cleanup -----\n    torch.save(model, f\"Spinal_ViT_Fold_{f}.pt\")\n    del model, df, tdf, vdf, tds, vds, tdl, vdl, dls, learn\n    gc.collect()\n\n# =================== FINAL SUMMARY ===================\n\ntotal_end = time.time()\nprint(f\"\\n================== 5-FOLD RESULTS ==================\")\n\nfor k in fold_results[0].keys():\n    if k == \"fold\": continue\n    values = [x[k] for x in fold_results]\n    print(f\"{k}: {np.mean(values):.4f} ± {np.std(values):.4f}\")\n\nprint(f\"\\nTotal training time: {(total_end-total_start_time)/3600:.2f} hours\")\n\n# =================== ADDED FINAL PAPER OUTPUTS ===================\n\n# Save CSV & LaTeX\ndf_res = pd.DataFrame(fold_results)\ndf_res.to_csv(\"kfold_results.csv\", index=False)\nwith open(\"kfold_results.tex\",\"w\") as f:\n    f.write(df_res.to_latex(index=False, float_format=\"%.4f\"))\n\n# Combined confusion matrix across folds\nlabels_all = np.concatenate(all_fold_labels)\npreds_all = np.concatenate(all_fold_preds)\ncm_all = confusion_matrix(labels_all, preds_all)\n\nplt.figure(figsize=(6,5))\nsns.heatmap(cm_all, annot=True, fmt=\"d\", cmap=\"Oranges\",\n            xticklabels=range(num_classes), yticklabels=range(num_classes))\nplt.title(\"Combined Confusion Matrix Across Folds\")\nplt.tight_layout()\nplt.savefig(\"combined_confusion_matrix.png\")\nplt.close()\n\n# Statistical test: macro F1 stability\nvals = [x[\"macro_f1\"] for x in fold_results]\nt,p = ttest_rel(vals, vals)  # self vs self = stability check placeholder\nprint(f\"\\nPaired t-test on macro_f1 stability: t={t:.3f}, p={p:.5f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T00:09:09.997858Z","iopub.execute_input":"2026-01-23T00:09:09.998253Z","iopub.status.idle":"2026-01-23T02:11:17.158273Z","shell.execute_reply.started":"2026-01-23T00:09:09.998224Z","shell.execute_reply":"2026-01-23T02:11:17.157264Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"nig, fix this error","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"----end of code---","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\n# ======================= FOLD LOOP =======================\n\nfold_results = []\ntotal_start_time = time.time()\n\nfor fold_idx, f in enumerate(FOLDS):\n    print(f\"\\n================== Fold {f} ({fold_idx+1}/{len(FOLDS)}) ==================\")\n    fold_start = time.time()\n\n    seed_everything(SEED)\n\n    # ----- Load discriminator -----\n    discriminator = torch.load(\n        \"/kaggle/input/lumbar-spine-keypoint-detection-models/Sagittal_T2_spine_discriminator_3\",\n        weights_only=False\n    )\n    model = Sagittal_T2_Spinal_ViT(discriminator.emb)\n\n    # ----- Split data -----\n    df = S\n    tdf = df[df.fold != f]\n    vdf = df[df.fold == f]\n\n    tds = Sagittal_T2_Spinal_Dataset(tdf)\n    vds = Sagittal_T2_Spinal_Dataset(vdf, VALID=True)\n\n    tdl = torch.utils.data.DataLoader(tds, batch_size=BS, shuffle=True, drop_last=True)\n    vdl = torch.utils.data.DataLoader(vds, batch_size=BS, shuffle=False)\n\n    dls = DataLoaders(tdl, vdl)\n\n    n_iter = len(tds)//BS\n\n    # ----- Learner -----\n    learn = Learner(\n        dls,\n        model,\n        loss_func=myLoss,\n        cbs=[ShowGraphCallback(), GradientClip(3.0), alpha_cb]\n    )\n\n    # ----- Train -----\n    print(\"Training...\")\n    train_start = time.time()\n    learn.fit_one_cycle(EPOCHS, lr_max=LR_MAX, wd=0.05, pct_start=0.02)\n    train_end = time.time()\n    print(f\"Training time: {(train_end - train_start):.2f} sec\")\n\n    # ----- Save loss curves -----\n    train_losses = learn.recorder.values\n    plt.figure()\n    plt.plot([x[0] for x in train_losses], label=\"train_loss\")\n    plt.plot([x[1] for x in train_losses], label=\"valid_loss\")\n    plt.title(f\"Loss Curve Fold {f}\")\n    plt.legend()\n    plt.savefig(f\"loss_curve_fold_{f}.png\")\n    plt.close()\n\n    # =================== VALIDATION ===================\n    \n    print(\"Validating...\")\n    model.eval()\n    all_preds = []\n    all_labels = []\n    all_probs_batches = []\n\n    val_start = time.time()\n\n    with torch.no_grad():\n        for (data, mask), labels in tqdm(vdl, desc=f\"Fold {f} Validation\", leave=True):\n            probs = model([data, mask]).softmax(-1)  # (B,5,3)\n            preds = probs.argmax(-1).cpu().numpy().reshape(-1)\n            true = labels.cpu().numpy().reshape(-1)\n\n            all_preds.append(preds)\n            all_labels.append(true)\n            all_probs_batches.append(probs.cpu().numpy().reshape(-1, 3))\n\n    val_end = time.time()\n    print(f\"Validation time: {(val_end - val_start):.2f} sec\")\n\n    # Merge batches\n    all_preds = np.concatenate(all_preds)\n    all_labels = np.concatenate(all_labels)\n    all_probs = np.concatenate(all_probs_batches)\n\n    # Remove ignored labels (-100)\n    valid_mask = all_labels != -100\n    all_preds = all_preds[valid_mask]\n    all_labels = all_labels[valid_mask]\n    all_probs = all_probs[valid_mask]\n\n    # =================== METRICS ===================\n\n    # Accuracy\n    acc = accuracy_score(all_labels, all_preds)\n\n    # Precision / Recall / F1\n    prec_macro, rec_macro, f1_macro, _ = precision_recall_fscore_support(\n        all_labels, all_preds, average='macro'\n    )\n    f1_weighted = precision_recall_fscore_support(\n        all_labels, all_preds, average='weighted'\n    )[2]\n\n    # Cohen Kappa & MCC\n    kappa = cohen_kappa_score(all_labels, all_preds)\n    mcc = matthews_corrcoef(all_labels, all_preds)\n\n    # AUC (OvR)\n    num_classes = 3\n    y_true_oh = np.eye(num_classes)[all_labels]  # one-hot\n\n    auc_macro = roc_auc_score(y_true_oh, all_probs, average='macro', multi_class='ovr')\n    auc_micro = roc_auc_score(y_true_oh, all_probs, average='micro', multi_class='ovr')\n\n    # Confusion Matrix\n    cm = confusion_matrix(all_labels, all_preds)\n    print(\"\\nConfusion Matrix:\")\n    print(cm)\n\n    # Classification report\n    print(\"\\nClassification Report:\")\n    print(classification_report(all_labels, all_preds, digits=3))\n\n    # Save fold metrics\n    fold_results.append({\n        \"fold\": f,\n        \"accuracy\": acc,\n        \"macro_f1\": f1_macro,\n        \"weighted_f1\": f1_weighted,\n        \"macro_auc\": auc_macro,\n        \"micro_auc\": auc_micro,\n        \"kappa\": kappa,\n        \"mcc\": mcc\n    })\n\n    # ----- Fold Timing -----\n    fold_end = time.time()\n    fold_duration = fold_end - fold_start\n    print(f\"Fold {f} duration: {fold_duration/60:.2f} min\")\n\n    # ----- ETA -----\n    folds_done = fold_idx + 1\n    folds_left = len(FOLDS) - folds_done\n    avg_fold_time = (fold_end - total_start_time) / folds_done\n\n    eta_min = (avg_fold_time * folds_left) / 60\n    eta_hr = (avg_fold_time * folds_left) / 3600\n    print(f\"ETA remaining: {eta_min:.1f} minutes ({eta_hr:.2f} hours)\")\n\n    # ----- Cleanup -----\n    torch.save(model, f\"Spinal_ViT_Fold_{f}.pt\")\n    del model, df, tdf, vdf, tds, vds, tdl, vdl, dls, learn\n    gc.collect()\n\n# =================== FINAL SUMMARY ===================\n\ntotal_end = time.time()\nprint(f\"\\n================== 5-FOLD RESULTS ==================\")\n\nfor k in fold_results[0].keys():\n    if k == \"fold\": continue\n    values = [x[k] for x in fold_results]\n    print(f\"{k}: {np.mean(values):.4f} ± {np.std(values):.4f}\")\n\nprint(f\"\\nTotal training time: {(total_end-total_start_time)/3600:.2f} hours\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T22:35:37.227782Z","iopub.execute_input":"2026-01-22T22:35:37.228146Z","iopub.status.idle":"2026-01-22T22:36:33.636447Z","shell.execute_reply.started":"2026-01-22T22:35:37.228111Z","shell.execute_reply":"2026-01-22T22:36:33.635349Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for f in FOLDS:\n    seed_everything(SEED)\n    model = Sagittal_T2_Spinal_ViT(\n        torch.load(\n            \"/kaggle/input/lumbar-spine-keypoint-detection-models/Sagittal_T2_spine_discriminator_3\",\n            weights_only=False\n        ).emb\n    )\n    df = S\n    tdf = df[df.fold != f]\n    vdf = df[df.fold == f]\n    tds = Sagittal_T2_Spinal_Dataset(tdf)\n    vds = Sagittal_T2_Spinal_Dataset(vdf,VALID=True)\n    tdl = torch.utils.data.DataLoader(tds, batch_size=BS, shuffle=True, drop_last=True)\n    vdl = torch.utils.data.DataLoader(vds, batch_size=BS, shuffle=False)\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            alpha_cb\n        ]\n    )\n    learn.fit_one_cycle(EPOCHS, lr_max=LR_MAX, wd=0.05, pct_start=0.02)\n    torch.save(model,'Sagittal_T2_Spinal_ViT_'+str(f))\n    del model,df,tdf,vdf,tds,vds,tdl,vdl,dls,learn\n    gc.collect()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T22:00:54.987248Z","iopub.execute_input":"2026-01-22T22:00:54.987608Z","iopub.status.idle":"2026-01-22T22:01:04.258597Z","shell.execute_reply.started":"2026-01-22T22:00:54.987582Z","shell.execute_reply":"2026-01-22T22:01:04.257534Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}