{"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":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":9054047,"sourceType":"datasetVersion","datasetId":5459348},{"sourceId":9136541,"sourceType":"datasetVersion","datasetId":5517612},{"sourceId":9140723,"sourceType":"datasetVersion","datasetId":5520621},{"sourceId":9173318,"sourceType":"datasetVersion","datasetId":5543734},{"sourceId":9470910,"sourceType":"datasetVersion","datasetId":5759498},{"sourceId":9472132,"sourceType":"datasetVersion","datasetId":5760411},{"sourceId":9474042,"sourceType":"datasetVersion","datasetId":5761786},{"sourceId":9478717,"sourceType":"datasetVersion","datasetId":5765285},{"sourceId":9481714,"sourceType":"datasetVersion","datasetId":5767517},{"sourceId":9481718,"sourceType":"datasetVersion","datasetId":5767520},{"sourceId":9495261,"sourceType":"datasetVersion","datasetId":5777871},{"sourceId":9495269,"sourceType":"datasetVersion","datasetId":5777877},{"sourceId":9534662,"sourceType":"datasetVersion","datasetId":5807130},{"sourceId":9562820,"sourceType":"datasetVersion","datasetId":5827761},{"sourceId":9562824,"sourceType":"datasetVersion","datasetId":5827765},{"sourceId":9562827,"sourceType":"datasetVersion","datasetId":5827767}],"dockerImageVersionId":30733,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"%pip -q install --no-deps /kaggle/input/segmentation-models-pytorch-0-3-3/efficientnet_pytorch-0.7.1-py3-none-any.whl\n%pip -q install --no-deps /kaggle/input/segmentation-models-pytorch-0-3-3/munch-4.0.0-py2.py3-none-any.whl\n%pip -q install --no-deps /kaggle/input/segmentation-models-pytorch-0-3-3/pretrainedmodels-0.7.4-py3-none-any.whl\n%pip -q install --no-deps /kaggle/input/segmentation-models-pytorch-0-3-3/segmentation_models_pytorch-0.3.3-py3-none-any.whl\n%pip -q install --no-deps /kaggle/input/segmentation-models-pytorch-0-3-3/timm-0.9.2-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2024-10-06T21:37:19.113577Z","iopub.execute_input":"2024-10-06T21:37:19.113933Z","iopub.status.idle":"2024-10-06T21:39:11.297688Z","shell.execute_reply.started":"2024-10-06T21:37:19.113903Z","shell.execute_reply":"2024-10-06T21:39:11.296633Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport random\nimport pandas as pd\nimport numpy as np\nimport pydicom\nimport cv2\nimport re\nimport timm\nimport json\nimport logging\nimport warnings\n\nfrom albumentations.pytorch import ToTensorV2\n\nimport torch\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn.functional as F\n\nimport segmentation_models_pytorch as smp\n\nlogging.getLogger('timm').setLevel(logging.WARNING)\n\n# Params\ndebug = False\ndebug_n = 100\n\ncoord_model_names = {\n    'sagt2': 'rsna-2024-glad-moon-593',\n    'sagt1': 'rsna-2024-leafy-cherry-654',\n    'axi': 'rsna-2024-scarlet-feather-603',\n}\n\nmodel_names = {\n    'spinal': 'rsna-2024-giddy-monkey-1266',\n    'foraminal': 'rsna-2024-hardy-voice-1244',\n    'subarticular': 'rsna-2024-fiery-meadow-1254',\n    'global': 'rsna-2024-dashing-spaceship-1252',\n    \n    'spinal_2': 'rsna-2024-leafy-river-1268',\n    'foraminal_2': 'rsna-2024-snowy-oath-1251',\n    'subarticular_2': 'rsna-2024-hearty-spaceship-1256',\n    'global_2': 'rsna-2024-cool-frost-1378',\n    \n    'spinal_3': 'rsna-2024-splendid-glade-1421',\n    'foraminal_3': 'rsna-2024-blooming-gorge-1250',\n    'subarticular_3': 'rsna-2024-smooth-resonance-1422',\n    'global_3': 'rsna-2024-radiant-tree-1423',\n}\n\n# Paths\ninput_dir = '/kaggle/input'\ndata_dir = 'rsna-2024-lumbar-spine-degenerative-classification'\nout_dir = '/kaggle/working'\n\nlevels = ['l1_l2', 'l2_l3', 'l3_l4', 'l4_l5', 'l5_s1']\nconditions = ['spinal_canal_stenosis', \n              'left_neural_foraminal_narrowing', 'right_neural_foraminal_narrowing', \n              'left_subarticular_stenosis', 'right_subarticular_stenosis']\nsides = ['left', 'right']\n\n# Functions\ndef seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n\n# Seed\nseed = 42\nseed_everything(seed)\ndevice = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')\n\n# Load descriptive data\nif debug:\n    # Use train data\n    test_series_filename = 'train_series_descriptions.csv'\n    img_dirname = 'train_images'\n    img_dir = os.path.join(input_dir, data_dir, img_dirname)\n    df = pd.read_csv(os.path.join(input_dir, data_dir, 'train.csv'), dtype={'study_id': 'str'}).drop_duplicates('study_id').sample(n=debug_n, random_state=seed)\n    df_series = pd.read_csv(os.path.join(input_dir, data_dir, test_series_filename), dtype={'study_id': 'str', 'series_id': 'str'})\n    phase = 'faketest'\nelse:\n    test_series_filename = 'test_series_descriptions.csv'\n    img_dirname = 'test_images'\n    img_dir = os.path.join(input_dir, data_dir, img_dirname)\n    df = pd.read_csv(os.path.join(input_dir, data_dir, test_series_filename), dtype={'study_id': 'str'}).drop_duplicates('study_id')[['study_id']]\n    for cond in conditions:\n        for level in levels:\n            df[f'{cond}_{level}'] = np.nan\n    df_series = pd.read_csv(os.path.join(input_dir, data_dir, test_series_filename), dtype={'study_id': 'str', 'series_id': 'str'})\n    phase = 'test'\ndf.head()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-10-06T21:39:11.299754Z","iopub.execute_input":"2024-10-06T21:39:11.300055Z","iopub.status.idle":"2024-10-06T21:39:19.896301Z","shell.execute_reply.started":"2024-10-06T21:39:11.300027Z","shell.execute_reply":"2024-10-06T21:39:19.895270Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Utils","metadata":{}},{"cell_type":"code","source":"def natural_sort(l):\n    convert = lambda text: int(text) if text.isdigit() else text.lower()\n    alphanum_key = lambda key: [convert(c) for c in re.split('([0-9]+)', key)]\n    return sorted(l, key=alphanum_key)\n\n\ndef get_series(df_series, study_id, series_description):\n    series_list = df_series[\n        (df_series['study_id'] == study_id)\n        & (df_series['series_description'] == series_description)\n    ]['series_id'].tolist()\n    if len(series_list) == 0:\n        return None\n    return series_list\n\n\ndef sagi_coord_to_axi_instance_number(sag_x_norm, sag_y_norm, mid_sag_ds, axi_ds_list):\n    # Calculate sagittal world coordinates based on middle slice\n    sag_affine = dcm_affine(mid_sag_ds)\n    sag_coord = np.array([sag_y_norm * mid_sag_ds.Rows, sag_x_norm * mid_sag_ds.Columns, 0, 1])\n    sag_world_coord = (sag_affine @ sag_coord)[:-1]\n\n    # Get closest axial slice\n    dist_list = []\n    for ds in axi_ds_list:\n        normal = np.cross(ds.ImageOrientationPatient[:3], ds.ImageOrientationPatient[3:])\n        normal /= np.linalg.norm(normal)\n        dist = np.abs(np.dot(sag_world_coord - ds.ImagePositionPatient, normal))\n        dist_list.append(dist)\n    axi_slice_idx = np.argmin(dist_list)\n    min_dist = dist_list[axi_slice_idx]\n    if min_dist > 5:\n        return None, np.nan\n    axi_series_id, axi_instance_number = axi_ds_list[axi_slice_idx].filename.split('/')[-2:]\n\n    return axi_series_id, int(axi_instance_number.replace('.dcm', ''))\n\n\ndef dcm_affine(ds):\n    F11, F21, F31 = ds.ImageOrientationPatient[3:]\n    F12, F22, F32 = ds.ImageOrientationPatient[:3]\n    dr, dc = ds.PixelSpacing\n    Sx, Sy, Sz = ds.ImagePositionPatient\n\n    return np.array(\n        [\n            [F11 * dr, F12 * dc, 0, Sx],\n            [F21 * dr, F22 * dc, 0, Sy],\n            [F31 * dr, F32 * dc, 0, Sz],\n            [0, 0, 0, 1],\n        ]\n    )\n\n### Kaggle specific utils ###\ndef load_config(config_path):\n    with open(config_path) as f:\n        return json.load(f)\n    \ndef get_model_params(model_name, fold_num=5):\n    cfg = load_config(os.path.join(input_dir, model_name + '-5fold', 'config.json'))\n    state_paths = [\n        os.path.join(input_dir, model_name + '-5fold', model_name + f'-cv{i}_best.pt')\n        for i in range(1, fold_num + 1)\n    ]\n    return cfg, state_paths\n\ndef get_instance(name, cfg, *args, **kwargs):        \n    return globals()[cfg[name]['type']](*args, **kwargs, **cfg[name]['args'])\n\ndef get_sagi2axi_data(study_id, series_id):\n    if pd.isnull(series_id):\n        return None, None\n    sag_dir = os.path.join(img_dir, str(study_id), str(series_id))\n    axi_series = get_series(df_series, study_id, 'Axial T2')\n    if axi_series is None:\n        return None, None\n    axi_dir_list = [\n        os.path.join(img_dir, str(study_id), str(axi_series_id))\n        for axi_series_id in axi_series\n    ]\n    if os.path.isdir(sag_dir) and len(axi_dir_list) != 0:\n        sag_file_list = natural_sort(os.listdir(sag_dir))\n        mid_sag_ds = pydicom.dcmread(os.path.join(sag_dir, sag_file_list[len(sag_file_list) // 2]))\n\n        axi_ds_list = []\n        for axi_dir in axi_dir_list:\n            axi_file_list = natural_sort(os.listdir(axi_dir))\n            axi_ds_list += [pydicom.dcmread(os.path.join(axi_dir, file)) for file in axi_file_list]\n        return mid_sag_ds, axi_ds_list\n    else:\n        return None, None\n\ndef get_coord_from_heatmap_pred(pred, idx):\n    predi = pred[0, idx, ...].squeeze()\n    y_coord, x_coord = np.unravel_index(predi.argmax(), predi.shape)\n    x_norm = x_coord / predi.shape[1]\n    y_norm = y_coord / predi.shape[0]\n    return x_norm, y_norm","metadata":{"execution":{"iopub.status.busy":"2024-10-06T21:39:19.897919Z","iopub.execute_input":"2024-10-06T21:39:19.898205Z","iopub.status.idle":"2024-10-06T21:39:19.919417Z","shell.execute_reply.started":"2024-10-06T21:39:19.898169Z","shell.execute_reply":"2024-10-06T21:39:19.918312Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Datasets","metadata":{}},{"cell_type":"code","source":"class DatasetBase(Dataset):\n    def __init__(\n        self,\n        df,\n        root_dir,\n        data_dir,\n        img_num,\n        resolution,\n        phase='train',\n        df_coordinates=None,\n        cleaning_rule=None,\n        coord_impute_file=None,\n        resample_slice_spacing=None,\n        interpolation='INTER_CUBIC',\n        standardize=True,\n        rand_instance_number_offsets=None,\n        transform=None,\n    ):\n        self.phase = phase\n        if self.phase in ['train', 'valid', 'predict', 'faketest', 'valid_check']:\n            self.img_subdir = 'train_images'\n            self.series_filename = 'train_series_descriptions.csv'\n        elif self.phase == 'test':\n            self.img_subdir = 'test_images'\n            self.series_filename = 'test_series_descriptions.csv'\n        if self.phase == 'faketest':\n            self.phase = 'test'\n\n        self.df = df\n        self.data_dir = os.path.join(root_dir, data_dir)\n        self.df_series = self.load_series_info()\n        self.coord_impute_file = coord_impute_file\n        if df_coordinates is None and self.phase != 'test':\n            self.df_coordinates = self.load_coordinates_info()\n        else:\n            self.df_coordinates = df_coordinates\n        if self.df_coordinates is not None:\n            self.df_coordinates = self.df_coordinates.merge(\n                self.df_series, how='left', on=['study_id', 'series_id']\n            )\n        self.img_dir = os.path.join(self.data_dir, self.img_subdir)\n        self.transform = transform\n        self.img_num = img_num\n        self.resolution = resolution\n        self.cleaning_rule = cleaning_rule\n        self.resample_slice_spacing = resample_slice_spacing\n        self.interpolation = interpolation\n        self.standardize = standardize\n        self.rand_instance_number_offsets = rand_instance_number_offsets\n\n        self.levels = ['l1_l2', 'l2_l3', 'l3_l4', 'l4_l5', 'l5_s1']\n        self.sides = ['left', 'right']\n\n        self.logger = logging.getLogger(__name__)\n        self.logger.setLevel(logging.ERROR)\n        self.handler = logging.FileHandler('level_roi_dataset.log')\n        self.logger.addHandler(self.handler)\n\n    def __len__(self):\n        return len(self.df)\n\n    def load_series_info(self):\n        return pd.read_csv(\n            os.path.join(self.data_dir, self.series_filename),\n            dtype={'study_id': 'str', 'series_id': 'str'},\n        )\n\n    def load_coordinates_info(self):\n        df_coordinates = pd.read_csv(\n            os.path.join(self.data_dir, '..', 'processed', 'train_label_coordinates.csv'),\n            dtype={'study_id': 'str', 'series_id': 'str'},\n        )\n        if self.coord_impute_file is not None:\n            df_coordinates_pred = pd.read_csv(\n                os.path.join(self.data_dir, '..', 'processed', self.coord_impute_file + '.csv'),\n                dtype={'study_id': 'str', 'series_id': 'str'},\n            )\n            bad_coordinates = pd.read_csv(\n                os.path.join(\n                    self.data_dir, '..', 'processed', 'bad_coords_spinal_canal_stenosis.csv'\n                ),\n                dtype={'study_id': 'str', 'series_id': 'str'},\n            )\n\n            # Remove bad coords\n            df_coordinates.set_index(['study_id', 'series_id', 'row_id'], inplace=True)\n            bad_coordinates.set_index(['study_id', 'series_id', 'row_id'], inplace=True)\n            df_coordinates = df_coordinates.drop(\n                df_coordinates.index.intersection(bad_coordinates.index)\n            ).reset_index()\n\n            # Impute with predicted coords\n            df_coordinates.set_index(['study_id', 'row_id'], inplace=True)\n            df_coordinates_pred.set_index(['study_id', 'row_id'], inplace=True)\n            df_coordinates = pd.concat(\n                [\n                    df_coordinates,\n                    df_coordinates_pred.drop(df_coordinates.index).dropna(subset='series_id'),\n                ],\n                axis=0,\n            ).reset_index()\n\n        return df_coordinates\n\n    def get_series(self, study_id, series_description):\n        series_list = self.df_series[\n            (self.df_series['study_id'] == study_id)\n            & (self.df_series['series_description'] == series_description)\n        ]['series_id'].tolist()\n\n        # Try to substitute T1 for T2 or vice versa\n        if len(series_list) == 0:\n            if series_description == 'Sagittal T2/STIR':\n                series_description = 'Sagittal T1'\n            elif series_description == 'Sagittal T1':\n                series_description = 'Sagittal T2/STIR'\n            series_list = self.df_series[\n                (self.df_series['study_id'] == study_id)\n                & (self.df_series['series_description'] == series_description)\n            ]['series_id'].tolist()\n\n        if len(series_list) == 0:\n            self.logger.warning('%s %s not found', study_id, series_description)\n            return None\n\n        if len(series_list) > 1:\n            self.logger.warning('%s %s multiple found', study_id, series_description)\n\n        return series_list\n\n    def __getitem__(self, idx):\n        raise NotImplementedError('Please implement this method in a subclass')\n\n\nclass CoordDataset(DatasetBase):\n    def __init__(\n        self,\n        df,\n        root_dir,\n        data_dir,\n        img_num,\n        resolution,\n        heatmap_std,\n        phase='train',\n        df_coordinates=None,\n        cleaning_rule=None,\n        coord_impute_file=None,\n        rand_instance_number_offsets=None,\n        transform=None,\n    ):\n        super().__init__(\n            df=df,\n            root_dir=root_dir,\n            data_dir=data_dir,\n            img_num=img_num,\n            resolution=resolution,\n            phase=phase,\n            df_coordinates=df_coordinates,\n            cleaning_rule=cleaning_rule,\n            coord_impute_file=coord_impute_file,\n            rand_instance_number_offsets=rand_instance_number_offsets,\n            transform=transform,\n        )\n        self.heatmap_std = heatmap_std\n\n    def gaussian_heatmap(self, width, height, center, std_dev):\n        \"\"\"\n        Args:\n        - width (int): Width of the heatmap\n        - height (int): Height of the heatmap\n        - center (tuple): The (x, y) coordinates of the Gaussian peak\n        - std_dev (int, optional): Standard deviation of the Gaussian\n\n        \"\"\"\n        x_axis = torch.arange(width).float() - center[0]\n        y_axis = torch.arange(height).float() - center[1]\n        x, y = torch.meshgrid(y_axis, x_axis, indexing='ij')\n\n        return torch.exp(-((x**2 + y**2) / (2 * std_dev**2)))\n\n    def create_heatmaps(self, coords):\n        heatmaps = []\n        for i in range(coords.shape[0]):\n            if np.isnan(coords[i, :]).any():\n                heatmaps.append(torch.zeros(self.resolution, self.resolution))\n            else:\n                heatmaps.append(\n                    self.gaussian_heatmap(\n                        self.resolution, self.resolution, coords[i, :], std_dev=self.heatmap_std\n                    )\n                )\n        return torch.stack(heatmaps, dim=-1).numpy()\n\n    def most_frequent(self, l):\n        l = [l_ for l_ in l if not pd.isnull(l_)]\n        if len(l) == 0:\n            return None\n        return max(l, key=l.count)\n\n    def get_series_with_coords(self, study_id, series_description, level=None, side=None):\n        # Get expected vars in standard order\n        if series_description == 'Sagittal T2/STIR':\n            expected_vars = [f'spinal_canal_stenosis_{lvl}' for lvl in self.levels]\n        elif series_description == 'Sagittal T1':\n            expected_vars = [f'{side}_neural_foraminal_narrowing_{lvl}' for lvl in self.levels]\n        elif series_description == 'Axial T2':\n            expected_vars = [f'{side}_subarticular_stenosis_{level}' for side in self.sides]\n\n        # Get coords in standard order, padded with nans\n        series_coords = self.df_coordinates[\n            (self.df_coordinates['study_id'] == study_id)\n            & (self.df_coordinates['series_description'] == series_description)\n        ]\n        series_coords = series_coords.merge(\n            pd.DataFrame({'row_id': expected_vars}), on='row_id', how='right'\n        )\n        coords = np.array(\n            [series_coords['x_norm'].values, series_coords['y_norm'].values], dtype=np.float32\n        ).T\n\n        series_list = series_coords['series_id'].unique().tolist()\n        series_id = self.most_frequent(series_list)\n\n        if pd.isnull(series_id):\n            instance_number = np.nan\n        else:\n            instance_number = self.most_frequent(\n                series_coords[series_coords['series_id'] == series_id][\n                    'instance_number'\n                ].values.tolist()\n            )\n            if pd.isnull(instance_number):\n                instance_number = np.nan\n\n        return series_id, coords, instance_number\n\n    def get_image(\n        self,\n        study_id,\n        series_id,\n        instance_number_type='middle',\n        instance_number=None,\n        interpolation=cv2.INTER_CUBIC,\n        standardize=True,\n    ):\n        # Zero image\n        x = np.zeros((self.resolution, self.resolution, self.img_num), dtype=np.float32)\n        if pd.isnull(series_id) or (instance_number_type != 'middle' and pd.isnull(instance_number)):\n            return x\n        \n        # Get all dcm files in series\n        series_dir = os.path.join(self.img_dir, str(study_id), str(series_id))\n        file_list = natural_sort(os.listdir(series_dir))\n        slice_num = len(file_list)\n        if slice_num == 0:\n            return x\n\n        # Fix direction\n        ds_first = pydicom.dcmread(os.path.join(series_dir, file_list[0]))\n        ds_last = pydicom.dcmread(os.path.join(series_dir, file_list[-1]))\n        pos_diff = np.array(ds_last.ImagePositionPatient) - np.array(ds_first.ImagePositionPatient)\n        pos_diff = pos_diff[np.abs(pos_diff).argmax()]\n        if pos_diff < 0:\n            file_list.reverse()\n            if instance_number_type == 'index':\n                instance_number = len(file_list) - instance_number - 1\n\n        if instance_number_type == 'middle':\n            start_index = (slice_num - self.img_num) // 2\n        elif instance_number_type == 'index':\n            start_index = instance_number - self.img_num // 2\n        elif instance_number_type == 'filename':\n            start_index = [int(file.rstrip('.dcm')) for file in file_list].index(\n                instance_number\n            ) - self.img_num // 2\n        elif instance_number_type == 'relative':\n            start_index = int(instance_number * slice_num) - self.img_num // 2\n        elif instance_number_type == 'centered_mm':\n            try:\n                start_index = int(\n                    np.ceil(\n                        (instance_number / float(ds_first.SpacingBetweenSlices) + slice_num / 2 - 1)\n                        - self.img_num / 2\n                    )\n                )\n            except Exception as e:\n                start_index = (slice_num - self.img_num) // 2\n        elif instance_number_type == 'centered_mm_old':\n            try:\n                start_index = (\n                    slice_num // 2\n                    + round(instance_number / float(ds_first.SpacingBetweenSlices))\n                    - self.img_num // 2\n                )\n            except Exception as e:\n                start_index = slice_num // 2 - self.img_num // 2\n                \n        # Augment instance number during training\n        if self.rand_instance_number_offsets is not None and self.phase in ['train']:\n            start_index += random.sample(self.rand_instance_number_offsets, 1)[0]\n        \n        start_index = min(max(start_index, 0), slice_num)\n        end_index = min(start_index + self.img_num, slice_num)\n        file_list = file_list[start_index:end_index]\n\n        for i, filename in enumerate(file_list):\n            ds = pydicom.dcmread(os.path.join(series_dir, filename))\n            img = ds.pixel_array.astype(np.float32)\n\n            # Resize\n            img = cv2.resize(img, (self.resolution, self.resolution), interpolation=interpolation)\n            x[..., i] = img\n\n        # Standardize image\n        if standardize and x.std() != 0:\n            x = (x - x.mean()) / x.std()\n\n        return x\n\n\nclass Sagt2CoordDataset(CoordDataset):\n    def __init__(self, *args, **kwargs):\n        super().__init__(*args, **kwargs)\n\n        # Clean data\n        if self.phase in ['train', 'valid', 'valid_check']:\n            if self.cleaning_rule == 'keep_only_complete':\n                coord_counts = (\n                    self.df_coordinates[\n                        self.df_coordinates['row_id'].str.contains('spinal_canal_stenosis')\n                    ]\n                    .groupby('study_id')\n                    .count()['series_id']\n                )\n                good_study_ids = coord_counts[coord_counts == 5].index.astype(str).tolist()\n                self.df = self.df[self.df['study_id'].isin(good_study_ids)]\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n\n        if self.phase in ['train', 'valid', 'predict', 'valid_check']:\n            series_id, coords, _ = self.get_series_with_coords(row.study_id, 'Sagittal T2/STIR')\n            img = self.get_image(row.study_id, series_id, instance_number_type='middle')\n\n            heatmaps = self.create_heatmaps(coords * self.resolution)\n            if self.transform:\n                t = self.transform(image=img, mask=heatmaps)\n                img, heatmaps = t['image'], t['mask']\n            heatmaps = np.transpose(heatmaps, (2, 0, 1))\n\n            if self.phase in ['predict', 'valid_check']:\n                if pd.isnull(series_id):\n                    series_id = ''\n                return img, heatmaps, row.study_id, series_id\n\n            return img, heatmaps\n\n        elif self.phase in ['test']:\n            series_id = self.get_series(row.study_id, 'Sagittal T2/STIR')\n            if series_id is not None:\n                series_id = series_id[0]  # get the first series_id\n            img = self.get_image(row.study_id, series_id, instance_number_type='middle')\n            if self.transform:\n                img = self.transform(image=img)['image']\n            if pd.isnull(series_id):\n                series_id = ''\n            return img, row.study_id, series_id\n\n\nclass Sagt1CoordDataset(CoordDataset):\n    def __init__(self, *args, **kwargs):\n        super().__init__(*args, **kwargs)\n\n        # Split for sides\n        def side_at_end(colname):\n            splits = colname.split('_')\n            if splits[0] in self.sides:\n                return '_'.join(splits[1:] + [splits[0]])\n            else:\n                return colname\n\n        self.df = self.df.rename(columns=side_at_end)\n\n        self.df = pd.wide_to_long(\n            self.df,\n            [\n                f'{cond}_{level}'\n                for cond in ['neural_foraminal_narrowing', 'subarticular_stenosis']\n                for level in self.levels\n            ],\n            i='study_id',\n            j='side',\n            sep='_',\n            suffix=r'\\w+',\n        ).reset_index()\n\n        # Clean data\n        if self.phase in ['train', 'valid', 'valid_check']:\n            if self.cleaning_rule == 'keep_only_complete':\n                coord_counts = (\n                    self.df_coordinates[\n                        self.df_coordinates['row_id'].str.contains('neural_foraminal_narrowing')\n                    ]\n                    .groupby('study_id')\n                    .count()['series_id']\n                )\n                good_study_ids = coord_counts[coord_counts == 10].index.astype(str).tolist()\n                self.df = self.df[self.df['study_id'].isin(good_study_ids)]\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n\n        # Get series_id (and coords)\n        if self.phase in ['train', 'valid', 'predict', 'valid_check']:\n            series_id, coords, _ = self.get_series_with_coords(\n                row.study_id, 'Sagittal T1', side=row.side\n            )\n        elif self.phase in ['test']:\n            series_id = self.get_series(row.study_id, 'Sagittal T1')\n            if series_id is not None:\n                series_id = series_id[0]  # get the first series_id\n\n        # Get image\n        instance_number_type = 'centered_mm'\n        instance_number = 16.8\n        if row.side == 'right':\n            instance_number = -1 * instance_number\n\n        img = self.get_image(\n            row.study_id,\n            series_id,\n            instance_number_type=instance_number_type,\n            instance_number=instance_number,\n        )\n\n        # Return data (and heatmaps)\n        if self.phase in ['train', 'valid', 'predict', 'valid_check']:\n            heatmaps = self.create_heatmaps(coords * self.resolution)\n            if self.transform:\n                t = self.transform(image=img, mask=heatmaps)\n                img, heatmaps = t['image'], t['mask']\n            heatmaps = np.transpose(heatmaps, (2, 0, 1))\n\n            if self.phase in ['predict', 'valid_check']:\n                if pd.isnull(series_id):\n                    series_id = ''\n                return img, heatmaps, row.study_id, series_id, row.side\n\n            return img, heatmaps\n\n        elif self.phase in ['test']:\n\n            if self.transform:\n                img = self.transform(image=img)['image']\n            if pd.isnull(series_id):\n                series_id = ''\n            return img, row.study_id, series_id, row.side\n\n\nclass AxiCoordDataset(CoordDataset):\n    def __init__(self, *args, **kwargs):\n        super().__init__(*args, **kwargs)\n\n        # Need to split for levels\n        self.df = pd.wide_to_long(\n            self.df,\n            [\n                'spinal_canal_stenosis',\n                'left_neural_foraminal_narrowing',\n                'right_neural_foraminal_narrowing',\n                'left_subarticular_stenosis',\n                'right_subarticular_stenosis',\n            ],\n            i='study_id',\n            j='level',\n            sep='_',\n            suffix=r'\\w+',\n        ).reset_index()\n\n        # Clean data\n        if self.phase in ['train', 'valid', 'valid_check']:\n            if self.cleaning_rule == 'keep_only_complete':\n                coord_counts = (\n                    self.df_coordinates[\n                        self.df_coordinates['row_id'].str.contains('subarticular_stenosis')\n                    ]\n                    .groupby('study_id')\n                    .count()['series_id']\n                )\n                good_study_ids = coord_counts[coord_counts == 10].index.astype(str).tolist()\n                self.df = self.df[self.df['study_id'].isin(good_study_ids)]\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n\n        series_id, coords, instance_number = self.get_series_with_coords(\n            row.study_id, 'Axial T2', level=row.level\n        )\n\n        img = self.get_image(\n            row.study_id,\n            series_id,\n            instance_number_type='filename',\n            instance_number=instance_number,\n        )\n\n        if self.phase in ['train', 'valid', 'predict', 'valid_check']:\n            heatmaps = self.create_heatmaps(coords * self.resolution)\n            if self.transform:\n                t = self.transform(image=img, mask=heatmaps)\n                img, heatmaps = t['image'], t['mask']\n            heatmaps = np.transpose(heatmaps, (2, 0, 1))\n\n            if self.phase in ['predict', 'valid_check']:\n                if pd.isnull(series_id):\n                    series_id = ''\n                return img, heatmaps, row.study_id, series_id, row.level\n\n            return img, heatmaps\n\n        elif self.phase in ['test']:\n            if self.transform:\n                img = self.transform(image=img)['image']\n            if pd.isnull(series_id):\n                series_id = ''\n            return img, row.study_id, series_id, row.level\n\n\nclass ROIDataset(DatasetBase):\n    def __init__(\n        self,\n        df,\n        root_dir,\n        data_dir,\n        img_num,\n        resolution,\n        roi_size,\n        phase='train',\n        df_coordinates=None,\n        interpolation='INTER_CUBIC',\n        standardize=True,\n        cleaning_rule=None,\n        coord_impute_file=None,\n        rand_instance_number_offsets=None,\n        resample_slice_spacing=None,\n        transform=None,\n    ):\n        super().__init__(\n            df=df,\n            root_dir=root_dir,\n            data_dir=data_dir,\n            img_num=img_num,\n            resolution=resolution,\n            phase=phase,\n            df_coordinates=df_coordinates,\n            cleaning_rule=cleaning_rule,\n            coord_impute_file=coord_impute_file,\n            rand_instance_number_offsets=rand_instance_number_offsets,\n            resample_slice_spacing=resample_slice_spacing,\n            transform=transform,\n        )\n            \n        self.df_orig = df\n        self.df = pd.wide_to_long(\n            df,\n            [\n                'spinal_canal_stenosis',\n                'left_neural_foraminal_narrowing',\n                'right_neural_foraminal_narrowing',\n                'left_subarticular_stenosis',\n                'right_subarticular_stenosis',\n            ],\n            i='study_id',\n            j='level',\n            sep='_',\n            suffix=r'\\w+',\n        ).reset_index()\n        self.roi_size = roi_size\n\n\n    def get_sagt2_coord(self, study_id, level):\n        coords = self.df_coordinates[\n            (self.df_coordinates['study_id'] == study_id)\n            & (self.df_coordinates['row_id'] == f'spinal_canal_stenosis_{level}')\n        ][['x_norm', 'y_norm', 'instance_number', 'series_id']].values\n\n        if coords.shape[0] == 0:\n            self.logger.warning('SAGT2 %s %s no coordinates found', study_id, level)\n            return (np.nan, np.nan, np.nan, None)\n\n        if coords.shape[0] > 1:\n            self.logger.warning('SAGT2 %s %s multiple coordinates found', study_id, level)\n\n        x_norm, y_norm, instance_number, series_id = coords[0]\n        if not pd.isnull(instance_number):\n            instance_number = int(instance_number)\n        return x_norm, y_norm, instance_number, series_id\n\n    def get_sagt1_coord(self, study_id, level, side):\n        coords = self.df_coordinates[\n            (self.df_coordinates['study_id'] == study_id)\n            & (self.df_coordinates['row_id'] == f'{side}_neural_foraminal_narrowing_{level}')\n        ][['x_norm', 'y_norm', 'instance_number', 'series_id']].values\n\n        if coords.shape[0] == 0:\n            self.logger.warning('SAGT1 %s %s %s no coordinates found', study_id, level, side)\n            return np.nan, np.nan, np.nan, None\n\n        if coords.shape[0] > 1:\n            self.logger.warning('SAGT1 %s %s %s multiple coordinates found', study_id, level, side)\n\n        x_norm, y_norm, instance_number, series_id = coords[0]\n        if not pd.isnull(instance_number):\n            instance_number = int(instance_number)\n        return x_norm, y_norm, instance_number, series_id\n\n    def get_axi_coord(self, study_id, level, side):\n        coords = self.df_coordinates[\n            (self.df_coordinates['study_id'] == study_id)\n            & (self.df_coordinates['row_id'] == f'{side}_subarticular_stenosis_{level}')\n        ][['x_norm', 'y_norm', 'instance_number', 'series_id']].values\n\n        if coords.shape[0] == 0:\n            self.logger.warning('AXI %s %s %s no coordinates found', study_id, level, side)\n            return np.nan, np.nan, np.nan, None\n\n        if coords.shape[0] > 1:\n            self.logger.warning('AXI %s %s %s multiple coordinates found', study_id, level, side)\n\n        x_norm, y_norm, instance_number, series_id = coords[0]\n        if not pd.isnull(instance_number):\n            instance_number = int(instance_number)\n        return x_norm, y_norm, instance_number, series_id\n\n    def get_label(self, idx, label_name):\n        label = self.df[label_name].iloc[idx]\n        if isinstance(label_name, list):\n            label = label.astype(float).values\n        else:\n            label = float(label)\n        label = np.nan_to_num(label, nan=-100).astype(np.int64)\n        return label\n\n    def get_level_onehot(self, level_name):\n        level = np.zeros(len(self.levels), dtype=np.float32)\n        level[self.levels.index(level_name)] = 1.0\n        return level\n\n    def get_side_onehot(self, side_name):\n        side = np.zeros(len(self.sides), dtype=np.float32)\n        side[self.sides.index(side_name)] = 1.0\n        return side\n\n    def get_roi(\n        self,\n        study_id,\n        series_id,\n        img_num,\n        resolution,\n        roi_size,\n        x_norm=np.nan,\n        y_norm=np.nan,\n        instance_number_type='middle',\n        instance_number=None,\n        resample_slice_spacing=None,\n    ):\n        # Zero image\n        x = np.zeros((resolution, resolution, img_num), dtype=np.float32)\n        if pd.isnull(series_id) or pd.isnull(x_norm) or pd.isnull(y_norm) or (instance_number_type != 'middle' and pd.isnull(instance_number)):\n            return x\n        \n        # Get all dcm files in series\n        series_dir = os.path.join(self.img_dir, str(study_id), str(series_id))\n        file_list = natural_sort(os.listdir(series_dir))\n        slice_num = len(file_list)\n        if slice_num == 0:\n            return x\n\n        # Fix direction\n        ds_first = pydicom.dcmread(os.path.join(series_dir, file_list[0]))\n        ds_last = pydicom.dcmread(os.path.join(series_dir, file_list[-1]))\n        pos_diff = np.array(ds_last.ImagePositionPatient) - np.array(ds_first.ImagePositionPatient)\n        pos_diff = pos_diff[np.abs(pos_diff).argmax()]\n        if pos_diff < 0:\n            file_list.reverse()\n            if instance_number_type == 'index':\n                instance_number = len(file_list) - instance_number - 1\n\n        if resample_slice_spacing is not None:\n            try:\n                slice_spacing = float(ds_first.SpacingBetweenSlices)\n            except Exception as e:\n                slice_spacing = 4.5\n            resample_factor = resample_slice_spacing / slice_spacing\n            img_num_final = img_num\n            img_num = int(img_num * resample_factor)\n            resample_factor = img_num / img_num_final\n            x = np.zeros((resolution, resolution, img_num), dtype=np.float32)\n\n        if instance_number_type == 'middle':\n            start_index = (slice_num - img_num) // 2\n        elif instance_number_type == 'index':\n            start_index = instance_number - img_num // 2\n        elif instance_number_type == 'filename':\n            start_index = [int(file.rstrip('.dcm')) for file in file_list].index(\n                instance_number\n            ) - img_num // 2\n        elif instance_number_type == 'relative':\n            start_index = int(instance_number * slice_num) - img_num // 2\n        elif instance_number_type == 'centered_index':\n            start_index = slice_num // 2 + instance_number - img_num // 2\n        elif instance_number_type == 'centered_mm':\n            try:\n                start_index = int(\n                    np.ceil(\n                        (instance_number / float(ds_first.SpacingBetweenSlices) + slice_num / 2 - 1)\n                        - img_num / 2\n                    )\n                )\n            except Exception as e:\n                start_index = (slice_num - img_num) // 2\n        elif instance_number_type == 'centered_mm_old':\n            try:\n                start_index = (\n                    slice_num // 2\n                    + round(instance_number / float(ds_first.SpacingBetweenSlices))\n                    - img_num // 2\n                )\n            except Exception as e:\n                start_index = slice_num // 2 - img_num // 2\n\n        # Augment instance number during training\n        if self.rand_instance_number_offsets is not None and self.phase in ['train']:\n            start_index += random.sample(self.rand_instance_number_offsets, 1)[0]\n\n        start_index = min(max(start_index, 0), slice_num)\n        end_index = min(start_index + img_num, slice_num)\n        file_list = file_list[start_index:end_index]\n\n        interpolation = getattr(cv2, self.interpolation)\n        for i, filename in enumerate(file_list):\n            ds = pydicom.dcmread(os.path.join(series_dir, filename))\n            img = ds.pixel_array.astype(np.float32)\n\n            if i == 0:\n                x_norm = x_norm * img.shape[1]\n                y_norm = y_norm * img.shape[0]\n\n            # Crop ROI\n            size_x = roi_size / ds.PixelSpacing[1]\n            size_y = roi_size / ds.PixelSpacing[0]\n            x1 = round(x_norm - (size_x / 2))\n            x2 = round(x_norm + (size_x / 2))\n            y1 = round(y_norm - (size_y / 2))\n            y2 = round(y_norm + (size_y / 2))\n            if any([x1 < 0, x2 > img.shape[1], y1 < 0, y2 > img.shape[0]]):\n                self.logger.warning('%s %s ROI out of bounds', study_id, series_id)\n                break\n            img = img[y1:y2, x1:x2]\n\n            # Resize\n            img = cv2.resize(img, (resolution, resolution), interpolation=interpolation)\n            x[..., i] = img\n\n        if resample_slice_spacing is not None:\n            x = scipy.ndimage.zoom(x, (1, 1, 1 / resample_factor), order=3)\n\n        # Standardize image\n        if self.standardize and x.std() != 0:\n            x = (x - x.mean()) / x.std()\n\n        return x\n\n\nclass SpinalROIDataset(ROIDataset):\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n\n        # Sagittal T2/STIR ROI\n        sagt2_x_norm, sagt2_y_norm, _, sagt2_series_id = self.get_sagt2_coord(\n            row.study_id, row.level\n        )\n        sagt2_roi = self.get_roi(\n            study_id=row.study_id,\n            series_id=sagt2_series_id,\n            img_num=self.img_num[0],\n            resolution=self.resolution,\n            roi_size=self.roi_size,\n            x_norm=sagt2_x_norm,\n            y_norm=sagt2_y_norm,\n            instance_number_type='middle',\n            resample_slice_spacing=self.resample_slice_spacing,\n        )\n\n        # Axial T2 central ROI\n        left_axi_x_norm, left_axi_y_norm, left_axi_instance_number, left_axi_series_id = (\n            self.get_axi_coord(row.study_id, row.level, 'left')\n        )\n        right_axi_x_norm, right_axi_y_norm, right_axi_instance_number, right_axi_series_id = (\n            self.get_axi_coord(row.study_id, row.level, 'right')\n        )\n\n        with warnings.catch_warnings():\n            warnings.simplefilter('ignore', category=RuntimeWarning)\n            axi_x_norm = np.nanmean([left_axi_x_norm, right_axi_x_norm])\n            axi_y_norm = np.nanmean([left_axi_y_norm, right_axi_y_norm])\n        if not pd.isnull(left_axi_instance_number):\n            axi_instance_number = left_axi_instance_number\n            axi_series_id = left_axi_series_id\n        else:\n            axi_instance_number = right_axi_instance_number\n            axi_series_id = right_axi_series_id\n\n        axi_roi = self.get_roi(\n            study_id=row.study_id,\n            series_id=axi_series_id,\n            img_num=self.img_num[1],\n            resolution=self.resolution,\n            roi_size=self.roi_size,\n            x_norm=axi_x_norm,\n            y_norm=axi_y_norm,\n            instance_number_type='filename',\n            instance_number=axi_instance_number,\n            resample_slice_spacing=self.resample_slice_spacing,\n        )\n\n        if self.transform:\n            if isinstance(self.transform, list):\n                print('multiple transforms used')\n                sagt2_roi = self.transform[0](image=sagt2_roi)['image']\n                axi_roi = self.transform[1](image=axi_roi)['image']\n            else:\n                sagt2_roi = self.transform(image=sagt2_roi)['image']\n                axi_roi = self.transform(image=axi_roi)['image']\n\n        level = self.get_level_onehot(row.level)\n        if self.phase in ['train', 'valid', 'predict', 'valid_check']:\n            label = self.get_label(idx, 'spinal_canal_stenosis')\n            if self.phase == 'valid_check':\n                sagt2_series_id = '' if sagt2_series_id is None else sagt2_series_id\n                axi_series_id = '' if axi_series_id is None else axi_series_id\n                return (\n                    sagt2_roi,\n                    axi_roi,\n                    level,\n                    label,\n                    row.study_id,\n                    sagt2_series_id,\n                    axi_series_id,\n                    row.level,\n                )\n            return sagt2_roi, axi_roi, level, label\n        elif self.phase in ['test']:\n            return sagt2_roi, axi_roi, level, 0\n\n\nclass ForaminalROIDataset(ROIDataset):\n    def __init__(self, *args, **kwargs):\n        super().__init__(*args, **kwargs)\n\n        self.df.rename(\n            columns={\n                'left_neural_foraminal_narrowing': 'neural_foraminal_narrowing_left',\n                'right_neural_foraminal_narrowing': 'neural_foraminal_narrowing_right',\n                'left_subarticular_stenosis': 'subarticular_stenosis_left',\n                'right_subarticular_stenosis': 'subarticular_stenosis_right',\n            },\n            inplace=True,\n        )\n        self.df = pd.wide_to_long(\n            self.df,\n            ['neural_foraminal_narrowing', 'subarticular_stenosis'],\n            i=['study_id', 'level'],\n            j='side',\n            sep='_',\n            suffix=r'\\w+',\n        ).reset_index()\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n\n        # Sagittal T1 ROI\n        sagt1_x_norm, sagt1_y_norm, _, sagt1_series_id = self.get_sagt1_coord(\n            row.study_id, row.level, row.side\n        )\n\n        instance_number_type = 'centered_mm'\n        level_instance_number = {\n            'l1_l2': 14.9,\n            'l2_l3': 15.7,\n            'l3_l4': 16.7,\n            'l4_l5': 17.8,\n            'l5_s1': 18.9,\n        }\n        instance_number = level_instance_number[row.level]\n        if row.side == 'right':\n            instance_number = -1 * instance_number\n\n        sagt1_roi = self.get_roi(\n            study_id=row.study_id,\n            series_id=sagt1_series_id,\n            img_num=self.img_num,\n            resolution=self.resolution,\n            roi_size=self.roi_size,\n            x_norm=sagt1_x_norm,\n            y_norm=sagt1_y_norm,\n            instance_number_type=instance_number_type,\n            instance_number=instance_number,\n            resample_slice_spacing=self.resample_slice_spacing,\n        )\n\n        if self.transform:\n            sagt1_roi = self.transform(image=sagt1_roi)['image']\n\n        level = self.get_level_onehot(row.level)\n        side = self.get_side_onehot(row.side)\n        if self.phase in ['train', 'valid', 'predict', 'valid_check']:\n            label = self.get_label(idx, 'neural_foraminal_narrowing')\n            if self.phase == 'valid_check':\n                sagt1_series_id = '' if sagt1_series_id is None else sagt1_series_id\n                return sagt1_roi, level, side, label, row.study_id, sagt1_series_id, row.level\n            return sagt1_roi, level, side, label\n        elif self.phase in ['test']:\n            return sagt1_roi, level, side, 0\n\n\nclass SubarticularROIDataset(ForaminalROIDataset):\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n\n        # Axial T2 central ROI\n        axi_x_norm, axi_y_norm, axi_instance_number, axi_series_id = self.get_axi_coord(\n            row.study_id, row.level, row.side\n        )\n\n        axi_roi = self.get_roi(\n            study_id=row.study_id,\n            series_id=axi_series_id,\n            img_num=self.img_num,\n            resolution=self.resolution,\n            roi_size=self.roi_size,\n            x_norm=axi_x_norm,\n            y_norm=axi_y_norm,\n            instance_number_type='filename',\n            instance_number=axi_instance_number,\n            resample_slice_spacing=self.resample_slice_spacing,\n        )\n\n        # Flip if right side\n        if row.side == 'right':\n            axi_roi = np.flip(axi_roi, axis=1).copy()\n\n        if self.transform:\n            axi_roi = self.transform(image=axi_roi)['image']\n\n        level = self.get_level_onehot(row.level)\n        side = self.get_side_onehot(row.side)\n        if self.phase in ['train', 'valid', 'predict', 'valid_check']:\n            label = self.get_label(idx, 'subarticular_stenosis')\n            if self.phase == 'valid_check':\n                axi_series_id = '' if axi_series_id is None else axi_series_id\n                return axi_roi, level, side, label, row.study_id, axi_series_id, row.level, row.side\n            return axi_roi, level, side, label\n        elif self.phase == 'test':\n            return axi_roi, level, side, 0\n\n\nclass GlobalROIDataset(ROIDataset):\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n\n        # Sagittal T2/STIR ROI\n        sagt2_x_norm, sagt2_y_norm, _, sagt2_series_id = self.get_sagt2_coord(\n            row.study_id, row.level\n        )\n        sagt2_roi = self.get_roi(\n            study_id=row.study_id,\n            series_id=sagt2_series_id,\n            img_num=self.img_num[0],\n            resolution=self.resolution,\n            roi_size=self.roi_size,\n            x_norm=sagt2_x_norm,\n            y_norm=sagt2_y_norm,\n            instance_number_type='middle',\n            resample_slice_spacing=self.resample_slice_spacing,\n        )\n\n        for side in ['left', 'right']:\n            # Sagittal T1 ROI\n            sagt1_x_norm, sagt1_y_norm, _, sagt1_series_id = self.get_sagt1_coord(\n                row.study_id, row.level, side\n            )\n\n            sagt1_level_instance_number = {\n                'l1_l2': 14.9,\n                'l2_l3': 15.7,\n                'l3_l4': 16.7,\n                'l4_l5': 17.8,\n                'l5_s1': 18.9,\n            }\n            sagt1_instance_number = sagt1_level_instance_number[row.level]\n            if side == 'right':\n                sagt1_instance_number = -1 * sagt1_instance_number\n\n            sagt1_roi = self.get_roi(\n                study_id=row.study_id,\n                series_id=sagt1_series_id,\n                img_num=self.img_num[1],\n                resolution=self.resolution,\n                roi_size=self.roi_size,\n                x_norm=sagt1_x_norm,\n                y_norm=sagt1_y_norm,\n                instance_number_type='centered_mm',\n                instance_number=sagt1_instance_number,\n                resample_slice_spacing=self.resample_slice_spacing,\n            )\n            if side == 'left':\n                sagt1_left_roi = sagt1_roi\n            elif side == 'right':\n                sagt1_right_roi = sagt1_roi\n\n            # Axial T2 ROI\n            axi_x_norm, axi_y_norm, axi_instance_number, axi_series_id = self.get_axi_coord(\n                row.study_id, row.level, side\n            )\n\n            axi_roi = self.get_roi(\n                study_id=row.study_id,\n                series_id=axi_series_id,\n                img_num=self.img_num[2],\n                resolution=self.resolution,\n                roi_size=self.roi_size,\n                x_norm=axi_x_norm,\n                y_norm=axi_y_norm,\n                instance_number_type='filename',\n                instance_number=axi_instance_number,\n                resample_slice_spacing=self.resample_slice_spacing,\n            )\n\n            # Flip if right side\n            if side == 'left':\n                axi_left_roi = axi_roi\n            elif side == 'right':\n                axi_right_roi = np.flip(axi_roi, axis=1).copy()\n\n        if self.transform:\n            sagt2_roi = self.transform[0](image=sagt2_roi)['image']\n            sagt1_left_roi = self.transform[1](image=sagt1_left_roi)['image']\n            sagt1_right_roi = self.transform[1](image=sagt1_right_roi)['image']\n            axi_left_roi = self.transform[2](image=axi_left_roi)['image']\n            axi_right_roi = self.transform[2](image=axi_right_roi)['image']\n\n        level = self.get_level_onehot(row.level)\n        if self.phase in ['train', 'valid', 'predict', 'valid_check']:\n            label = self.get_label(\n                idx,\n                [\n                    'spinal_canal_stenosis',\n                    'left_neural_foraminal_narrowing',\n                    'right_neural_foraminal_narrowing',\n                    'left_subarticular_stenosis',\n                    'right_subarticular_stenosis',\n                ],\n            )\n            return (\n                sagt2_roi,\n                sagt1_left_roi,\n                sagt1_right_roi,\n                axi_left_roi,\n                axi_right_roi,\n                level,\n                label,\n            )\n        elif self.phase in ['test']:\n            return sagt2_roi, sagt1_left_roi, sagt1_right_roi, axi_left_roi, axi_right_roi, level, 0","metadata":{"execution":{"iopub.status.busy":"2024-10-06T21:39:19.922059Z","iopub.execute_input":"2024-10-06T21:39:19.922367Z","iopub.status.idle":"2024-10-06T21:39:20.259719Z","shell.execute_reply.started":"2024-10-06T21:39:19.922342Z","shell.execute_reply":"2024-10-06T21:39:20.258744Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Models","metadata":{}},{"cell_type":"code","source":"class CoordModel(nn.Module):\n    def __init__(\n        self,\n        base_model,\n        encoder_name,\n        num_classes=None,\n        in_channels=None,\n        encoder_weights='imagenet',\n    ):\n        super().__init__()\n        self.base_model = base_model\n        self.encoder_name = encoder_name\n        self.num_classes = num_classes\n        self.in_channels = in_channels\n        self.encoder_weights = encoder_weights\n\n        self.unet = getattr(smp, self.base_model)(\n            encoder_name=self.encoder_name,\n            classes=self.num_classes,\n            in_channels=self.in_channels,\n            encoder_weights=self.encoder_weights,\n        )\n\n    def forward(self, x):\n        return self.unet(x)\n\n\nclass SplitROIFeatures(nn.Module):\n    def __init__(\n        self,\n        base_model,\n        in_channels=None,\n        pretrained=True,\n        rnn_hidden_size=512,\n        rnn_num_layers=2,\n        rnn_dropout=0,\n        rnn_bidirectional=False,\n    ):\n        super().__init__()\n        self.in_channels = in_channels\n        self.model = timm.create_model(\n            model_name=base_model,\n            pretrained=pretrained,\n            in_chans=self.in_channels,\n            num_classes=0,\n        )\n        self.feature_num = self.model.num_features\n        self.global_pool_1d = nn.AdaptiveAvgPool1d(1)\n        self.rnn = nn.GRU(\n            input_size=self.feature_num,\n            hidden_size=rnn_hidden_size,\n            num_layers=rnn_num_layers,\n            dropout=rnn_dropout,\n            bidirectional=rnn_bidirectional,\n            batch_first=True,\n        )\n\n    def forward(self, x):\n        x_subset_list = []\n        start_size = self.in_channels // 2\n        end_size = self.in_channels - start_size\n        for i in range(start_size, x.shape[1] - end_size + 1):\n            x_subset = x[:, i - start_size : i + end_size, :, :]\n            x_subset = self.model(x_subset)\n            x_subset_list.append(x_subset)\n        x_split = torch.stack(x_subset_list, dim=1)\n\n        # RNN\n        x_split, _ = self.rnn(x_split)\n        x_split = x_split.permute(0, 2, 1)\n        x_split = self.global_pool_1d(x_split).squeeze()\n        if x_split.dim() == 1:\n            x_split = x_split.unsqueeze(0)\n\n        return x_split\n\n\nclass SplitROIFeaturesV2(SplitROIFeatures):\n    pass\n\n\nclass SpinalROIModel(nn.Module):\n    def __init__(\n        self,\n        base_model,\n        num_classes,\n        in_channels=None,\n        pretrained=True,\n        rnn_hidden_size=512,\n        rnn_num_layers=2,\n        rnn_dropout=0,\n        rnn_bidirectional=False,\n    ):\n\n        super().__init__()\n        self.model_sagt2 = SplitROIFeatures(\n            base_model,\n            in_channels=in_channels[0],\n            pretrained=pretrained,\n            rnn_hidden_size=rnn_hidden_size,\n            rnn_num_layers=rnn_num_layers,\n            rnn_dropout=rnn_dropout,\n            rnn_bidirectional=rnn_bidirectional,\n        )\n        self.model_axi = SplitROIFeatures(\n            base_model,\n            in_channels=in_channels[1],\n            pretrained=pretrained,\n            rnn_hidden_size=rnn_hidden_size,\n            rnn_num_layers=rnn_num_layers,\n            rnn_dropout=rnn_dropout,\n            rnn_bidirectional=rnn_bidirectional,\n        )\n        if rnn_bidirectional:\n            fc_in_size = rnn_hidden_size * 2\n        else:\n            fc_in_size = rnn_hidden_size\n        self.classifier = nn.Linear(2 * fc_in_size, num_classes)\n\n    def forward(self, x_sagt2, x_axi, level):\n        x_sagt2 = self.model_sagt2(x_sagt2)\n        x_axi = self.model_axi(x_axi)\n        return self.classifier(torch.cat((x_sagt2, x_axi), dim=1))\n\n\nclass ForaminalROIModel(nn.Module):\n    def __init__(\n        self,\n        base_model,\n        num_classes,\n        in_channels=None,\n        pretrained=True,\n        rnn_hidden_size=512,\n        rnn_num_layers=2,\n        rnn_dropout=0,\n        rnn_bidirectional=False,\n    ):\n        super().__init__()\n        self.model = SplitROIFeatures(\n            base_model,\n            in_channels=in_channels,\n            pretrained=pretrained,\n            rnn_hidden_size=rnn_hidden_size,\n            rnn_num_layers=rnn_num_layers,\n            rnn_dropout=rnn_dropout,\n            rnn_bidirectional=rnn_bidirectional,\n        )\n        if rnn_bidirectional:\n            fc_in_size = rnn_hidden_size * 2\n        else:\n            fc_in_size = rnn_hidden_size\n        self.classifier = nn.Linear(fc_in_size, num_classes)\n\n    def forward(self, x, level, side):\n        x = self.model(x)\n        return self.classifier(x)\n\n\nclass SubarticularROIModel(ForaminalROIModel):\n    pass\n\n\nclass GlobalROIModel(nn.Module):\n    def __init__(\n        self,\n        base_model,\n        num_classes,\n        in_channels=None,\n        pretrained=True,\n        rnn_hidden_size=512,\n        rnn_num_layers=2,\n        rnn_dropout=0,\n        rnn_bidirectional=False,\n    ):\n\n        super().__init__()\n        self.model_sagt2 = SplitROIFeatures(\n            base_model,\n            in_channels=in_channels[0],\n            pretrained=pretrained,\n            rnn_hidden_size=rnn_hidden_size,\n            rnn_num_layers=rnn_num_layers,\n            rnn_dropout=rnn_dropout,\n            rnn_bidirectional=rnn_bidirectional,\n        )\n        self.model_sagt1 = SplitROIFeatures(\n            base_model,\n            in_channels=in_channels[1],\n            pretrained=pretrained,\n            rnn_hidden_size=rnn_hidden_size,\n            rnn_num_layers=rnn_num_layers,\n            rnn_dropout=rnn_dropout,\n            rnn_bidirectional=rnn_bidirectional,\n        )\n        self.model_axi = SplitROIFeatures(\n            base_model,\n            in_channels=in_channels[2],\n            pretrained=pretrained,\n            rnn_hidden_size=rnn_hidden_size,\n            rnn_num_layers=rnn_num_layers,\n            rnn_dropout=rnn_dropout,\n            rnn_bidirectional=rnn_bidirectional,\n        )\n        if rnn_bidirectional:\n            fc_in_size = rnn_hidden_size * 2\n        else:\n            fc_in_size = rnn_hidden_size\n        self.classifier = nn.Linear(5 * fc_in_size, num_classes)\n\n    def forward(self, x_sagt2, x_sagt1_left, x_sagt1_right, x_axi_left, x_axi_right, level):\n        x_sagt2 = self.model_sagt2(x_sagt2)\n        x_sagt1_left = self.model_sagt1(x_sagt1_left)\n        x_sagt1_right = self.model_sagt1(x_sagt1_right)\n        x_axi_left = self.model_axi(x_axi_left)\n        x_axi_right = self.model_axi(x_axi_right)\n        features = torch.cat((x_sagt2, x_sagt1_left, x_sagt1_right, x_axi_left, x_axi_right), dim=1)\n        return self.classifier(features)","metadata":{"execution":{"iopub.status.busy":"2024-10-06T21:39:20.261155Z","iopub.execute_input":"2024-10-06T21:39:20.261516Z","iopub.status.idle":"2024-10-06T21:39:20.289666Z","shell.execute_reply.started":"2024-10-06T21:39:20.261475Z","shell.execute_reply":"2024-10-06T21:39:20.288758Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Generate coordinates","metadata":{}},{"cell_type":"code","source":"def coord_model_init(model_type, df_coordinates=None):\n    cfg, state_paths = get_model_params(coord_model_names[model_type])\n    cfg['dataset']['args']['data_dir'] = data_dir\n    transform = ToTensorV2()\n    dataset = get_instance('dataset', cfg, df=df, phase=phase, root_dir=input_dir, df_coordinates=df_coordinates, transform=transform)\n    dataloader = DataLoader(dataset, batch_size=1, num_workers=4, shuffle=False)\n    model = get_instance('model', cfg, encoder_weights=None)\n    \n    return model, state_paths, dataloader\n\ndef ensemble_predict(x, model, state_paths, device):\n    ensemble_pred = []\n    for model_path in state_paths:\n        model.load_state_dict(torch.load(model_path, map_location=device));\n        model.to(device);\n        model.eval();\n\n        pred = model(x)\n        ensemble_pred.append(pred)\n    return torch.stack(ensemble_pred, dim=0).mean(0)","metadata":{"execution":{"iopub.status.busy":"2024-10-06T21:39:20.290987Z","iopub.execute_input":"2024-10-06T21:39:20.291426Z","iopub.status.idle":"2024-10-06T21:39:20.305165Z","shell.execute_reply.started":"2024-10-06T21:39:20.291392Z","shell.execute_reply":"2024-10-06T21:39:20.304303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Sagt2 (+ Axi instance_number)","metadata":{}},{"cell_type":"code","source":"model, state_paths, dataloader = coord_model_init('sagt2')\n\ncoord_df_sagt2_list = []\ncoord_df_axi_list = []\nwith torch.no_grad():\n    for i, (x, study_id, series_id) in enumerate(dataloader):\n        x = x.to(device)\n        study_id = study_id[0]\n        series_id = series_id[0]\n        if series_id == '':\n            series_id = np.nan\n        ensemble_pred = ensemble_predict(x, model, state_paths, device).cpu().numpy()\n\n        # Get data for sagi to axi coord projection\n        mid_sag_ds, axi_ds_list = get_sagi2axi_data(study_id, series_id)\n    \n        for level_idx, level in enumerate(levels):\n            x_norm, y_norm = get_coord_from_heatmap_pred(ensemble_pred, level_idx)\n            coord_df_sagt2_list.append(\n                {\n                    'study_id': study_id,\n                    'series_id': series_id,\n                    'condition': 'Spinal Canal Stenosis',\n                    'level': level.replace('_', '/').upper(),\n                    'x_norm': x_norm,\n                    'y_norm': y_norm,\n                    'instance_number': np.nan,\n                    'row_id': 'spinal_canal_stenosis_' + level,\n                }\n            )\n\n            # Generate axial instance numbers\n            if mid_sag_ds is None or axi_ds_list is None or np.isnan(x_norm) or np.isnan(y_norm):\n                axi_series_id, axi_instance_number = None, np.nan\n            else:\n                axi_series_id, axi_instance_number = sagi_coord_to_axi_instance_number(\n                    x_norm, y_norm, mid_sag_ds, axi_ds_list\n                )\n            for side in sides:\n                coord_df_axi_list.append(\n                    {\n                        'study_id': study_id,\n                        'series_id': axi_series_id,\n                        'condition': side.capitalize() + ' Subarticular Stenosis',\n                        'level': level.replace('_', '/').upper(),\n                        'x_norm': np.nan,\n                        'y_norm': np.nan,\n                        'instance_number': axi_instance_number,\n                        'row_id': side + '_subarticular_stenosis_' + level,\n                    }\n                )\ncoord_df = pd.concat((pd.DataFrame(coord_df_sagt2_list), pd.DataFrame(coord_df_axi_list)))","metadata":{"execution":{"iopub.status.busy":"2024-10-06T21:39:20.308357Z","iopub.execute_input":"2024-10-06T21:39:20.308622Z","iopub.status.idle":"2024-10-06T21:39:27.113171Z","shell.execute_reply.started":"2024-10-06T21:39:20.308599Z","shell.execute_reply":"2024-10-06T21:39:27.112032Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Sagt1","metadata":{}},{"cell_type":"code","source":"model, state_paths, dataloader = coord_model_init('sagt1')\n\ncoord_df_sagt1_list = []\nwith torch.no_grad():\n    for i, (x, study_id, series_id, side) in enumerate(dataloader):\n        x = x.to(device)\n        study_id = study_id[0]\n        series_id = series_id[0]\n        if series_id == '':\n            series_id = np.nan\n        side = side[0]\n        ensemble_pred = ensemble_predict(x, model, state_paths, device).cpu().numpy()\n        \n        for level_idx, level in enumerate(levels):\n            x_norm, y_norm = get_coord_from_heatmap_pred(ensemble_pred, level_idx)\n            \n            coord_df_sagt1_list.append(\n                {\n                    'study_id': study_id,\n                    'series_id': series_id,\n                    'condition': side.capitalize() + ' Neural Foraminal Narrowing',\n                    'level': level.replace('_', '/').upper(),\n                    'x_norm': x_norm,\n                    'y_norm': y_norm,\n                    'instance_number': np.nan,\n                    'row_id': side\n                    + '_neural_foraminal_narrowing_'\n                    + level,\n                }\n            )\ncoord_df = pd.concat((coord_df, pd.DataFrame(coord_df_sagt1_list)))","metadata":{"execution":{"iopub.status.busy":"2024-10-06T21:39:27.114904Z","iopub.execute_input":"2024-10-06T21:39:27.115766Z","iopub.status.idle":"2024-10-06T21:39:32.700042Z","shell.execute_reply.started":"2024-10-06T21:39:27.115735Z","shell.execute_reply":"2024-10-06T21:39:32.698777Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Axi","metadata":{}},{"cell_type":"code","source":"model, state_paths, dataloader = coord_model_init('axi', df_coordinates=coord_df)\n\nwith torch.no_grad():\n    for i, (x, study_id, series_id, level) in enumerate(dataloader):\n        x = x.to(device)\n        study_id = study_id[0]\n        series_id = series_id[0]\n        level = level[0]\n        ensemble_pred = ensemble_predict(x, model, state_paths, device).cpu().numpy()\n        \n        for side_idx, side in enumerate(sides):\n            series_id = coord_df.loc[\n                (coord_df['study_id'] == study_id)\n                & (coord_df['row_id'] == side + '_subarticular_stenosis_' + level),\n                'series_id',\n            ].values[0]\n                \n            x_norm, y_norm = get_coord_from_heatmap_pred(ensemble_pred, side_idx)\n            \n            # Insert into coord_df\n            coord_df.loc[\n                (coord_df['study_id'] == study_id)\n                & (coord_df['series_id'] == series_id)\n                & (coord_df['row_id'] == side + '_subarticular_stenosis_' + level),\n                ['x_norm', 'y_norm'],\n            ] = [x_norm, y_norm]","metadata":{"execution":{"iopub.status.busy":"2024-10-06T21:39:32.701788Z","iopub.execute_input":"2024-10-06T21:39:32.702232Z","iopub.status.idle":"2024-10-06T21:39:39.882293Z","shell.execute_reply.started":"2024-10-06T21:39:32.702168Z","shell.execute_reply":"2024-10-06T21:39:39.881232Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Modeling","metadata":{}},{"cell_type":"code","source":"def roi_ensemble_preds(model_type, df_coordinates, transform=ToTensorV2()):\n    cfg, state_paths = get_model_params(model_names[model_type])\n    cfg['dataset']['args']['data_dir'] = data_dir\n    cfg['model']['args']['pretrained'] = False\n    dataset = get_instance('dataset', cfg, df=df, phase=phase, root_dir=input_dir, df_coordinates=df_coordinates, transform=transform)\n    dataloader = DataLoader(dataset, batch_size=16, num_workers=4, shuffle=False)\n\n    ensemble_preds = []\n    model_list = []\n    for model_path in state_paths:\n        model = get_instance('model', cfg)\n        model.load_state_dict(torch.load(model_path, map_location=device));\n        model.to(device);\n        model.eval();\n        model_list.append(model)\n\n    preds = []\n    with torch.no_grad():\n        for *x, _ in dataloader:\n            x = [x_.to(device) for x_ in x]\n            pred = torch.stack([model_list[i](*x) for i in range(5)], dim=0)\n            if pred.dim() == 2:\n                pred = pred.unsqueeze(1)\n            if pred.shape[2] > 4:\n                pred = torch.unflatten(pred, 2, [3, -1])\n            pred = pred.softmax(2).mean(0)\n            preds.append(pred)\n    preds = torch.cat(preds, dim=0)\n    \n    return preds","metadata":{"execution":{"iopub.status.busy":"2024-10-06T21:39:39.885590Z","iopub.execute_input":"2024-10-06T21:39:39.885912Z","iopub.status.idle":"2024-10-06T21:39:39.897033Z","shell.execute_reply.started":"2024-10-06T21:39:39.885878Z","shell.execute_reply":"2024-10-06T21:39:39.895987Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Spinal","metadata":{}},{"cell_type":"code","source":"spinal_ensemble_preds = roi_ensemble_preds(model_type='spinal', df_coordinates=coord_df)\nspinal_ensemble_preds = spinal_ensemble_preds.reshape(5, -1, 3).moveaxis(0, -1)\n\nspinal_ensemble_preds_2 = roi_ensemble_preds(model_type='spinal_2', df_coordinates=coord_df)\nspinal_ensemble_preds_2 = spinal_ensemble_preds_2.reshape(5, -1, 3).moveaxis(0, -1)\n\nspinal_ensemble_preds_3 = roi_ensemble_preds(model_type='spinal_3', df_coordinates=coord_df)\nspinal_ensemble_preds_3 = spinal_ensemble_preds_3.reshape(5, -1, 3).moveaxis(0, -1)","metadata":{"execution":{"iopub.status.busy":"2024-10-06T21:39:39.898127Z","iopub.execute_input":"2024-10-06T21:39:39.898446Z","iopub.status.idle":"2024-10-06T21:40:19.631627Z","shell.execute_reply.started":"2024-10-06T21:39:39.898413Z","shell.execute_reply":"2024-10-06T21:40:19.630274Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Foraminal","metadata":{}},{"cell_type":"code","source":"foraminal_ensemble_preds = roi_ensemble_preds(model_type='foraminal', df_coordinates=coord_df)\nforaminal_ensemble_preds = foraminal_ensemble_preds.reshape(5, -1, 2, 3).moveaxis(2, -1).moveaxis(0, -1).reshape(-1, 3, 10)\n\nforaminal_ensemble_preds_2 = roi_ensemble_preds(model_type='foraminal_2', df_coordinates=coord_df)\nforaminal_ensemble_preds_2 = foraminal_ensemble_preds_2.reshape(5, -1, 2, 3).moveaxis(2, -1).moveaxis(0, -1).reshape(-1, 3, 10)\n\nforaminal_ensemble_preds_3 = roi_ensemble_preds(model_type='foraminal_3', df_coordinates=coord_df)\nforaminal_ensemble_preds_3 = foraminal_ensemble_preds_3.reshape(5, -1, 2, 3).moveaxis(2, -1).moveaxis(0, -1).reshape(-1, 3, 10)","metadata":{"execution":{"iopub.status.busy":"2024-10-06T21:40:19.633275Z","iopub.execute_input":"2024-10-06T21:40:19.633620Z","iopub.status.idle":"2024-10-06T21:40:43.871064Z","shell.execute_reply.started":"2024-10-06T21:40:19.633591Z","shell.execute_reply":"2024-10-06T21:40:43.869927Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Subarticular","metadata":{}},{"cell_type":"code","source":"subarticular_ensemble_preds = roi_ensemble_preds(model_type='subarticular', df_coordinates=coord_df)\nsubarticular_ensemble_preds = subarticular_ensemble_preds.reshape(5, -1, 2, 3).moveaxis(2, -1).moveaxis(0, -1).reshape(-1, 3, 10)\n\nsubarticular_ensemble_preds_2 = roi_ensemble_preds(model_type='subarticular_2', df_coordinates=coord_df)\nsubarticular_ensemble_preds_2 = subarticular_ensemble_preds_2.reshape(5, -1, 2, 3).moveaxis(2, -1).moveaxis(0, -1).reshape(-1, 3, 10)\n\nsubarticular_ensemble_preds_3 = roi_ensemble_preds(model_type='subarticular_3', df_coordinates=coord_df)\nsubarticular_ensemble_preds_3 = subarticular_ensemble_preds_3.reshape(5, -1, 2, 3).moveaxis(2, -1).moveaxis(0, -1).reshape(-1, 3, 10)","metadata":{"execution":{"iopub.status.busy":"2024-10-06T21:40:43.872611Z","iopub.execute_input":"2024-10-06T21:40:43.872934Z","iopub.status.idle":"2024-10-06T21:41:03.715587Z","shell.execute_reply.started":"2024-10-06T21:40:43.872906Z","shell.execute_reply":"2024-10-06T21:41:03.714492Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Global","metadata":{}},{"cell_type":"code","source":"global_ensemble_preds = roi_ensemble_preds(model_type='global', df_coordinates=coord_df, transform=[ToTensorV2()] * 3)\nglobal_ensemble_preds = global_ensemble_preds.reshape(5, -1, 3, 5).moveaxis(0, -1).reshape(-1, 3, 25)\n\nglobal_ensemble_preds_2 = roi_ensemble_preds(model_type='global_2', df_coordinates=coord_df, transform=[ToTensorV2()] * 3)\nglobal_ensemble_preds_2 = global_ensemble_preds_2.reshape(5, -1, 3, 5).moveaxis(0, -1).reshape(-1, 3, 25)\n\nglobal_ensemble_preds_3 = roi_ensemble_preds(model_type='global_3', df_coordinates=coord_df, transform=[ToTensorV2()] * 3)\nglobal_ensemble_preds_3 = global_ensemble_preds_3.reshape(5, -1, 3, 5).moveaxis(0, -1).reshape(-1, 3, 25)","metadata":{"execution":{"iopub.status.busy":"2024-10-06T21:41:03.717021Z","iopub.execute_input":"2024-10-06T21:41:03.717349Z","iopub.status.idle":"2024-10-06T21:42:05.406972Z","shell.execute_reply.started":"2024-10-06T21:41:03.717321Z","shell.execute_reply":"2024-10-06T21:42:05.405684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Combine predictions","metadata":{}},{"cell_type":"code","source":"split_ensemble_preds = torch.concatenate([spinal_ensemble_preds, foraminal_ensemble_preds, subarticular_ensemble_preds], axis=-1)\nsplit_ensemble_preds_2 = torch.concatenate([spinal_ensemble_preds_2, foraminal_ensemble_preds_2, subarticular_ensemble_preds_2], axis=-1)\nsplit_ensemble_preds_3 = torch.concatenate([spinal_ensemble_preds_3, foraminal_ensemble_preds_3, subarticular_ensemble_preds_3], axis=-1)\n\nensemble_preds = (split_ensemble_preds + global_ensemble_preds + \n                  split_ensemble_preds_2 + global_ensemble_preds_2 + \n                  split_ensemble_preds_3 + global_ensemble_preds_3) / 6\n\nrow_ids = [f'{study_id}_{cond}_{level}' for study_id in df['study_id'].tolist() for cond in conditions for level in levels]\ncombined_preds = ensemble_preds.swapaxes(1, 2).flatten(0, 1).cpu().numpy().copy()","metadata":{"execution":{"iopub.status.busy":"2024-10-06T21:42:05.408909Z","iopub.execute_input":"2024-10-06T21:42:05.409241Z","iopub.status.idle":"2024-10-06T21:42:05.432048Z","shell.execute_reply.started":"2024-10-06T21:42:05.409212Z","shell.execute_reply":"2024-10-06T21:42:05.431350Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Calculate loss for debug","metadata":{}},{"cell_type":"code","source":"if debug:\n    pd.set_option('future.no_silent_downcasting', True)\n    def print_loss(preds, ys):\n        loss = nn.NLLLoss(weight=torch.tensor([1.0, 2.0, 4.0]).to(device))(torch.log(preds), ys).item()\n        print(loss)\n    ys = torch.tensor(\n        np.nan_to_num(\n            df.iloc[:,1:].replace({'Normal/Mild': 0, 'Moderate': 1, 'Severe': 2}).to_numpy(dtype=np.float32)\n        ), dtype=torch.int64).to(device)\n    \n    spinal_idx = np.arange(5)\n    foraminal_idx = np.arange(5, 15)\n    subarticular_idx = np.arange(15, 25)\n    print_loss(ensemble_preds[..., spinal_idx], ys[..., spinal_idx])\n    print_loss(ensemble_preds[..., foraminal_idx], ys[..., foraminal_idx])\n    print_loss(ensemble_preds[..., subarticular_idx], ys[..., subarticular_idx])\n    print()\n    print_loss(ensemble_preds, ys)","metadata":{"execution":{"iopub.status.busy":"2024-10-06T21:42:05.433076Z","iopub.execute_input":"2024-10-06T21:42:05.433381Z","iopub.status.idle":"2024-10-06T21:42:05.441529Z","shell.execute_reply.started":"2024-10-06T21:42:05.433357Z","shell.execute_reply":"2024-10-06T21:42:05.440517Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Submit results","metadata":{}},{"cell_type":"code","source":"sample_sub = pd.read_csv(os.path.join(input_dir, data_dir, 'sample_submission.csv'))\nlabels = list(sample_sub.columns[1:])\n\nsubmission = pd.DataFrame()\nsubmission['row_id'] = row_ids\nsubmission[labels] = combined_preds\nsubmission.head()\n\nif debug:\n    display(pd.melt(df, id_vars='study_id').sort_values(['study_id', 'variable']).head(50))\n    display(submission.sort_values('row_id').head(50))\nsubmission.to_csv('submission.csv', index=False)\npd.read_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2024-10-06T21:42:05.442595Z","iopub.execute_input":"2024-10-06T21:42:05.442847Z","iopub.status.idle":"2024-10-06T21:42:05.480544Z","shell.execute_reply.started":"2024-10-06T21:42:05.442825Z","shell.execute_reply":"2024-10-06T21:42:05.479723Z"},"trusted":true},"execution_count":null,"outputs":[]}]}