{"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"}],"dockerImageVersionId":30732,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\n\nbase_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'\ntrain_desc_path = os.path.join(base_path, 'train_series_descriptions.csv')\ntrain_label_path = os.path.join(base_path, 'train_label_coordinates.csv')\ntrain_csv_path = os.path.join(base_path, 'train.csv')\ntrain_folder = os.path.join(base_path, 'train_images')\ntest_folder = os.path.join(base_path, 'test_images')","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:49.395481Z","iopub.execute_input":"2024-08-13T05:03:49.395741Z","iopub.status.idle":"2024-08-13T05:03:49.408042Z","shell.execute_reply.started":"2024-08-13T05:03:49.395718Z","shell.execute_reply":"2024-08-13T05:03:49.407199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\n\ntrain_desc = pd.read_csv(train_desc_path)\ntrain_label = pd.read_csv(train_label_path)\ntrain = pd.read_csv(train_csv_path)","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:49.409669Z","iopub.execute_input":"2024-08-13T05:03:49.410019Z","iopub.status.idle":"2024-08-13T05:03:50.680085Z","shell.execute_reply.started":"2024-08-13T05:03:49.409989Z","shell.execute_reply":"2024-08-13T05:03:50.679257Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_desc.head()","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:50.681699Z","iopub.execute_input":"2024-08-13T05:03:50.682048Z","iopub.status.idle":"2024-08-13T05:03:50.700406Z","shell.execute_reply.started":"2024-08-13T05:03:50.682021Z","shell.execute_reply":"2024-08-13T05:03:50.699602Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_label","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:50.701559Z","iopub.execute_input":"2024-08-13T05:03:50.702137Z","iopub.status.idle":"2024-08-13T05:03:50.719011Z","shell.execute_reply.started":"2024-08-13T05:03:50.702103Z","shell.execute_reply":"2024-08-13T05:03:50.718117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.head()","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:50.721492Z","iopub.execute_input":"2024-08-13T05:03:50.722097Z","iopub.status.idle":"2024-08-13T05:03:50.744832Z","shell.execute_reply.started":"2024-08-13T05:03:50.722067Z","shell.execute_reply":"2024-08-13T05:03:50.743921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data = pd.merge(train_label, train_desc, on=['study_id', 'series_id'], how='inner')","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:50.745921Z","iopub.execute_input":"2024-08-13T05:03:50.746664Z","iopub.status.idle":"2024-08-13T05:03:50.771667Z","shell.execute_reply.started":"2024-08-13T05:03:50.746634Z","shell.execute_reply":"2024-08-13T05:03:50.771011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data ","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:50.772544Z","iopub.execute_input":"2024-08-13T05:03:50.772798Z","iopub.status.idle":"2024-08-13T05:03:50.786939Z","shell.execute_reply.started":"2024-08-13T05:03:50.772767Z","shell.execute_reply":"2024-08-13T05:03:50.786083Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data[train_data.series_description == 'Axial T2'].head()","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:50.787927Z","iopub.execute_input":"2024-08-13T05:03:50.788216Z","iopub.status.idle":"2024-08-13T05:03:50.809727Z","shell.execute_reply.started":"2024-08-13T05:03:50.788193Z","shell.execute_reply":"2024-08-13T05:03:50.808952Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Sagittal T2 ---> Spinal Canal Stenosis\n# Sagittal T1 ---> Neural Foraminal Narrowing\n# Axial T2 ---> Subarticular Stenosis","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:50.810717Z","iopub.execute_input":"2024-08-13T05:03:50.811017Z","iopub.status.idle":"2024-08-13T05:03:50.814688Z","shell.execute_reply.started":"2024-08-13T05:03:50.810993Z","shell.execute_reply":"2024-08-13T05:03:50.813732Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data['image_id'] = train_data.condition.apply(lambda x: x.lower().replace(' ', '_')) + '_' + train_data.level.apply(lambda x: x.lower().replace('/', '_'))","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:50.815664Z","iopub.execute_input":"2024-08-13T05:03:50.815940Z","iopub.status.idle":"2024-08-13T05:03:50.891474Z","shell.execute_reply.started":"2024-08-13T05:03:50.815918Z","shell.execute_reply":"2024-08-13T05:03:50.890651Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data.head()","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:50.895524Z","iopub.execute_input":"2024-08-13T05:03:50.895999Z","iopub.status.idle":"2024-08-13T05:03:50.907549Z","shell.execute_reply.started":"2024-08-13T05:03:50.895976Z","shell.execute_reply":"2024-08-13T05:03:50.906646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_folder","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:50.908815Z","iopub.execute_input":"2024-08-13T05:03:50.909185Z","iopub.status.idle":"2024-08-13T05:03:50.918356Z","shell.execute_reply.started":"2024-08-13T05:03:50.909154Z","shell.execute_reply":"2024-08-13T05:03:50.917558Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data['image_path'] = train_folder + '/' + train_data.study_id.astype(str) + '/' + train_data.series_id.astype(str) + '/' + train_data.instance_number.astype(str) + '.dcm'","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:50.920261Z","iopub.execute_input":"2024-08-13T05:03:50.920522Z","iopub.status.idle":"2024-08-13T05:03:51.031336Z","shell.execute_reply.started":"2024-08-13T05:03:50.920501Z","shell.execute_reply":"2024-08-13T05:03:51.030459Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data.head()","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:51.032259Z","iopub.execute_input":"2024-08-13T05:03:51.032495Z","iopub.status.idle":"2024-08-13T05:03:51.045871Z","shell.execute_reply.started":"2024-08-13T05:03:51.032473Z","shell.execute_reply":"2024-08-13T05:03:51.044731Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_mapped = train.set_index(['study_id']).stack().reset_index()\ntrain_mapped.columns = ['study_id', 'image_id', 'severity_value']\n\n# Merge train_data with this mapped DataFrame\nmerged = train_data.merge(train_mapped, on=['study_id', 'image_id'], how='left')\n\n# The 'severity_value' column in the merged DataFrame contains the desired severity values\ntrain_data['severity'] = merged['severity_value']\n\ntrain_data.head()","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:51.047489Z","iopub.execute_input":"2024-08-13T05:03:51.047838Z","iopub.status.idle":"2024-08-13T05:03:51.126862Z","shell.execute_reply.started":"2024-08-13T05:03:51.047803Z","shell.execute_reply":"2024-08-13T05:03:51.125918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data['x'] = train_data.pop('x')\ntrain_data['y'] = train_data.pop('y')","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:51.128095Z","iopub.execute_input":"2024-08-13T05:03:51.128477Z","iopub.status.idle":"2024-08-13T05:03:51.135278Z","shell.execute_reply.started":"2024-08-13T05:03:51.128442Z","shell.execute_reply":"2024-08-13T05:03:51.134484Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data['image_path'] = train_data.pop('image_path')","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:51.136342Z","iopub.execute_input":"2024-08-13T05:03:51.136647Z","iopub.status.idle":"2024-08-13T05:03:51.145787Z","shell.execute_reply.started":"2024-08-13T05:03:51.136617Z","shell.execute_reply":"2024-08-13T05:03:51.144813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data.head()","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:51.147179Z","iopub.execute_input":"2024-08-13T05:03:51.147496Z","iopub.status.idle":"2024-08-13T05:03:51.165146Z","shell.execute_reply.started":"2024-08-13T05:03:51.147467Z","shell.execute_reply":"2024-08-13T05:03:51.164008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pydicom\nimport matplotlib.pyplot as plt\n\ndef get_plots(image_paths: pd.Series, count_of_images: int=5, x: pd.Series|None=None, y: pd.Series|None=None):\n    for path in image_paths[:count_of_images]:\n        image_data = pydicom.dcmread(path)\n        plt.imshow(image_data.pixel_array, cmap='gray')\n        if x is not None and y is not None:\n            plt.scatter(x[:count_of_images], y[:count_of_images], c='red')\n        plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:51.166548Z","iopub.execute_input":"2024-08-13T05:03:51.166893Z","iopub.status.idle":"2024-08-13T05:03:51.367071Z","shell.execute_reply.started":"2024-08-13T05:03:51.166863Z","shell.execute_reply":"2024-08-13T05:03:51.366254Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"get_plots(train_data.image_path, x=train_data.x, y=train_data.y)","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:51.368320Z","iopub.execute_input":"2024-08-13T05:03:51.368682Z","iopub.status.idle":"2024-08-13T05:03:53.091078Z","shell.execute_reply.started":"2024-08-13T05:03:51.368631Z","shell.execute_reply":"2024-08-13T05:03:53.089985Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data.head()","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.092334Z","iopub.execute_input":"2024-08-13T05:03:53.092691Z","iopub.status.idle":"2024-08-13T05:03:53.108157Z","shell.execute_reply.started":"2024-08-13T05:03:53.092640Z","shell.execute_reply":"2024-08-13T05:03:53.107049Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data['image_id'] = train_data.study_id.astype(str) + '_' + train_data.image_id","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.109386Z","iopub.execute_input":"2024-08-13T05:03:53.109720Z","iopub.status.idle":"2024-08-13T05:03:53.176667Z","shell.execute_reply.started":"2024-08-13T05:03:53.109688Z","shell.execute_reply":"2024-08-13T05:03:53.175573Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data.severity = train_data.severity.astype(str).apply(lambda x: x.strip().lower().replace('/', '_'))","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.177804Z","iopub.execute_input":"2024-08-13T05:03:53.178199Z","iopub.status.idle":"2024-08-13T05:03:53.213467Z","shell.execute_reply.started":"2024-08-13T05:03:53.178162Z","shell.execute_reply":"2024-08-13T05:03:53.212628Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data.head()","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.214490Z","iopub.execute_input":"2024-08-13T05:03:53.214966Z","iopub.status.idle":"2024-08-13T05:03:53.229204Z","shell.execute_reply.started":"2024-08-13T05:03:53.214943Z","shell.execute_reply":"2024-08-13T05:03:53.228139Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for col in train_data.columns:\n    train_data[col] = train_data[col].apply(lambda x: None if x == 'nan' else x)","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.230547Z","iopub.execute_input":"2024-08-13T05:03:53.230868Z","iopub.status.idle":"2024-08-13T05:03:53.524765Z","shell.execute_reply.started":"2024-08-13T05:03:53.230842Z","shell.execute_reply":"2024-08-13T05:03:53.523838Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_data), len(train_data.columns)","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.526443Z","iopub.execute_input":"2024-08-13T05:03:53.526806Z","iopub.status.idle":"2024-08-13T05:03:53.532882Z","shell.execute_reply.started":"2024-08-13T05:03:53.526770Z","shell.execute_reply":"2024-08-13T05:03:53.532025Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data.dropna(inplace=True)","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.533969Z","iopub.execute_input":"2024-08-13T05:03:53.534342Z","iopub.status.idle":"2024-08-13T05:03:53.590602Z","shell.execute_reply.started":"2024-08-13T05:03:53.534307Z","shell.execute_reply":"2024-08-13T05:03:53.589787Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_data), len(train_data.columns)","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.592160Z","iopub.execute_input":"2024-08-13T05:03:53.592485Z","iopub.status.idle":"2024-08-13T05:03:53.605461Z","shell.execute_reply.started":"2024-08-13T05:03:53.592455Z","shell.execute_reply":"2024-08-13T05:03:53.600816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"conditions = {\n    'Sagittal T1': {\n        'left': 'left_neural_foraminal_narrowing',\n        'right': 'right_neural_foraminal_narrowing'\n    },\n    'Sagittal T2/STIR': 'spinal_canal_stenosis',\n    'Axial T2': {\n        'left': 'left_subarticular_stenosis',\n        'right': 'right_subarticular_stenosis'\n    }\n}","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.618773Z","iopub.execute_input":"2024-08-13T05:03:53.619251Z","iopub.status.idle":"2024-08-13T05:03:53.627571Z","shell.execute_reply.started":"2024-08-13T05:03:53.619214Z","shell.execute_reply":"2024-08-13T05:03:53.626678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_desc = pd.read_csv(os.path.join(base_path, 'test_series_descriptions.csv'))","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.629067Z","iopub.execute_input":"2024-08-13T05:03:53.629829Z","iopub.status.idle":"2024-08-13T05:03:53.641216Z","shell.execute_reply.started":"2024-08-13T05:03:53.629798Z","shell.execute_reply":"2024-08-13T05:03:53.640152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_desc.head()","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.642890Z","iopub.execute_input":"2024-08-13T05:03:53.643650Z","iopub.status.idle":"2024-08-13T05:03:53.654933Z","shell.execute_reply.started":"2024-08-13T05:03:53.643619Z","shell.execute_reply":"2024-08-13T05:03:53.653962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_study = []\ntest_series = []\ntest_image_paths = []\ninstance_numbers = []\n\nfor study in os.listdir(test_folder):\n    study_path = os.path.join(test_folder, study)\n    \n    # Iterate over the series within each study\n    for each_series in os.listdir(study_path):\n        series_path = os.path.join(study_path, each_series)\n        \n        # Iterate over the images within each series\n        for image in os.listdir(series_path):\n            image_path = os.path.join(series_path, image)\n            \n            # Append the data to the lists\n            test_study.append(int(study))\n            test_series.append(int(each_series))\n            test_image_paths.append(image_path)\n            instance_numbers.append(int(image.rstrip('.dcm')))","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.656133Z","iopub.execute_input":"2024-08-13T05:03:53.656387Z","iopub.status.idle":"2024-08-13T05:03:53.686903Z","shell.execute_reply.started":"2024-08-13T05:03:53.656365Z","shell.execute_reply":"2024-08-13T05:03:53.685987Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data = pd.DataFrame({'study_id': test_study, 'series_id': test_series, 'instance_number': instance_numbers, 'image_path': test_image_paths})","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.688101Z","iopub.execute_input":"2024-08-13T05:03:53.688438Z","iopub.status.idle":"2024-08-13T05:03:53.694004Z","shell.execute_reply.started":"2024-08-13T05:03:53.688406Z","shell.execute_reply":"2024-08-13T05:03:53.692942Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data.head()","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.695463Z","iopub.execute_input":"2024-08-13T05:03:53.696191Z","iopub.status.idle":"2024-08-13T05:03:53.708551Z","shell.execute_reply.started":"2024-08-13T05:03:53.696161Z","shell.execute_reply":"2024-08-13T05:03:53.707516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data = test_data.merge(test_desc[['study_id', 'series_id', 'series_description']], on=['study_id', 'series_id'], how='inner')","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.709878Z","iopub.execute_input":"2024-08-13T05:03:53.710200Z","iopub.status.idle":"2024-08-13T05:03:53.720713Z","shell.execute_reply.started":"2024-08-13T05:03:53.710171Z","shell.execute_reply":"2024-08-13T05:03:53.719642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data.head()","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.722014Z","iopub.execute_input":"2024-08-13T05:03:53.722336Z","iopub.status.idle":"2024-08-13T05:03:53.735935Z","shell.execute_reply.started":"2024-08-13T05:03:53.722306Z","shell.execute_reply":"2024-08-13T05:03:53.734923Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"new_rows = []\n\n# Iterate over the DataFrame rows\nfor _, row in test_data.iterrows():\n    series_desc = row['series_description']\n    if series_desc in conditions:\n        condition_value = conditions[series_desc]\n        if isinstance(condition_value, dict):\n            for side, condition in condition_value.items():\n                new_row = row.copy()\n                new_row['row_id'] = f'{new_row.study_id}_{condition}'\n                new_rows.append(new_row)\n        else:\n            new_row = row.copy()\n            new_row['row_id'] = f'{new_row.study_id}_{condition_value}'\n            new_rows.append(new_row)","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.737388Z","iopub.execute_input":"2024-08-13T05:03:53.737720Z","iopub.status.idle":"2024-08-13T05:03:53.859708Z","shell.execute_reply.started":"2024-08-13T05:03:53.737690Z","shell.execute_reply":"2024-08-13T05:03:53.858803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data = pd.DataFrame(new_rows).reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.860811Z","iopub.execute_input":"2024-08-13T05:03:53.861088Z","iopub.status.idle":"2024-08-13T05:03:53.882900Z","shell.execute_reply.started":"2024-08-13T05:03:53.861065Z","shell.execute_reply":"2024-08-13T05:03:53.881879Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_condition(row_id):\n    return '_'.join(row_id.split('_')[1:])\n\ntest_data['condition'] = test_data.row_id.apply(lambda row_id: get_condition(row_id))","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.884003Z","iopub.execute_input":"2024-08-13T05:03:53.884256Z","iopub.status.idle":"2024-08-13T05:03:53.894585Z","shell.execute_reply.started":"2024-08-13T05:03:53.884235Z","shell.execute_reply":"2024-08-13T05:03:53.893729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data.head()","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.896276Z","iopub.execute_input":"2024-08-13T05:03:53.896537Z","iopub.status.idle":"2024-08-13T05:03:53.911891Z","shell.execute_reply.started":"2024-08-13T05:03:53.896515Z","shell.execute_reply":"2024-08-13T05:03:53.910818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"levels = ['l1_l2', 'l2_l3', 'l3_l4', 'l4_l5', 'l5_s1']","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.913059Z","iopub.execute_input":"2024-08-13T05:03:53.913306Z","iopub.status.idle":"2024-08-13T05:03:53.920589Z","shell.execute_reply.started":"2024-08-13T05:03:53.913285Z","shell.execute_reply":"2024-08-13T05:03:53.919655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def update_level(row, levels):\n    level = levels[row.name % len(levels)]\n    return f'{row.row_id}_{level}'\n\ntest_data['row_id'] = test_data.apply(lambda x: update_level(x, levels), axis=1)","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.922420Z","iopub.execute_input":"2024-08-13T05:03:53.922674Z","iopub.status.idle":"2024-08-13T05:03:53.934877Z","shell.execute_reply.started":"2024-08-13T05:03:53.922645Z","shell.execute_reply":"2024-08-13T05:03:53.934033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data['image_path'] = test_data.pop('image_path')","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.935944Z","iopub.execute_input":"2024-08-13T05:03:53.936354Z","iopub.status.idle":"2024-08-13T05:03:53.948040Z","shell.execute_reply.started":"2024-08-13T05:03:53.936332Z","shell.execute_reply":"2024-08-13T05:03:53.947200Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data.head()","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.949143Z","iopub.execute_input":"2024-08-13T05:03:53.949951Z","iopub.status.idle":"2024-08-13T05:03:53.965408Z","shell.execute_reply.started":"2024-08-13T05:03:53.949918Z","shell.execute_reply":"2024-08-13T05:03:53.964419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"get_plots(test_data.image_path, count_of_images=4)","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:53.966924Z","iopub.execute_input":"2024-08-13T05:03:53.967308Z","iopub.status.idle":"2024-08-13T05:03:55.487010Z","shell.execute_reply.started":"2024-08-13T05:03:53.967246Z","shell.execute_reply":"2024-08-13T05:03:55.486075Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data.head()","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:55.488246Z","iopub.execute_input":"2024-08-13T05:03:55.488568Z","iopub.status.idle":"2024-08-13T05:03:55.501276Z","shell.execute_reply.started":"2024-08-13T05:03:55.488539Z","shell.execute_reply":"2024-08-13T05:03:55.500369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\n\ndef load_DICOM(path):\n    file = pydicom.dcmread(path)\n    data = file.pixel_array\n    \n    # Make the least pixel value to be 0\n    data -= np.min(data)\n    \n    # Bring pixels to range 0-1\n    if np.max(data) != 0:\n        data = data / np.max(data)\n    \n    # Make pixels to range 0-255 with dtype uint8\n    data = (data * 255).astype(np.uint8)\n    \n    return data","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:55.502707Z","iopub.execute_input":"2024-08-13T05:03:55.502999Z","iopub.status.idle":"2024-08-13T05:03:55.512967Z","shell.execute_reply.started":"2024-08-13T05:03:55.502976Z","shell.execute_reply":"2024-08-13T05:03:55.512024Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import Dataset\nimport torchvision.transforms as transforms\n\nclass CustomDataset(Dataset):\n    def __init__(self, dataframe, transform=None):\n        self.dataframe = dataframe\n        self.transform = transform\n        self.labels = {'normal_mild': 0, 'moderate': 1, 'severe': 2}\n        \n    def __len__(self):\n        return len(self.dataframe)\n    \n    def __getitem__(self, index):\n        row = self.dataframe.iloc[index]\n        image_path = row.image_path\n        status = self.labels[row.severity]\n        \n        data = load_DICOM(image_path)\n        \n        if self.transform:\n            data = self.transform(data)\n        \n        return data, status","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:03:55.514032Z","iopub.execute_input":"2024-08-13T05:03:55.514318Z","iopub.status.idle":"2024-08-13T05:04:00.762284Z","shell.execute_reply.started":"2024-08-13T05:03:55.514294Z","shell.execute_reply":"2024-08-13T05:04:00.761348Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nfrom torch.utils.data import DataLoader\n\ndef prepare_data(df, series_description, transform, batch_size):\n    final_df = df[df.series_description == series_description]\n    \n    train_df, test_df = train_test_split(final_df, test_size=0.2, random_state=84)\n    train_df = train_df.reset_index(drop=True)\n    test_df = test_df.reset_index(drop=True)\n    \n    train_dataset = CustomDataset(train_df, transform)\n    test_dataset = CustomDataset(test_df, transform)\n    \n    train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)\n    test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=True)\n    \n    return train_loader, test_loader","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:04:00.763635Z","iopub.execute_input":"2024-08-13T05:04:00.764460Z","iopub.status.idle":"2024-08-13T05:04:01.424274Z","shell.execute_reply.started":"2024-08-13T05:04:00.764414Z","shell.execute_reply":"2024-08-13T05:04:01.423310Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torchvision.transforms as transforms\n\ntransform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.Resize((224, 224)),\n    transforms.Grayscale(num_output_channels=3),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:04:01.425494Z","iopub.execute_input":"2024-08-13T05:04:01.425810Z","iopub.status.idle":"2024-08-13T05:04:01.431389Z","shell.execute_reply.started":"2024-08-13T05:04:01.425779Z","shell.execute_reply":"2024-08-13T05:04:01.430574Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualize_batch(dataloader):\n    images, labels = next(iter(dataloader))\n    batch_size = len(images)\n    \n    # Define number of rows and columns for the subplot grid\n    grid_size = int(np.ceil(np.sqrt(batch_size)))\n    \n    fig, axes = plt.subplots(grid_size, grid_size, figsize=(15, 15))\n    axes = axes.flatten()  # Flatten the array of axes for easier iteration\n    \n    for i, (img, lbl) in enumerate(zip(images, labels)):\n        if i >= grid_size * grid_size:\n            break  # Avoids indexing issues if batch_size > grid_size * grid_size\n        \n        img = img.permute(1, 2, 0)  # Convert to HWC for visualization\n        axes[i].imshow(img)\n        axes[i].set_title(f\"Label: {lbl}\")\n        axes[i].axis('off')\n    \n    # Hide any remaining subplots if batch_size < grid_size * grid_size\n    for j in range(i + 1, grid_size * grid_size):\n        axes[j].axis('off')\n    \n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:04:01.432557Z","iopub.execute_input":"2024-08-13T05:04:01.432838Z","iopub.status.idle":"2024-08-13T05:04:01.444260Z","shell.execute_reply.started":"2024-08-13T05:04:01.432815Z","shell.execute_reply":"2024-08-13T05:04:01.443498Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.optim as optim\nimport tqdm\nimport torch.nn as nn\nimport torch\nfrom datetime import datetime as dt\n\ndef train_model(model, train_loader, val_loader, **kwargs):\n    criterion = kwargs['criterion']\n    optimizer = kwargs['optimizer']\n    num_epochs = kwargs['num_epochs']\n    \n    train_accuracies = []\n    train_losses = []\n    val_accuracies = []\n    val_losses = []\n    \n    device = 'cuda' if torch.cuda.is_available() else 'cpu'\n    \n    st = dt.now()\n    \n    for epoch in range(num_epochs):\n        model.train()\n        train_loss = 0\n        correct = 0\n        total = 0\n        \n        for inputs, labels in tqdm.tqdm(train_loader, unit='batch', desc=f'Epoch {epoch+1}/{num_epochs}: Training'):\n            inputs, labels = inputs.to(device), labels.to(device)\n            optimizer.zero_grad()\n            outputs = model(inputs)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            optimizer.step()\n            train_loss += loss.item() * inputs.size(0)\n            \n            _, predicted = torch.max(outputs.data, 1)\n            total += labels.size(0)\n            correct += (predicted == labels).sum().item()\n        \n        train_acc = 100 * correct / total\n        train_loss /= len(train_loader.dataset)\n        \n        model.eval()\n        val_loss = 0\n        correct = 0\n        total = 0\n\n        with torch.no_grad():\n            for inputs, labels in tqdm.tqdm(val_loader, desc=f'Epoch {epoch+1}/{num_epochs}: Validation'):\n                inputs, labels = inputs.to(device), labels.to(device)\n                outputs = model(inputs)\n                loss = criterion(outputs, labels)\n                val_loss += loss.item() * inputs.size(0)\n\n                _, predicted = torch.max(outputs.data, 1)\n                correct += (predicted == labels).sum().item()\n                total += labels.size(0)\n\n        val_acc = 100 * correct / total\n        val_loss /= len(val_loader.dataset)\n        \n        train_accuracies.append(train_acc)\n        val_accuracies.append(val_acc)\n        train_losses.append(train_loss)\n        val_losses.append(val_loss)\n        \n        print(f'Epoch {epoch+1}/{num_epochs}: Train loss: {train_loss:.4f}, Train acc: {train_acc:.4f}%, Val loss: {val_loss:.4f}, Val acc: {val_acc:.4f}%')\n    \n    end = dt.now()\n    \n    return {'train': [train_losses, train_accuracies], 'val': [val_losses, val_accuracies], 'time': end-st}","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:04:01.445249Z","iopub.execute_input":"2024-08-13T05:04:01.445580Z","iopub.status.idle":"2024-08-13T05:04:01.460735Z","shell.execute_reply.started":"2024-08-13T05:04:01.445550Z","shell.execute_reply":"2024-08-13T05:04:01.459837Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\ndef plot_losses_and_accuracies(train, val):\n    plt.figure(figsize=(10, 5))\n    \n    plt.subplot(1, 2, 1)\n    plt.plot(train[1], label='Training Accuracy')\n    plt.plot(val[1], label='Validation Accuracy')\n    plt.xlabel('Epoch')\n    plt.ylabel('Accuracy')\n    plt.title('Training and Validation Accuracy')\n    plt.legend()\n    \n    plt.subplot(1, 2, 2)\n    plt.plot(train[0], label='Training Loss')\n    plt.plot(val[0], label='Validation Loss')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.title('Training and Validation Loss')\n    plt.legend()\n    \n    plt.tight_layout()\n    plt.show()    ","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:04:01.461891Z","iopub.execute_input":"2024-08-13T05:04:01.462147Z","iopub.status.idle":"2024-08-13T05:04:01.473801Z","shell.execute_reply.started":"2024-08-13T05:04:01.462125Z","shell.execute_reply":"2024-08-13T05:04:01.472939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\n\nclass SimpleCNN(nn.Module):\n    def __init__(self, num_classes=3):\n        super(SimpleCNN, self).__init__()\n        self.conv1 = nn.Conv2d(in_channels=3, out_channels=32, kernel_size=3, stride=1, padding=1)\n        self.conv2 = nn.Conv2d(in_channels=32, out_channels=64, kernel_size=3, stride=1, padding=1)\n        self.conv3 = nn.Conv2d(in_channels=64, out_channels=128, kernel_size=3, stride=1, padding=1)\n        self.pool = nn.MaxPool2d(kernel_size=2, stride=2, padding=0)\n        self.fc1 = nn.Linear(128 * 28 * 28, 512)\n        self.fc2 = nn.Linear(512, num_classes)\n        self.dropout = nn.Dropout(0.3)\n        self.batch_norm1 = nn.BatchNorm2d(32)\n        self.batch_norm2 = nn.BatchNorm2d(64)\n        self.batch_norm3 = nn.BatchNorm2d(128)\n        self.batch_norm4 = nn.BatchNorm1d(512)\n\n    def forward(self, x):\n        x = self.pool(F.relu(self.batch_norm1(self.conv1(x))))\n        x = self.pool(F.relu(self.batch_norm2(self.conv2(x))))\n        x = self.pool(F.relu(self.batch_norm3(self.conv3(x))))\n        x = x.view(-1, 128 * 28 * 28)  # Flatten the tensor\n        x = F.relu(self.batch_norm4(self.fc1(x)))\n        x = self.dropout(x)\n        x = self.fc2(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:04:01.474896Z","iopub.execute_input":"2024-08-13T05:04:01.475396Z","iopub.status.idle":"2024-08-13T05:04:01.486690Z","shell.execute_reply.started":"2024-08-13T05:04:01.475367Z","shell.execute_reply":"2024-08-13T05:04:01.485633Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torchvision","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:04:01.487772Z","iopub.execute_input":"2024-08-13T05:04:01.488054Z","iopub.status.idle":"2024-08-13T05:04:01.499564Z","shell.execute_reply.started":"2024-08-13T05:04:01.488032Z","shell.execute_reply":"2024-08-13T05:04:01.498729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_classes = 3","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:04:01.500516Z","iopub.execute_input":"2024-08-13T05:04:01.500843Z","iopub.status.idle":"2024-08-13T05:04:01.509714Z","shell.execute_reply.started":"2024-08-13T05:04:01.500820Z","shell.execute_reply":"2024-08-13T05:04:01.508950Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sagittal_t1_model = SimpleCNN(num_classes)\nsagittal_t2_model = SimpleCNN(num_classes)\naxial_t2_model = SimpleCNN(num_classes)","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:04:01.510768Z","iopub.execute_input":"2024-08-13T05:04:01.511029Z","iopub.status.idle":"2024-08-13T05:04:03.035180Z","shell.execute_reply.started":"2024-08-13T05:04:01.511008Z","shell.execute_reply":"2024-08-13T05:04:03.034362Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def modify_model(model, num_classes):\n    num_ftrs = model.fc.in_features\n    model.fc = nn.Linear(num_ftrs, num_classes)\n    return model","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:04:03.036222Z","iopub.execute_input":"2024-08-13T05:04:03.036509Z","iopub.status.idle":"2024-08-13T05:04:03.041780Z","shell.execute_reply.started":"2024-08-13T05:04:03.036486Z","shell.execute_reply":"2024-08-13T05:04:03.040783Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optimizer_sagittal_t2_model = torch.optim.Adam(sagittal_t2_model.parameters(), lr=0.01)\noptimizer_sagittal_t1_model = torch.optim.Adam(sagittal_t1_model.parameters(), lr=0.01)\noptimizer_axial_t2_model = torch.optim.Adam(axial_t2_model.parameters(), lr=0.01)","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:04:03.042927Z","iopub.execute_input":"2024-08-13T05:04:03.043331Z","iopub.status.idle":"2024-08-13T05:04:03.052095Z","shell.execute_reply.started":"2024-08-13T05:04:03.043301Z","shell.execute_reply":"2024-08-13T05:04:03.051241Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader_axial, val_loader_axial = prepare_data(train_data, 'Axial T2', transform=transform, batch_size=32)\ntrain_loader_t1, val_loader_t1 = prepare_data(train_data, 'Sagittal T1', transform=transform, batch_size=32)\ntrain_loader_t2_stir, val_loader_t2_stir = prepare_data(train_data, 'Sagittal T2/STIR', transform=transform, batch_size=32)","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:04:03.053034Z","iopub.execute_input":"2024-08-13T05:04:03.054623Z","iopub.status.idle":"2024-08-13T05:04:03.137423Z","shell.execute_reply.started":"2024-08-13T05:04:03.054598Z","shell.execute_reply":"2024-08-13T05:04:03.136601Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Assuming train_loader_axial is a DataLoader object\nsample_batch = next(iter(train_loader_axial))\ndata, labels = sample_batch\n\n# To get the shape of the data (assuming data is a tensor)\ndata_shape = data.shape\nprint(data_shape)","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:04:03.138594Z","iopub.execute_input":"2024-08-13T05:04:03.138919Z","iopub.status.idle":"2024-08-13T05:04:03.970699Z","shell.execute_reply.started":"2024-08-13T05:04:03.138892Z","shell.execute_reply":"2024-08-13T05:04:03.969800Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Assuming train_loader_axial is a DataLoader object\nsample_batch = next(iter(train_loader_t1))\ndata, labels = sample_batch\n\n# To get the shape of the data (assuming data is a tensor)\ndata_shape = data.shape\nprint(data_shape)","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:04:03.971839Z","iopub.execute_input":"2024-08-13T05:04:03.972119Z","iopub.status.idle":"2024-08-13T05:04:04.911681Z","shell.execute_reply.started":"2024-08-13T05:04:03.972096Z","shell.execute_reply":"2024-08-13T05:04:04.910777Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Assuming train_loader_axial is a DataLoader object\nsample_batch = next(iter(train_loader_t2_stir))\ndata, labels = sample_batch\n\n# To get the shape of the data (assuming data is a tensor)\ndata_shape = data.shape\nprint(data_shape)","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:04:04.912858Z","iopub.execute_input":"2024-08-13T05:04:04.913203Z","iopub.status.idle":"2024-08-13T05:04:05.721049Z","shell.execute_reply.started":"2024-08-13T05:04:04.913171Z","shell.execute_reply":"2024-08-13T05:04:05.720080Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'\n\nmodels = {\n    'Sagittal T1': sagittal_t1_model.to(device),\n    'Sagittal T2/STIR': sagittal_t2_model.to(device),\n    'Axial T2': axial_t2_model.to(device)\n}\n\ntrain_loaders = {\n    'Sagittal T1': [train_loader_t1, val_loader_t1],\n    'Sagittal T2/STIR': [train_loader_t2_stir, val_loader_t2_stir],\n    'Axial T2': [train_loader_axial, val_loader_axial]\n}\n\noptimizers = {\n    'Sagittal T1': optimizer_sagittal_t1_model,\n    'Sagittal T2/STIR': optimizer_sagittal_t2_model,\n    'Axial T2': optimizer_axial_t2_model\n}","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:04:05.722523Z","iopub.execute_input":"2024-08-13T05:04:05.722937Z","iopub.status.idle":"2024-08-13T05:04:06.110861Z","shell.execute_reply.started":"2024-08-13T05:04:05.722904Z","shell.execute_reply":"2024-08-13T05:04:06.109813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\nnum_epochs = 5\nresults = {}\n\nfor key in models.keys():\n    model = models[key]\n    optimizer = optimizers[key]\n    results[key] = train_model(model, *train_loaders[key], criterion=criterion, num_epochs=num_epochs, optimizer=optimizer)","metadata":{"execution":{"iopub.status.busy":"2024-08-13T05:04:06.112030Z","iopub.execute_input":"2024-08-13T05:04:06.112314Z","iopub.status.idle":"2024-08-13T06:25:37.256252Z","shell.execute_reply.started":"2024-08-13T05:04:06.112288Z","shell.execute_reply":"2024-08-13T06:25:37.255349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"results","metadata":{"execution":{"iopub.status.busy":"2024-08-13T06:25:37.257395Z","iopub.execute_input":"2024-08-13T06:25:37.257673Z","iopub.status.idle":"2024-08-13T06:25:37.264977Z","shell.execute_reply.started":"2024-08-13T06:25:37.257650Z","shell.execute_reply":"2024-08-13T06:25:37.264127Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for key in models.keys():\n    plot_losses_and_accuracies(results[key]['train'], results[key]['val'])","metadata":{"execution":{"iopub.status.busy":"2024-08-13T06:25:37.266155Z","iopub.execute_input":"2024-08-13T06:25:37.266553Z","iopub.status.idle":"2024-08-13T06:25:39.013854Z","shell.execute_reply.started":"2024-08-13T06:25:37.266520Z","shell.execute_reply":"2024-08-13T06:25:39.012942Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for key in results.keys():\n    print(f'Time taken for {key} = {results[key][\"time\"]}')","metadata":{"execution":{"iopub.status.busy":"2024-08-13T06:25:39.014960Z","iopub.execute_input":"2024-08-13T06:25:39.015242Z","iopub.status.idle":"2024-08-13T06:25:39.020652Z","shell.execute_reply.started":"2024-08-13T06:25:39.015217Z","shell.execute_reply":"2024-08-13T06:25:39.019819Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import Dataset\nimport torchvision.transforms as transforms\n\nclass CustomTestDataset(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, index):\n        row = self.dataframe.iloc[index]\n        image_path = row.image_path\n        \n        img = load_DICOM(image_path)\n        \n        if self.transform:\n            img = self.transform(img)\n        \n        return img","metadata":{"execution":{"iopub.status.busy":"2024-08-13T06:25:39.021742Z","iopub.execute_input":"2024-08-13T06:25:39.022075Z","iopub.status.idle":"2024-08-13T06:25:39.033839Z","shell.execute_reply.started":"2024-08-13T06:25:39.022045Z","shell.execute_reply":"2024-08-13T06:25:39.033102Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = CustomTestDataset(test_data, transform=transform)\ntest_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2024-08-13T06:25:39.034819Z","iopub.execute_input":"2024-08-13T06:25:39.035077Z","iopub.status.idle":"2024-08-13T06:25:39.043039Z","shell.execute_reply.started":"2024-08-13T06:25:39.035056Z","shell.execute_reply":"2024-08-13T06:25:39.042209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def evaluate_model(test_loader, test_data, models):\n    predictions = []\n    normal_mild_probs = []\n    moderate_probs = []\n    severe_probs = []\n    \n    device = 'cuda' if torch.cuda.is_available() else 'cpu'\n    \n    for model in models.values():\n        model.to(device)\n    \n    \n    with torch.no_grad():\n        for idx, img in enumerate(tqdm.tqdm(test_loader, desc='Evaluating:')):\n            img = img.to(device)\n            series_desc = test_data.iloc[idx]['series_description']\n            \n            model = models[series_desc]\n            model.eval()\n            \n            outputs = model(img)\n            probs = torch.softmax(outputs, dim=1).squeeze(0)\n            normal_mild_probs.append(probs[0].item())\n            moderate_probs.append(probs[1].item())\n            severe_probs.append(probs[2].item())\n            predictions.append(probs)\n    \n    return normal_mild_probs, moderate_probs, severe_probs, predictions","metadata":{"execution":{"iopub.status.busy":"2024-08-13T06:25:39.043976Z","iopub.execute_input":"2024-08-13T06:25:39.044519Z","iopub.status.idle":"2024-08-13T06:25:39.053320Z","shell.execute_reply.started":"2024-08-13T06:25:39.044495Z","shell.execute_reply":"2024-08-13T06:25:39.052482Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data['normal_mild'], test_data['moderate'], test_data['severe'], _ = evaluate_model(test_loader, test_data, models)","metadata":{"execution":{"iopub.status.busy":"2024-08-13T06:25:39.054323Z","iopub.execute_input":"2024-08-13T06:25:39.054569Z","iopub.status.idle":"2024-08-13T06:25:43.198219Z","shell.execute_reply.started":"2024-08-13T06:25:39.054548Z","shell.execute_reply":"2024-08-13T06:25:43.197306Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"final_submission = test_data[['row_id', 'normal_mild', 'moderate', 'severe']]","metadata":{"execution":{"iopub.status.busy":"2024-08-13T06:25:43.199497Z","iopub.execute_input":"2024-08-13T06:25:43.199988Z","iopub.status.idle":"2024-08-13T06:25:43.205815Z","shell.execute_reply.started":"2024-08-13T06:25:43.199954Z","shell.execute_reply":"2024-08-13T06:25:43.204779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pd.reset_option('max_rows')","metadata":{"execution":{"iopub.status.busy":"2024-08-13T06:55:51.073130Z","iopub.execute_input":"2024-08-13T06:55:51.073489Z","iopub.status.idle":"2024-08-13T06:55:51.078233Z","shell.execute_reply.started":"2024-08-13T06:55:51.073459Z","shell.execute_reply":"2024-08-13T06:55:51.077204Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"final_submission","metadata":{"execution":{"iopub.status.busy":"2024-08-13T06:55:52.551609Z","iopub.execute_input":"2024-08-13T06:55:52.552328Z","iopub.status.idle":"2024-08-13T06:55:52.564538Z","shell.execute_reply.started":"2024-08-13T06:55:52.552296Z","shell.execute_reply":"2024-08-13T06:55:52.563682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"grouped_submission = final_submission.groupby('row_id').max().reset_index()","metadata":{"execution":{"iopub.status.busy":"2024-08-13T06:25:43.224836Z","iopub.execute_input":"2024-08-13T06:25:43.225094Z","iopub.status.idle":"2024-08-13T06:25:43.240211Z","shell.execute_reply.started":"2024-08-13T06:25:43.225072Z","shell.execute_reply":"2024-08-13T06:25:43.239511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"grouped_submission.head(5)","metadata":{"execution":{"iopub.status.busy":"2024-08-13T07:03:07.195786Z","iopub.execute_input":"2024-08-13T07:03:07.196454Z","iopub.status.idle":"2024-08-13T07:03:07.207003Z","shell.execute_reply.started":"2024-08-13T07:03:07.196425Z","shell.execute_reply":"2024-08-13T07:03:07.206201Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"grouped_submission.to_csv('/kaggle/working/submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2024-08-13T07:06:02.795143Z","iopub.execute_input":"2024-08-13T07:06:02.795506Z","iopub.status.idle":"2024-08-13T07:06:02.802914Z","shell.execute_reply.started":"2024-08-13T07:06:02.795476Z","shell.execute_reply":"2024-08-13T07:06:02.802052Z"},"trusted":true},"execution_count":null,"outputs":[]}]}