{"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.7.12"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"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\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":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.execute_input":"2024-07-22T15:36:47.391937Z","iopub.status.busy":"2024-07-22T15:36:47.389385Z","iopub.status.idle":"2024-07-22T15:37:02.967144Z","shell.execute_reply":"2024-07-22T15:37:02.965578Z","shell.execute_reply.started":"2024-07-22T15:36:47.391871Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEED = 1337\nFOLDS = [1,2,3,4,5]\nPATH = 'C:/Users/Angel/kaggle/'# Main path\nTRAIN_PATH = 'C:/Users/Angel/kaggle/train/'# Training images folder\nENCODER_NAME = \"resnet18\"\nPATCH_H = 512\nPATCH_W = 512\nANGLE = 30\npatch_size = 64\nLmax = 10\nLR_MAX = 5e-6\nBS = 16\nEPOCHS = 16","metadata":{"execution":{"iopub.execute_input":"2024-07-22T15:37:25.201958Z","iopub.status.busy":"2024-07-22T15:37:25.201502Z","iopub.status.idle":"2024-07-22T15:37:25.208848Z","shell.execute_reply":"2024-07-22T15:37:25.207312Z","shell.execute_reply.started":"2024-07-22T15:37:25.201926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.execute_input":"2024-07-22T15:37:02.982642Z","iopub.status.busy":"2024-07-22T15:37:02.982119Z","iopub.status.idle":"2024-07-22T15:37:03.364342Z","shell.execute_reply":"2024-07-22T15:37:03.363214Z","shell.execute_reply.started":"2024-07-22T15:37:02.982516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lforaminal = [\n    'left_neural_foraminal_narrowing_l1_l2',\n    'left_neural_foraminal_narrowing_l2_l3',\n    'left_neural_foraminal_narrowing_l3_l4',\n    'left_neural_foraminal_narrowing_l4_l5',\n    'left_neural_foraminal_narrowing_l5_s1'\n]\nrforaminal = [\n    'right_neural_foraminal_narrowing_l1_l2',\n    'right_neural_foraminal_narrowing_l2_l3',\n    'right_neural_foraminal_narrowing_l3_l4',\n    'right_neural_foraminal_narrowing_l4_l5',\n    'right_neural_foraminal_narrowing_l5_s1'\n]\nlabels = {\n    'Normal/Mild':0,\n    'Moderate':1,\n    'Severe':2,\n    'UNK':-100\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":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv(PATH + 'train_split.csv')\ntrain = train[['study_id','fold']+lforaminal+rforaminal][train[lforaminal+rforaminal].isna().sum(1) < len(lforaminal+rforaminal)].reset_index(drop=True)\ntrain.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = train.fillna('UNK')\ntrain[(train[lforaminal+rforaminal] == 'UNK').sum(1)>0].reset_index(drop=True).tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_coor = pd.read_csv(PATH + 'train_label_coordinates.csv')\ndf_coor.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"L_coor = df_coor[df_coor.condition == 'Left Neural Foraminal Narrowing'].groupby(['study_id','series_id','level']).mean()\nL_coor.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"R_coor = df_coor[df_coor.condition == 'Right Neural Foraminal Narrowing'].groupby(['study_id','series_id','level']).mean()\nR_coor.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"F_coor =L_coor.merge(\n    R_coor,\n    left_on=L_coor.index.to_numpy(),\n    right_on=R_coor.index.to_numpy(),\n    suffixes=('_L','_R')\n)\nF_coor['study_id'] = F_coor['key_0'].apply(lambda v:v[0])\nF_coor['series_id'] = F_coor['key_0'].apply(lambda v:v[1])\nF_coor['level'] = F_coor['key_0'].apply(lambda v:v[2])\nF_coor = F_coor.drop(columns=['key_0'])\nF_coor.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"F_coor['instance_number_L'] = F_coor['instance_number_L'] - 1\nF_coor['instance_number_R'] = F_coor['instance_number_R'] - 1\nF_coor['instance_number'] = (F_coor['instance_number_L'] + F_coor['instance_number_R'])/2\nF_coor['x'] = (F_coor['x_L'] + F_coor['x_R'])/2\nF_coor['y'] = (F_coor['y_L'] + F_coor['y_R'])/2\nF_coor.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"coordinates = {}\nfor study_id,df in F_coor.groupby('study_id'):\n    coordinates[study_id] = {}\nfor (study_id,series_id),df in tqdm(F_coor.groupby(['study_id','series_id'])):\n    coordinates[study_id][series_id] = {\n                'L1/L2':{\n                    'x':torch.nan,\n                    'y':torch.nan,\n                    'L':torch.nan,\n                    'R':torch.nan,\n                    'instance_number':torch.nan\n                },\n                    'L2/L3':{\n                    'x':torch.nan,\n                    'y':torch.nan,\n                    'L':torch.nan,\n                    'R':torch.nan,\n                    'instance_number':torch.nan\n                },\n                'L3/L4':{\n                    'x':torch.nan,\n                    'y':torch.nan,\n                    'L':torch.nan,\n                    'R':torch.nan,\n                    'instance_number':torch.nan\n                },\n                'L4/L5':{\n                    'x':torch.nan,\n                    'y':torch.nan,\n                    'L':torch.nan,\n                    'R':torch.nan,\n                    'instance_number':torch.nan\n                },\n                'L5/S1':{\n                    'x':torch.nan,\n                    'y':torch.nan,\n                    'L':torch.nan,\n                    'R':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']]['L'] = row['instance_number_L']\n        coordinates[row['study_id']][row['series_id']][row['level']]['R'] = row['instance_number_R']\n        coordinates[row['study_id']][row['series_id']][row['level']]['instance_number'] = row['instance_number']","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"F_coor =  F_coor[[\n    'study_id',\n    'series_id'\n]].groupby([\n    'study_id',\n    'series_id'\n]).count().reset_index()\nF_coor.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"v = np.zeros((len(F_coor),25))\nfor i in tqdm(range(len(F_coor))):\n    row = F_coor.iloc[i]\n    k = 0\n    for level in coordinates[row['study_id']][row['series_id']]:\n        v[i,k:k+5] = list(coordinates[row['study_id']][row['series_id']][level].values())\n        k += 5","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"F_coor[[\n    'x_L1L2',\n    'y_L1L2',\n    'L_L1L2',\n    'R_L1L2',\n    'M_L1L2',\n    'x_L2L3',\n    'y_L2L3',\n    'L_L2L3',\n    'R_L2L3',\n    'M_L2L3',\n    'x_L3L4',\n    'y_L3L4',\n    'L_L3L4',\n    'R_L3L4',\n    'M_L3L4',\n    'x_L4L5',\n    'y_L4L5',\n    'L_L4L5',\n    'R_L4L5',\n    'M_L4L5',\n    'x_L5S1',\n    'y_L5S1',\n    'L_L5S1',\n    'R_L5S1',\n    'M_L5S1',\n]] = v\nF_coor.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"F = F_coor.merge(train,left_on='study_id',right_on='study_id')\nF.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mask = F[[\n    'x_L1L2',\n    'x_L2L3',\n    'x_L3L4',\n    'x_L4L5',\n    'x_L5S1'\n]].isna().values\nmask += F[[\n    'y_L1L2',\n    'y_L2L3',\n    'y_L3L4',\n    'y_L4L5',\n    'y_L5S1'\n]].isna().values\nmask += F[[\n    'L_L1L2',\n    'L_L2L3',\n    'L_L3L4',\n    'L_L4L5',\n    'L_L5S1'\n]].isna().values\nmask += F[[\n    'R_L1L2',\n    'R_L2L3',\n    'R_L3L4',\n    'R_L4L5',\n    'R_L5S1'\n]].isna().values\nmask += F[[\n    'M_L1L2',\n    'M_L2L3',\n    'M_L3L4',\n    'M_L4L5',\n    'M_L5S1'\n]].isna().values\nmask = mask > 0\nmask[-5:]\nv = F[lforaminal].values\nv[mask] = 'UNK'\nF[lforaminal] = v\nv = F[rforaminal].values\nv[mask] = 'UNK'\nF[rforaminal] = v","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"F['flipped_L1L2'] = F['L_L1L2'] < F['R_L1L2']\nF['flipped_L2L3'] = F['L_L2L3'] < F['R_L2L3']\nF['flipped_L3L4'] = F['L_L3L4'] < F['R_L3L4']\nF['flipped_L4L5'] = F['L_L4L5'] < F['R_L4L5']\nF['flipped_L5S1'] = F['L_L5S1'] < F['R_L5S1']\nF['flipped'] = F[[\n    'flipped_L1L2',\n    'flipped_L2L3',\n    'flipped_L3L4',\n    'flipped_L4L5',\n    'flipped_L5S1'\n]].sum(1)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"F[[\n    'L_L1L2','R_L1L2',\n    'L_L2L3','R_L2L3',\n    'L_L3L4','R_L3L4',\n    'L_L4L5','R_L4L5',\n    'L_L5S1','R_L5S1',\n    'flipped'\n]][~F['flipped'].isin([0,5])]","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"F = F[F['flipped'].isin([0,5])].drop(columns=[\n    'flipped_L1L2',\n    'flipped_L2L3',\n    'flipped_L3L4',\n    'flipped_L4L5',\n    'flipped_L5S1'\n])\nF['flipped'] = F['flipped'] > 0\nF.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"F.groupby('fold').count()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Sagittal_T1_foraminal_Dataset(Dataset):\n    def __init__(self, df, VALID=False, P=patch_size, alpha=0):\n        self.data = df\n        self.VALID = VALID\n        self.P = P\n        self.alpha = alpha\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        LR = row[[\n            'L_L1L2','R_L1L2',\n            'L_L2L3','R_L2L3',\n            'L_L3L4','R_L3L4',\n            'L_L4L5','R_L4L5',\n            'L_L5S1','R_L5S1'\n        ]].values.reshape(5,2)\n        M = row[[\n            'M_L1L2',\n            'M_L2L3',\n            'M_L3L4',\n            'M_L4L5',\n            'M_L5S1'\n        ]].values\n\n        label = torch.as_tensor([labels[x] for x in row[lforaminal+rforaminal]]).view(2,5).to(device)\n\n        if row['flipped']:\n            LR = np.flip(LR,1)\n            label = label.flip(0)\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] = 0\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(2,5,Lmax,2*self.P,2*self.P).to(device)\n        slices_mask = torch.ones(2,5,Lmax).bool().to(device)\n        for i in range(5):\n            if ~missing[i]:\n                left_start = int(M[i]) + 1\n                left_end = min([D,left_start + Lmax])\n                left_crop = crops[i,left_start:left_end]\n\n                right_end = int(M[i])\n                right_start = max([0,right_end - Lmax])\n                right_crop = crops[i,right_start:right_end]\n\n                image[0,i,:len(left_crop)] = left_crop\n                image[1,i,:len(right_crop)] = right_crop.flip(0)\n\n                slices_mask[0,i,:len(left_crop)] = False\n                slices_mask[1,i,:len(right_crop)] = False\n\n                if not self.VALID:\n                    image[:,i] = augment_image(image[:,i].reshape(-1,2*self.P,2*self.P)).reshape(2,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":{"execution":{"iopub.execute_input":"2024-07-22T15:37:03.377957Z","iopub.status.busy":"2024-07-22T15:37:03.376752Z","iopub.status.idle":"2024-07-22T15:37:03.413904Z","shell.execute_reply":"2024-07-22T15:37:03.412746Z","shell.execute_reply.started":"2024-07-22T15:37:03.377915Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds = Sagittal_T1_foraminal_Dataset(F)","metadata":{},"execution_count":null,"outputs":[]},{"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(2, Lmax, figsize=(10,2))\n    for i in range(2):\n        for j in range(Lmax):\n            axes[i,j].imshow(sample.cpu()[i,k,j])\n    plt.show()\n\nplt.imshow(mask.cpu().view(-1,Lmax))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del ds\ngc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Sagittal_T1_Foramina_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":{},"execution_count":null,"outputs":[]},{"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_T1_Foraminal_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        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,10)\n        x = x.view(-1,10,512) + torch.concat([self.pos_enc,self.pos_enc],1)\n        x = self.transformer(x,src_key_padding_mask=level_mask)\n        x = self.proj_out(x.view(-1,512)).view(-1,2,5,3)\n        return x","metadata":{"execution":{"iopub.execute_input":"2024-07-22T15:37:03.416625Z","iopub.status.busy":"2024-07-22T15:37:03.415818Z","iopub.status.idle":"2024-07-22T15:37:03.443753Z","shell.execute_reply":"2024-07-22T15:37:03.442621Z","shell.execute_reply.started":"2024-07-22T15:37:03.416584Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.execute_input":"2024-07-22T15:37:03.446270Z","iopub.status.busy":"2024-07-22T15:37:03.445500Z","iopub.status.idle":"2024-07-22T15:37:03.459352Z","shell.execute_reply":"2024-07-22T15:37:03.457940Z","shell.execute_reply.started":"2024-07-22T15:37:03.446231Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for f in FOLDS:\n    seed_everything(SEED)\n    model = Sagittal_T1_Foraminal_ViT(\n        torch.load(PATH + 'Sagittal_T1/foraminal_prediction/discriminator/Sagittal_T1_foramina_discriminator_'+str(f)).emb\n    )\n    df = F\n    tdf = df[df.fold != f]\n    vdf = df[df.fold == f]\n    tds = Sagittal_T1_foraminal_Dataset(tdf)\n    vds = Sagittal_T1_foraminal_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_T1_pretrained_foraminal_ViT_'+str(f))\n    del model,df,tdf,vdf,tds,vds,tdl,vdl,dls,learn\n    gc.collect()","metadata":{"execution":{"iopub.execute_input":"2024-07-22T15:39:58.974867Z","iopub.status.busy":"2024-07-22T15:39:58.974392Z","iopub.status.idle":"2024-07-22T15:40:10.837835Z","shell.execute_reply":"2024-07-22T15:40:10.835528Z","shell.execute_reply.started":"2024-07-22T15:39:58.974828Z"},"trusted":true},"execution_count":null,"outputs":[]}]}