{"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":9076668,"sourceType":"datasetVersion","datasetId":5475657}],"dockerImageVersionId":30732,"isInternetEnabled":false,"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 matplotlib.pyplot as plt\nimport warnings\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset\nfrom tqdm import tqdm\nfrom PIL import Image\nimport pydicom\nimport cv2\nimport os\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport timm\nimport random\nfrom sklearn.model_selection import KFold","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:24.852460Z","iopub.execute_input":"2024-09-10T20:25:24.852878Z","iopub.status.idle":"2024-09-10T20:25:32.496525Z","shell.execute_reply.started":"2024-09-10T20:25:24.852838Z","shell.execute_reply":"2024-09-10T20:25:32.495546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.cpu_count()","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:32.498239Z","iopub.execute_input":"2024-09-10T20:25:32.498562Z","iopub.status.idle":"2024-09-10T20:25:32.505872Z","shell.execute_reply.started":"2024-09-10T20:25:32.498530Z","shell.execute_reply":"2024-09-10T20:25:32.504986Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_size = 512\nin_channels = 30\nnum_classes = 75\nbatch_size = 16\nn_workers = os.cpu_count()\nmodel_name = 'tf_efficientnet_b3.ns_jft_in1k'\nlr = 1e-4\nepochs = 30\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\ncomm_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:32.506882Z","iopub.execute_input":"2024-09-10T20:25:32.507119Z","iopub.status.idle":"2024-09-10T20:25:32.588325Z","shell.execute_reply.started":"2024-09-10T20:25:32.507098Z","shell.execute_reply":"2024-09-10T20:25:32.587515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_random_seed(seed: int = 42, deterministic: bool = False):\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)  \n    torch.backends.cudnn.benchmark = True\n    torch.backends.cudnn.deterministic = deterministic \n\nset_random_seed(42)","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:32.589787Z","iopub.execute_input":"2024-09-10T20:25:32.590153Z","iopub.status.idle":"2024-09-10T20:25:32.602737Z","shell.execute_reply.started":"2024-09-10T20:25:32.590118Z","shell.execute_reply":"2024-09-10T20:25:32.601819Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv\")\ndf","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:32.605208Z","iopub.execute_input":"2024-09-10T20:25:32.605523Z","iopub.status.idle":"2024-09-10T20:25:32.672385Z","shell.execute_reply.started":"2024-09-10T20:25:32.605477Z","shell.execute_reply":"2024-09-10T20:25:32.671458Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"desc_df = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv\")\ndesc_df","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:32.673477Z","iopub.execute_input":"2024-09-10T20:25:32.673772Z","iopub.status.idle":"2024-09-10T20:25:32.693634Z","shell.execute_reply.started":"2024-09-10T20:25:32.673748Z","shell.execute_reply":"2024-09-10T20:25:32.692735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label_df = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_label_coordinates.csv\")\nlabel_df","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:32.694888Z","iopub.execute_input":"2024-09-10T20:25:32.695635Z","iopub.status.idle":"2024-09-10T20:25:32.842079Z","shell.execute_reply.started":"2024-09-10T20:25:32.695600Z","shell.execute_reply":"2024-09-10T20:25:32.841171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"figure, axis = plt.subplots(1,3, figsize=(20,5)) \nfor idx, d in enumerate(['foraminal', 'subarticular', 'canal']):\n    diagnosis = list(filter(lambda x: x.find(d) > -1, df.columns))\n    dff = df[diagnosis]\n    with warnings.catch_warnings():\n        warnings.simplefilter(action='ignore', category=FutureWarning)\n        value_counts = dff.apply(pd.value_counts).fillna(0).T\n    value_counts.plot(kind='bar', stacked=True, ax=axis[idx])\n    axis[idx].set_title(f'{d} distribution')","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:32.843173Z","iopub.execute_input":"2024-09-10T20:25:32.843438Z","iopub.status.idle":"2024-09-10T20:25:34.047610Z","shell.execute_reply.started":"2024-09-10T20:25:32.843417Z","shell.execute_reply":"2024-09-10T20:25:34.046686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.fillna(-100, inplace=True)","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:34.048808Z","iopub.execute_input":"2024-09-10T20:25:34.049089Z","iopub.status.idle":"2024-09-10T20:25:34.061453Z","shell.execute_reply.started":"2024-09-10T20:25:34.049064Z","shell.execute_reply":"2024-09-10T20:25:34.060749Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = {'Normal/Mild' : 0, \n          'Moderate' : 1,\n          'Severe' : 2}\ndf.replace(labels, inplace=True)\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:34.062641Z","iopub.execute_input":"2024-09-10T20:25:34.063077Z","iopub.status.idle":"2024-09-10T20:25:34.117509Z","shell.execute_reply.started":"2024-09-10T20:25:34.063044Z","shell.execute_reply":"2024-09-10T20:25:34.116618Z"},"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-09-10T20:25:34.118772Z","iopub.execute_input":"2024-09-10T20:25:34.119070Z","iopub.status.idle":"2024-09-10T20:25:34.123672Z","shell.execute_reply.started":"2024-09-10T20:25:34.119045Z","shell.execute_reply":"2024-09-10T20:25:34.122814Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class_counts = [0, 0, 0]\nfor cond in CONDITIONS:\n    for level in LEVELS:\n        f_id = f'{cond}_{level}'\n        temp_df = df[f_id].value_counts()\n        class_counts[0] += temp_df[0]\n        class_counts[1] += temp_df[1]\n        class_counts[2] += temp_df[2]","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:34.124871Z","iopub.execute_input":"2024-09-10T20:25:34.125142Z","iopub.status.idle":"2024-09-10T20:25:34.147209Z","shell.execute_reply.started":"2024-09-10T20:25:34.125119Z","shell.execute_reply":"2024-09-10T20:25:34.146336Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class_weights = [1.0 / count for count in class_counts]\nclass_weights = torch.tensor(class_weights, dtype=torch.float)\nclass_weights = class_weights / class_weights.sum()\nclass_weights = class_weights.to(device)\nprint(class_weights)","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:34.148422Z","iopub.execute_input":"2024-09-10T20:25:34.148866Z","iopub.status.idle":"2024-09-10T20:25:34.618956Z","shell.execute_reply.started":"2024-09-10T20:25:34.148836Z","shell.execute_reply":"2024-09-10T20:25:34.618003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"st_ids = df['study_id'].unique()\nprint(st_ids.size)\nprint(st_ids)","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:34.623293Z","iopub.execute_input":"2024-09-10T20:25:34.623620Z","iopub.status.idle":"2024-09-10T20:25:34.630308Z","shell.execute_reply.started":"2024-09-10T20:25:34.623592Z","shell.execute_reply":"2024-09-10T20:25:34.629426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"desc_ids = desc_df['series_description'].unique()\ndesc_ids","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:34.631535Z","iopub.execute_input":"2024-09-10T20:25:34.631875Z","iopub.status.idle":"2024-09-10T20:25:34.645967Z","shell.execute_reply.started":"2024-09-10T20:25:34.631842Z","shell.execute_reply":"2024-09-10T20:25:34.645027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/100206310/1012284084'\nfig, axes = plt.subplots(2, 2, figsize=(70, 70))\naxes = axes.flatten()\nfor file_name, ax in zip(os.listdir(path), axes):\n    path = f'/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/100206310/1012284084/{file_name}'\n    dcm_data = pydicom.dcmread(path)\n    img = dcm_data.pixel_array\n    img = (img - img.min()) / (img.max() - img.min() + 1e-6) * 255\n    img = cv2.resize(img, (512, 512), interpolation=cv2.INTER_CUBIC)\n    ax.imshow(img, cmap = plt.cm.bone)\n    ax.axis('off')  # Hide the axis\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:34.647152Z","iopub.execute_input":"2024-09-10T20:25:34.647473Z","iopub.status.idle":"2024-09-10T20:25:39.536103Z","shell.execute_reply.started":"2024-09-10T20:25:34.647449Z","shell.execute_reply":"2024-09-10T20:25:39.531475Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/100206310/1792451510'\nfig, axes = plt.subplots(2, 2, figsize=(70, 70))\naxes = axes.flatten()\nfor file_name, ax in zip(os.listdir(path), axes):\n    src_path = f'{path}/{file_name}'\n    dcm_data = pydicom.dcmread(src_path)\n    img = dcm_data.pixel_array\n    img = (img - img.min()) / (img.max() - img.min() + 1e-6) * 255\n    img = cv2.resize(img, (512, 512), interpolation=cv2.INTER_CUBIC)\n    ax.imshow(img, cmap = plt.cm.bone)\n    ax.axis('off')  # Hide the axis\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:39.537212Z","iopub.execute_input":"2024-09-10T20:25:39.537541Z","iopub.status.idle":"2024-09-10T20:25:44.397356Z","shell.execute_reply.started":"2024-09-10T20:25:39.537508Z","shell.execute_reply":"2024-09-10T20:25:44.395743Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/100206310/2092806862'\nfig, axes = plt.subplots(2, 2, figsize=(70, 70))\naxes = axes.flatten()\nfor file_name, ax in zip(os.listdir(path), axes):\n    src_path = f'{path}/{file_name}'\n    dcm_data = pydicom.dcmread(src_path)\n    img = dcm_data.pixel_array\n    img = (img - img.min()) / (img.max() - img.min() + 1e-6) * 255\n    img = cv2.resize(img, (512, 512), interpolation=cv2.INTER_CUBIC)\n    ax.imshow(img, cmap = plt.cm.bone)\n    ax.axis('off')  # Hide the axis\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:44.398785Z","iopub.execute_input":"2024-09-10T20:25:44.399127Z","iopub.status.idle":"2024-09-10T20:25:49.254095Z","shell.execute_reply.started":"2024-09-10T20:25:44.399092Z","shell.execute_reply":"2024-09-10T20:25:49.253135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transforms = A.Compose([\n    A.RandomBrightnessContrast(brightness_limit=(-0.2, 0.2), contrast_limit=(-0.2, 0.2), p=0.75),\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=0.75),\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=0.75),\n    A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=15, border_mode=0, p=0.75),\n    A.CoarseDropout(max_holes=16, max_height=64, max_width=64, min_holes=1, min_height=8, min_width=8, p=0.75),    \n    A.Normalize(mean=0.5, std=0.5)\n])\ntransforms_val = A.Compose([\n    A.Normalize(mean=0.5, std=0.5)\n])","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:49.255264Z","iopub.execute_input":"2024-09-10T20:25:49.255676Z","iopub.status.idle":"2024-09-10T20:25:49.266668Z","shell.execute_reply.started":"2024-09-10T20:25:49.255644Z","shell.execute_reply":"2024-09-10T20:25:49.265708Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x_df = desc_df[desc_df['study_id']==4003253]","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:49.267751Z","iopub.execute_input":"2024-09-10T20:25:49.268001Z","iopub.status.idle":"2024-09-10T20:25:49.277363Z","shell.execute_reply.started":"2024-09-10T20:25:49.267979Z","shell.execute_reply":"2024-09-10T20:25:49.276518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i, row in x_df.iterrows():\n    print(i, row['series_id'])","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:49.278563Z","iopub.execute_input":"2024-09-10T20:25:49.278884Z","iopub.status.idle":"2024-09-10T20:25:49.289807Z","shell.execute_reply.started":"2024-09-10T20:25:49.278862Z","shell.execute_reply":"2024-09-10T20:25:49.288964Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNA_Dataset(Dataset):\n    def __init__(self, df, desc_df, transforms):\n        super().__init__()\n        self.df = df\n        self.desc_df = desc_df\n        self.transforms = transforms\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        X = np.zeros((img_size, img_size, in_channels), dtype=np.uint8)\n        i = self.df.iloc[idx]\n        study_id = i['study_id']\n        labels = i[1 : ].values.astype(np.int64)\n        temp_df = self.desc_df[self.desc_df['study_id']==study_id]\n        for i, row in temp_df.iterrows():\n            study_id = row['study_id']\n            series_id = row['series_id']\n            series_desc = row['series_description']\n            if series_desc == 'Axial T2':\n                try:\n                    j = 0\n                    for img_file in os.listdir(f'{comm_path}/train_images/{study_id}/{series_id}'):\n                        src_path  = f'{comm_path}/train_images/{study_id}/{series_id}/{img_file}'\n                        dcm_data = pydicom.dcmread(src_path)\n                        img = dcm_data.pixel_array\n                        img = (img - img.min()) / (img.max() - img.min() + 1e-6) * 255\n                        img = cv2.resize(img, (img_size, img_size), interpolation=cv2.INTER_CUBIC)\n                        img = np.array(img)           \n                        X[:, :, j] = img.astype(np.uint8)\n                        j=j+1\n                        if j==10:\n                            break\n                except:\n                    print(f'failed to load on {study_id}, Axial T2')\n                    pass \n                \n            if series_desc == 'Sagittal T1':\n                try:\n                    j = 0\n                    for img_file in os.listdir(f'{comm_path}/train_images/{study_id}/{series_id}'):\n                        src_path  = f'{comm_path}/train_images/{study_id}/{series_id}/{img_file}'\n                        dcm_data = pydicom.dcmread(src_path)\n                        img = dcm_data.pixel_array\n                        img = (img - img.min()) / (img.max() - img.min() + 1e-6) * 255\n                        img = cv2.resize(img, (img_size, img_size), interpolation=cv2.INTER_CUBIC)\n                        img = np.array(img)                    \n                        X[:, :, j+10] = img.astype(np.uint8)\n                        j=j+1\n                        if j==10:\n                            break\n                except:\n                    print(f'failed to load on {study_id}, Sagittal T1')\n                    pass \n            \n            if series_desc == 'Sagittal T2/STIR':\n                try:\n                    j = 0\n                    for img_file in os.listdir(f'{comm_path}/train_images/{study_id}/{series_id}'):\n                        src_path  = f'{comm_path}/train_images/{study_id}/{series_id}/{img_file}'\n                        dcm_data = pydicom.dcmread(src_path)\n                        img = dcm_data.pixel_array\n                        img = (img - img.min()) / (img.max() - img.min() + 1e-6) * 255\n                        img = cv2.resize(img, (img_size, img_size), interpolation=cv2.INTER_CUBIC)\n                        img = np.array(img)           \n                        X[:, :, j+20] = img.astype(np.uint8)\n                        j=j+1\n                        if j==10:\n                            break\n                except:\n                    print(f'failed to load on {study_id}, Sagittal T2/STIR')\n                    pass \n        \n        X = self.transforms(image=X)[\"image\"]\n        \n        X = X.transpose(2, 0, 1)\n        \n        return X, labels     ","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:49.290970Z","iopub.execute_input":"2024-09-10T20:25:49.291237Z","iopub.status.idle":"2024-09-10T20:25:49.313036Z","shell.execute_reply.started":"2024-09-10T20:25:49.291214Z","shell.execute_reply":"2024-09-10T20:25:49.312019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"split_ratio = 0.8\nsplit_index = int(len(df) * split_ratio)\n\ntrain_df = df[:split_index]\nval_df = df[split_index:]","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:49.314116Z","iopub.execute_input":"2024-09-10T20:25:49.314390Z","iopub.status.idle":"2024-09-10T20:25:49.327157Z","shell.execute_reply.started":"2024-09-10T20:25:49.314365Z","shell.execute_reply":"2024-09-10T20:25:49.326308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_df), len(val_df)","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:49.328461Z","iopub.execute_input":"2024-09-10T20:25:49.328744Z","iopub.status.idle":"2024-09-10T20:25:49.339155Z","shell.execute_reply.started":"2024-09-10T20:25:49.328721Z","shell.execute_reply":"2024-09-10T20:25:49.338311Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds = RSNA_Dataset(train_df, desc_df, transforms)\ntrain_dl = DataLoader(train_ds, batch_size=batch_size, shuffle=False, pin_memory=True, drop_last=True, num_workers=n_workers)\nval_ds = RSNA_Dataset(val_df, desc_df, transforms_val)\nval_dl = DataLoader(val_ds, batch_size=batch_size, shuffle=False, pin_memory=True, drop_last=True, num_workers=n_workers)","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:49.340359Z","iopub.execute_input":"2024-09-10T20:25:49.340919Z","iopub.status.idle":"2024-09-10T20:25:49.349041Z","shell.execute_reply.started":"2024-09-10T20:25:49.340886Z","shell.execute_reply":"2024-09-10T20:25:49.348196Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i, (x, t) in enumerate(train_dl):\n    if i==1:\n        break\n    print(x.shape)\n    z = x.numpy().transpose(0,2,3,1)[0, ..., i]\n    print(z.shape)\n    plt.imshow(z, plt.cm.bone)\n    plt.show()\nplt.close()","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:25:49.350123Z","iopub.execute_input":"2024-09-10T20:25:49.350371Z","iopub.status.idle":"2024-09-10T20:26:29.203971Z","shell.execute_reply.started":"2024-09-10T20:25:49.350348Z","shell.execute_reply":"2024-09-10T20:26:29.202755Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNAModel(nn.Module):\n    def __init__(self, model_name, in_channels, num_classes):\n        super().__init__()\n        self.model = timm.create_model(model_name, in_chans=in_channels, num_classes=num_classes, global_pool='avg')\n    \n    def forward(self, x):\n        y = self.model(x)\n        return y","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:26:29.206482Z","iopub.execute_input":"2024-09-10T20:26:29.207253Z","iopub.status.idle":"2024-09-10T20:26:29.214888Z","shell.execute_reply.started":"2024-09-10T20:26:29.207189Z","shell.execute_reply":"2024-09-10T20:26:29.213864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = RSNAModel(model_name, in_channels, num_classes).to(device)\ncriterion = nn.CrossEntropyLoss(weight=class_weights)\noptimizer = optim.AdamW(list(model.parameters()), lr=lr, betas=(0.5, 0.999))\n\nautocast = torch.cuda.amp.autocast(enabled=True, dtype=torch.half)\nscaler = torch.cuda.amp.GradScaler(enabled=True, init_scale=4096)","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:26:29.219425Z","iopub.execute_input":"2024-09-10T20:26:29.221938Z","iopub.status.idle":"2024-09-10T20:26:29.520827Z","shell.execute_reply.started":"2024-09-10T20:26:29.221901Z","shell.execute_reply":"2024-09-10T20:26:29.520019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(model, criterion, optimizer, train_dl, val_dl):\n    model.train()\n    total_loss = 0\n    loop = tqdm(train_dl, leave=True)\n    for X, y in loop:\n        X = X.to(device)\n        y = y.reshape((batch_size*25))\n        y = y.to(device)\n        with autocast:\n            y_pred = model(X)\n            y_pred = y_pred.reshape((batch_size*25, 3))\n            loss = criterion(y_pred, y)\n        \n        loop.set_description(f\"Loss: {loss.item():.4f}\")\n        total_loss += loss.item()\n        \n        optimizer.zero_grad()\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n    \n    train_loss = total_loss / len(train_dl)\n    print(f\"Train Loss: {train_loss:.4f}\")\n    \n    model.eval()\n    total_loss = 0\n    loop = tqdm(val_dl, leave=True)\n    with torch.no_grad():\n        for X, y in loop:\n            X = X.to(device)\n            y = y.reshape((batch_size*25))\n            y = y.to(device)\n            with autocast:\n                y_pred = model(X)\n                y_pred = y_pred.reshape((batch_size*25, 3))\n                loss = criterion(y_pred, y)\n            \n            loop.set_description(f\"Loss: {loss.item():.4f}\")\n            total_loss += loss.item()\n    \n    val_loss = total_loss / len(val_dl)\n    print(f\"Val Loss: {val_loss:.4f}\")\n    \n    #fname = f'/kaggle/working/model_fo88   ld{fold}.pt'\n    #torch.save(model.state_dict(), fname)","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:26:29.521929Z","iopub.execute_input":"2024-09-10T20:26:29.522216Z","iopub.status.idle":"2024-09-10T20:26:29.532824Z","shell.execute_reply.started":"2024-09-10T20:26:29.522192Z","shell.execute_reply":"2024-09-10T20:26:29.531844Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def main():\n    for num in range(epochs):\n        print(\"Epoch:\", num+1)\n        train(model, criterion, optimizer, train_dl, val_dl)\n            \n\nif __name__ == \"__main__\":\n    main()","metadata":{"execution":{"iopub.status.busy":"2024-09-10T20:26:29.533968Z","iopub.execute_input":"2024-09-10T20:26:29.534224Z","iopub.status.idle":"2024-09-10T22:57:10.050474Z","shell.execute_reply.started":"2024-09-10T20:26:29.534202Z","shell.execute_reply":"2024-09-10T22:57:10.049318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tdesc_df = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_series_descriptions.csv\")","metadata":{"execution":{"iopub.status.busy":"2024-09-10T22:57:10.052441Z","iopub.execute_input":"2024-09-10T22:57:10.053372Z","iopub.status.idle":"2024-09-10T22:57:10.064849Z","shell.execute_reply.started":"2024-09-10T22:57:10.053325Z","shell.execute_reply":"2024-09-10T22:57:10.063981Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tdesc_df","metadata":{"execution":{"iopub.status.busy":"2024-09-10T22:57:10.066066Z","iopub.execute_input":"2024-09-10T22:57:10.066383Z","iopub.status.idle":"2024-09-10T22:57:10.077680Z","shell.execute_reply.started":"2024-09-10T22:57:10.066354Z","shell.execute_reply":"2024-09-10T22:57:10.076665Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"study_ids = list(tdesc_df['study_id'].unique())","metadata":{"execution":{"iopub.status.busy":"2024-09-10T22:57:10.078818Z","iopub.execute_input":"2024-09-10T22:57:10.079182Z","iopub.status.idle":"2024-09-10T22:57:10.088172Z","shell.execute_reply.started":"2024-09-10T22:57:10.079148Z","shell.execute_reply":"2024-09-10T22:57:10.087324Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"study_ids","metadata":{"execution":{"iopub.status.busy":"2024-09-10T22:57:10.089305Z","iopub.execute_input":"2024-09-10T22:57:10.089610Z","iopub.status.idle":"2024-09-10T22:57:10.099974Z","shell.execute_reply.started":"2024-09-10T22:57:10.089585Z","shell.execute_reply":"2024-09-10T22:57:10.099017Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNATestDataset(Dataset):\n    def __init__(self, tdesc_df, study_ids, transform=None):\n        self.tdesc_df = tdesc_df\n        self.study_ids = study_ids\n        self.transforms = transform\n    \n    def __len__(self):\n        return len(self.study_ids)\n    \n    def __getitem__(self, idx):\n        X_test = np.zeros((img_size, img_size, in_channels), dtype=np.uint8)\n        st_id = self.study_ids[idx]\n        temp_df = tdesc_df[tdesc_df['study_id']==st_id]\n        for i, row in temp_df.iterrows():\n            study_id = row['study_id']\n            series_id = row['series_id']\n            series_desc = row['series_description']\n            if series_desc == 'Axial T2':\n                try:\n                    j = 0\n                    for img_file in os.listdir(f'{comm_path}/test_images/{study_id}/{series_id}'):\n                        src_path  = f'{comm_path}/test_images/{study_id}/{series_id}/{img_file}'\n                        dcm_data = pydicom.dcmread(src_path)\n                        img = dcm_data.pixel_array\n                        img = (img - img.min()) / (img.max() - img.min() + 1e-6) * 255\n                        img = cv2.resize(img, (img_size, img_size), interpolation=cv2.INTER_CUBIC)\n                        img = np.array(img)           \n                        X_test[:, :, j] = img.astype(np.uint8)\n                        j=j+1\n                        if j==10:\n                            break\n                except:\n                    print(f'failed to load on {study_id}, Axial T2')\n                    pass \n                \n            if series_desc == 'Sagittal T1': \n                try:\n                    j = 0\n                    for img_file in os.listdir(f'{comm_path}/test_images/{study_id}/{series_id}'):\n                        src_path  = f'{comm_path}/test_images/{study_id}/{series_id}/{img_file}'\n                        dcm_data = pydicom.dcmread(src_path)\n                        img = dcm_data.pixel_array\n                        img = (img - img.min()) / (img.max() - img.min() + 1e-6) * 255\n                        img = cv2.resize(img, (img_size, img_size), interpolation=cv2.INTER_CUBIC)\n                        img = np.array(img)                    \n                        X_test[:, :, j+10] = img.astype(np.uint8)\n                        j=j+1\n                        if j==10:\n                            break\n                except Exception as e:\n                    print(f'failed to load on {study_id}, Sagittal T1, {e}')\n                    pass \n            \n            if series_desc == 'Sagittal T2/STIR':\n                try:\n                    j = 0\n                    for img_file in os.listdir(f'{comm_path}/test_images/{study_id}/{series_id}'):\n                        src_path  = f'{comm_path}/test_images/{study_id}/{series_id}/{img_file}'\n                        dcm_data = pydicom.dcmread(src_path)\n                        img = dcm_data.pixel_array\n                        img = (img - img.min()) / (img.max() - img.min() + 1e-6) * 255\n                        img = cv2.resize(img, (img_size, img_size), interpolation=cv2.INTER_CUBIC)\n                        img = np.array(img)           \n                        X_test[:, :, j+20] = img.astype(np.uint8)\n                        j=j+1\n                        if j==10:\n                            break\n                except:\n                    print(f'failed to load on {study_id}, Sagittal T2/STIR')\n                    pass \n        \n        \n        X_test = self.transforms(image=X_test)[\"image\"]\n        X_test = X_test.transpose(2, 0, 1)\n        \n        return X_test, st_id","metadata":{"execution":{"iopub.status.busy":"2024-09-10T22:57:10.101089Z","iopub.execute_input":"2024-09-10T22:57:10.101384Z","iopub.status.idle":"2024-09-10T22:57:10.121358Z","shell.execute_reply.started":"2024-09-10T22:57:10.101348Z","shell.execute_reply":"2024-09-10T22:57:10.120535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_ds = RSNATestDataset(tdesc_df, study_ids, transform=transforms_val)\ntest_dl = DataLoader(test_ds, batch_size=1, shuffle=False, pin_memory=True, drop_last=False)","metadata":{"execution":{"iopub.status.busy":"2024-09-10T22:57:10.122713Z","iopub.execute_input":"2024-09-10T22:57:10.123040Z","iopub.status.idle":"2024-09-10T22:57:10.135633Z","shell.execute_reply.started":"2024-09-10T22:57:10.123010Z","shell.execute_reply":"2024-09-10T22:57:10.134814Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(test_dl)","metadata":{"execution":{"iopub.status.busy":"2024-09-10T22:57:10.136705Z","iopub.execute_input":"2024-09-10T22:57:10.137044Z","iopub.status.idle":"2024-09-10T22:57:10.147951Z","shell.execute_reply.started":"2024-09-10T22:57:10.137015Z","shell.execute_reply":"2024-09-10T22:57:10.147210Z"},"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]\nrow_names = []\ny_preds = []\nautocast = torch.cuda.amp.autocast(enabled=True, dtype=torch.half)\nwith tqdm(test_dl, leave=True) as pbar:\n    with torch.no_grad():\n        for idx, (X, st_id) in enumerate(pbar):\n            m = nn.Softmax(dim=1)\n            X = X.to(device)\n            for cond in CONDITIONS:\n                for level in LEVELS:\n                    row_names.append(f'{st_id.item()}_{cond}_{level}')\n            with autocast:\n                y_pred = model(X).reshape((25, 3))\n                y_pred = m(y_pred)\n                y_pred = y_pred.cpu().numpy()\n                y_preds.append(y_pred)","metadata":{"execution":{"iopub.status.busy":"2024-09-10T22:57:10.149087Z","iopub.execute_input":"2024-09-10T22:57:10.149404Z","iopub.status.idle":"2024-09-10T22:57:12.733344Z","shell.execute_reply.started":"2024-09-10T22:57:10.149372Z","shell.execute_reply":"2024-09-10T22:57:12.732448Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_preds = np.concatenate(y_preds, axis=0)","metadata":{"execution":{"iopub.status.busy":"2024-09-10T22:57:12.734560Z","iopub.execute_input":"2024-09-10T22:57:12.734842Z","iopub.status.idle":"2024-09-10T22:57:12.739008Z","shell.execute_reply.started":"2024-09-10T22:57:12.734818Z","shell.execute_reply":"2024-09-10T22:57:12.738099Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/sample_submission.csv\")\nLABELS = list(sub_df.columns[1:])\nLABELS","metadata":{"execution":{"iopub.status.busy":"2024-09-10T22:57:12.740199Z","iopub.execute_input":"2024-09-10T22:57:12.740552Z","iopub.status.idle":"2024-09-10T22:57:12.757959Z","shell.execute_reply.started":"2024-09-10T22:57:12.740520Z","shell.execute_reply":"2024-09-10T22:57:12.757098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"row_names","metadata":{"execution":{"iopub.status.busy":"2024-09-10T22:57:12.764797Z","iopub.execute_input":"2024-09-10T22:57:12.765080Z","iopub.status.idle":"2024-09-10T22:57:12.771208Z","shell.execute_reply.started":"2024-09-10T22:57:12.765055Z","shell.execute_reply":"2024-09-10T22:57:12.770376Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_preds","metadata":{"execution":{"iopub.status.busy":"2024-09-10T22:57:12.772448Z","iopub.execute_input":"2024-09-10T22:57:12.773042Z","iopub.status.idle":"2024-09-10T22:57:12.785905Z","shell.execute_reply.started":"2024-09-10T22:57:12.773003Z","shell.execute_reply":"2024-09-10T22:57:12.785011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.DataFrame()\nsubmission['row_id'] = row_names\nsubmission[LABELS] = y_preds\nsubmission.head(25)","metadata":{"execution":{"iopub.status.busy":"2024-09-10T22:57:12.787372Z","iopub.execute_input":"2024-09-10T22:57:12.787774Z","iopub.status.idle":"2024-09-10T22:57:12.809770Z","shell.execute_reply.started":"2024-09-10T22:57:12.787744Z","shell.execute_reply":"2024-09-10T22:57:12.808881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv', index=False)\npd.read_csv('submission.csv').head()","metadata":{"execution":{"iopub.status.busy":"2024-09-10T22:57:12.810843Z","iopub.execute_input":"2024-09-10T22:57:12.811097Z","iopub.status.idle":"2024-09-10T22:57:12.831085Z","shell.execute_reply.started":"2024-09-10T22:57:12.811075Z","shell.execute_reply":"2024-09-10T22:57:12.830253Z"},"trusted":true},"execution_count":null,"outputs":[]}]}