{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.10.13","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":9334184,"sourceType":"datasetVersion","datasetId":5656063},{"sourceId":183439689,"sourceType":"kernelVersion"}],"dockerImageVersionId":30733,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!unzip -q /kaggle/input/rsna2024-lsdc-making-dataset/_output_.zip ","metadata":{"execution":{"iopub.status.busy":"2024-09-08T13:57:35.564557Z","iopub.execute_input":"2024-09-08T13:57:35.564860Z","iopub.status.idle":"2024-09-08T14:01:36.563813Z","shell.execute_reply.started":"2024-09-08T13:57:35.564835Z","shell.execute_reply":"2024-09-08T14:01:36.562639Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# DINO Contrastive Pretraining\n## Testing this idea of pretraining an already successful model on the task at hand using the self supervised learning DINO emploty\n### Things to try: Pretrain different models, possibly a vision transformer architecture, experiment with using projection layers for contrastive loss, hyperparameter tuning, etc.\n##","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nimport sys\nfrom PIL import Image\nimport cv2\nimport math, random\nimport numpy as np\nimport pandas as pd\nfrom glob import glob\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import KFold\n\nfrom collections import OrderedDict\n\nimport torch\nimport torch.nn.functional as F\nfrom torch import nn\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim import AdamW\n\nimport timm\nfrom transformers import get_cosine_schedule_with_warmup\n\nimport albumentations as A\n\nfrom sklearn.model_selection import KFold","metadata":{"execution":{"iopub.status.busy":"2024-09-08T14:01:36.565923Z","iopub.execute_input":"2024-09-08T14:01:36.566250Z","iopub.status.idle":"2024-09-08T14:01:44.957543Z","shell.execute_reply.started":"2024-09-08T14:01:36.566222Z","shell.execute_reply":"2024-09-08T14:01:44.956771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rd = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'","metadata":{"execution":{"iopub.status.busy":"2024-09-08T14:01:44.958552Z","iopub.execute_input":"2024-09-08T14:01:44.958998Z","iopub.status.idle":"2024-09-08T14:01:44.963124Z","shell.execute_reply.started":"2024-09-08T14:01:44.958972Z","shell.execute_reply":"2024-09-08T14:01:44.962128Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"NOT_DEBUG = True # True -> run naormally, False -> debug mode, with lesser computing cost\n\nOUTPUT_DIR = f'rsna24-results'\ndevice = 'cuda:0' if torch.cuda.is_available() else 'cpu'\nN_WORKERS = os.cpu_count() \nUSE_AMP = True # can change True if using T4 or newer than Ampere\nSEED = 8620\n\nIMG_SIZE = [224, 224]\nIN_CHANS = 42\nN_LABELS = 25\nN_CLASSES = 3 * N_LABELS\n\nAUG_PROB = 0.75\n\nN_FOLDS = 5 if NOT_DEBUG else 2\nEPOCHS = 20 if NOT_DEBUG else 2\nMODEL_NAME = \"tf_efficientnet_b3.ns_jft_in1k\" if NOT_DEBUG else \"tf_efficientnet_b0.ns_jft_in1k\"\n\nGRAD_ACC = 2\nTGT_BATCH_SIZE = 32\nBATCH_SIZE = TGT_BATCH_SIZE // GRAD_ACC\nMAX_GRAD_NORM = None\nEARLY_STOPPING_EPOCH = 3\n\nLR = 2e-4 * TGT_BATCH_SIZE / 32\nWD = 1e-2\nAUG = True","metadata":{"execution":{"iopub.status.busy":"2024-09-08T14:01:44.965234Z","iopub.execute_input":"2024-09-08T14:01:44.965499Z","iopub.status.idle":"2024-09-08T14:01:45.056745Z","shell.execute_reply.started":"2024-09-08T14:01:44.965476Z","shell.execute_reply":"2024-09-08T14:01:45.055766Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.makedirs(OUTPUT_DIR, exist_ok=True)","metadata":{"execution":{"iopub.status.busy":"2024-09-08T14:01:45.058201Z","iopub.execute_input":"2024-09-08T14:01:45.058650Z","iopub.status.idle":"2024-09-08T14:01:45.067530Z","shell.execute_reply.started":"2024-09-08T14:01:45.058622Z","shell.execute_reply":"2024-09-08T14:01:45.066697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_random_seed(seed: int = 8620, deterministic: bool = False):\n    \"\"\"Set seeds\"\"\"\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)  # type: ignore\n    torch.backends.cudnn.benchmark = True\n    torch.backends.cudnn.deterministic = deterministic  # type: ignore\n\nset_random_seed(SEED)","metadata":{"execution":{"iopub.status.busy":"2024-09-08T14:01:45.068725Z","iopub.execute_input":"2024-09-08T14:01:45.069085Z","iopub.status.idle":"2024-09-08T14:01:45.079773Z","shell.execute_reply.started":"2024-09-08T14:01:45.069055Z","shell.execute_reply":"2024-09-08T14:01:45.078917Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Open Dataframes","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv(f'{rd}/train.csv')\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-09-08T14:01:45.080849Z","iopub.execute_input":"2024-09-08T14:01:45.081482Z","iopub.status.idle":"2024-09-08T14:01:45.145415Z","shell.execute_reply.started":"2024-09-08T14:01:45.081447Z","shell.execute_reply":"2024-09-08T14:01:45.144559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Change the state to Label.\n\nThe dataframe contains some Nans, which we will replace with -100 so that We and function can ignore them when calculating the loss and score.","metadata":{}},{"cell_type":"code","source":"df = df.fillna(-100)","metadata":{"execution":{"iopub.status.busy":"2024-09-08T14:01:45.146340Z","iopub.execute_input":"2024-09-08T14:01:45.146577Z","iopub.status.idle":"2024-09-08T14:01:45.160317Z","shell.execute_reply.started":"2024-09-08T14:01:45.146556Z","shell.execute_reply":"2024-09-08T14:01:45.159428Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label2id = {'Normal/Mild': 0, 'Moderate':1, 'Severe':2}\ndf = df.replace(label2id)\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-09-08T14:01:45.161528Z","iopub.execute_input":"2024-09-08T14:01:45.162320Z","iopub.status.idle":"2024-09-08T14:01:45.220965Z","shell.execute_reply.started":"2024-09-08T14:01:45.162293Z","shell.execute_reply":"2024-09-08T14:01:45.220023Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CONDITIONS = [\n    'Spinal Canal Stenosis', \n    'Left Neural Foraminal Narrowing', \n    'Right Neural Foraminal Narrowing',\n    'Left Subarticular Stenosis',\n    'Right Subarticular Stenosis'\n]\n\nLEVELS = [\n    'L1/L2',\n    'L2/L3',\n    'L3/L4',\n    'L4/L5',\n    'L5/S1',\n]","metadata":{"execution":{"iopub.status.busy":"2024-09-08T14:01:45.224675Z","iopub.execute_input":"2024-09-08T14:01:45.225009Z","iopub.status.idle":"2024-09-08T14:01:45.229289Z","shell.execute_reply.started":"2024-09-08T14:01:45.224986Z","shell.execute_reply":"2024-09-08T14:01:45.228407Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define Dataset\n\nThis implementation is very slow and leaves a lot of room for improvement.","metadata":{}},{"cell_type":"code","source":"class RSNA24Dataset(Dataset):\n    def __init__(self, df, phase='train', transform=None):\n        self.df = df\n        self.transform = transform\n        self.phase = phase\n    \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        x = np.zeros((512, 512, IN_CHANS), dtype=np.uint8)\n        t = self.df.iloc[idx]\n        st_id = int(t['study_id'])\n        label = t[1:].values.astype(np.int64)\n        \n        # Sagittal T1\n        for i in range(0, 10, 1):\n            try:\n                p = f'./cvt_png/{st_id}/Sagittal T1/{i:03d}.png'\n                img = Image.open(p).convert('L')\n                img = np.array(img)\n                x[..., i] = img.astype(np.uint8)\n            except:\n                #print(f'failed to load on {st_id}, Sagittal T1')\n                pass\n            \n        # Sagittal T2/STIR\n        for i in range(0, 10, 1):\n            try:\n                p = f'./cvt_png/{st_id}/Sagittal T2_STIR/{i:03d}.png'\n                img = Image.open(p).convert('L')\n                img = np.array(img)\n                x[..., i+10] = img.astype(np.uint8)\n            except:\n                #print(f'failed to load on {st_id}, Sagittal T2/STIR')\n                pass\n            \n        # Axial T2\n        axt2 = glob(f'./cvt_png/{st_id}/Axial T2/*.png')\n        axt2 = sorted(axt2)\n    \n        step = len(axt2) / 10.0\n        st = len(axt2)/2.0 - 4.0*step\n        end = len(axt2)+0.0001\n                \n        for i, j in enumerate(np.arange(st, end, step)):\n            try:\n                p = axt2[max(0, int((j-0.5001).round()))]\n                img = Image.open(p).convert('L')\n                img = np.array(img)\n                x[..., i+20] = img.astype(np.uint8)\n            except:\n                #print(f'failed to load on {st_id}, Sagittal T2/STIR')\n                pass  \n            \n        assert np.sum(x)>0\n            \n        if self.transform is not None:\n            x = self.transform(image=x)['image']\n\n        x = x.transpose(2, 0, 1)\n                \n        return x, label","metadata":{"execution":{"iopub.status.busy":"2024-09-08T14:01:45.230514Z","iopub.execute_input":"2024-09-08T14:01:45.230875Z","iopub.status.idle":"2024-09-08T14:01:45.244523Z","shell.execute_reply.started":"2024-09-08T14:01:45.230845Z","shell.execute_reply":"2024-09-08T14:01:45.243568Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\ndef global_augment(images):\n    size = 224\n    \n    # Define the full transformation pipeline\n    transform = A.Compose([\n        A.RandomResizedCrop(size, size, scale=(0.4, 1.0)),\n        A.HorizontalFlip(),\n        A.ColorJitter(brightness=0.4, contrast=0.4, saturation=0.4, hue=0.1),\n        A.ToGray(p=0.2),\n        A.GaussianBlur(blur_limit=(23, 23), sigma_limit=(0.1, 2.0)),\n        A.Normalize(mean=[0.5], std=[0.5]),\n        ToTensorV2()\n    ])\n\n    transformed_images = []\n\n    for img in images:\n        # Split the 42-channel image into 14 groups of 3 channels\n        channel_splits = torch.split(img, 3, dim=0)\n        transformed_splits = []\n        \n        for split in channel_splits:\n            split_np = split.numpy().transpose(1, 2, 0)  # Convert to HWC format\n            transformed_split = transform(image=split_np)[\"image\"]\n            transformed_splits.append(transformed_split)\n        \n        # Concatenate the transformed splits back into a 42-channel image\n        img = torch.cat(transformed_splits, dim=0)\n        transformed_images.append(img)\n    \n    return torch.stack(transformed_images)\n\ndef multiple_local_augments(images, num_crops=6):\n    size = 96  # Smaller crops for local\n    \n    # Define the full transformation pipeline\n    transform = A.Compose([\n        A.RandomResizedCrop(size, size, scale=(0.05, 0.4)),\n        A.HorizontalFlip(),\n        A.ColorJitter(brightness=0.4, contrast=0.4, saturation=0.4, hue=0.1),\n        A.ToGray(p=0.2),\n        A.GaussianBlur(blur_limit=(23, 23), sigma_limit=(0.1, 2.0)),\n        A.Normalize(mean=[0.5], std=[0.5]),\n        ToTensorV2()\n    ])\n\n    transformed_images = []\n\n    for img in images:\n        # Split the 42-channel image into 14 groups of 3 channels\n        channel_splits = torch.split(img, 3, dim=0)\n        transformed_splits = []\n        \n        for split in channel_splits:\n            split_np = split.numpy().transpose(1, 2, 0)  # Convert to HWC format\n            transformed_split = transform(image=split_np)[\"image\"]\n            transformed_splits.append(transformed_split)\n        \n        # Concatenate the transformed splits back into a 42-channel image\n        img = torch.cat(transformed_splits, dim=0)\n        transformed_images.append(img)\n    \n    return torch.stack(transformed_images)","metadata":{"execution":{"iopub.status.busy":"2024-09-08T14:01:45.245616Z","iopub.execute_input":"2024-09-08T14:01:45.245865Z","iopub.status.idle":"2024-09-08T14:01:45.261912Z","shell.execute_reply.started":"2024-09-08T14:01:45.245844Z","shell.execute_reply":"2024-09-08T14:01:45.261074Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import timm\nfrom collections import OrderedDict\n\nstudent_model_name = \"edgenext_base.in21k_ft_in1k\"\nteacher_model_name = \"edgenext_base.in21k_ft_in1k\" \n\nclass OG_DINO(nn.Module):\n    def __init__(self, num_classes, in_channels=42, pretrained=True):\n        super().__init__()\n        \n        self.teacher = timm.create_model(teacher_model_name,\n                                         in_chans=in_channels,\n                                         num_classes=num_classes,\n                                         pretrained=False,\n                                         features_only=False, \n                                         global_pool='avg')\n        self.t_dict = torch.load('/kaggle/input/marjorie-wc/model_fold-1.pt')\n        new_state_dict = {}\n        for key, value in self.t_dict.items():\n            new_key = key.replace(\"model.\", \"\")  # Remove \"model.\" from the key\n            new_state_dict[new_key] = value\n        new_state_dict = OrderedDict(new_state_dict)\n        self.teacher.load_state_dict(new_state_dict)\n        self.student = timm.create_model(student_model_name,\n                                         in_chans=in_channels,\n                                         num_classes=num_classes,\n                                         pretrained=pretrained,\n                                         features_only=False,\n                                         global_pool='avg')\n        self.register_buffer('center', torch.zeros(1, self.student.head.fc.out_features))\n\n        # Ensure the teacher parameters do not get updated during backprop\n        for param in self.teacher.parameters():\n            param.requires_grad = False\n    @staticmethod\n    def distillation_loss(student_output, teacher_output, center, tau_s, tau_t):\n        \"\"\"\n        Calculates distillation loss with centering and sharpening (function H in pseudocode).\n        \"\"\"\n        # Detach teacher output to stop gradients.\n        teacher_output = teacher_output.detach()\n\n        # Center and sharpen teacher's outputs\n        teacher_probs = F.softmax((teacher_output - center) / tau_t, dim=1)\n\n        # Sharpen student's outputs\n        student_probs = F.log_softmax(student_output / tau_s, dim=1)\n\n        # Calculate cross-entropy loss between student's and teacher's probabilities.\n        loss = - (teacher_probs * student_probs).sum(dim=1).mean()\n        return loss\n    \n    def teacher_update(self, beta: float):\n        for teacher_params, student_params in zip(self.teacher.parameters(), self.student.parameters()):\n            teacher_params.data.mul_(beta).add_(student_params.data, alpha=(1 - beta))","metadata":{"execution":{"iopub.status.busy":"2024-09-08T14:01:45.263155Z","iopub.execute_input":"2024-09-08T14:01:45.263586Z","iopub.status.idle":"2024-09-08T14:01:45.280638Z","shell.execute_reply.started":"2024-09-08T14:01:45.263553Z","shell.execute_reply":"2024-09-08T14:01:45.279650Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DINO(nn.Module):\n    def __init__(self, num_classes, in_channels=42, pretrained=True):\n        super().__init__()\n        \n        self.teacher = timm.create_model(teacher_model_name,\n                                         in_chans=in_channels,\n                                         num_classes=num_classes,\n                                         pretrained=False,\n                                         features_only=False, \n                                         global_pool='avg')\n        self.t_dict = torch.load('/kaggle/input/marjorie-wc/model_fold-1.pt')\n        new_state_dict = {key.replace(\"model.\", \"\"): value for key, value in self.t_dict.items()}\n        new_state_dict = OrderedDict(new_state_dict)\n        self.teacher.load_state_dict(new_state_dict)\n        \n        self.student = timm.create_model(student_model_name,\n                                         in_chans=in_channels,\n                                         num_classes=num_classes,\n                                         pretrained=pretrained,\n                                         features_only=False,\n                                         global_pool='avg')\n        \n        self.register_buffer('center', torch.zeros(1, self.student.head.fc.out_features))\n        \n        # Freeze teacher parameters\n        for param in self.teacher.parameters():\n            param.requires_grad = False\n        \n        # Initialize projection head for student only\n        self.student_projection = nn.Sequential(\n            nn.Linear(num_classes, 256),\n            nn.ReLU(),\n            nn.Linear(256, 128)\n        )\n        \n        # Initialize teacher projection as a copy of student projection\n        self.teacher_projection = nn.Sequential(\n            nn.Linear(num_classes, 256),\n            nn.ReLU(),\n            nn.Linear(256, 128)\n        )\n        self.teacher_projection.load_state_dict(self.student_projection.state_dict())\n        \n        # Ensure teacher projection parameters don't require gradients\n        for param in self.teacher_projection.parameters():\n            param.requires_grad = False\n            \n    @staticmethod\n    def distillation_loss(student_output, teacher_output, center, tau_s, tau_t):\n        teacher_output = teacher_output.detach()\n        teacher_probs = F.softmax((teacher_output - center) / tau_t, dim=1)\n        student_probs = F.log_softmax(student_output / tau_s, dim=1)\n        loss = - (teacher_probs * student_probs).sum(dim=1).mean()\n        return loss\n    \n    @staticmethod\n    def contrastive_loss(q, k, temperature=0.1):\n        q = F.normalize(q, dim=1)\n        k = F.normalize(k, dim=1)\n        logits = torch.einsum('nc,mc->nm', [q, k]) / temperature\n        N = logits.shape[0]\n        labels = torch.arange(N, device=logits.device)\n        return F.cross_entropy(logits, labels) * (2 * temperature)\n    \n    def forward(self, x1, x2):\n        student_output1 = self.student(x1)\n        student_output2 = self.student(x2)\n        \n        with torch.no_grad():\n            teacher_output1 = self.teacher(x1)\n            teacher_output2 = self.teacher(x2)\n        \n        # Project outputs for contrastive loss\n        q1 = self.student_projection(student_output1)\n        q2 = self.student_projection(student_output2)\n        k1 = self.teacher_projection(teacher_output1)\n        k2 = self.teacher_projection(teacher_output2)\n        \n        return student_output1, student_output2, teacher_output1, teacher_output2, q1, q2, k1, k2\n    \n        \n    def teacher_update(self, beta: float):\n        # Update teacher network parameters\n        for t_params, s_params in zip(self.teacher.parameters(), self.student.parameters()):\n            t_params.data.mul_(beta).add_(s_params.data, alpha=(1 - beta))\n        \n        # Update teacher projection parameters\n        for t_params, s_params in zip(self.teacher_projection.parameters(), self.student_projection.parameters()):\n            t_params.data.mul_(beta).add_(s_params.data, alpha=(1 - beta))","metadata":{"execution":{"iopub.status.busy":"2024-09-08T14:01:45.281827Z","iopub.execute_input":"2024-09-08T14:01:45.282102Z","iopub.status.idle":"2024-09-08T14:01:45.754517Z","shell.execute_reply.started":"2024-09-08T14:01:45.282079Z","shell.execute_reply":"2024-09-08T14:01:45.752796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_dino(dino,\n                        data_loader,\n                        optimizer,\n                        device,\n                        num_epochs,\n                        tps=0.1,\n                        tpt=0.04,\n                        beta=0.996,\n                        m=0.9,\n                        weight_decay=1e-6):\n    dino.to(device)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=num_epochs)\n    \n    for epoch in range(num_epochs):\n        print(f\"Epoch: {epoch+1}/{num_epochs}\")\n        total_loss = 0\n        for n, (x, _) in enumerate(data_loader):\n            x1 = global_augment(x).to(device)\n            x2 = multiple_local_augments(x).to(device)\n            \n            student_output1, student_output2, teacher_output1, teacher_output2, q1, q2, k1, k2 = dino(x1, x2)\n            \n            # Compute losses\n            distillation_loss = (dino.distillation_loss(student_output1, teacher_output2, dino.center, tps, tpt) +\n                                 dino.distillation_loss(student_output2, teacher_output1, dino.center, tps, tpt)) / 2\n            \n            contrastive_loss = (dino.contrastive_loss(q1, k2) + dino.contrastive_loss(q2, k1)) / 2\n            \n            loss = distillation_loss + contrastive_loss\n            \n            # Backpropagation\n            optimizer.zero_grad()\n            loss.backward()\n            \n            # Gradient clipping\n            torch.nn.utils.clip_grad_norm_(dino.student.parameters(), max_norm=1.0)\n            \n            optimizer.step()\n            \n            if n % 10 == 0:\n                print(f\"Batch {n}, Loss: {loss.item():.4f}\")\n            total_loss += loss.item()\n            \n            # Update the teacher network parameters\n            dino.teacher_update(beta)\n            \n            # Update the center\n            with torch.no_grad():\n                dino.center = m * dino.center + (1 - m) * torch.cat([teacher_output1, teacher_output2], dim=0).mean(dim=0)\n        \n        print(f\"Epoch {epoch+1}, Average Loss: {total_loss/len(data_loader):.4f}\")\n        scheduler.step()\n    \n    return dino.student","metadata":{"execution":{"iopub.status.busy":"2024-09-08T14:01:45.757041Z","iopub.execute_input":"2024-09-08T14:01:45.757395Z","iopub.status.idle":"2024-09-08T14:01:45.776139Z","shell.execute_reply.started":"2024-09-08T14:01:45.757362Z","shell.execute_reply":"2024-09-08T14:01:45.774926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"        \"\"\"\n        Args:\n        dino: DINO Module\n        data_loader (nn.Module): Dataloader for training\n        optimizer (nn.optimizer): Optimizer for optimization (SGD etc.)\n        defice (torch.device): 'cuda', 'cpu'\n        num_epochs: Number of Epochs\n        tps (float): tau for sharpening student logits\n        tpt: for sharpening teacher logits\n        beta (float): moving average decay \n        m (float): center moveing average decay\n        \"\"\"","metadata":{"execution":{"iopub.status.busy":"2024-09-08T14:01:45.777778Z","iopub.execute_input":"2024-09-08T14:01:45.778235Z","iopub.status.idle":"2024-09-08T14:01:45.792105Z","shell.execute_reply.started":"2024-09-08T14:01:45.778198Z","shell.execute_reply":"2024-09-08T14:01:45.791124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nimport torchvision.transforms as transforms\n\n# Split the data into training and validation sets\ntrain_idx, val_idx = train_test_split(range(len(df)), test_size=0.05, random_state=SEED)\n\n# Print the split information\nprint('#' * 30)\nprint('Start training and validation split')\nprint('#' * 30)\nprint(len(train_idx), len(val_idx))\n\n# Create DataFrames for training and validation sets\ndf_train = df.iloc[train_idx]\n\n# Create the dataset and dataloader for training\ntrain_ds = RSNA24Dataset(df_train, phase='train')\n\ndino = DINO(75)\ntrain_dl = DataLoader(\n    train_ds,\n    batch_size=32,\n    shuffle=True,\n    pin_memory=True,\n    drop_last=True,\n    num_workers=N_WORKERS\n)\noptimizer = torch.optim.AdamW(dino.student.parameters(), lr=1e-4)","metadata":{"execution":{"iopub.status.busy":"2024-09-08T14:01:45.795850Z","iopub.execute_input":"2024-09-08T14:01:45.796228Z","iopub.status.idle":"2024-09-08T14:01:48.869217Z","shell.execute_reply.started":"2024-09-08T14:01:45.796192Z","shell.execute_reply":"2024-09-08T14:01:48.868424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lumbar_dino = train_dino(dino,\n           train_dl,\n           optimizer,\n           device,\n           100)","metadata":{"execution":{"iopub.status.busy":"2024-09-08T14:01:48.870349Z","iopub.execute_input":"2024-09-08T14:01:48.870620Z","iopub.status.idle":"2024-09-08T21:05:46.984176Z","shell.execute_reply.started":"2024-09-08T14:01:48.870596Z","shell.execute_reply":"2024-09-08T21:05:46.983016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(dino.student.state_dict()'dino.pth')","metadata":{"execution":{"iopub.status.busy":"2024-09-08T21:15:14.607898Z","iopub.execute_input":"2024-09-08T21:15:14.608649Z","iopub.status.idle":"2024-09-08T21:15:15.026614Z","shell.execute_reply.started":"2024-09-08T21:15:14.608618Z","shell.execute_reply":"2024-09-08T21:15:15.025344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(dino.student.state_dict)","metadata":{"execution":{"iopub.status.busy":"2024-09-08T21:13:57.259837Z","iopub.execute_input":"2024-09-08T21:13:57.260188Z","iopub.status.idle":"2024-09-08T21:13:57.267284Z","shell.execute_reply.started":"2024-09-08T21:13:57.260161Z","shell.execute_reply":"2024-09-08T21:13:57.266367Z"},"trusted":true},"execution_count":null,"outputs":[]}]}