{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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":"nvidiaTeslaT4","dataSources":[{"sourceId":36363,"databundleVersionId":4050810,"sourceType":"competition"}],"dockerImageVersionId":30761,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"from glob import glob\nimport matplotlib.pylab as plt\n\nimport pydicom as dicom","metadata":{"execution":{"iopub.status.busy":"2024-09-13T08:15:25.929181Z","iopub.execute_input":"2024-09-13T08:15:25.929559Z","iopub.status.idle":"2024-09-13T08:15:26.120125Z","shell.execute_reply.started":"2024-09-13T08:15:25.929516Z","shell.execute_reply":"2024-09-13T08:15:26.119244Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim","metadata":{"execution":{"iopub.status.busy":"2024-09-13T08:15:26.411853Z","iopub.execute_input":"2024-09-13T08:15:26.41226Z","iopub.status.idle":"2024-09-13T08:15:29.848029Z","shell.execute_reply.started":"2024-09-13T08:15:26.412224Z","shell.execute_reply":"2024-09-13T08:15:29.847097Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_images = glob(\"/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train_images/1.2.826.0.1.3680043.10001/*\")\nlen(train_images)","metadata":{"execution":{"iopub.status.busy":"2024-09-04T13:07:57.072841Z","iopub.execute_input":"2024-09-04T13:07:57.073319Z","iopub.status.idle":"2024-09-04T13:07:57.084992Z","shell.execute_reply.started":"2024-09-04T13:07:57.073277Z","shell.execute_reply":"2024-09-04T13:07:57.083535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.style.use('default')\nfig, axes = plt.subplots(4,4, figsize=(12,12))\ntrain_images\nfor i, ax in enumerate(axes.reshape(-1)):\n    img_path = train_images[i]\n    img = dicom.dcmread(img_path)  \n    ax.imshow(img.pixel_array)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-09-04T13:07:57.763162Z","iopub.execute_input":"2024-09-04T13:07:57.763594Z","iopub.status.idle":"2024-09-04T13:08:00.790141Z","shell.execute_reply.started":"2024-09-04T13:07:57.763553Z","shell.execute_reply":"2024-09-04T13:08:00.788849Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data processing\nshould make images as a tensor for each patient","metadata":{}},{"cell_type":"code","source":"import os\nimport pydicom\nimport torch\nimport numpy as np\nfrom torch.utils.data import Dataset\nfrom pathlib import Path\nimport random\n\nclass SpineCTDataset(Dataset):\n    def __init__(self, root_dir, patients_per_load=2, seed=None):\n        self.root_dir = Path(root_dir)\n        self.patients_per_load = patients_per_load\n        self.all_patients = [d for d in self.root_dir.iterdir() if d.is_dir()]\n        self.total_patients = len(self.all_patients)\n        \n        if seed is not None:\n            random.seed(seed)\n        \n        self.load_new_subset()\n\n    def load_new_subset(self):\n        self.current_patients = random.sample(self.all_patients, min(self.patients_per_load, self.total_patients))\n        self.data = self._load_patient_subset()\n\n    def _load_patient_subset(self):\n        patient_subset = {}\n        for patient_dir in self.current_patients:\n            patient_id = patient_dir.name\n            patient_scans = self._load_patient_scans(patient_dir)\n            patient_subset[patient_id] = patient_scans\n        return patient_subset\n\n    def _load_patient_scans(self, patient_dir):\n        scans = []\n        for dcm_file in patient_dir.glob('*.dcm'):\n            try:\n                dicom = pydicom.dcmread(str(dcm_file))\n                image = self._get_pixel_data(dicom)\n                scans.append(image)\n            except RuntimeError as e:\n                print(f\"Error reading {dcm_file}: {str(e)}\")\n                # You might want to skip this file or use an alternative method\n        return np.array(scans)\n\n    def _get_pixel_data(self, dicom):\n        try:\n            return dicom.pixel_array\n        except RuntimeError:\n            # If pydicom fails, try using SimpleITK as an alternative\n            import SimpleITK as sitk\n            img = sitk.ReadImage(dicom.filename)\n            return sitk.GetArrayFromImage(img)\n\n    def __len__(self):\n        return len(self.current_patients)\n\n    def __getitem__(self, idx):\n        patient_id = self.current_patients[idx].name\n        scans = self.data[patient_id]\n        return {'patient_id': patient_id, 'scans': torch.from_numpy(scans).float()}\n\n# Usage example\nroot_directory = '/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train_images'\ndataset = SpineCTDataset(root_directory, patients_per_load=2)\n\n# Access data for a specific patient in the current subset\npatient_data = dataset[0]\nprint(f\"Patient ID: {patient_data['patient_id']}\")\nprint(f\"Number of scans: {patient_data['scans'].shape[0]}\")\nprint(f\"Scan dimensions: {patient_data['scans'].shape[1:]}\")","metadata":{"execution":{"iopub.status.busy":"2024-09-13T08:15:31.011905Z","iopub.execute_input":"2024-09-13T08:15:31.012434Z","iopub.status.idle":"2024-09-13T08:15:41.671356Z","shell.execute_reply.started":"2024-09-13T08:15:31.012394Z","shell.execute_reply":"2024-09-13T08:15:41.670321Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**MODEL DESIGN**","metadata":{}},{"cell_type":"code","source":"patient_data[\"scans\"].shape","metadata":{"execution":{"iopub.status.busy":"2024-09-13T08:15:41.673408Z","iopub.execute_input":"2024-09-13T08:15:41.673892Z","iopub.status.idle":"2024-09-13T08:15:41.681306Z","shell.execute_reply.started":"2024-09-13T08:15:41.673846Z","shell.execute_reply":"2024-09-13T08:15:41.680337Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_size=(512,512)","metadata":{"execution":{"iopub.status.busy":"2024-09-13T08:15:41.682403Z","iopub.execute_input":"2024-09-13T08:15:41.682747Z","iopub.status.idle":"2024-09-13T08:15:41.691758Z","shell.execute_reply.started":"2024-09-13T08:15:41.682708Z","shell.execute_reply":"2024-09-13T08:15:41.691002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MultiHeadAttention(nn.Module):\n    def __init__(self, emb_size, heads):\n        super(MultiHeadAttention,self).__init__()\n        self.heads=heads\n        self.emb_size=emb_size\n        self.head_dim=self.emb_size//self.heads\n        self.w_k=nn.Linear(emb_size,emb_size)\n        self.w_q=nn.Linear(emb_size,emb_size)\n        self.w_v=nn.Linear(emb_size,emb_size)\n        self.out=nn.Linear(emb_size,emb_size)\n\n        assert(self.head_dim * heads == emb_size),\"embeding size is not divisible by number of heads\"\n\n\n    def forward(self,k,q,v,mask=None):\n        N=q.shape[0]  # batch size\n        K=self.w_k(k)\n        Q=self.w_q(q)\n        V=self.w_v(v)\n\n        K=K.view(N,K.shape[1],self.heads,self.head_dim).transpose(1,2)    # (batch size, sequence len, heads, head dimention)\n        Q=Q.view(N,Q.shape[1],self.heads,self.head_dim).transpose(1,2)    # transposed to give(batch size, heads, sequence len, head dimention)\n        V=V.view(N,V.shape[1],self.heads,self.head_dim).transpose(1,2)\n\n        attention=(torch.matmul(Q,K.transpose(-2,-1)))/torch.tensor(self.head_dim**0.5)\n        \n        if mask is not None:\n            mask=mask.reshape(-1,1,1,128)\n            attention.masked_fill_(mask==0, -1e9)\n\n        attention_scores=F.softmax(attention, dim=-1)\n        output=torch.matmul(attention_scores,V)\n        output = output.transpose(1, 2).reshape(N, -1, self.emb_size)\n        output=self.out(output)\n\n        return output","metadata":{"execution":{"iopub.status.busy":"2024-09-13T08:15:41.695306Z","iopub.execute_input":"2024-09-13T08:15:41.696153Z","iopub.status.idle":"2024-09-13T08:15:41.707659Z","shell.execute_reply.started":"2024-09-13T08:15:41.696108Z","shell.execute_reply":"2024-09-13T08:15:41.706826Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Positional embedings need to be changed\n# **need to recheck it**\nneed to use positional embedings","metadata":{}},{"cell_type":"code","source":"def pos_embedding(seq_len, emb_size, n=10000):\n    P = np.zeros((seq_len, emb_size))\n    for pos in range(seq_len):\n        for i in range(emb_size // 2):\n            denominator = np.power(n, 2 * i / emb_size)\n            P[pos, 2 * i] = np.sin(pos / denominator)\n            P[pos, 2 * i + 1] = np.cos(pos / denominator)\n    return torch.tensor(P, dtype=torch.float32)","metadata":{"execution":{"iopub.status.busy":"2024-09-13T08:15:41.708632Z","iopub.execute_input":"2024-09-13T08:15:41.708884Z","iopub.status.idle":"2024-09-13T08:15:41.720265Z","shell.execute_reply.started":"2024-09-13T08:15:41.708856Z","shell.execute_reply":"2024-09-13T08:15:41.719378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Encoder(nn.Module):\n    def __init__(self, heads, emb_size):\n        super(Encoder, self).__init__()\n        self.mha=MultiHeadAttention(emb_size, heads)\n        self.ff1=nn.Linear(emb_size,4*emb_size)\n        self.ff2=nn.Linear(4*emb_size, emb_size)\n        self.norm1=nn.LayerNorm(emb_size)\n        self.norm2=nn.LayerNorm(emb_size)\n        self.dropout=nn.Dropout(p=0.2)\n\n    def forward(self, x, mask=None):\n        attention_out=self.mha(x,x,x,mask)\n        attention_out = self.dropout(attention_out)\n        out1=self.norm1(x+attention_out)\n\n        ff_out=F.relu(self.ff1(out1))\n        ff_out=self.ff2(ff_out)\n        out2=self.dropout(ff_out)\n        encoder_out=self.norm2(out1+out2)\n        return encoder_out","metadata":{"execution":{"iopub.status.busy":"2024-09-13T08:15:41.721404Z","iopub.execute_input":"2024-09-13T08:15:41.72172Z","iopub.status.idle":"2024-09-13T08:15:41.730465Z","shell.execute_reply.started":"2024-09-13T08:15:41.721686Z","shell.execute_reply":"2024-09-13T08:15:41.729511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"there are a lot of images so making it using vision transformers with patches of individual images and take patches from all images is useful","metadata":{}},{"cell_type":"code","source":"128 * (512/8) * (512/8)","metadata":{"execution":{"iopub.status.busy":"2024-09-13T08:15:41.731601Z","iopub.execute_input":"2024-09-13T08:15:41.731876Z","iopub.status.idle":"2024-09-13T08:15:41.746031Z","shell.execute_reply.started":"2024-09-13T08:15:41.731846Z","shell.execute_reply":"2024-09-13T08:15:41.745233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CNNTransformer(nn.Module):\n    def __init__(self, num_encoder_layers):\n        super(CNNTransformer, self).__init__()\n        self.conv_block = nn.ModuleList([\n            nn.Conv2d(in_chan, out_chan, kernel_size=(3, 3), stride=2, padding=1)\n            for in_chan, out_chan in [(1, 32), (32, 64), (64, 128)]                           #rgb or grayscale\n        ])\n        \n        self.fc_block = nn.Linear(524288, 1024)\n        \n        self.encoder_block = nn.ModuleList([\n            Encoder(4, 1024)\n            for _ in range(num_encoder_layers)\n        ])\n        \n        self.fc_end_block = nn.ModuleList([\n            nn.Linear(in_dim, out_dim)\n            for in_dim, out_dim in [(1024,256),(256,16)]\n        ])\n        \n        self.out_fc = nn.Linear(16,2)\n        self.flatten = nn.Flatten()\n        #self.position_encodings = pos_embedding(seq_len, emb_size)\n        \n    def forward(self, num_slits, images):\n        input_ids=[]\n        image_list = []\n\n        # Loop through each slice (image) and append it to the list\n        for i in range(images.shape[0]):\n            image = images[i]\n            image_list.append(image)\n\n        # Now `image_list` contains 258 tensors of shape [512, 512]\n\n        for i in range(num_slits):\n            x=image_list[i]\n            x=x.reshape((1,1,512,512))       #(batchsize, channels, img_h, img_w)\n            for conv in self.conv_block:\n                x = F.relu(conv(x))\n\n            x=self.flatten(x)\n            \n            #print(x.shape)\n\n            x = F.relu(self.fc_block(x))\n\n            input_ids.append(x)\n\n        x = torch.stack(input_ids)  # stack slices along a new dimension\n\n        for encoder in self.encoder_block:\n            x = encoder(x)\n\n        # polling done then useful \n\n        for fc in self.fc_end_block:\n            x = F.relu(fc(x))\n\n        x = self.out_fc(x)\n\n        return x\n        \nmodel = CNNTransformer(num_encoder_layers=2)","metadata":{"execution":{"iopub.status.busy":"2024-09-13T08:15:41.747595Z","iopub.execute_input":"2024-09-13T08:15:41.74827Z","iopub.status.idle":"2024-09-13T08:15:46.258691Z","shell.execute_reply.started":"2024-09-13T08:15:41.748229Z","shell.execute_reply":"2024-09-13T08:15:46.257616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\noptimizer = optim.Adam(model.parameters(), lr=0.002)     #might need to reduce lr","metadata":{"execution":{"iopub.status.busy":"2024-09-13T08:15:46.25996Z","iopub.execute_input":"2024-09-13T08:15:46.260278Z","iopub.status.idle":"2024-09-13T08:15:47.177645Z","shell.execute_reply.started":"2024-09-13T08:15:46.260245Z","shell.execute_reply":"2024-09-13T08:15:47.176843Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"output=model(patient_data['scans'].shape[0],patient_data[\"scans\"])","metadata":{"execution":{"iopub.status.busy":"2024-09-13T08:17:19.212283Z","iopub.execute_input":"2024-09-13T08:17:19.212696Z","iopub.status.idle":"2024-09-13T08:17:50.992615Z","shell.execute_reply.started":"2024-09-13T08:17:19.21266Z","shell.execute_reply":"2024-09-13T08:17:50.991704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"output.shape","metadata":{"execution":{"iopub.status.busy":"2024-09-13T08:18:06.500031Z","iopub.execute_input":"2024-09-13T08:18:06.500758Z","iopub.status.idle":"2024-09-13T08:18:06.506693Z","shell.execute_reply.started":"2024-09-13T08:18:06.500719Z","shell.execute_reply":"2024-09-13T08:18:06.505763Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"patient_data['scans'].shape[0]","metadata":{"execution":{"iopub.status.busy":"2024-09-13T08:18:19.772323Z","iopub.execute_input":"2024-09-13T08:18:19.77271Z","iopub.status.idle":"2024-09-13T08:18:19.778801Z","shell.execute_reply.started":"2024-09-13T08:18:19.772674Z","shell.execute_reply":"2024-09-13T08:18:19.777851Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"patient_data['scans'].shape","metadata":{"execution":{"iopub.status.busy":"2024-09-13T07:34:19.368658Z","iopub.execute_input":"2024-09-13T07:34:19.369041Z","iopub.status.idle":"2024-09-13T07:34:19.375762Z","shell.execute_reply.started":"2024-09-13T07:34:19.369004Z","shell.execute_reply":"2024-09-13T07:34:19.374759Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_list[0].shape","metadata":{"execution":{"iopub.status.busy":"2024-09-13T07:52:54.805573Z","iopub.execute_input":"2024-09-13T07:52:54.805964Z","iopub.status.idle":"2024-09-13T07:52:54.834792Z","shell.execute_reply.started":"2024-09-13T07:52:54.805928Z","shell.execute_reply":"2024-09-13T07:52:54.833427Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}