{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":183439689,"sourceType":"kernelVersion"}],"dockerImageVersionId":30732,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pydicom\nimport glob, os\nimport matplotlib.pyplot as plt\nimport pandas as pd\nimport numpy as np\nimport cv2\nfrom tqdm import tqdm\nimport re\nimport timm\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\nfrom PIL import Image\n","metadata":{"execution":{"iopub.status.busy":"2024-07-07T18:56:54.871956Z","iopub.execute_input":"2024-07-07T18:56:54.872863Z","iopub.status.idle":"2024-07-07T18:56:54.878209Z","shell.execute_reply.started":"2024-07-07T18:56:54.872830Z","shell.execute_reply":"2024-07-07T18:56:54.877396Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!unzip -q /kaggle/input/rsna2024-lsdc-making-dataset/_output_.zip ","metadata":{"execution":{"iopub.status.busy":"2024-07-07T18:56:55.914240Z","iopub.execute_input":"2024-07-07T18:56:55.914774Z","iopub.status.idle":"2024-07-07T19:01:45.579938Z","shell.execute_reply.started":"2024-07-07T18:56:55.914743Z","shell.execute_reply":"2024-07-07T19:01:45.578723Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rd = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-07-07T19:15:20.910304Z","iopub.execute_input":"2024-07-07T19:15:20.911106Z","iopub.status.idle":"2024-07-07T19:15:20.915577Z","shell.execute_reply.started":"2024-07-07T19:15:20.911067Z","shell.execute_reply":"2024-07-07T19:15:20.914646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def atoi(text):\n    return int(text) if text.isdigit() else text\n\ndef natural_keys(text):\n    return [ atoi(c) for c in re.split(r'(\\d+)', text) ]","metadata":{"execution":{"iopub.status.busy":"2024-07-07T19:15:21.106213Z","iopub.execute_input":"2024-07-07T19:15:21.106613Z","iopub.status.idle":"2024-07-07T19:15:21.111694Z","shell.execute_reply.started":"2024-07-07T19:15:21.106583Z","shell.execute_reply":"2024-07-07T19:15:21.110828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dfc = pd.read_csv(f'{rd}/train_label_coordinates.csv')","metadata":{"execution":{"iopub.status.busy":"2024-07-07T19:15:21.241612Z","iopub.execute_input":"2024-07-07T19:15:21.241886Z","iopub.status.idle":"2024-07-07T19:15:21.397782Z","shell.execute_reply.started":"2024-07-07T19:15:21.241862Z","shell.execute_reply":"2024-07-07T19:15:21.397048Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(f'{rd}/train_series_descriptions.csv')\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-07-07T19:15:21.399176Z","iopub.execute_input":"2024-07-07T19:15:21.399467Z","iopub.status.idle":"2024-07-07T19:15:21.458064Z","shell.execute_reply.started":"2024-07-07T19:15:21.399442Z","shell.execute_reply":"2024-07-07T19:15:21.457290Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Visualization","metadata":{}},{"cell_type":"markdown","source":"Data preprocessing completely copied from: https://www.kaggle.com/code/itsuki9180/rsna2024-lsdc-making-dataset","metadata":{}},{"cell_type":"code","source":"df['series_description'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2024-06-29T20:35:18.713222Z","iopub.status.idle":"2024-06-29T20:35:18.713664Z","shell.execute_reply.started":"2024-06-29T20:35:18.713471Z","shell.execute_reply":"2024-06-29T20:35:18.713488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df[df['study_id']==4096820034]","metadata":{"execution":{"iopub.status.busy":"2024-06-29T20:35:18.714712Z","iopub.status.idle":"2024-06-29T20:35:18.715123Z","shell.execute_reply.started":"2024-06-29T20:35:18.714923Z","shell.execute_reply":"2024-06-29T20:35:18.714939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dfc[dfc['study_id']==4096820034]","metadata":{"execution":{"iopub.status.busy":"2024-06-29T20:35:18.719191Z","iopub.status.idle":"2024-06-29T20:35:18.719689Z","shell.execute_reply.started":"2024-06-29T20:35:18.719478Z","shell.execute_reply":"2024-06-29T20:35:18.719500Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def imread_and_imwrite(src_path, dst_path):\n    dicom_data = pydicom.dcmread(src_path)\n    image = dicom_data.pixel_array\n    image = (image - image.min()) / (image.max() - image.min() +1e-6) * 255\n    img = cv2.resize(image, (512, 512),interpolation=cv2.INTER_CUBIC)\n    assert img.shape==(512,512)\n    cv2.imwrite(dst_path, img)","metadata":{"execution":{"iopub.status.busy":"2024-06-25T19:28:42.434635Z","iopub.execute_input":"2024-06-25T19:28:42.435523Z","iopub.status.idle":"2024-06-25T19:28:42.441343Z","shell.execute_reply.started":"2024-06-25T19:28:42.435486Z","shell.execute_reply":"2024-06-25T19:28:42.440471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"st_ids = df['study_id'].unique()\nst_ids[:3], len(st_ids)","metadata":{"execution":{"iopub.status.busy":"2024-06-25T19:28:42.576504Z","iopub.execute_input":"2024-06-25T19:28:42.576786Z","iopub.status.idle":"2024-06-25T19:28:42.586316Z","shell.execute_reply.started":"2024-06-25T19:28:42.576762Z","shell.execute_reply":"2024-06-25T19:28:42.585445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"desc = list(df['series_description'].unique())\ndesc","metadata":{"execution":{"iopub.status.busy":"2024-06-25T19:28:42.939333Z","iopub.execute_input":"2024-06-25T19:28:42.940190Z","iopub.status.idle":"2024-06-25T19:28:42.948357Z","shell.execute_reply.started":"2024-06-25T19:28:42.940159Z","shell.execute_reply":"2024-06-25T19:28:42.947486Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for idx, si in enumerate(tqdm(st_ids, total=len(st_ids))):\n#     pdf = df[df['study_id']==si]\n#     for ds in desc:\n#         ds_ = ds.replace('/', '_')\n#         pdf_ = pdf[pdf['series_description']==ds]\n#         os.makedirs(f'cvt_png/{si}/{ds_}', exist_ok=True)\n#         allimgs = []\n#         for i, row in pdf_.iterrows():\n#             pimgs = glob.glob(f'{rd}/train_images/{row[\"study_id\"]}/{row[\"series_id\"]}/*.dcm')\n#             pimgs = sorted(pimgs, key=natural_keys)\n#             allimgs.extend(pimgs)\n            \n#         if len(allimgs)==0:\n#             print(si, ds, 'has no images')\n#             continue\n\n#         if ds == 'Axial T2':\n#             for j, impath in enumerate(allimgs):\n#                 dst = f'cvt_png/{si}/{ds}/{j:03d}.png'\n#                 imread_and_imwrite(impath, dst)\n                \n#         elif ds == 'Sagittal T2/STIR':\n            \n#             step = len(allimgs) / 10.0\n#             st = len(allimgs)/2.0 - 4.0*step\n#             end = len(allimgs)+0.0001\n#             for j, i in enumerate(np.arange(st, end, step)):\n#                 dst = f'cvt_png/{si}/{ds_}/{j:03d}.png'\n#                 ind2 = max(0, int((i-0.5001).round()))\n#                 imread_and_imwrite(allimgs[ind2], dst)\n                \n#             assert len(glob.glob(f'cvt_png/{si}/{ds_}/*.png'))==10\n                \n#         elif ds == 'Sagittal T1':\n#             step = len(allimgs) / 10.0\n#             st = len(allimgs)/2.0 - 4.0*step\n#             end = len(allimgs)+0.0001\n#             for j, i in enumerate(np.arange(st, end, step)):\n#                 dst = f'cvt_png/{si}/{ds}/{j:03d}.png'\n#                 ind2 = max(0, int((i-0.5001).round()))\n#                 imread_and_imwrite(allimgs[ind2], dst)\n                \n#             assert len(glob.glob(f'cvt_png/{si}/{ds}/*.png'))==10","metadata":{"execution":{"iopub.status.busy":"2024-06-25T03:26:19.647143Z","iopub.execute_input":"2024-06-25T03:26:19.647729Z","iopub.status.idle":"2024-06-25T04:21:52.441792Z","shell.execute_reply.started":"2024-06-25T03:26:19.647701Z","shell.execute_reply":"2024-06-25T04:21:52.440870Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"","metadata":{"execution":{"iopub.status.busy":"2024-06-29T21:17:44.563061Z","iopub.execute_input":"2024-06-29T21:17:44.563685Z","iopub.status.idle":"2024-06-29T21:17:44.568642Z","shell.execute_reply.started":"2024-06-29T21:17:44.563657Z","shell.execute_reply":"2024-06-29T21:17:44.567713Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_size = [512, 512]\nin_channels = 30\nn_labels = 25\nn_classes = 3 * n_labels\nnum_workers = os.cpu_count()\n\nAUG_PROB = 0.75\n\n#TGT_BATCH_SIZE = 16\n\nlr = 2e-4\n\nsave_dir = \"rsna_results\"\nos.makedirs(save_dir, exist_ok=True)","metadata":{"execution":{"iopub.status.busy":"2024-07-07T19:21:25.127131Z","iopub.execute_input":"2024-07-07T19:21:25.127761Z","iopub.status.idle":"2024-07-07T19:21:25.133233Z","shell.execute_reply.started":"2024-07-07T19:21:25.127728Z","shell.execute_reply":"2024-07-07T19:21:25.132395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(num_workers)","metadata":{"execution":{"iopub.status.busy":"2024-07-07T19:21:25.311877Z","iopub.execute_input":"2024-07-07T19:21:25.312144Z","iopub.status.idle":"2024-07-07T19:21:25.316779Z","shell.execute_reply.started":"2024-07-07T19:21:25.312120Z","shell.execute_reply":"2024-07-07T19:21:25.315894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Net(nn.Module):\n    def __init__(self, in_channels, n_classes):\n        super(Net, self).__init__()\n        # need to align in channels with image - 1 channel if grayscale, 3 for rgb... acutally in this case its 30 because the 3d scans are 512x512x30 \n        self.conv1 = nn.Conv2d(in_channels, 32, kernel_size=7, padding='same')  # set the size of the convolution to 5x5, and have 12 of them\n        self.conv2 = nn.Conv2d(32, 64, kernel_size=5, padding='same')\n        self.conv3 = nn.Conv2d(64, 128, kernel_size=3, padding='same')\n        self.conv4 = nn.Conv2d(128, 256, kernel_size=3, padding='same')\n        self.fc1 = nn.Linear(256, 100) # changed from 50 to 100 since 75 classes... dont know if this actually matteres\n        self.fc2 = nn.Linear(100, n_classes)\n        \n        self.global_avg_pool = nn.AdaptiveAvgPool2d((1, 1))\n        \n        self.bn1 = nn.BatchNorm2d(32)\n        self.bn2 = nn.BatchNorm2d(64)\n        self.bn3 = nn.BatchNorm2d(128)\n        self.bn4 = nn.BatchNorm2d(256)\n    \n    def forward(self, x):\n        x = F.relu(self.bn1(F.max_pool2d(self.conv1(x), 2)))\n        x = F.relu(self.bn2(F.max_pool2d(self.conv2(x), 2)))\n        x = F.relu(self.bn3(F.max_pool2d(self.conv3(x), 2)))\n        x = F.relu(self.bn4(F.max_pool2d(self.conv4(x), 2)))\n        x = self.global_avg_pool(x)\n        x = x.view(x.size(0), -1)\n        x = F.relu(self.fc1(x))\n        x = self.fc2(x)\n        return x\n\n    \nnetwork = Net(in_channels, n_classes)","metadata":{"execution":{"iopub.status.busy":"2024-07-07T19:21:25.460810Z","iopub.execute_input":"2024-07-07T19:21:25.461110Z","iopub.status.idle":"2024-07-07T19:21:25.484236Z","shell.execute_reply.started":"2024-07-07T19:21:25.461086Z","shell.execute_reply":"2024-07-07T19:21:25.483322Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import zipfile\n\n# zip_file_path = \"/kaggle/input/rsna2024-lsdc-making-dataset/_output_.zip\"\n# extraction_dir = \"/kaggle/working/rsna2024_dataset\"\n\n# with zipfile.ZipFile(zip_file_path, 'r') as zip_ref:\n#     zip_ref.extractall(extraction_dir)\n\n# # Assuming there's a CSV file inside the extracted contents, load it\n# csv_file_path = os.path.join(extraction_dir, 'your_csv_file.csv')  # Replace with the actual CSV file name\n# df = pd.read_csv(csv_file_path)\n\n","metadata":{"execution":{"iopub.status.busy":"2024-07-07T19:21:25.569700Z","iopub.execute_input":"2024-07-07T19:21:25.569988Z","iopub.status.idle":"2024-07-07T19:21:25.573980Z","shell.execute_reply.started":"2024-07-07T19:21:25.569964Z","shell.execute_reply":"2024-07-07T19:21:25.573188Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(f'{rd}/train.csv')\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-07-07T19:21:25.689849Z","iopub.execute_input":"2024-07-07T19:21:25.690122Z","iopub.status.idle":"2024-07-07T19:21:25.727278Z","shell.execute_reply.started":"2024-07-07T19:21:25.690099Z","shell.execute_reply":"2024-07-07T19:21:25.726108Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = df.fillna(-100)\nlabel2id = {'Normal/Mild': 0, 'Moderate':1, 'Severe':2}\ndf = df.replace(label2id)\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-07-07T19:21:25.929187Z","iopub.execute_input":"2024-07-07T19:21:25.929551Z","iopub.status.idle":"2024-07-07T19:21:25.999600Z","shell.execute_reply.started":"2024-07-07T19:21:25.929524Z","shell.execute_reply":"2024-07-07T19:21:25.998657Z"},"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-07-07T19:21:26.079955Z","iopub.execute_input":"2024-07-07T19:21:26.080459Z","iopub.status.idle":"2024-07-07T19:21:26.084926Z","shell.execute_reply.started":"2024-07-07T19:21:26.080434Z","shell.execute_reply":"2024-07-07T19:21:26.084046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nclass RSNA24Dataset(Dataset):\n    def __init__(self, df, indices, phase='train', transform=None):\n        self.df = df.iloc[indices].reset_index(drop=True)\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_channels), 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.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        # convert from torch.cuda.ByteTensor to torch.cuda.FloatTensor\n        x = x.astype(np.float32)\n        x = x.transpose(2, 0, 1)\n                \n        return x, label","metadata":{"execution":{"iopub.status.busy":"2024-07-07T19:21:26.222888Z","iopub.execute_input":"2024-07-07T19:21:26.223357Z","iopub.status.idle":"2024-07-07T19:21:26.236666Z","shell.execute_reply.started":"2024-07-07T19:21:26.223331Z","shell.execute_reply":"2024-07-07T19:21:26.235800Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import albumentations as A\ntransforms_train = A.Compose([\n    A.RandomBrightnessContrast(brightness_limit=(-0.2, 0.2), contrast_limit=(-0.2, 0.2), p=AUG_PROB),\n    A.OneOf([\n        A.MotionBlur(blur_limit=5),\n        A.MedianBlur(blur_limit=5),\n        A.GaussianBlur(blur_limit=5),\n        A.GaussNoise(var_limit=(5.0, 30.0)),\n    ], p=AUG_PROB),\n\n    A.OneOf([\n        A.OpticalDistortion(distort_limit=1.0),\n        A.GridDistortion(num_steps=5, distort_limit=1.),\n        A.ElasticTransform(alpha=3),\n    ], p=AUG_PROB),\n\n    A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=15, border_mode=0, p=AUG_PROB),\n    A.Resize(img_size[0], img_size[1]),\n    A.CoarseDropout(max_holes=16, max_height=64, max_width=64, min_holes=1, min_height=8, min_width=8, p=AUG_PROB),    \n    A.Normalize(mean=0.5, std=0.5)\n])\n\ntransforms_val = A.Compose([\n    A.Resize(img_size[0], img_size[1]),\n    A.Normalize(mean=0.5, std=0.5)\n])","metadata":{"execution":{"iopub.status.busy":"2024-07-07T19:21:26.640220Z","iopub.execute_input":"2024-07-07T19:21:26.640605Z","iopub.status.idle":"2024-07-07T19:21:26.649346Z","shell.execute_reply.started":"2024-07-07T19:21:26.640577Z","shell.execute_reply":"2024-07-07T19:21:26.648548Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import DataLoader, random_split\n\nfrom torch.utils.data import random_split\n\n# Define the proportions for each set\ntrain_ratio = 0.85\nval_ratio = 0.15\n\n# Ensure the ratios sum to 1\nassert train_ratio + val_ratio == 1, \"The ratios must sum to 1!\"\n\ntotal_size = len(df)\ntrain_size = int(train_ratio * total_size + 1) # add 1 bc doesnt split evenly\nval_size = int(val_ratio * total_size)\n\nprint(f'total: {total_size} ---- train: {train_size} ----- val: {val_size} ---- train+val: {train_size+val_size}')\n\n# Split the dataset\ntrain_subset, val_subset = random_split(df, [train_size, val_size])\n\ntrain_indices = train_subset.indices\nval_indices = val_subset.indices\n\ntrain_dataset = RSNA24Dataset(df, train_indices, phase='train', transform=transforms_train) #, transform=transforms_train)\nval_dataset = RSNA24Dataset(df, val_indices, phase='valid', transform=transforms_val)\n\n# Create the data loaders for each set\ntrain_loader = DataLoader(train_dataset, batch_size=16, shuffle=True, num_workers=num_workers)\nval_loader = DataLoader(val_dataset, batch_size=16, shuffle=False, num_workers=num_workers)","metadata":{"execution":{"iopub.status.busy":"2024-07-07T19:21:27.517024Z","iopub.execute_input":"2024-07-07T19:21:27.517366Z","iopub.status.idle":"2024-07-07T19:21:27.550790Z","shell.execute_reply.started":"2024-07-07T19:21:27.517336Z","shell.execute_reply":"2024-07-07T19:21:27.549964Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_model(network, train_loader, val_loader, optimizer, num_epochs=10):\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    # Lists to store losses for plotting\n    epoch_train_losses = []\n    epoch_val_losses = []\n    \n    weights = torch.tensor([1.0, 2.0, 4.0])\n    criterion = nn.CrossEntropyLoss(weight=weights.to(device))\n    \n    if torch.cuda.device_count() > 1:\n        print(f'Using {torch.cuda.device_count()} GPUs for training.')\n        network = nn.DataParallel(network, device_ids = [0,1]) # THIS USES BOTH GPUS WHEN TRAINING ON KAGGLE's \"GPU T4 x2\"\n        \n    network = network.to(device)\n    \n    \n    for epoch in range(num_epochs):\n        network.train()\n        train_loss = 0\n\n\n    \n        for data, target in train_loader:\n            data, target = data.to(device), target.to(device)\n            optimizer.zero_grad()\n            output = network(data)\n            \n            # split loss calculation by category - loss calculated separately for each mild/moderate/severe prediction in each of the 25 sections of the spine. maybe because different weights assigned for mild/moderate/severe? \n            loss = 0\n            for col in range(n_labels):\n                pred = output[:, col * 3: col * 3 + 3]\n                gt = target[:, col]\n                loss += criterion(pred, gt) # ACCUMULATE LOSS! \n            loss.backward()\n            optimizer.step()\n            train_loss += loss.item()\n\n        avg_train_loss = train_loss / len(train_loader)\n        epoch_train_losses.append(avg_train_loss)\n        print(f'Epoch {epoch+1}/{num_epochs}, Train Loss: {avg_train_loss:.4f}')\n\n        # Validation\n        network.eval()\n        val_loss = 0\n        with torch.no_grad():\n            for data, target in val_loader:\n                data, target = data.to(device), target.to(device)\n                output = network(data)\n                loss = 0\n                for col in range(n_labels):\n                    pred = output[:, col * 3: col * 3 + 3]\n                    gt = target[:, col]\n                    loss += criterion(pred, gt)\n                val_loss += loss.item()\n\n        avg_val_loss = val_loss / len(val_loader)\n        epoch_val_losses.append(avg_val_loss)\n        print(f'Epoch {epoch+1}/{num_epochs}, Validation Loss: {avg_val_loss:.4f}')\n\n    plt.figure(figsize=(10, 5))\n    plt.plot(epoch_train_losses, label='Training Loss')\n    plt.plot(epoch_val_losses, label='Validation Loss')\n    plt.title('Training and Validation Losses Over Epochs')\n    plt.xlabel('Epochs')\n    plt.ylabel('Cross Entropy Loss')\n    plt.legend()\n    plt.grid(True)\n    plt.show()\n\n    return epoch_train_losses, epoch_val_losses","metadata":{"execution":{"iopub.status.busy":"2024-07-07T19:21:29.467711Z","iopub.execute_input":"2024-07-07T19:21:29.468037Z","iopub.status.idle":"2024-07-07T19:21:29.482538Z","shell.execute_reply.started":"2024-07-07T19:21:29.468013Z","shell.execute_reply":"2024-07-07T19:21:29.481652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# network = Net(in_channels, n_classes)\n# optimizer = torch.optim.Adam(network.parameters(), lr=0.001)\n\n# print(torch.cuda.is_available()) # -------------------------------------------------------------------------TO JULIET KERN:  THESE LINES OF CODE BELOW ARE TO SET DEVICE TO GPU, and NETWORK TO GPU\n# device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n# network = network.to(device)","metadata":{"execution":{"iopub.status.busy":"2024-07-07T19:21:31.805110Z","iopub.execute_input":"2024-07-07T19:21:31.805942Z","iopub.status.idle":"2024-07-07T19:21:31.810226Z","shell.execute_reply.started":"2024-07-07T19:21:31.805898Z","shell.execute_reply":"2024-07-07T19:21:31.809457Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model 1","metadata":{}},{"cell_type":"code","source":"\nclass RSNA24Model(nn.Module):\n    def __init__(self, model_name, in_c=30, n_classes=75, pretrained=True, features_only=False):\n        super().__init__()\n        self.model = timm.create_model(\n                                    model_name,\n                                    pretrained=pretrained, \n                                    features_only=features_only,\n                                    in_chans=in_c,\n                                    num_classes=n_classes,\n                                    global_pool='avg'\n                                    )\n    \n    def forward(self, x):\n        y = self.model(x)\n        return y","metadata":{"execution":{"iopub.status.busy":"2024-07-07T19:21:35.721391Z","iopub.execute_input":"2024-07-07T19:21:35.722327Z","iopub.status.idle":"2024-07-07T19:21:35.733669Z","shell.execute_reply.started":"2024-07-07T19:21:35.722275Z","shell.execute_reply":"2024-07-07T19:21:35.732665Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"network = RSNA24Model(model_name = 'densenet201', pretrained=False)\noptimizer = torch.optim.AdamW(network.parameters(), lr=2e-4, weight_decay=1e-2)\n\nprint(torch.cuda.device_count())  # Should output 2 if both GPUs are available\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n#network = network.to(device)","metadata":{"execution":{"iopub.status.busy":"2024-07-07T19:21:39.134546Z","iopub.execute_input":"2024-07-07T19:21:39.134979Z","iopub.status.idle":"2024-07-07T19:21:39.701421Z","shell.execute_reply.started":"2024-07-07T19:21:39.134946Z","shell.execute_reply":"2024-07-07T19:21:39.700457Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_losses, val_losses = train_model(network, train_loader, val_loader, optimizer, num_epochs=20)","metadata":{"execution":{"iopub.status.busy":"2024-06-30T01:16:26.981819Z","iopub.execute_input":"2024-06-30T01:16:26.982632Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_losses, val_losses = train_model(network, train_loader, val_loader, optimizer, num_epochs=20)","metadata":{"execution":{"iopub.status.busy":"2024-07-07T19:21:47.872889Z","iopub.execute_input":"2024-07-07T19:21:47.873729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f'train losses: {train_losses}\\nval losses: {val_losses}')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\n\ndef clear_cuda_memory():\n    torch.cuda.empty_cache()\n    torch.cuda.synchronize()\n\nclear_cuda_memory()","metadata":{"execution":{"iopub.status.busy":"2024-06-29T21:57:48.628118Z","iopub.execute_input":"2024-06-29T21:57:48.628483Z","iopub.status.idle":"2024-06-29T21:57:48.633215Z","shell.execute_reply.started":"2024-06-29T21:57:48.628455Z","shell.execute_reply":"2024-06-29T21:57:48.632319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model 2","metadata":{}},{"cell_type":"markdown","source":"same architecture as model 1 except correct train function and increased batch size (16 -> 32)","metadata":{}},{"cell_type":"code","source":"train_losses, val_losses = train_model(network, train_loader, val_loader, optimizer)","metadata":{"execution":{"iopub.status.busy":"2024-06-26T04:10:27.369794Z","iopub.execute_input":"2024-06-26T04:10:27.370610Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model 3","metadata":{}},{"cell_type":"code","source":"# trying slightly larger net to see if it makes any difference\n\nclass Net(nn.Module):\n    def __init__(self, in_channels, n_classes):\n        super(Net, self).__init__()\n        # need to align in channels with image - 1 channel if grayscale, 3 for rgb... acutally in this case its 30 because the 3d scans are 512x512x30 \n        self.conv1 = nn.Conv2d(in_channels, 32, kernel_size=7, padding='same')  # set the size of the convolution to 5x5, and have 12 of them\n        self.conv2 = nn.Conv2d(32, 64, kernel_size=5, padding='same')\n        self.conv3 = nn.Conv2d(64, 128, kernel_size=3, padding='same')\n        self.conv4 = nn.Conv2d(128, 256, kernel_size=3, padding='same')\n        self.conv5 = nn.Conv2d(256, 512, kernel_size=3, padding='same')\n        self.fc1 = nn.Linear(512, 100) # changed from 50 to 100 since 75 classes... dont know if this actually matteres\n        self.fc2 = nn.Linear(100, n_classes)\n        \n        self.global_avg_pool = nn.AdaptiveAvgPool2d((1, 1))\n        \n        self.bn1 = nn.BatchNorm2d(32)\n        self.bn2 = nn.BatchNorm2d(64)\n        self.bn3 = nn.BatchNorm2d(128)\n        self.bn4 = nn.BatchNorm2d(256)\n        self.bn5 = nn.BatchNorm2d(512)\n    \n    def forward(self, x):\n        x = F.relu(self.bn1(F.max_pool2d(self.conv1(x), 2)))\n        x = F.relu(self.bn2(F.max_pool2d(self.conv2(x), 2)))\n        x = F.relu(self.bn3(F.max_pool2d(self.conv3(x), 2)))\n        x = F.relu(self.bn4(F.max_pool2d(self.conv4(x), 2)))\n        x = F.relu(self.bn5(F.max_pool2d(self.conv5(x), 2)))\n        x = self.global_avg_pool(x)\n        x = x.view(x.size(0), -1)\n        x = F.relu(self.fc1(x))\n        x = self.fc2(x)\n        return x\n\n    \nnetwork = Net(in_channels, n_classes)","metadata":{"execution":{"iopub.status.busy":"2024-06-26T03:28:57.959122Z","iopub.execute_input":"2024-06-26T03:28:57.959848Z","iopub.status.idle":"2024-06-26T03:28:57.991580Z","shell.execute_reply.started":"2024-06-26T03:28:57.959812Z","shell.execute_reply":"2024-06-26T03:28:57.990738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_losses, val_losses = train_model(network, train_loader, val_loader, optimizer)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}