{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"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":"none","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"}],"dockerImageVersionId":30746,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Import libraries & set variables ","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport pydicom\nimport torch","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-07-30T08:29:13.438198Z","iopub.execute_input":"2024-07-30T08:29:13.438591Z","iopub.status.idle":"2024-07-30T08:29:13.443875Z","shell.execute_reply.started":"2024-07-30T08:29:13.438559Z","shell.execute_reply":"2024-07-30T08:29:13.442523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"base_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'\ntrain_path = f'{base_path}/train.csv'\ntrain_coords_path = f'{base_path}/train_label_coordinates.csv'\ntrain_desc_path = f'{base_path}/train_series_descriptions.csv'\ntest_desc_path = f'{base_path}/test_series_descriptions.csv'\nsample_submission_path = f'{base_path}/sample_submission.csv'\n\ntrain_data = pd.read_csv(train_path)\ntrain_coords = pd.read_csv(train_coords_path)\ntrain_desc = pd.read_csv(train_desc_path)\ntest_desc = pd.read_csv(test_desc_path)\nsample_submission = pd.read_csv(sample_submission_path)\n\ntrain_img_path = f'{base_path}/train_images'\ntest_img_path = f'{base_path}/test_images'","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:29:13.445617Z","iopub.execute_input":"2024-07-30T08:29:13.446073Z","iopub.status.idle":"2024-07-30T08:29:13.571331Z","shell.execute_reply.started":"2024-07-30T08:29:13.446032Z","shell.execute_reply":"2024-07-30T08:29:13.570177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create Dataset","metadata":{}},{"cell_type":"markdown","source":"For our Dataset, we want to map \nImages (which can be found by study_id and series_id)\nand series description (sagittal, axial etc)\nto \none-hot encoded labels (normal/moderate/severe)\n\nThis is done in 2 steps\n1. Generate the full DataFrame (study_id, series_id, description, all one-hot encoded labels)\n2. Split into 2 smaller DataFrames, filtered by description. \nSagittal images are mapped to Foraminal labels, axial images are mapped to Subarticular and Canal labels ","metadata":{"execution":{"iopub.status.busy":"2024-07-28T08:00:45.633746Z","iopub.execute_input":"2024-07-28T08:00:45.634271Z","iopub.status.idle":"2024-07-28T08:00:45.643080Z","shell.execute_reply.started":"2024-07-28T08:00:45.634232Z","shell.execute_reply":"2024-07-28T08:00:45.641467Z"}}},{"cell_type":"code","source":"df_merged = pd.merge(train_data, train_desc, on='study_id', how='inner')\ndf_merged.insert(1, 'series_description', df_merged.pop('series_description'))\ndf_merged.insert(1, 'series_id', df_merged.pop('series_id'))\ndf_merged","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:29:13.572841Z","iopub.execute_input":"2024-07-30T08:29:13.573205Z","iopub.status.idle":"2024-07-30T08:29:13.620989Z","shell.execute_reply.started":"2024-07-30T08:29:13.573176Z","shell.execute_reply":"2024-07-30T08:29:13.619803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Extract labels \ndf_index = df_merged.columns[3:]\ndf_id = df_merged.columns[:3]\nsagittal_cols = df_index.str.contains(\"foraminal\")\naxial_cols = ~sagittal_cols\nsagittal_cols = df_id.append(df_index[sagittal_cols])\naxial_cols = df_id.append(df_index[axial_cols])\n\ndf_sagittal = df_merged[df_merged['series_description'].str.contains(\"Sagittal\")][sagittal_cols]\ndf_axial = df_merged[df_merged['series_description'].str.contains(\"Axial\")][axial_cols]","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:29:13.622697Z","iopub.execute_input":"2024-07-30T08:29:13.623147Z","iopub.status.idle":"2024-07-30T08:29:13.647898Z","shell.execute_reply.started":"2024-07-30T08:29:13.623102Z","shell.execute_reply":"2024-07-30T08:29:13.646652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_sagittal.head()","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:29:13.649485Z","iopub.execute_input":"2024-07-30T08:29:13.649934Z","iopub.status.idle":"2024-07-30T08:29:13.675132Z","shell.execute_reply.started":"2024-07-30T08:29:13.649897Z","shell.execute_reply":"2024-07-30T08:29:13.673899Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_axial.head()","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:29:13.676532Z","iopub.execute_input":"2024-07-30T08:29:13.676903Z","iopub.status.idle":"2024-07-30T08:29:13.704844Z","shell.execute_reply.started":"2024-07-30T08:29:13.676873Z","shell.execute_reply":"2024-07-30T08:29:13.703554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# One-hot encoding of respective labels\n# Recall that (Sagittal -> Foraminal) and (Axial -> Subarticular and Canal)\ndf_sagittal = pd.get_dummies(df_sagittal, columns=df_sagittal.columns[3:])\ndf_axial = pd.get_dummies(df_axial, columns=df_axial.columns[3:])","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:29:13.706665Z","iopub.execute_input":"2024-07-30T08:29:13.707055Z","iopub.status.idle":"2024-07-30T08:29:13.761955Z","shell.execute_reply.started":"2024-07-30T08:29:13.707019Z","shell.execute_reply":"2024-07-30T08:29:13.760690Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_sagittal.head()","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:29:13.763260Z","iopub.execute_input":"2024-07-30T08:29:13.763670Z","iopub.status.idle":"2024-07-30T08:29:13.796365Z","shell.execute_reply.started":"2024-07-30T08:29:13.763638Z","shell.execute_reply":"2024-07-30T08:29:13.794679Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_axial.head()","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:29:13.798034Z","iopub.execute_input":"2024-07-30T08:29:13.798556Z","iopub.status.idle":"2024-07-30T08:29:13.835977Z","shell.execute_reply.started":"2024-07-30T08:29:13.798508Z","shell.execute_reply":"2024-07-30T08:29:13.834590Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Small test to ensure everything works\ntemp = df_sagittal.iloc[0]\ntemp_study = temp['study_id']\ntemp_series = temp['series_id']\ntemp_path = f\"{train_img_path}/{temp_study}/{temp_series}\"\ntemp_path = f\"{temp_path}/{os.listdir(temp_path)[0]}\"\ndicom_image = pydicom.dcmread(temp_path)\nimg_arr = torch.tensor(np.array(dicom_image.pixel_array, dtype=np.float32))\nimg_arr","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:29:13.837774Z","iopub.execute_input":"2024-07-30T08:29:13.838194Z","iopub.status.idle":"2024-07-30T08:29:13.885129Z","shell.execute_reply.started":"2024-07-30T08:29:13.838162Z","shell.execute_reply":"2024-07-30T08:29:13.883844Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_arr.max()","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:29:13.887128Z","iopub.execute_input":"2024-07-30T08:29:13.887627Z","iopub.status.idle":"2024-07-30T08:29:13.896910Z","shell.execute_reply.started":"2024-07-30T08:29:13.887589Z","shell.execute_reply":"2024-07-30T08:29:13.895523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_arr.min()","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:29:13.898273Z","iopub.execute_input":"2024-07-30T08:29:13.898685Z","iopub.status.idle":"2024-07-30T08:29:13.910524Z","shell.execute_reply.started":"2024-07-30T08:29:13.898652Z","shell.execute_reply":"2024-07-30T08:29:13.908687Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%matplotlib inline\nimport matplotlib.pyplot as plt\n\nplt.imshow(img_arr, cmap='gray')","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:29:13.911930Z","iopub.execute_input":"2024-07-30T08:29:13.912297Z","iopub.status.idle":"2024-07-30T08:29:14.255344Z","shell.execute_reply.started":"2024-07-30T08:29:13.912267Z","shell.execute_reply":"2024-07-30T08:29:14.253912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# What's the max number of slices? \n# It's 192, every image will be padded to 192?\n# Update, 192 is an outlier, the max_slices we are going to use is 70, discard images with > 70 slices\n\nmax_slices = 0\npath = \"\"\n\nfor study_id in os.listdir(train_img_path):\n    for series_id in os.listdir(f\"{train_img_path}/{study_id}\"):\n        new_slice = len(os.listdir(f\"{train_img_path}/{study_id}/{series_id}\"))\n        if new_slice > max_slices:\n            max_slices = new_slice\n            path = f\"{train_img_path}/{study_id}/{series_id}\"\n        \nmax_slices","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:29:14.262888Z","iopub.execute_input":"2024-07-30T08:29:14.263628Z","iopub.status.idle":"2024-07-30T08:29:59.034551Z","shell.execute_reply.started":"2024-07-30T08:29:14.263564Z","shell.execute_reply":"2024-07-30T08:29:59.033275Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_lens = [len(os.listdir(f\"{train_img_path}/{study_id}/{series_id}\")) for study_id in os.listdir(train_img_path) for series_id in os.listdir(f\"{train_img_path}/{study_id}\")]","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:36:25.302444Z","iopub.execute_input":"2024-07-30T08:36:25.302870Z","iopub.status.idle":"2024-07-30T08:36:30.057387Z","shell.execute_reply.started":"2024-07-30T08:36:25.302838Z","shell.execute_reply":"2024-07-30T08:36:30.056386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_lens = np.array(all_lens)\nall_lens.shape","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:36:44.263579Z","iopub.execute_input":"2024-07-30T08:36:44.264077Z","iopub.status.idle":"2024-07-30T08:36:44.273155Z","shell.execute_reply.started":"2024-07-30T08:36:44.264035Z","shell.execute_reply":"2024-07-30T08:36:44.271942Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pd.Series(all_lens).describe()","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:37:06.271772Z","iopub.execute_input":"2024-07-30T08:37:06.272249Z","iopub.status.idle":"2024-07-30T08:37:06.288083Z","shell.execute_reply.started":"2024-07-30T08:37:06.272208Z","shell.execute_reply":"2024-07-30T08:37:06.286745Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"B = plt.boxplot(all_lens)","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:47:08.744533Z","iopub.execute_input":"2024-07-30T08:47:08.744991Z","iopub.status.idle":"2024-07-30T08:47:08.996164Z","shell.execute_reply.started":"2024-07-30T08:47:08.744954Z","shell.execute_reply":"2024-07-30T08:47:08.994824Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"[item.get_ydata() for item in B['whiskers']]\n","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:47:27.665715Z","iopub.execute_input":"2024-07-30T08:47:27.666131Z","iopub.status.idle":"2024-07-30T08:47:27.675645Z","shell.execute_reply.started":"2024-07-30T08:47:27.666099Z","shell.execute_reply":"2024-07-30T08:47:27.673764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.sum([all_lens > 100])","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:49:14.904990Z","iopub.execute_input":"2024-07-30T08:49:14.905434Z","iopub.status.idle":"2024-07-30T08:49:14.913815Z","shell.execute_reply.started":"2024-07-30T08:49:14.905389Z","shell.execute_reply":"2024-07-30T08:49:14.912505Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.sum([all_lens > 25])","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:42:53.775541Z","iopub.execute_input":"2024-07-30T08:42:53.776704Z","iopub.status.idle":"2024-07-30T08:42:53.784597Z","shell.execute_reply.started":"2024-07-30T08:42:53.776661Z","shell.execute_reply":"2024-07-30T08:42:53.783264Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.sum([all_lens < 5])","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:48:36.125543Z","iopub.execute_input":"2024-07-30T08:48:36.125997Z","iopub.status.idle":"2024-07-30T08:48:36.134602Z","shell.execute_reply.started":"2024-07-30T08:48:36.125960Z","shell.execute_reply":"2024-07-30T08:48:36.133378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Use the images less than 100, find the max_slices, it's 70\n(all_lens[all_lens < 100]).max()","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:52:06.694944Z","iopub.execute_input":"2024-07-30T08:52:06.696105Z","iopub.status.idle":"2024-07-30T08:52:06.703914Z","shell.execute_reply.started":"2024-07-30T08:52:06.696055Z","shell.execute_reply":"2024-07-30T08:52:06.702813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import Dataset\n\n\nclass CustomImageDataset(Dataset):\n    def __init__(self, df, img_dir, transform=None, target_transform=None):\n        self.img_labels = df.iloc[:, 3:]  # The first 3 columns are [study_id, series_id, series_description]\n        self.img_dir = img_dir\n        self.img_paths = df.iloc[:, :3]\n        self.transform = transform\n        self.target_transform = target_transform\n\n    def __len__(self):\n        return len(self.img_labels)\n\n    def __getitem__(self, idx):\n        slices = []\n        # Each path at .../<study_id>/<series_id>/ contains multiple .dcm files\n        # We concat the .dcm slices to form a 3D image\n        directory = os.path.join(self.img_dir, str(self.img_paths.iloc[idx, 0]), str(self.img_paths.iloc[idx, 1]))\n        for filename in sorted(os.listdir(directory), key=lambda x: int(x.split('.')[0])):  \n            if filename.endswith(\".dcm\"):\n                filepath = os.path.join(directory, filename)\n                dicom_image = np.array(pydicom.dcmread(filepath).pixel_array, dtype=np.float32)\n                \n                if self.transform:\n                    dicom_image = self.transform(dicom_image)\n                    dicom_image = torch.squeeze(dicom_image)\n                \n                slices.append(dicom_image)\n        \n        slices = torch.tensor(np.array(slices, dtype=np.float32))\n\n        label = torch.tensor(self.img_labels.iloc[idx].astype(int).values)\n#         if self.transform:\n#             image = self.transform(image)\n#         if self.target_transform:\n#             label = self.target_transform(label)\n        return slices, label","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:02:44.881860Z","iopub.execute_input":"2024-07-30T08:02:44.882226Z","iopub.status.idle":"2024-07-30T08:02:44.893308Z","shell.execute_reply.started":"2024-07-30T08:02:44.882188Z","shell.execute_reply":"2024-07-30T08:02:44.892018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torchvision.transforms as transforms\n\n\n# Example transform\ntransform = transforms.Compose([\n    transforms.ToTensor(),\n    transforms.Resize((256, 256), antialias=True)  # Resize height and width to 256\n])","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:04:40.608965Z","iopub.execute_input":"2024-07-30T08:04:40.609385Z","iopub.status.idle":"2024-07-30T08:04:42.132426Z","shell.execute_reply.started":"2024-07-30T08:04:40.609354Z","shell.execute_reply":"2024-07-30T08:04:42.131388Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset_sagittal = CustomImageDataset(df_sagittal, train_img_path, transform)\nimg, labels = next(iter(dataset_sagittal))","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:04:43.539902Z","iopub.execute_input":"2024-07-30T08:04:43.541074Z","iopub.status.idle":"2024-07-30T08:04:44.009265Z","shell.execute_reply.started":"2024-07-30T08:04:43.541039Z","shell.execute_reply":"2024-07-30T08:04:44.008059Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img.shape","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:04:44.011399Z","iopub.execute_input":"2024-07-30T08:04:44.011981Z","iopub.status.idle":"2024-07-30T08:04:44.018184Z","shell.execute_reply.started":"2024-07-30T08:04:44.011949Z","shell.execute_reply":"2024-07-30T08:04:44.017034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Uncomment to see 2D plots of the images\n# fig, axes = plt.subplots(3, 5, figsize=(15, 15))\n# axes = axes.flatten()  # Flatten the 2D array of axes to easily iterate\n\n# for i, ax in enumerate(axes):\n#     if i < img.shape[0]:\n#         ax.imshow(img[i], cmap='gray')\n#         ax.axis('off')\n#     else:\n#         fig.delaxes(ax)  # Remove any extra subplots\n\n# plt.tight_layout()\n# plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:04:45.478751Z","iopub.execute_input":"2024-07-30T08:04:45.479752Z","iopub.status.idle":"2024-07-30T08:04:45.484062Z","shell.execute_reply.started":"2024-07-30T08:04:45.479712Z","shell.execute_reply":"2024-07-30T08:04:45.483044Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:04:45.830123Z","iopub.execute_input":"2024-07-30T08:04:45.830925Z","iopub.status.idle":"2024-07-30T08:04:45.838076Z","shell.execute_reply.started":"2024-07-30T08:04:45.830888Z","shell.execute_reply":"2024-07-30T08:04:45.837038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset_axial = CustomImageDataset(df_axial, train_img_path, transform)\nimg, labels = next(iter(dataset_axial))","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:04:46.168785Z","iopub.execute_input":"2024-07-30T08:04:46.169165Z","iopub.status.idle":"2024-07-30T08:04:46.722579Z","shell.execute_reply.started":"2024-07-30T08:04:46.169135Z","shell.execute_reply":"2024-07-30T08:04:46.721590Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Uncomment to see 2D plots of the images\n# fig, axes = plt.subplots(9, 5, figsize=(15, 15))\n# axes = axes.flatten()  # Flatten the 2D array of axes to easily iterate\n\n# for i, ax in enumerate(axes):\n#     if i < img.shape[0]:\n#         ax.imshow(img[i], cmap='gray')\n#         ax.axis('off')\n#     else:\n#         fig.delaxes(ax)  # Remove any extra subplots\n\n# plt.tight_layout()\n# plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:04:46.962010Z","iopub.execute_input":"2024-07-30T08:04:46.962780Z","iopub.status.idle":"2024-07-30T08:04:46.967150Z","shell.execute_reply.started":"2024-07-30T08:04:46.962743Z","shell.execute_reply":"2024-07-30T08:04:46.966011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:04:47.858722Z","iopub.execute_input":"2024-07-30T08:04:47.859101Z","iopub.status.idle":"2024-07-30T08:04:47.867095Z","shell.execute_reply.started":"2024-07-30T08:04:47.859069Z","shell.execute_reply":"2024-07-30T08:04:47.866037Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Split into train and test\nfrom torch.utils.data import random_split\n\ngenerator = torch.Generator().manual_seed(42)\n\ntrain_dataset_sagittal, test_dataset_sagittal = random_split(dataset_sagittal, [0.8, 0.2], generator=generator)\ntrain_dataset_axial, test_dataset_axial = random_split(dataset_axial, [0.8, 0.2], generator=generator)","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:05:06.971083Z","iopub.execute_input":"2024-07-30T08:05:06.971484Z","iopub.status.idle":"2024-07-30T08:05:06.983116Z","shell.execute_reply.started":"2024-07-30T08:05:06.971433Z","shell.execute_reply":"2024-07-30T08:05:06.982022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Put into DataLoader for training\nfrom torch.utils.data import DataLoader\n\ntrain_dataloader_sagittal = DataLoader(train_dataset_sagittal, batch_size=1, shuffle=False)  # Change shuffle to True later\ntrain_dataloader_axial = DataLoader(train_dataset_axial, batch_size=64, shuffle=True)\ntest_dataloader_sagittal = DataLoader(test_dataset_sagittal, batch_size=64, shuffle=False)\ntest_dataloader_axial = DataLoader(test_dataset_axial, batch_size=64, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:05:07.128747Z","iopub.execute_input":"2024-07-30T08:05:07.129126Z","iopub.status.idle":"2024-07-30T08:05:07.134842Z","shell.execute_reply.started":"2024-07-30T08:05:07.129095Z","shell.execute_reply":"2024-07-30T08:05:07.133797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# next(iter(train_dataloader_sagittal))","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:05:07.328573Z","iopub.execute_input":"2024-07-30T08:05:07.328945Z","iopub.status.idle":"2024-07-30T08:05:07.333624Z","shell.execute_reply.started":"2024-07-30T08:05:07.328916Z","shell.execute_reply":"2024-07-30T08:05:07.332229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# TODO pad the images to same number of slices\nmy_iter = iter(train_dataloader_sagittal)\nimg1, label1 = next(my_iter)\nimg2, label2 = next(my_iter)","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:05:08.398856Z","iopub.execute_input":"2024-07-30T08:05:08.399719Z","iopub.status.idle":"2024-07-30T08:05:09.261408Z","shell.execute_reply.started":"2024-07-30T08:05:08.399679Z","shell.execute_reply":"2024-07-30T08:05:09.260317Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img1.shape","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:05:09.269021Z","iopub.execute_input":"2024-07-30T08:05:09.269771Z","iopub.status.idle":"2024-07-30T08:05:09.275654Z","shell.execute_reply.started":"2024-07-30T08:05:09.269737Z","shell.execute_reply":"2024-07-30T08:05:09.274659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img2.shape","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:05:09.968652Z","iopub.execute_input":"2024-07-30T08:05:09.969362Z","iopub.status.idle":"2024-07-30T08:05:09.975831Z","shell.execute_reply.started":"2024-07-30T08:05:09.969323Z","shell.execute_reply":"2024-07-30T08:05:09.974648Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Uncomment to see 2D plots of the images\nfig, axes = plt.subplots(4, 5, figsize=(15, 15))\naxes = axes.flatten()  # Flatten the 2D array of axes to easily iterate\n\nfor i, ax in enumerate(axes):\n    if i < img1[0].shape[0]:\n        ax.imshow(img1[0][i], cmap='gray')\n        ax.axis('off')\n    else:\n        fig.delaxes(ax)  # Remove any extra subplots\n\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:05:12.759899Z","iopub.execute_input":"2024-07-30T08:05:12.760703Z","iopub.status.idle":"2024-07-30T08:05:14.278356Z","shell.execute_reply.started":"2024-07-30T08:05:12.760668Z","shell.execute_reply":"2024-07-30T08:05:14.277120Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Uncomment to see 2D plots of the images\nfig, axes = plt.subplots(4, 5, figsize=(15, 15))\naxes = axes.flatten()  # Flatten the 2D array of axes to easily iterate\n\nfor i, ax in enumerate(axes):\n    if i < img2[0].shape[0]:\n        ax.imshow(img2[0][i], cmap='gray')\n        ax.axis('off')\n    else:\n        fig.delaxes(ax)  # Remove any extra subplots\n\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-07-30T08:05:14.280027Z","iopub.execute_input":"2024-07-30T08:05:14.280351Z","iopub.status.idle":"2024-07-30T08:05:15.693946Z","shell.execute_reply.started":"2024-07-30T08:05:14.280322Z","shell.execute_reply":"2024-07-30T08:05:15.692719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train\nWe train 2 models.\n1. Sagittal model trains on dataset_sagittal\n2. Axial model trains on dataset_axial","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}