{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","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":114536,"sourceType":"modelInstanceVersion","modelInstanceId":64905,"modelId":89293}],"dockerImageVersionId":30762,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"This is the training companion notebook to the inference notebook I have published here:\n","metadata":{}},{"cell_type":"markdown","source":"## Dependencies","metadata":{}},{"cell_type":"code","source":"%pip install --quiet /kaggle/input/timm_3d_deps/other/initial/10/pydicom/pydicom/pydicom-2.4.4-py3-none-any.whl\n%pip install timm_3d --no-index --quiet --find-links=/kaggle/input/timm_3d_deps/other/initial/10/timm_3d/\n%pip install torchio --no-index --quiet --find-links=/kaggle/input/timm_3d_deps/other/initial/10/torchio/\n%pip install itk --no-index --quiet --find-links=/kaggle/input/timm_3d_deps/other/initial/10/itk/itk\n%pip install skorch --no-index --quiet --find-links=/kaggle/input/timm_3d_deps/other/initial/10/skorch/skorch\n%pip install spacecutter --no-index --quiet --find-links=/kaggle/input/timm_3d_deps/other/initial/10/spacecutter/\n%pip install open3d --no-index --quiet --find-links=/kaggle/input/timm_3d_deps/other/initial/10/open3d\n%pip install pgzip --no-index --quiet --find-links=/kaggle/input/timm_3d_deps/other/initial/10/pgzip","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-09-28T11:13:51.653720Z","iopub.execute_input":"2024-09-28T11:13:51.654200Z","iopub.status.idle":"2024-09-28T11:16:18.289084Z","shell.execute_reply.started":"2024-09-28T11:13:51.654160Z","shell.execute_reply":"2024-09-28T11:16:18.288014Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Config","metadata":{}},{"cell_type":"code","source":"data_path = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/\"","metadata":{"execution":{"iopub.status.busy":"2024-09-28T11:21:13.716080Z","iopub.execute_input":"2024-09-28T11:21:13.716514Z","iopub.status.idle":"2024-09-28T11:21:13.721666Z","shell.execute_reply.started":"2024-09-28T11:21:13.716471Z","shell.execute_reply":"2024-09-28T11:21:13.720685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\n\nCONFIG = dict(\n    n_levels=5,\n    num_classes=25,\n    num_conditions=5,\n    image_interpolation=\"nearest\",\n    # backbone=\"maxvit_rmlp_tiny_rw_256\",\n    backbone=\"efficientnet_b0\",\n    # vol_size=(256, 256, 256),\n    vol_size=(64, 64, 64),\n    num_workers=4,\n    gradient_acc_steps=1,\n    drop_rate=0.4,\n    drop_rate_last=0.,\n    drop_path_rate=0.4,\n    aug_prob=0.9,\n    out_dim=3,\n    # epochs=40\n    epochs=5,\n    batch_size=16,\n    split_k=5,\n    device=torch.device(\"cuda\") if torch.cuda.is_available() else \"cpu\",\n    seed=42\n)","metadata":{"execution":{"iopub.status.busy":"2024-09-28T11:21:17.600478Z","iopub.execute_input":"2024-09-28T11:21:17.600862Z","iopub.status.idle":"2024-09-28T11:21:20.762152Z","shell.execute_reply.started":"2024-09-28T11:21:17.600826Z","shell.execute_reply":"2024-09-28T11:21:20.761334Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CONDITIONS = {\n    \"Sagittal T2/STIR\": [\"Spinal Canal Stenosis\"],\n    \"Axial T2\": [\"Left Subarticular Stenosis\", \"Right Subarticular Stenosis\"],\n    \"Sagittal T1\": [\"Left Neural Foraminal Narrowing\", \"Right Neural Foraminal Narrowing\"],\n}\nLABEL_MAP = {'normal_mild': 0, 'moderate': 1, 'severe': 2}","metadata":{"execution":{"iopub.status.busy":"2024-09-28T11:21:20.763676Z","iopub.execute_input":"2024-09-28T11:21:20.764043Z","iopub.status.idle":"2024-09-28T11:21:20.768909Z","shell.execute_reply.started":"2024-09-28T11:21:20.764011Z","shell.execute_reply":"2024-09-28T11:21:20.768012Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Loading","metadata":{}},{"cell_type":"markdown","source":"## Training Metadata","metadata":{}},{"cell_type":"code","source":"def retrieve_coordinate_training_data(train_path):\n    def reshape_row(row):\n        data = {'study_id': [], 'condition': [], 'level': [], 'severity': []}\n\n        for column, value in row.items():\n            if column not in ['study_id', 'series_id', 'instance_number', 'x', 'y', 'series_description']:\n                parts = column.split('_')\n                condition = ' '.join([word.capitalize() for word in parts[:-2]])\n                level = parts[-2].capitalize() + '/' + parts[-1].capitalize()\n                data['study_id'].append(row['study_id'])\n                data['condition'].append(condition)\n                data['level'].append(level)\n                data['severity'].append(value)\n\n        return pd.DataFrame(data)\n\n    train = pd.read_csv(train_path + 'train.csv')\n    label = pd.read_csv(train_path + 'train_label_coordinates.csv')\n    train_desc = pd.read_csv(train_path + 'train_series_descriptions.csv')\n    test_desc = pd.read_csv(train_path + 'test_series_descriptions.csv')\n    sub = pd.read_csv(train_path + 'sample_submission.csv')\n\n    new_train_df = pd.concat([reshape_row(row) for _, row in train.iterrows()], ignore_index=True)\n    merged_df = pd.merge(new_train_df, label, on=['study_id', 'condition', 'level'], how='inner')\n    final_merged_df = pd.merge(merged_df, train_desc, on=['series_id', 'study_id'], how='inner')\n    final_merged_df['severity'] = final_merged_df['severity'].map(\n        {'Normal/Mild': 'normal_mild', 'Moderate': 'moderate', 'Severe': 'severe'})\n\n    final_merged_df['row_id'] = (\n            final_merged_df['study_id'].astype(str) + '_' +\n            final_merged_df['condition'].str.lower().str.replace(' ', '_') + '_' +\n            final_merged_df['level'].str.lower().str.replace('/', '_')\n    )\n\n    # Create the image_path column\n    final_merged_df['image_path'] = (\n            f'{train_path}/train_images/' +\n            final_merged_df['study_id'].astype(str) + '/' +\n            final_merged_df['series_id'].astype(str) + '/' +\n            final_merged_df['instance_number'].astype(str) + '.dcm'\n    )\n\n    return final_merged_df","metadata":{"execution":{"iopub.status.busy":"2024-09-28T11:21:38.412845Z","iopub.execute_input":"2024-09-28T11:21:38.413619Z","iopub.status.idle":"2024-09-28T11:21:38.428324Z","shell.execute_reply.started":"2024-09-28T11:21:38.413580Z","shell.execute_reply":"2024-09-28T11:21:38.427431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\n\ntrain_data = retrieve_coordinate_training_data(data_path)","metadata":{"execution":{"iopub.status.busy":"2024-09-28T11:21:41.530881Z","iopub.execute_input":"2024-09-28T11:21:41.531329Z","iopub.status.idle":"2024-09-28T11:21:43.743363Z","shell.execute_reply.started":"2024-09-28T11:21:41.531285Z","shell.execute_reply":"2024-09-28T11:21:43.742533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data[:5]","metadata":{"execution":{"iopub.status.busy":"2024-09-28T11:21:59.723473Z","iopub.execute_input":"2024-09-28T11:21:59.723961Z","iopub.status.idle":"2024-09-28T11:21:59.752391Z","shell.execute_reply.started":"2024-09-28T11:21:59.723921Z","shell.execute_reply":"2024-09-28T11:21:59.751285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset Definitions\n\nNote the mirror trick -- by flipping across the X axis, we get any right labels as left also and vice versa","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader\n\n\nclass StudyLevelDataset(Dataset):\n    def __init__(self,\n                 base_path: str,\n                 dataframe: pd.DataFrame,\n                 transform_3d=None,\n                 is_train=False,\n                 vol_size=(192, 192, 192),\n                 use_mirror_trick=False):\n        self.base_path = base_path\n        self.is_train = is_train\n        self.use_mirror_trick = use_mirror_trick\n\n        self.dataframe = (dataframe[['study_id', \"series_id\", \"series_description\", \"condition\", \"severity\", \"level\"]]\n                          .drop_duplicates())\n\n        self.subjects = self.dataframe[['study_id']].drop_duplicates().reset_index(drop=True)\n        self.series = self.dataframe[[\"study_id\", \"series_id\"]].drop_duplicates().groupby(\"study_id\")[\n            \"series_id\"].apply(list).to_dict()\n        self.series_descs = {e[0]: e[1] for e in\n                             self.dataframe[[\"series_id\", \"series_description\"]].drop_duplicates().values}\n\n        self.transform_3d = transform_3d\n\n        self.levels = sorted(self.dataframe[\"level\"].unique())\n        self.labels = self._get_labels()\n        self.vol_size = vol_size\n\n    def __len__(self):\n        return len(self.subjects) * (2 if self.use_mirror_trick else 1)\n\n    def __getitem__(self, index):\n        is_mirror = index >= len(self.subjects)\n        curr = self.subjects.iloc[index % len(self.subjects)]\n\n        label = np.array(self.labels[(curr[\"study_id\"])])\n        study_path = os.path.join(self.base_path, str(curr[\"study_id\"]))\n\n        study_images = read_study_as_voxel_grid_v2(study_path,\n                                                   curr[\"study_id\"],\n                                                   series_type_dict=self.series_descs,\n                                                   img_size=(self.vol_size[0], self.vol_size[1]))\n\n        if is_mirror:\n            temp = label[:10].copy()\n            label[:10] = label[10:20].copy()\n            label[10:20] = temp\n\n        if self.transform_3d is not None:\n            study_images = torch.FloatTensor(study_images)\n\n            if is_mirror:\n                study_images = torch.flip(study_images, [1])\n\n            study_images = self.transform_3d(study_images)  # .data\n            return study_images.to(torch.half), torch.tensor(label, dtype=torch.long)\n\n        print(\"loaded\")\n        return torch.HalfTensor(study_images.copy()), torch.tensor(label, dtype=torch.long)\n\n    def _get_labels(self):\n        labels = dict()\n        for name, group in self.dataframe.groupby([\"study_id\"]):\n            group = group[[\"condition\", \"level\", \"severity\"]].drop_duplicates().sort_values([\"condition\", \"level\"])\n            label_indices = []\n            for index, row in group.iterrows():\n                if row[\"severity\"] in LABEL_MAP:\n                    label_indices.append(LABEL_MAP[row[\"severity\"]])\n                else:\n                    raise ValueError()\n\n            study_id = name[0]\n\n            labels[study_id] = label_indices\n\n        return labels","metadata":{"execution":{"iopub.status.busy":"2024-09-28T11:24:17.634832Z","iopub.execute_input":"2024-09-28T11:24:17.635217Z","iopub.status.idle":"2024-09-28T11:24:17.652772Z","shell.execute_reply.started":"2024-09-28T11:24:17.635181Z","shell.execute_reply":"2024-09-28T11:24:17.651820Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3D Data Loading\n\nExplained in a little more detail on my inference notebook. Note the caching part -- this makes a huge difference in the training time since the transforms are a massive bottleneck otherwise","metadata":{}},{"cell_type":"code","source":"def read_study_as_voxel_grid_v2(dir_path,\n                                study_id,\n                                series_type_dict=None, \n                                downsampling_factor=1, \n                                img_size=(256, 256), \n                                caching=True,\n                                cache_base_path=\"/kaggle/working/3d_cache\"):\n    if caching:\n        os.makedirs(os.path.join(cache_base_path, str(study_id)), exist_ok=True)\n        cache_path = os.path.join(cache_base_path, str(study_id), f\"cached_grid_v2_{img_size[0]}.npy.gz\")\n        f = None\n        if os.path.exists(cache_path):\n            try:\n                f = pgzip.PgzipFile(cache_path, \"r\")\n                ret = np.load(f, allow_pickle=True)\n                f.close()\n                return ret\n            except Exception as e:\n                print(dir_path, \"\\n\", e)\n                if f:\n                    f.close()\n                os.remove(cache_path)\n\n    pcd_overall = read_study_as_pcd(dir_path,\n                                    series_types_dict=series_type_dict,\n                                    downsampling_factor=downsampling_factor,\n                                    img_size=img_size)\n    box = pcd_overall.get_axis_aligned_bounding_box()\n\n    max_b = np.array(box.get_max_bound())\n    min_b = np.array(box.get_min_bound())\n\n    pts = (np.array(pcd_overall.points) - (min_b)) * (\n            (img_size[0] - 1, img_size[0] - 1, img_size[0] - 1) / (max_b - min_b))\n    coords = np.round(pts).astype(np.int32)\n    vals = np.array(pcd_overall.colors, dtype=np.float16)\n\n    grid = np.zeros((3, img_size[0], img_size[0], img_size[0]), dtype=np.float16)\n    indices = coords[:, 0], coords[:, 1], coords[:, 2]\n\n    np.maximum.at(grid[0], indices, vals[:, 0])\n    np.maximum.at(grid[1], indices, vals[:, 1])\n    np.maximum.at(grid[2], indices, vals[:, 2])\n    \n    if caching:\n        f = pgzip.PgzipFile(cache_path, \"w\")\n        np.save(f, grid)\n        f.close()\n\n    return grid\n\n\ndef read_study_as_pcd(dir_path,\n                      series_types_dict=None,\n                      downsampling_factor=1,\n                      resize_slices=True,\n                      resize_method=\"nearest\",\n                      stack_slices_thickness=True,\n                      img_size=(256, 256)):\n    pcd_overall = o3d.geometry.PointCloud()\n\n    for path in glob.glob(os.path.join(dir_path, \"**/*.dcm\"), recursive=True):\n        dicom_slice = dcmread(path)\n\n        series_id = os.path.basename(os.path.dirname(path))\n        study_id = os.path.basename(os.path.dirname(os.path.dirname(path)))\n        if series_types_dict is None or int(series_id) not in series_types_dict:\n            series_desc = dicom_slice.SeriesDescription\n        else:\n            series_desc = series_types_dict[int(series_id)]\n            series_desc = series_desc.split(\" \")[-1]\n\n        x_orig, y_orig = dicom_slice.pixel_array.shape\n        if resize_slices:\n            if resize_method == \"nearest\":\n                img = np.expand_dims(cv2.resize(dicom_slice.pixel_array, img_size, interpolation=cv2.INTER_AREA), -1)\n            elif resize_method == \"maxpool\":\n                img_tensor = torch.tensor(dicom_slice.pixel_array).float()\n                img = F.adaptive_max_pool2d(img_tensor.unsqueeze(0), img_size).numpy()\n            else:\n                raise ValueError(f\"Invalid resize_method {resize_method}\")\n        else:\n            img = np.expand_dims(np.array(dicom_slice.pixel_array), -1)\n        x, y, z = np.where(img)\n\n        downsampling_factor_iter = max(downsampling_factor, int(math.ceil(len(x) / 6e6)))\n\n        index_voxel = np.vstack((x, y, z))[:, ::downsampling_factor_iter]\n        grid_index_array = index_voxel.T\n        pcd = o3d.geometry.PointCloud(o3d.utility.Vector3dVector(grid_index_array.astype(np.float64)))\n\n        vals = np.expand_dims(img[x, y, z][::downsampling_factor_iter], -1)\n        if series_desc == \"T1\":\n            vals = np.pad(vals, ((0, 0), (0, 2)))\n        elif series_desc == \"T2\":\n            vals = np.pad(vals, ((0, 0), (1, 1)))\n        elif series_desc == \"T2/STIR\":\n            vals = np.pad(vals, ((0, 0), (2, 0)))\n        else:\n            raise ValueError(f\"Unknown series desc: {series_desc}\")\n\n        pcd.colors = o3d.utility.Vector3dVector(vals.astype(np.float64))\n\n        if resize_slices:\n            transform_matrix_factor = np.matrix(\n                [[0, y_orig / img_size[1], 0, 0],\n                 [x_orig / img_size[0], 0, 0, 0],\n                 [0, 0, 1, 0],\n                 [0, 0, 0, 1]]\n            )\n        else:\n            transform_matrix_factor = np.matrix(\n                [[0, 1, 0, 0],\n                 [1, 0, 0, 0],\n                 [0, 0, 1, 0],\n                 [0, 0, 0, 1]]\n            )\n\n        dX, dY = dicom_slice.PixelSpacing\n        dZ = dicom_slice.SliceThickness\n\n        X = np.array(list(dicom_slice.ImageOrientationPatient[:3]) + [0]) * dX\n        Y = np.array(list(dicom_slice.ImageOrientationPatient[3:]) + [0]) * dY\n\n        S = np.array(list(dicom_slice.ImagePositionPatient) + [1])\n\n        transform_matrix = np.array([X, Y, np.zeros(len(X)), S]).T\n        transform_matrix = transform_matrix @ transform_matrix_factor\n\n        if stack_slices_thickness:\n            for z in range(int(dZ)):\n                pos = list(dicom_slice.ImagePositionPatient)\n                if series_desc == \"T2\":\n                    pos[-1] += z\n                else:\n                    pos[0] += z\n                S = np.array(pos + [1])\n\n                transform_matrix = np.array([X, Y, np.zeros(len(X)), S]).T\n                transform_matrix = transform_matrix @ transform_matrix_factor\n\n                pcd_overall += copy.deepcopy(pcd).transform(transform_matrix)\n\n        else:\n            pcd_overall += copy.deepcopy(pcd).transform(transform_matrix)\n\n    return pcd_overall","metadata":{"execution":{"iopub.status.busy":"2024-09-28T11:25:05.801223Z","iopub.execute_input":"2024-09-28T11:25:05.801608Z","iopub.status.idle":"2024-09-28T11:25:05.831915Z","shell.execute_reply.started":"2024-09-28T11:25:05.801574Z","shell.execute_reply":"2024-09-28T11:25:05.830814Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset Utility\n\nSee the cleaning here -- I am simply dropping the data points where any of the 25 labels are missing","metadata":{}},{"cell_type":"code","source":"def create_study_level_datasets_and_loaders_k_fold(df: pd.DataFrame,\n                                                   base_path: str,\n                                                   transform_3d_train=None,\n                                                   transform_3d_val=None,\n                                                   vol_size=None,\n                                                   split_k=5,\n                                                   random_seed=42,\n                                                   batch_size=1,\n                                                   num_workers=0,\n                                                   pin_memory=True,\n                                                   use_mirroring_trick=True):\n    df = df.dropna()\n    # This drops any subjects with nans\n\n    filtered_df = pd.DataFrame(columns=df.columns)\n    for series_desc in CONDITIONS.keys():\n        subset = df[df['series_description'] == series_desc]\n        if series_desc == \"Sagittal T2/STIR\":\n            subset = subset[subset.groupby([\"study_id\"]).transform('size') == 5]\n        else:\n            subset = subset[subset.groupby([\"study_id\"]).transform('size') == 10]\n        filtered_df = pd.concat([filtered_df, subset])\n\n    filtered_df = filtered_df[filtered_df.groupby([\"study_id\"]).transform('size') == 25]\n\n    np.random.seed(random_seed)\n    ids = filtered_df[\"study_id\"].unique()\n    np.random.shuffle(ids)\n\n    ret = []\n    folds = np.array_split(ids, split_k)\n\n    for index, fold in enumerate(folds):\n        val_studies = fold\n\n        train_df = filtered_df[~filtered_df[\"study_id\"].isin(val_studies)]\n        val_df = filtered_df[filtered_df[\"study_id\"].isin(val_studies)]\n\n        train_df = train_df.reset_index(drop=True)\n        val_df = val_df.reset_index(drop=True)\n\n        train_dataset = StudyLevelDataset(base_path, train_df,\n                                          transform_3d=transform_3d_train,\n                                          is_train=True,\n                                          use_mirror_trick=use_mirroring_trick,\n                                          vol_size=vol_size\n                                          )\n        val_dataset = StudyLevelDataset(base_path, val_df,\n                                        transform_3d=transform_3d_val,\n                                        vol_size=vol_size\n                                        )\n\n        train_loader = DataLoader(train_dataset,\n                                  batch_size=batch_size,\n                                  shuffle=True,\n                                  num_workers=num_workers,\n                                  pin_memory=pin_memory,\n                                  persistent_workers=num_workers > 0)\n        val_loader = DataLoader(val_dataset,\n                                batch_size=batch_size,\n                                shuffle=False,\n                                num_workers=num_workers,\n                                pin_memory=pin_memory,\n                                persistent_workers=num_workers > 0)\n\n        ret.append((train_loader, val_loader, train_dataset, val_dataset))\n\n    return ret","metadata":{"execution":{"iopub.status.busy":"2024-09-28T11:27:43.827234Z","iopub.execute_input":"2024-09-28T11:27:43.827937Z","iopub.status.idle":"2024-09-28T11:27:43.841302Z","shell.execute_reply.started":"2024-09-28T11:27:43.827897Z","shell.execute_reply":"2024-09-28T11:27:43.840216Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Augmentations","metadata":{}},{"cell_type":"code","source":"import torchio as tio\n\ntransform_3d_train = tio.Compose([\n    tio.ZNormalization(),\n    tio.RandomAffine(translation=10, p=CONFIG[\"aug_prob\"]),\n    tio.RandomNoise(p=CONFIG[\"aug_prob\"]),\n    tio.RandomSpike(1, intensity=(-0.5, 0.5), p=CONFIG[\"aug_prob\"]),\n    tio.RescaleIntensity((0, 1)),\n])\n\ntransform_3d_val = tio.Compose([\n    tio.RescaleIntensity((0, 1)),\n])","metadata":{"execution":{"iopub.status.busy":"2024-09-28T11:27:55.904448Z","iopub.execute_input":"2024-09-28T11:27:55.904820Z","iopub.status.idle":"2024-09-28T11:27:56.782257Z","shell.execute_reply.started":"2024-09-28T11:27:55.904786Z","shell.execute_reply":"2024-09-28T11:27:56.781346Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Load the dataset","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\n\ndataset_folds = create_study_level_datasets_and_loaders_k_fold(train_data,\n                                                               transform_3d_train=transform_3d_train,\n                                                               transform_3d_val=transform_3d_val,\n                                                               base_path=os.path.join(\n                                                                data_path,\n                                                                \"train_images\"),\n                                                               vol_size=CONFIG[\"vol_size\"],\n                                                               num_workers=CONFIG[\"num_workers\"],\n                                                               split_k=CONFIG[\"split_k\"],\n                                                               batch_size=CONFIG[\"batch_size\"],\n                                                               pin_memory=True,\n                                                               use_mirroring_trick=True\n                                                               )","metadata":{"execution":{"iopub.status.busy":"2024-09-28T11:30:16.465953Z","iopub.execute_input":"2024-09-28T11:30:16.467083Z","iopub.status.idle":"2024-09-28T11:30:50.977812Z","shell.execute_reply.started":"2024-09-28T11:30:16.467040Z","shell.execute_reply":"2024-09-28T11:30:50.976955Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset_folds[1]","metadata":{"execution":{"iopub.status.busy":"2024-09-28T11:31:37.984407Z","iopub.execute_input":"2024-09-28T11:31:37.985312Z","iopub.status.idle":"2024-09-28T11:31:37.991170Z","shell.execute_reply.started":"2024-09-28T11:31:37.985251Z","shell.execute_reply":"2024-09-28T11:31:37.990296Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Sanity Check","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport open3d as o3d\nimport glob\nfrom pydicom import dcmread\nimport cv2\nimport math \nimport copy\nimport pgzip\n\nfor index, fold in enumerate(dataset_folds):\n    trainloader, valloader, trainset, testset = fold\n\n    plt.imshow(np.mean(trainset[0][0].numpy()[0, 31:34], axis=0), cmap = \"gray\")\n    plt.show()\n    \n    break","metadata":{"execution":{"iopub.status.busy":"2024-09-28T11:32:37.637177Z","iopub.execute_input":"2024-09-28T11:32:37.637925Z","iopub.status.idle":"2024-09-28T11:32:42.829295Z","shell.execute_reply.started":"2024-09-28T11:32:37.637882Z","shell.execute_reply":"2024-09-28T11:32:42.828221Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Prepopulate the cache\nOptionally, you can prepopulate the cache first before calling the training loop","metadata":{}},{"cell_type":"code","source":"from tqdm import tqdm\n\nfor fold in dataset_folds:\n    trainloader, valloader, trainset, testset = fold\n    \n    for train_point in tqdm(trainset):\n        pass\n        break\n    print(\"------------------\")\n    for val_point in tqdm(testset):\n        pass\n        break\n    \n    break","metadata":{"execution":{"iopub.status.busy":"2024-09-28T11:34:14.601222Z","iopub.execute_input":"2024-09-28T11:34:14.602195Z","iopub.status.idle":"2024-09-28T11:34:17.442203Z","shell.execute_reply.started":"2024-09-28T11:34:14.602151Z","shell.execute_reply":"2024-09-28T11:34:17.441209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"import timm_3d\nimport torch.nn as nn\nfrom spacecutter.losses import CumulativeLinkLoss\nfrom spacecutter.models import LogisticCumulativeLink\nfrom spacecutter.callbacks import AscensionCallback\n\n\nclass Classifier3dMultihead(nn.Module):\n    def __init__(self,\n                 backbone=\"efficientnet_lite0\",\n                 in_chans=1,\n                 out_classes=5,\n                 cutpoint_margin=0.15,\n                 pretrained=False):\n        super(Classifier3dMultihead, self).__init__()\n        self.out_classes = out_classes\n\n        self.backbone = timm_3d.create_model(\n            backbone,\n            features_only=False,\n            drop_rate=CONFIG[\"drop_rate\"],\n            drop_path_rate=CONFIG[\"drop_path_rate\"],\n            pretrained=pretrained,\n            in_chans=in_chans,\n            global_pool=\"max\",\n        )\n        if \"efficientnet\" in backbone:\n            head_in_dim = self.backbone.classifier.in_features\n            self.backbone.classifier = nn.Sequential(\n                nn.LayerNorm(head_in_dim),\n                nn.Dropout(CONFIG[\"drop_rate_last\"]),\n            )\n\n        elif \"vit\" in backbone:\n            self.backbone.head.drop = nn.Dropout(p=CONFIG[\"drop_rate_last\"])\n            head_in_dim = self.backbone.head.fc.in_features\n            self.backbone.head.fc = nn.Identity()\n\n        self.heads = nn.ModuleList(\n            [nn.Sequential(\n                nn.Linear(head_in_dim, 1),\n                LogisticCumulativeLink(CONFIG[\"out_dim\"])\n            ) for i in range(out_classes)]\n        )\n\n        self.ascension_callback = AscensionCallback(margin=cutpoint_margin)\n\n    def forward(self, x):\n        feat = self.backbone(x)\n        return torch.swapaxes(torch.stack([head(feat) for head in self.heads]), 0, 1)\n\n    def _ascension_callback(self):\n        for head in self.heads:\n            self.ascension_callback.clip(head[-1])","metadata":{"execution":{"iopub.status.busy":"2024-09-28T11:43:15.475852Z","iopub.execute_input":"2024-09-28T11:43:15.476567Z","iopub.status.idle":"2024-09-28T11:43:18.443934Z","shell.execute_reply.started":"2024-09-28T11:43:15.476525Z","shell.execute_reply":"2024-09-28T11:43:18.443092Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"markdown","source":"## Training util functions","metadata":{}},{"cell_type":"code","source":"import os.path\n\nfrom tqdm import tqdm\nfrom torch.cuda.amp import autocast, GradScaler\nimport torch.nn as nn\nimport matplotlib.pyplot as plt\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n\ndef model_validation_loss(model, val_loader, loss_fns, epoch):\n    val_loss = 0\n    unweighted_val_loss = 0\n    alt_val_loss = 0\n    unweighted_alt_val_loss = 0\n\n    with torch.no_grad():\n        model.eval()\n\n        for images, label in tqdm(val_loader, desc=f\"Validating epoch {epoch}\"):\n            label = label.to(device).unsqueeze(-1)\n\n            with autocast(enabled=device != \"cpu\", dtype=torch.bfloat16):\n                output = model(images.to(device))\n\n                for index, loss_fn in enumerate(loss_fns[\"train\"]):\n                    if len(loss_fns[\"train\"]) > 1:\n                        loss = loss_fn(output[:, index], label[:, index]) / len(loss_fns[\"train\"])\n                    else:\n                        loss = loss_fn(output, label) / len(loss_fns[\"train\"])\n                    val_loss += loss.cpu().item()\n\n                for index, loss_fn in enumerate(loss_fns[\"unweighted_val\"]):\n                    if len(loss_fns[\"unweighted_val\"]) > 1:\n                        loss = loss_fn(output[:, index], label[:, index]) / len(loss_fns[\"unweighted_val\"])\n                    else:\n                        loss = loss_fn(output, label) / len(loss_fns[\"unweighted_val\"])\n                    unweighted_val_loss += loss.cpu().item()\n\n                for index, loss_fn in enumerate(loss_fns[\"alt_val\"]):\n                    if len(loss_fns[\"alt_val\"]) > 1:\n                        # !TODO: Label squeezed for CE loss\n                        loss = loss_fn(output[:, index], label.squeeze(-1)[:, index]) / len(loss_fns[\"alt_val\"])\n                    else:\n                        loss = loss_fn(output, label) / len(loss_fns[\"alt_val\"])\n                    alt_val_loss += loss.cpu().item()\n\n                for index, loss_fn in enumerate(loss_fns[\"unweighted_alt_val\"]):\n                    if len(loss_fns[\"unweighted_alt_val\"]) > 1:\n                        # !TODO: Label squeezed for CE loss\n                        loss = loss_fn(output[:, index], label.squeeze(-1)[:, index]) / len(\n                            loss_fns[\"unweighted_alt_val\"])\n                    else:\n                        loss = loss_fn(output, label) / len(loss_fns[\"alt_val\"])\n                    unweighted_alt_val_loss += loss.cpu().item()\n\n                del output\n            # torch.cuda.empty_cache()\n\n        val_loss = val_loss / len(val_loader)\n        unweighted_val_loss = unweighted_val_loss / len(val_loader)\n        alt_val_loss = alt_val_loss / len(val_loader)\n        unweighted_alt_val_loss = unweighted_alt_val_loss / len(val_loader)\n\n        return val_loss, unweighted_val_loss, alt_val_loss, unweighted_alt_val_loss\n\n\ndef dump_plots_for_loss_and_acc(losses,\n                                val_losses,\n                                unweighted_val_losses,\n                                alt_val_losses,\n                                unweighted_alt_val_losses,\n                                data_subset_label,\n                                model_label):\n    plt.plot(np.log(losses), label=\"train\")\n    plt.plot(np.log(val_losses), label=\"weighted_val\")\n    plt.plot(np.log(unweighted_val_losses), label=\"unweighted_val\")\n    plt.plot(np.log(alt_val_losses), label=\"alt_val\")\n    plt.plot(np.log(unweighted_alt_val_losses), label=\"unweighted_alt_val\")\n    plt.legend(loc=\"center right\")\n    plt.title(data_subset_label)\n#     plt.savefig(f'./figures/{model_label}_loss.png')\n#     plt.close()\n    plt.show()\n\ndef train_model_with_validation(model,\n                                optimizers,\n                                schedulers,\n                                loss_fns,\n                                train_loader,\n                                val_loader,\n                                train_loader_desc=None,\n                                model_desc=\"my_model\",\n                                gradient_accumulation_per=1,\n                                epochs=10,\n                                freeze_backbone_initial_epochs=0,\n                                empty_cache_every_n_iterations=0,\n                                loss_weights=None,\n                                callbacks=None):\n    epoch_losses = []\n    epoch_validation_losses = []\n    epoch_unweighted_validation_losses = []\n    epoch_alt_validation_losses = []\n    epoch_unweighted_alt_validation_losses = []\n\n    scaler = GradScaler(init_scale=4096)\n\n    if freeze_backbone_initial_epochs > 0:\n        freeze_model_backbone(model)\n\n    for epoch in tqdm(range(epochs), desc=train_loader_desc):\n        epoch_loss = 0\n        model.train()\n\n        if freeze_backbone_initial_epochs > 0 and epoch == freeze_backbone_initial_epochs:\n            unfreeze_model_backbone(model)\n\n        for index, val in enumerate(tqdm(train_loader, desc=f\"Epoch {epoch}\")):\n            images, label = val\n            label = label.to(device).unsqueeze(-1)\n\n            with autocast(enabled=device != \"cpu\", dtype=torch.bfloat16):\n                output = model(images.to(device))\n\n                del images\n\n                if len(loss_fns[\"train\"]) > 1:\n                    loss = sum([(loss_fn(output[:, loss_index], label[:, loss_index]) / gradient_accumulation_per) for\n                                loss_index, loss_fn in enumerate(loss_fns[\"train\"])]) / len(loss_fns[\"train\"])\n                else:\n                    loss = loss_fns[\"train\"][0](output, label) / gradient_accumulation_per\n                epoch_loss += loss.detach().cpu().item() * gradient_accumulation_per  # / len(loss_fns[\"train\"])\n\n                del label\n\n            scaler.scale(loss).backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1e9)\n\n            del output\n\n            # Per gradient accumulation batch or if the last iter\n            if index % gradient_accumulation_per == 0 or index == len(train_loader) - 1:\n                for optimizer in optimizers:\n                    scaler.step(optimizer)\n                    optimizer.zero_grad(set_to_none=True)\n                scaler.update()\n\n            if callbacks:\n                for callback in callbacks:\n                    callback()\n\n            # prof.step()\n            if empty_cache_every_n_iterations > 0 and index % empty_cache_every_n_iterations == 0:\n                torch.cuda.empty_cache()\n\n            while os.path.exists(\".pause\"):\n                pass\n\n        epoch_loss = epoch_loss / len(train_loader)\n        epoch_validation_loss, epoch_unweighted_validation_loss, epoch_alt_validation_loss, epoch_unweighted_alt_validation_loss = (\n            model_validation_loss(model, val_loader, loss_fns, epoch)\n        )\n\n        for scheduler in schedulers:\n            scheduler.step()\n\n        if (epoch % 5 == 0\n            or len(epoch_validation_losses) == 0\n            or epoch_validation_loss < min(epoch_validation_losses)) \\\n                or epoch_unweighted_validation_loss < min(epoch_unweighted_validation_losses) \\\n                or epoch_alt_validation_loss < min(epoch_alt_validation_losses):\n            os.makedirs(f'/kaggle/working/models/{model_desc}', exist_ok=True)\n            torch.save(model.state_dict(),\n                       # torch.jit.script(model),\n                       f'/kaggle/working/models/{model_desc}/{model_desc}' + \"_\" + str(epoch) + \".pt\")\n\n        epoch_validation_losses.append(epoch_validation_loss)\n        epoch_unweighted_validation_losses.append(epoch_unweighted_validation_loss)\n        epoch_alt_validation_losses.append(epoch_alt_validation_loss)\n        epoch_unweighted_alt_validation_losses.append(epoch_unweighted_alt_validation_loss)\n\n        epoch_losses.append(epoch_loss)\n\n        dump_plots_for_loss_and_acc(epoch_losses,\n                                    epoch_validation_losses,\n                                    epoch_unweighted_validation_losses,\n                                    epoch_alt_validation_losses,\n                                    epoch_unweighted_alt_validation_losses,\n                                    train_loader_desc, model_desc)\n        print(f\"Training Loss for epoch {epoch}: {epoch_loss:.6f}\")\n        print(f\"Validation Loss for epoch {epoch}: {epoch_validation_loss:.6f}\")\n        print(f\"Unweighted Validation Loss for epoch {epoch}: {epoch_unweighted_validation_loss:.6f}\")\n        print(f\"Alt Validation Loss for epoch {epoch}: {epoch_alt_validation_loss:.6f}\")\n        print(f\"Unweighted Alt Validation Loss for epoch {epoch}: {epoch_unweighted_alt_validation_loss:.6f}\")\n\n    return epoch_losses, epoch_validation_losses","metadata":{"execution":{"iopub.status.busy":"2024-09-28T11:45:21.756130Z","iopub.execute_input":"2024-09-28T11:45:21.756603Z","iopub.status.idle":"2024-09-28T11:45:21.792008Z","shell.execute_reply.started":"2024-09-28T11:45:21.756565Z","shell.execute_reply":"2024-09-28T11:45:21.791011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loss Functions","metadata":{}},{"cell_type":"code","source":"schedulers = [\n]\ncriteria = {\n    \"train\": [\n        CumulativeLinkLoss(class_weights=[1,2,4]) for i in range(CONFIG[\"num_classes\"])\n    ],\n    \"unweighted_val\": [\n        CumulativeLinkLoss() for i in range(CONFIG[\"num_classes\"])\n    ],\n    \"alt_val\": [\n        nn.CrossEntropyLoss(weight=torch.Tensor([1,2,4])).to(CONFIG[\"device\"]) for i in range(CONFIG[\"num_classes\"])\n    ],\n    \"unweighted_alt_val\": [\n        nn.CrossEntropyLoss().to(device) for i in range(CONFIG[\"num_classes\"])\n    ]\n}","metadata":{"execution":{"iopub.status.busy":"2024-09-28T11:47:52.815236Z","iopub.execute_input":"2024-09-28T11:47:52.815663Z","iopub.status.idle":"2024-09-28T11:47:52.973869Z","shell.execute_reply.started":"2024-09-28T11:47:52.815623Z","shell.execute_reply":"2024-09-28T11:47:52.972834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Putting it all together","metadata":{}},{"cell_type":"code","source":"for index, fold in enumerate(dataset_folds):\n    # !TODO: Remove me, I am here to save some GPU time\n    if index > 0:\n        break\n    model = Classifier3dMultihead(backbone=CONFIG[\"backbone\"], in_chans=3, out_classes=CONFIG[\"num_classes\"]).to(CONFIG[\"device\"])\n    optimizers = [\n        torch.optim.AdamW(model.parameters(), lr=3e-4),\n    ]\n\n    trainloader, valloader, trainset, testset = fold\n\n    train_model_with_validation(model,\n                                optimizers,\n                                schedulers,\n                                criteria,\n                                trainloader,\n                                valloader,\n                                model_desc=CONFIG[\"backbone\"] + f\"_fold_{index}\",\n                                train_loader_desc=f\"Training {CONFIG['backbone']} fold {index}\",\n                                epochs=CONFIG[\"epochs\"],\n                                freeze_backbone_initial_epochs=0,\n                                callbacks=[model._ascension_callback],\n                                gradient_accumulation_per=CONFIG[\"gradient_acc_steps\"]\n                                )","metadata":{"execution":{"iopub.status.busy":"2024-09-28T11:49:06.038943Z","iopub.execute_input":"2024-09-28T11:49:06.039351Z"},"trusted":true},"execution_count":null,"outputs":[]}]}