{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":99552,"databundleVersionId":13441085}],"dockerImageVersionId":31089,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np \nimport pandas as pd \nimport seaborn as sns\nimport matplotlib.pyplot as plt\nimport pydicom\n\nimport os\nimport torch\nimport torch.nn as nn\nimport numpy as np\n\nfrom pathlib import Path\nfrom skimage import io, transform\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms, utils, models\n\nfrom PIL import Image\nfrom sklearn.metrics import confusion_matrix, classification_report\nfrom sklearn.model_selection import train_test_split\n\nimport time\n\n# Ignore warnings\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:04:25.053815Z","iopub.execute_input":"2025-08-27T12:04:25.054061Z","iopub.status.idle":"2025-08-27T12:04:36.578223Z","shell.execute_reply.started":"2025-08-27T12:04:25.054035Z","shell.execute_reply":"2025-08-27T12:04:36.577598Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BASE_FOLDER = '/kaggle/input'\nDETECTION_FOLDER = BASE_FOLDER + '/rsna-intracranial-aneurysm-detection'\nSERIES_FOLDER = DETECTION_FOLDER + '/series'\nTRAINING_FILE = DETECTION_FOLDER + '/train.csv'\n\nCATEGORIES = [\n    'Left Infraclinoid Internal Carotid Artery',\n    'Right Infraclinoid Internal Carotid Artery',\n    'Left Supraclinoid Internal Carotid Artery',\n    'Right Supraclinoid Internal Carotid Artery',\n    'Left Middle Cerebral Artery',\n    'Right Middle Cerebral Artery',\n    'Anterior Communicating Artery',\n    'Left Anterior Cerebral Artery',\n    'Right Anterior Cerebral Artery',\n    'Left Posterior Communicating Artery',\n    'Right Posterior Communicating Artery',\n    'Basilar Tip',\n    'Other Posterior Circulation',\n    'Aneurysm Present'\n]\n\n# Model selection - Change this to select which model to use for inference\n# Options: 'tf_efficientnetv2_s', 'convnext_small', 'swin_small_patch4_window7_224', 'ensemble'\nSELECTED_MODEL = 'tf_efficientnetv2_s' \n\n# Model paths configuration\nMODEL_PATHS = {\n    'tf_efficientnetv2_s': '/kaggle/input/rsna-iad-trained-models/models/tf_efficientnetv2_s_fold0_best.pth',\n    'convnext_small': '/kaggle/input/rsna-iad-trained-models/models/convnext_small_fold0_best.pth',\n    'swin_small_patch4_window7_224': '/kaggle/input/rsna-iad-trained-models/models/swin_small_patch4_window7_224_fold0_best.pth'\n}\n\nBATCH_SIZE = 32\nNUM_EPOCHS = 20\nFILE_NAME = 'best_vit_rsna_model'\nTARGET_SIZE = (224,224)\nNUM_TRAINING_ROWS = 100","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:04:36.579904Z","iopub.execute_input":"2025-08-27T12:04:36.580282Z","iopub.status.idle":"2025-08-27T12:04:36.585243Z","shell.execute_reply.started":"2025-08-27T12:04:36.580263Z","shell.execute_reply":"2025-08-27T12:04:36.584548Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_train = pd.read_csv(TRAINING_FILE)\n\ndf_train.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:04:36.585996Z","iopub.execute_input":"2025-08-27T12:04:36.586268Z","iopub.status.idle":"2025-08-27T12:04:36.63474Z","shell.execute_reply.started":"2025-08-27T12:04:36.586247Z","shell.execute_reply":"2025-08-27T12:04:36.634194Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_train.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:04:36.635385Z","iopub.execute_input":"2025-08-27T12:04:36.635589Z","iopub.status.idle":"2025-08-27T12:04:36.659212Z","shell.execute_reply.started":"2025-08-27T12:04:36.635571Z","shell.execute_reply":"2025-08-27T12:04:36.658486Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_train.columns","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:04:36.66Z","iopub.execute_input":"2025-08-27T12:04:36.660193Z","iopub.status.idle":"2025-08-27T12:04:36.666255Z","shell.execute_reply.started":"2025-08-27T12:04:36.660176Z","shell.execute_reply":"2025-08-27T12:04:36.665408Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"gender_counts = df_train['PatientSex'].value_counts()\n\ncolors = sns.color_palette('Set1')[0:len(gender_counts)]\n\nplt.figure(figsize=(15,8))\nfig, ax = plt.subplots()\nax.pie(gender_counts, labels=gender_counts.index, autopct='%1.1f%%', colors=colors, textprops={'fontsize': 12})\nax.set_title('Patients by Gender', fontsize=14)\nplt.axis('equal')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:04:36.666993Z","iopub.execute_input":"2025-08-27T12:04:36.667196Z","iopub.status.idle":"2025-08-27T12:04:36.832985Z","shell.execute_reply.started":"2025-08-27T12:04:36.667181Z","shell.execute_reply":"2025-08-27T12:04:36.832219Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"unique_modalities = df_train['Modality'].unique()\n\nprint(\"Unique values in 'Modality' column:\")\nprint(unique_modalities)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:04:36.835186Z","iopub.execute_input":"2025-08-27T12:04:36.835408Z","iopub.status.idle":"2025-08-27T12:04:36.842083Z","shell.execute_reply.started":"2025-08-27T12:04:36.835391Z","shell.execute_reply":"2025-08-27T12:04:36.841382Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"modality_counts = df_train['Modality'].value_counts()\n\ncolors = sns.color_palette('Set1')[0:len(modality_counts)]\n\nplt.figure(figsize=(15,8))\nfig, ax = plt.subplots()\nax.pie(modality_counts, labels=modality_counts.index, autopct='%1.1f%%', colors=colors, textprops={'fontsize': 12})\nax.set_title('Modality Breakdown', fontsize=14)\nplt.axis('equal')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:04:36.842971Z","iopub.execute_input":"2025-08-27T12:04:36.843249Z","iopub.status.idle":"2025-08-27T12:04:36.948836Z","shell.execute_reply.started":"2025-08-27T12:04:36.843224Z","shell.execute_reply":"2025-08-27T12:04:36.948187Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"bins = [19, 36, 50, 65, 100]\nlabels = ['19-35', '36-49', '50-64', '65+']\n\ndf = df_train.copy()\n\ndf['AgeGroup'] = pd.cut(df['PatientAge'], bins=bins, labels=labels, right=False)\n\nage_group_counts = df['AgeGroup'].value_counts().sort_index()\n\ncolors = sns.color_palette('Set1')[0:len(age_group_counts)]\n\nplt.figure(figsize=(15,8))\nfig, ax = plt.subplots()\nax.pie(age_group_counts, labels=age_group_counts.index, autopct='%1.1f%%', colors=colors, textprops={'fontsize': 12})\nax.set_title('Age Group Breakdown', fontsize=14)\nplt.axis('equal')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:04:36.949522Z","iopub.execute_input":"2025-08-27T12:04:36.949708Z","iopub.status.idle":"2025-08-27T12:04:37.051492Z","shell.execute_reply.started":"2025-08-27T12:04:36.949688Z","shell.execute_reply":"2025-08-27T12:04:37.050543Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def convert_binary_to_yn(ax):\n    legend = ax.get_legend()\n\n    new_labels = ['No', 'Yes']\n\n    for t, l in zip(legend.get_texts(), new_labels):\n        t.set_text(l)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:04:37.052311Z","iopub.execute_input":"2025-08-27T12:04:37.052658Z","iopub.status.idle":"2025-08-27T12:04:37.057136Z","shell.execute_reply.started":"2025-08-27T12:04:37.052624Z","shell.execute_reply":"2025-08-27T12:04:37.056272Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(10,6))\nax = sns.histplot(\n    data=df_train,\n    x='PatientAge',\n    hue='Aneurysm Present',\n    bins=30,\n    kde=True,\n    palette={0: '#00BFC4', 1: '#C77CFF'}\n)\n\nconvert_binary_to_yn(ax)\n\nplt.title(\"Age Distribution by Aneurysm Presence\", fontsize=14)\nplt.xlabel(\"Age\")\nplt.ylabel(\"Frequency\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:04:37.057981Z","iopub.execute_input":"2025-08-27T12:04:37.058189Z","iopub.status.idle":"2025-08-27T12:04:37.489702Z","shell.execute_reply.started":"2025-08-27T12:04:37.058172Z","shell.execute_reply":"2025-08-27T12:04:37.489094Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = df_train.copy()\nlocation_cols = df.columns[4:-1]  # skip UID, Age, Sex, Modality, and skip final label\nlocation_df = df[location_cols].astype(int) \n\n# Co-occurrence matrix\nco_matrix = location_df.T.dot(location_df)\n\n# Plot heatmap\nplt.figure(figsize=(12, 10))\nax = sns.heatmap(co_matrix, cmap=\"magma\", annot=True, fmt=\".0f\", linewidths=0.5)\nax.tick_params(axis='x', colors='white')\nax.tick_params(axis='y', colors='white')\nplt.title(\"Aneurysm Co-occurrence Matrix\", fontsize=16, color='white')\nplt.xticks(rotation=45, ha='right', fontsize=9)\nplt.yticks(rotation=0, fontsize=9)\nplt.gca().set_facecolor('black')\nplt.gcf().set_facecolor('#111111')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:04:37.490534Z","iopub.execute_input":"2025-08-27T12:04:37.490815Z","iopub.status.idle":"2025-08-27T12:04:38.16536Z","shell.execute_reply.started":"2025-08-27T12:04:37.490771Z","shell.execute_reply":"2025-08-27T12:04:38.164532Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RSNADataset(Dataset):\n    def __init__(self, csv_file, series_dir=SERIES_FOLDER, incoming_df=None, transform=None):\n        \n        self.series_dir = series_dir\n        self.transform = transform\n        \n        if incoming_df is None:\n            self.df = pd.read_csv(csv_file)\n        else:\n            self.df = incoming_df\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        if torch.is_tensor(idx):\n            idx = idx.tolist()\n\n        series_path = Path(self.series_dir) / self.df.iloc[idx, 0]\n        images = list(series_path.glob('**/*.dcm')) \n        imgs = [pydicom.dcmread(str(f)).pixel_array for f in sorted(images)]\n        volume = np.stack(imgs) \n\n        labels = self.df.iloc[idx, 4:]\n        sample = {'images': volume, 'labels': labels}\n        \n        return sample\n        ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:04:38.166123Z","iopub.execute_input":"2025-08-27T12:04:38.166372Z","iopub.status.idle":"2025-08-27T12:04:38.172525Z","shell.execute_reply.started":"2025-08-27T12:04:38.166353Z","shell.execute_reply":"2025-08-27T12:04:38.171802Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_ds = RSNADataset(TRAINING_FILE)\nsample = train_ds[0]\nprint(sample['images'].shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:04:38.173261Z","iopub.execute_input":"2025-08-27T12:04:38.173509Z","iopub.status.idle":"2025-08-27T12:04:41.962536Z","shell.execute_reply.started":"2025-08-27T12:04:38.173482Z","shell.execute_reply":"2025-08-27T12:04:41.9619Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print (len(train_ds))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:04:41.96317Z","iopub.execute_input":"2025-08-27T12:04:41.963362Z","iopub.status.idle":"2025-08-27T12:04:41.967643Z","shell.execute_reply.started":"2025-08-27T12:04:41.963348Z","shell.execute_reply":"2025-08-27T12:04:41.966906Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(sample['labels'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:04:41.968292Z","iopub.execute_input":"2025-08-27T12:04:41.968522Z","iopub.status.idle":"2025-08-27T12:04:41.980778Z","shell.execute_reply.started":"2025-08-27T12:04:41.968506Z","shell.execute_reply":"2025-08-27T12:04:41.980173Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fig = plt.figure(figsize=(15, 10))\ncolumns = 5; rows = 4\nfor i in range(20):\n    fig.add_subplot(rows, columns, i + 1)\n    plt.imshow(sample['images'][i], cmap=plt.cm.bone)\n    plt.axis('off')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:04:41.981631Z","iopub.execute_input":"2025-08-27T12:04:41.981887Z","iopub.status.idle":"2025-08-27T12:04:43.114317Z","shell.execute_reply.started":"2025-08-27T12:04:41.981862Z","shell.execute_reply":"2025-08-27T12:04:43.113479Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Use the most power possible\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:04:43.115184Z","iopub.execute_input":"2025-08-27T12:04:43.115425Z","iopub.status.idle":"2025-08-27T12:04:43.19433Z","shell.execute_reply.started":"2025-08-27T12:04:43.115406Z","shell.execute_reply":"2025-08-27T12:04:43.193492Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_transform = transforms.Compose([\n    #Resize to the target size defined above, if necessary\n    transforms.Resize(TARGET_SIZE),\n\n    #Flip it horizontally... maybe\n    transforms.RandomHorizontalFlip(),\n\n    #More random transformation\n    #Read more here: https://medium.com/@MarkAiCode/random-affine-transformations-in-pytorch-c45a290e44d0\n    transforms.RandomAffine(degrees=0, translate=(0.1, 0.1)),\n\n    #Have fun with the color brightness\n    transforms.ColorJitter(brightness=(0.8, 1.2)),\n\n    #Rotate the image by a max of 10 degrees in either direction\n    transforms.RandomRotation(10),\n\n    #Tranform it to a tensor\n    transforms.ToTensor(),\n\n    #Standard normalization value here, feel free to experiment\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\ntest_transform = transforms.Compose([\n    transforms.Resize(TARGET_SIZE),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:04:43.19531Z","iopub.execute_input":"2025-08-27T12:04:43.195657Z","iopub.status.idle":"2025-08-27T12:04:43.209483Z","shell.execute_reply.started":"2025-08-27T12:04:43.195631Z","shell.execute_reply":"2025-08-27T12:04:43.208823Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ImageDataset(Dataset):\n    def __init__(self, dataframe, transform=None):\n        self.dataframe = dataframe\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, idx):\n        img_path = self.dataframe.iloc[idx, 0]\n        labels = self.dataframe.iloc[idx, 1]\n        labels_tensor = torch.tensor(labels, dtype=torch.float32) \n        img_array = pydicom.dcmread(img_path).pixel_array\n\n        img = Image.fromarray(img_array).convert('RGB')\n\n        if self.transform:\n            img = self.transform(img)\n\n        return img, labels_tensor","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:04:43.210176Z","iopub.execute_input":"2025-08-27T12:04:43.210372Z","iopub.status.idle":"2025-08-27T12:04:43.227018Z","shell.execute_reply.started":"2025-08-27T12:04:43.210357Z","shell.execute_reply":"2025-08-27T12:04:43.226137Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ImageDataLoader():\n    #Note that this class allows you to send in custom transformers, here, we're just using\n    #the transformers we created above\n    #Also, we pass in the BASE series directory here\n    def __init__(self, df_train, series_dir, train_transform=train_transform, test_transform=test_transform):\n\n        self.df_train = df_train\n        \n        #Set the images dir as an instance variable\n        self.images_dir = series_dir\n\n        #Gotta create the datasets before we create the loaders\n        #Note that we datasets here are instance variables as a convenience for debugging\n        self.create_datasets(train_transform, test_transform)\n\n        #And, finally, the loaders\n        self.create_data_loaders()\n\n    #Creates all datasets\n    def create_datasets(self, train_transform, test_transform):\n        #Create the training dataframe\n        full_dataset = self.create_dataset()\n\n        train, test_full = train_test_split(full_dataset)\n        test, val = train_test_split(test_full)\n        \n        # Reset indices\n        test = test.reset_index(drop=True)\n        val = val.reset_index(drop=True)\n\n        #Now create three (3) datasets and store them as instance variables so\n        #developers can debug\n        self.train_dataset = self.create_image_dataset(train, transform=train_transform)\n        self.val_dataset = self.create_image_dataset(val, transform=test_transform)\n        self.test_dataset = self.create_image_dataset(test, transform=test_transform)\n\n    def create_data_loaders(self):\n\n        #We shuffule this one because it's used for training, don't want to send everything in\n        #in the same order\n        self.train_loader = self.create_data_loader(self.train_dataset, shuffle=True)\n        \n        self.val_loader = self.create_data_loader(self.val_dataset)\n        self.test_loader = self.create_data_loader(self.test_dataset)\n\n    #Convenience function for creating the data loader. It's small now but left\n    #in its own function for scalability purposes.\n    def create_data_loader(self, dataset, batch_size=BATCH_SIZE, shuffle=False):\n        return DataLoader(dataset, batch_size, shuffle)\n\n    #Convenience function for creating the dataset. Uses the ImageDataset class defined above. \n    #What's the difference between the standard dataset and the ImageDataSet?\n    #The standard dataset stores the class plus the path to the image.\n    #The ImageDataSet allows retrieval of the image itself.\n    def create_image_dataset(self, df, transform=None):\n        return ImageDataset(df, transform)\n\n    #Here's where we create the standard dataset.\n    #Note: this only stores the class plus the path to the image, not the image itself.\n    def create_dataset(self):\n        my_list = []\n        \n        #for i,row in self.df_train.iterrows():\n        for i in range(0,NUM_TRAINING_ROWS):\n            labels = self.df_train.iloc[i, 4:]\n            full_path = os.path.join(self.images_dir, self.df_train.iloc[i,0])\n            \n            for file_name in os.listdir(full_path):\n                file_path = os.path.join(full_path, file_name)\n                if file_path.endswith('.dcm'):\n                    try:\n                        img_array = pydicom.dcmread(file_path).pixel_array\n                        img = Image.fromarray(img_array).convert('RGB')\n                        my_list.append([file_path, labels])\n                    except TypeError as e:\n                        print(f\"Caught TypeError for file: {file_path}\")\n                        \n        return pd.DataFrame(my_list, columns=['file_path', 'labels'])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:04:43.227924Z","iopub.execute_input":"2025-08-27T12:04:43.228153Z","iopub.status.idle":"2025-08-27T12:04:43.239857Z","shell.execute_reply.started":"2025-08-27T12:04:43.22813Z","shell.execute_reply":"2025-08-27T12:04:43.239213Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_loader = ImageDataLoader(df_train, SERIES_FOLDER)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:04:43.242593Z","iopub.execute_input":"2025-08-27T12:04:43.242984Z","iopub.status.idle":"2025-08-27T12:10:59.523474Z","shell.execute_reply.started":"2025-08-27T12:04:43.242953Z","shell.execute_reply":"2025-08-27T12:10:59.522443Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(len(image_loader.train_dataset))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:10:59.52455Z","iopub.execute_input":"2025-08-27T12:10:59.524837Z","iopub.status.idle":"2025-08-27T12:10:59.529177Z","shell.execute_reply.started":"2025-08-27T12:10:59.524816Z","shell.execute_reply":"2025-08-27T12:10:59.528605Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"first_batch_data, first_batch_labels = next(iter(image_loader.train_loader))\n\nprint(f\"First batch data shape: {first_batch_data.shape}\")\nprint(f\"First batch labels shape: {first_batch_labels.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:10:59.529921Z","iopub.execute_input":"2025-08-27T12:10:59.530688Z","iopub.status.idle":"2025-08-27T12:11:00.248837Z","shell.execute_reply.started":"2025-08-27T12:10:59.530662Z","shell.execute_reply":"2025-08-27T12:11:00.248086Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"    # Note that model is not optional here. You must pass in a model.\n    # Note that optimizer is also not optional. It usually requires the model as input\n    # so it's best to pass it in this way.\n    # Gotta add the data loader as well.\n    # Note that we default the optimizer to CrossEntropyLoss here. But you can use any loss function\n    # That suits your fancy.\n    # The num_epochs parameters defaults to 10 and tells the trainer how many times to run\n    # through the training.\n    # The patience parameter tells the model how many epochs to go through when it's not getting\n    # any better.\n    def train_model(model, optimizer, image_loader, \n                    criterion=nn.CrossEntropyLoss(), num_epochs=100, patience=6, file_name='best_model'):\n\n        #Routine code here\n        model.to(device)\n\n        #Start with best value loss of infinity\n        best_val_loss = float(\"inf\")\n\n        #We'll use to check for early stoppage\n        tolerance = 0\n\n        #This is what we'll actually return here, giving developers insight into training success\n        history = {'train_loss': [], 'train_acc': [], 'val_loss': [], 'val_acc': []}\n\n        for epoch in range(num_epochs):\n            # Boilerplate code again - training mode\n            model.train()\n\n            # Set counters\n            running_loss = 0.0\n            correct_train = 0\n            total_train = 0\n\n            for images, labels in image_loader.train_loader:\n                images, labels = images.to(device), labels.to(device)\n\n                #This resets the gradients of optimized tensors\n                optimizer.zero_grad()\n\n                #This is where we get actual output from training\n                outputs = model(images)\n\n                #This is what determins the loss based on the outputs\n                loss = criterion(outputs, labels)\n\n                #Initiates back propagation\n                #As the name implies, the logic traverses through the computational\n                #graph that it created during the forward pass.\n                #It computes the loss with respect to each intermediate tensor and weights and biases.\n                #Note: this just calculates but does not optimize. The next line optimizes.\n                loss.backward()\n\n                #Handles the process of minimizing the loss function.\n                #This is, in fact, the optimization step.\n                optimizer.step()\n\n                #This gets the scalar value of the loss\n                running_loss += loss.item()\n\n                predicted = outputs > 0.5\n\n                #This just increases the total_train value by the batch size\n                total_train += (len(CATEGORIES) * BATCH_SIZE)\n\n                # The (predicted == labels) part performs an element-wise tensor comparison\n                # It creates a new boolean tensor\n                # The sum() method calculates the number of True values in the tensor\n                # item() gets the boolean value as a number\n                # Basically it's the number of right answers\n                correct_train += (predicted == labels).sum().item()\n\n            train_loss = running_loss / len(image_loader.train_loader)\n            train_acc = 100 * correct_train / total_train\n\n            # Now that we've trained, let's evaluate\n            # First, set the model to eval() because we're no longer training\n            model.eval()\n\n            # Once again, establish the counters\n            val_loss = 0.0\n            correct_val = 0\n            total_val = 0\n\n            # No need for gradient here because we're not correcting/training\n            with torch.no_grad():\n                for images, labels in image_loader.val_loader:\n                    images, labels = images.to(device), labels.to(device)\n\n                    # Once again, the actual output\n                    outputs = model(images)\n                    \n                    # Now let's see how right we were (or weren't)\n                    loss = criterion(outputs, labels)\n                    val_loss += loss.item()\n                    \n                    predicted = outputs > 0.5\n                    \n                    total_val += (len(CATEGORIES) * BATCH_SIZE)\n                    correct_val += (predicted == labels).sum().item()\n\n            val_loss = val_loss / len(image_loader.val_loader)\n            val_acc = 100 * correct_val / total_val\n\n            # Again, this is what this function returns\n            history['train_loss'].append(train_loss)\n            history['train_acc'].append(train_acc)\n            history['val_loss'].append(val_loss)\n            history['val_acc'].append(val_acc)\n\n            # Print out some helpful info to the console\n            print(f\"Epoch [{epoch + 1}/{num_epochs}]\")\n            print(f\"Training Loss: {train_loss:.4f}, Training Accuracy: {train_acc:.2f}%\")\n            print(f\"Evaluation Loss: {val_loss:.4f}, Evaluation Accuracy: {val_acc:.2f}%\")\n            print(\"#\" * 80)\n\n            # Save the model if we got the best score so far\n            if val_loss < best_val_loss:\n                best_val_loss = val_loss\n                print('Saving model')\n                torch.save(model.state_dict(), f'{file_name}.pth')\n                tolerance = 0\n            else:\n                tolerance += 1\n                if tolerance >= patience:\n                    print(f\"Jumping out early, we can't do any better after {epoch + 1} epochs.\")\n                    break\n\n        return history","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:11:00.249813Z","iopub.execute_input":"2025-08-27T12:11:00.250098Z","iopub.status.idle":"2025-08-27T12:11:00.261042Z","shell.execute_reply.started":"2025-08-27T12:11:00.250074Z","shell.execute_reply":"2025-08-27T12:11:00.260228Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"    def test_model(model, image_loader, file_name='best_model'):\n        #According to PyTorch docs, loading state_dict is the preferred way to load the model from disk\n        model.load_state_dict(torch.load(f'{file_name}.pth'))\n        \n        #Moves all parameters and buffers to a specific device\n        #On my laptop, it's always CPU\n        model.to(device)\n\n        #Setting to eval disables dropout that's used during training to prevent overfitting\n        #It also adjusts batch normalization\n        model.eval()\n\n        correct = 0\n        total = 0\n\n        all_preds = []\n        all_labels = []\n        all_images = []\n\n        #no_grad() is for memory efficiency\n        #It disables gradient calculation so PyTorch won't store the computation graph\n        with torch.no_grad():\n            for images, labels in image_loader.test_loader:\n                images, labels = images.to(device), labels.to(device)\n\n                #Run the inputs through the model to get the outputs\n                outputs = model(images)\n                predicted = outputs > 0.5\n                \n                total += (len(CATEGORIES) * BATCH_SIZE)\n                \n                #The (predicted == labels) part performs an element-wise tensor comparison\n                #It creates a new boolean tensor\n                #The sum() method calculates the number of True values in the tensor\n                #item() gets the boolean value as a number\n                correct += (predicted == labels).sum().item()\n        \n        test_acc = 100 * correct / total\n\n        print(f\"Test Accuracy: {test_acc:.2f}%\\n\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:11:00.261858Z","iopub.execute_input":"2025-08-27T12:11:00.262175Z","iopub.status.idle":"2025-08-27T12:11:00.277346Z","shell.execute_reply.started":"2025-08-27T12:11:00.262152Z","shell.execute_reply":"2025-08-27T12:11:00.276607Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class VitModel(nn.Module):\n    def __init__(self, num_classes):\n        super(VitModel, self).__init__()\n        self.pretrained = models.vit_b_16(weights=models.ViT_B_16_Weights.DEFAULT)\n        self.pretrained.head = nn.Identity()\n        self.new_head = nn.Sequential(\n            nn.Linear(1000, num_classes),\n            nn.Sigmoid()\n        )\n\n    def forward(self, x):\n        x = self.pretrained(x)\n        x = self.new_head(x)\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:11:00.278092Z","iopub.execute_input":"2025-08-27T12:11:00.278371Z","iopub.status.idle":"2025-08-27T12:11:00.291564Z","shell.execute_reply.started":"2025-08-27T12:11:00.278355Z","shell.execute_reply":"2025-08-27T12:11:00.290936Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"vit_model = VitModel(num_classes=len(CATEGORIES))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:11:00.292255Z","iopub.execute_input":"2025-08-27T12:11:00.292527Z","iopub.status.idle":"2025-08-27T12:11:03.573196Z","shell.execute_reply.started":"2025-08-27T12:11:00.292477Z","shell.execute_reply":"2025-08-27T12:11:03.572572Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"optimizer = torch.optim.SGD(vit_model.parameters(), lr=0.001, momentum=0.9)\n\nfile_name = 'best_vit_rsna_model'\n\nstart_time = time.perf_counter()\n\nprint(\"Launching ViT training...\")\nhistory = train_model(vit_model, optimizer, image_loader, criterion=torch.nn.BCELoss(), num_epochs=NUM_EPOCHS, file_name=FILE_NAME)\n\nend_time = time.perf_counter()\nelapsed_time = end_time - start_time\n\nprint(f\"Elapsed time: {elapsed_time:.1f} seconds\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T12:11:03.574192Z","iopub.execute_input":"2025-08-27T12:11:03.574465Z","execution_failed":"2025-08-27T17:26:25.875Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_model(vit_model, image_loader, file_name)","metadata":{"trusted":true,"execution":{"iopub.status.idle":"2025-08-27T16:45:01.711527Z","shell.execute_reply.started":"2025-08-27T16:43:18.362797Z","shell.execute_reply":"2025-08-27T16:45:01.710928Z"}},"outputs":[],"execution_count":null}]}