{"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":9187072,"sourceType":"datasetVersion","datasetId":5504483},{"sourceId":12997065,"sourceType":"datasetVersion","datasetId":8227093},{"sourceId":12997102,"sourceType":"datasetVersion","datasetId":8227121},{"sourceId":562995,"sourceType":"modelInstanceVersion","modelInstanceId":425992,"modelId":443476}],"dockerImageVersionId":30747,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Data processing","metadata":{}},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport torch\nfrom torch.utils.data import Dataset\nimport pydicom\nimport numpy as np\nfrom collections import defaultdict\nimport cv2\nimport matplotlib.pyplot as plt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:12.661490Z","iopub.execute_input":"2025-09-10T14:37:12.662345Z","iopub.status.idle":"2025-09-10T14:37:12.667281Z","shell.execute_reply.started":"2025-09-10T14:37:12.662302Z","shell.execute_reply":"2025-09-10T14:37:12.666177Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_train = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv')\nlabel_train = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_label_coordinates.csv')\ntrain_desc = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv')\ncoords_improved = pd.read_csv('/kaggle/input/lumbar-coordinate-pretraining-dataset/coords_rsna_improved.csv')\ncoords_improved.drop('Unnamed: 0', axis=1, inplace=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:12.668947Z","iopub.execute_input":"2025-09-10T14:37:12.669338Z","iopub.status.idle":"2025-09-10T14:37:12.833439Z","shell.execute_reply.started":"2025-09-10T14:37:12.669309Z","shell.execute_reply":"2025-09-10T14:37:12.832602Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sagittal_t1 = train_desc[train_desc['series_description'] == 'Sagittal T1']\nsagittal_t2 = train_desc[train_desc['series_description'] == 'Sagittal T2/STIR']\naxial_t2 = train_desc[train_desc['series_description'] == 'Axial T2']\n\nsagittal_t1 = sagittal_t1.merge(coords_improved)\nsagittal_t2 = sagittal_t2.merge(coords_improved)\naxial_t2 = axial_t2.merge(coords_improved)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:12.834995Z","iopub.execute_input":"2025-09-10T14:37:12.835262Z","iopub.status.idle":"2025-09-10T14:37:12.873298Z","shell.execute_reply.started":"2025-09-10T14:37:12.835240Z","shell.execute_reply":"2025-09-10T14:37:12.872601Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_train.fillna('N', inplace=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:12.874364Z","iopub.execute_input":"2025-09-10T14:37:12.874711Z","iopub.status.idle":"2025-09-10T14:37:12.883041Z","shell.execute_reply.started":"2025-09-10T14:37:12.874688Z","shell.execute_reply":"2025-09-10T14:37:12.882270Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"temp1 = sagittal_t1.groupby(['study_id', 'series_id', 'side']).apply(lambda x: list(zip(x['relative_x'], x['relative_y']))).reset_index(name='coords')\ntemp2 = sagittal_t1.groupby(['study_id', 'series_id', 'side']).apply(lambda x: sum(x['instance_number'])//len(x)).reset_index(name='instance_number')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:12.884925Z","iopub.execute_input":"2025-09-10T14:37:12.885201Z","iopub.status.idle":"2025-09-10T14:37:13.343472Z","shell.execute_reply.started":"2025-09-10T14:37:12.885169Z","shell.execute_reply":"2025-09-10T14:37:13.342514Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"temp2.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:13.344435Z","iopub.execute_input":"2025-09-10T14:37:13.344693Z","iopub.status.idle":"2025-09-10T14:37:13.353269Z","shell.execute_reply.started":"2025-09-10T14:37:13.344672Z","shell.execute_reply":"2025-09-10T14:37:13.352368Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"temp3 = temp1.merge(temp2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:13.354674Z","iopub.execute_input":"2025-09-10T14:37:13.355040Z","iopub.status.idle":"2025-09-10T14:37:13.365282Z","shell.execute_reply.started":"2025-09-10T14:37:13.355008Z","shell.execute_reply":"2025-09-10T14:37:13.364635Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"temp3","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:13.366234Z","iopub.execute_input":"2025-09-10T14:37:13.366443Z","iopub.status.idle":"2025-09-10T14:37:13.389901Z","shell.execute_reply.started":"2025-09-10T14:37:13.366427Z","shell.execute_reply":"2025-09-10T14:37:13.389066Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"scs_severity = df_train.iloc[:, :6]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:13.390893Z","iopub.execute_input":"2025-09-10T14:37:13.391148Z","iopub.status.idle":"2025-09-10T14:37:13.396782Z","shell.execute_reply.started":"2025-09-10T14:37:13.391130Z","shell.execute_reply":"2025-09-10T14:37:13.396133Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"scs_severity.columns = ['study_id', 'scs_l1_l2', 'scs_l2_l3', 'scs_l3_l4', 'scs_l4_l5', 'scs_l5_s1']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:13.397689Z","iopub.execute_input":"2025-09-10T14:37:13.397915Z","iopub.status.idle":"2025-09-10T14:37:13.407776Z","shell.execute_reply.started":"2025-09-10T14:37:13.397898Z","shell.execute_reply":"2025-09-10T14:37:13.406960Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"scs_severity","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:13.410615Z","iopub.execute_input":"2025-09-10T14:37:13.410880Z","iopub.status.idle":"2025-09-10T14:37:13.423909Z","shell.execute_reply.started":"2025-09-10T14:37:13.410854Z","shell.execute_reply":"2025-09-10T14:37:13.423104Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"temp3.merge(scs_severity)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:13.425052Z","iopub.execute_input":"2025-09-10T14:37:13.425919Z","iopub.status.idle":"2025-09-10T14:37:13.454229Z","shell.execute_reply.started":"2025-09-10T14:37:13.425888Z","shell.execute_reply":"2025-09-10T14:37:13.453548Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"temp3['scs'] = temp3.merge(scs_severity).apply(lambda x: list(zip(x['scs_l1_l2'][0], x['scs_l2_l3'][0], x['scs_l3_l4'][0], x['scs_l4_l5'][0], x['scs_l5_s1'][0])), axis=1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:13.454935Z","iopub.execute_input":"2025-09-10T14:37:13.455119Z","iopub.status.idle":"2025-09-10T14:37:13.521471Z","shell.execute_reply.started":"2025-09-10T14:37:13.455104Z","shell.execute_reply":"2025-09-10T14:37:13.520644Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"temp3","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:13.522395Z","iopub.execute_input":"2025-09-10T14:37:13.522664Z","iopub.status.idle":"2025-09-10T14:37:13.546984Z","shell.execute_reply.started":"2025-09-10T14:37:13.522645Z","shell.execute_reply":"2025-09-10T14:37:13.546111Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"temp1 = sagittal_t1.groupby(['study_id', 'series_id', 'side']).apply(lambda x: list(zip(x['relative_x'], x['relative_y']))).reset_index(name='coords')\ntemp2 = sagittal_t1.groupby(['study_id', 'series_id', 'side']).apply(lambda x: sum(x['instance_number'])//len(x)).reset_index(name='instance_number')\n\nscs = df_train.iloc[:, :6]\nl_nfn = df_train.iloc[:, [0, 6, 7, 8, 9, 10]]\nr_nfn = df_train.iloc[:, [0, 11, 12, 13, 14, 15]]\nl_ss = df_train.iloc[:, [0, 16, 17, 18, 19, 20]]\nr_ss = df_train.iloc[:, [0, 21, 22, 23, 24, 25]]\n\nscs.columns = ['study_id', '0', '1', '2', '3', '4']\nl_nfn.columns = ['study_id', '0', '1', '2', '3', '4']\nr_nfn.columns = ['study_id', '0', '1', '2', '3', '4']\nl_ss.columns = ['study_id', '0', '1', '2', '3', '4']\nr_ss.columns = ['study_id', '0', '1', '2', '3', '4']\n\ntemp3 = temp1.merge(temp2)\ntemp3['file_path'] = temp3.apply(lambda x: os.path.join(str(x['study_id']), str(x['series_id']), str(x['instance_number']) + '.dcm'), axis=1)\n\ntemp3['scs'] = temp3.merge(scs).apply(lambda x: list(zip(x['0'][0], x['1'][0], x['2'][0], x['3'][0], x['4'][0])), axis=1)\ntemp3['l_nfn'] = temp3.merge(l_nfn).apply(lambda x: list(zip(x['0'][0], x['1'][0], x['2'][0], x['3'][0], x['4'][0])), axis=1)\ntemp3['r_nfn'] = temp3.merge(r_nfn).apply(lambda x: list(zip(x['0'][0], x['1'][0], x['2'][0], x['3'][0], x['4'][0])), axis=1)\ntemp3['l_ss'] = temp3.merge(l_ss).apply(lambda x: list(zip(x['0'][0], x['1'][0], x['2'][0], x['3'][0], x['4'][0])), axis=1)\ntemp3['r_ss'] = temp3.merge(r_ss).apply(lambda x: list(zip(x['0'][0], x['1'][0], x['2'][0], x['3'][0], x['4'][0])), axis=1)\n\nsagittal_t1_crop = temp3","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:13.548435Z","iopub.execute_input":"2025-09-10T14:37:13.548766Z","iopub.status.idle":"2025-09-10T14:37:14.333477Z","shell.execute_reply.started":"2025-09-10T14:37:13.548737Z","shell.execute_reply":"2025-09-10T14:37:14.332546Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sagittal_t1_crop.head()\n# sagittal_t1_crop_l = sagittal_t1_crop[sagittal_t1_crop['side'] == 'L']\n# sagittal_t1_crop_r = sagittal_t1_crop[sagittal_t1_crop['side'] == 'R']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:14.334445Z","iopub.execute_input":"2025-09-10T14:37:14.334702Z","iopub.status.idle":"2025-09-10T14:37:14.361545Z","shell.execute_reply.started":"2025-09-10T14:37:14.334682Z","shell.execute_reply":"2025-09-10T14:37:14.360639Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"temp1 = sagittal_t2.groupby(['study_id', 'series_id', 'side']).apply(lambda x: list(zip(x['relative_x'], x['relative_y']))).reset_index(name='coords')\ntemp2 = sagittal_t2[sagittal_t2['side'] == 'R']\ntemp2 = temp2.groupby(['study_id', 'series_id', 'side']).apply(lambda x: sum(x['instance_number'])//len(x)).reset_index(name='instance_number')\n\nscs = df_train.iloc[:, :6]\nl_nfn = df_train.iloc[:, [0, 6, 7, 8, 9, 10]]\nr_nfn = df_train.iloc[:, [0, 11, 12, 13, 14, 15]]\nl_ss = df_train.iloc[:, [0, 16, 17, 18, 19, 20]]\nr_ss = df_train.iloc[:, [0, 21, 22, 23, 24, 25]]\n\nscs.columns = ['study_id', '0', '1', '2', '3', '4']\nl_nfn.columns = ['study_id', '0', '1', '2', '3', '4']\nr_nfn.columns = ['study_id', '0', '1', '2', '3', '4']\nl_ss.columns = ['study_id', '0', '1', '2', '3', '4']\nr_ss.columns = ['study_id', '0', '1', '2', '3', '4']\n\ntemp3 = temp1.merge(temp2)\ntemp3['file_path'] = temp3.apply(lambda x: os.path.join(str(x['study_id']), str(x['series_id']), str(x['instance_number']) + '.dcm'), axis=1)\n\ntemp3['scs'] = temp3.merge(scs).apply(lambda x: list(zip(x['0'][0], x['1'][0], x['2'][0], x['3'][0], x['4'][0])), axis=1)\ntemp3['l_nfn'] = temp3.merge(l_nfn).apply(lambda x: list(zip(x['0'][0], x['1'][0], x['2'][0], x['3'][0], x['4'][0])), axis=1)\ntemp3['r_nfn'] = temp3.merge(r_nfn).apply(lambda x: list(zip(x['0'][0], x['1'][0], x['2'][0], x['3'][0], x['4'][0])), axis=1)\ntemp3['l_ss'] = temp3.merge(l_ss).apply(lambda x: list(zip(x['0'][0], x['1'][0], x['2'][0], x['3'][0], x['4'][0])), axis=1)\ntemp3['r_ss'] = temp3.merge(r_ss).apply(lambda x: list(zip(x['0'][0], x['1'][0], x['2'][0], x['3'][0], x['4'][0])), axis=1)\n\nsagittal_t2_crop = temp3","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:14.362467Z","iopub.execute_input":"2025-09-10T14:37:14.362696Z","iopub.status.idle":"2025-09-10T14:37:15.186343Z","shell.execute_reply.started":"2025-09-10T14:37:14.362678Z","shell.execute_reply":"2025-09-10T14:37:15.185635Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sagittal_t2_crop.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:15.187409Z","iopub.execute_input":"2025-09-10T14:37:15.187736Z","iopub.status.idle":"2025-09-10T14:37:15.213880Z","shell.execute_reply.started":"2025-09-10T14:37:15.187707Z","shell.execute_reply":"2025-09-10T14:37:15.213056Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"temp1 = axial_t2.groupby(['study_id', 'series_id', 'level']).apply(lambda x: list(zip(x['relative_x'], x['relative_y']))).reset_index(name='coords')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:15.215108Z","iopub.execute_input":"2025-09-10T14:37:15.215437Z","iopub.status.idle":"2025-09-10T14:37:15.886521Z","shell.execute_reply.started":"2025-09-10T14:37:15.215408Z","shell.execute_reply":"2025-09-10T14:37:15.885656Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"temp1.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:15.887903Z","iopub.execute_input":"2025-09-10T14:37:15.888493Z","iopub.status.idle":"2025-09-10T14:37:15.899798Z","shell.execute_reply.started":"2025-09-10T14:37:15.888460Z","shell.execute_reply":"2025-09-10T14:37:15.898932Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TRAIN_DIR = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images\"\nBASE_DIR = \"/kaggle/input/axial-cropped/kaggle/working/axial_cropped_dataset\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:15.900910Z","iopub.execute_input":"2025-09-10T14:37:15.901244Z","iopub.status.idle":"2025-09-10T14:37:15.908879Z","shell.execute_reply.started":"2025-09-10T14:37:15.901217Z","shell.execute_reply":"2025-09-10T14:37:15.908065Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"import pydicom\nimport numpy as np\nfrom PIL import Image","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:15.909943Z","iopub.execute_input":"2025-09-10T14:37:15.910208Z","iopub.status.idle":"2025-09-10T14:37:15.917373Z","shell.execute_reply.started":"2025-09-10T14:37:15.910188Z","shell.execute_reply":"2025-09-10T14:37:15.916646Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_axis_cropped = pd.read_csv(\"/kaggle/input/axial-with-box-df/axial_with_box_df.csv\")\nlen(df_axis_cropped)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:15.918448Z","iopub.execute_input":"2025-09-10T14:37:15.918778Z","iopub.status.idle":"2025-09-10T14:37:15.941943Z","shell.execute_reply.started":"2025-09-10T14:37:15.918751Z","shell.execute_reply":"2025-09-10T14:37:15.941239Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_axis_cropped.groupby(\"study_id\").size().unique()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:15.942847Z","iopub.execute_input":"2025-09-10T14:37:15.943084Z","iopub.status.idle":"2025-09-10T14:37:15.950025Z","shell.execute_reply.started":"2025-09-10T14:37:15.943065Z","shell.execute_reply":"2025-09-10T14:37:15.949266Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset\nimport numpy as np\nimport cv2\nfrom scipy.interpolate import splprep, splev\nimport pydicom\n\n\ndef load_dcm_image(path):\n    ds = pydicom.dcmread(path)\n    img = ds.pixel_array.astype(np.float32)\n    img -= img.min()\n    img /= (img.max() + 1e-5)\n    return img\n\ndef crop_centered_patches(image, keypoints, box_width=150, box_height=100):\n    \"\"\"\n    Crop các patch hình chữ nhật không xoay, centered tại các keypoint.\n    \"\"\"\n    h, w = image.shape\n    patches = []\n\n    for px, py in keypoints:\n        cx = int(px)\n        cy = int(py)\n\n        x1 = max(cx - box_width // 2, 0)\n        y1 = max(cy - box_height // 2, 0)\n        x2 = min(cx + box_width // 2, w)\n        y2 = min(cy + box_height // 2, h)\n\n        patch = image[y1:y2, x1:x2]\n\n        # Resize lại nếu patch nhỏ hơn box (do bị cắt ở biên)\n        patch = cv2.resize(patch, (box_width, box_height))\n        patches.append(patch)\n\n    return patches\n\ndef crop_patch_by_spline(image, all_coords, disc_idx, box_width=150, box_height=100):\n    h, w = image.shape\n    keypoints = [(x * w, y * h) for x, y in all_coords]\n\n    patches = crop_centered_patches(image, keypoints, box_width, box_height)\n    return patches[disc_idx]\n\ndef load_png_image(path, box_width=150, box_height=100):\n    img = cv2.imread(path, cv2.IMREAD_GRAYSCALE).astype(np.float32)\n    img -= img.min()\n    img /= (img.max() + 1e-5)\n\n    # Resize cho cùng kích thước với crop từ DICOM\n    img = cv2.resize(img, (box_width, box_height))\n    return img\n\nclass SpineMultiViewDiscDataset(Dataset):\n    def __init__(self, df1, df2, df3, resize=(150, 100), base_dir=TRAIN_DIR, base_dir_axial=BASE_DIR, transform=None, iloc=0, condition='scs'):\n        self.samples = []\n        self.resize = resize\n        self.base_dir = base_dir\n        self.base_dir_axial = base_dir_axial\n        self.transform = transform\n        self.iloc = iloc\n        self.condition = condition\n        label_map = {'N': 0, 'M': 1, 'S': 2}\n\n        # Lấy danh sách các study_id duy nhất\n        study_ids = df1['study_id'].unique()\n\n        for sid in study_ids:\n            rows1 = df1[df1['study_id'] == sid]\n            rows2 = df2[df2['study_id'] == sid]\n            rows3 = df3[df3['study_id'] == sid]\n            if len(rows1) == 0 or len(rows2) == 0 or len(rows3) != 5:\n                continue\n\n            # Chọn một dòng duy nhất từ mỗi dataframe\n            row1 = rows1.iloc[self.iloc]\n            row2 = rows2.iloc[0]\n\n            coords1 = row1['coords']\n            coords2 = row2['coords']\n            labels = row1[condition]\n\n            for i in range(5):  # 5 disc levels\n                row3 = rows3[rows3[\"pred_level\"] == i + 1].iloc[0]\n                label_char = labels[0][i] if isinstance(labels, list) else labels[i]\n                label = label_map.get(label_char)\n                sample = {\n                    'file1': f\"{self.base_dir}/{row1['file_path']}\",\n                    'file2': f\"{self.base_dir}/{row2['file_path']}\",\n                    'file3': f\"{self.base_dir_axial}/{row3['study_id']}___{row3['series_id']}___{row3['instance_number']}.png\",\n                    'all_coords1': coords1,\n                    'all_coords2': coords2,\n                    'label': label,\n                    'disc_level': i\n                }\n                self.samples.append(sample)\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n\n        img1 = load_dcm_image(sample['file1'])\n        img2 = load_dcm_image(sample['file2'])\n        img3 = load_png_image(sample['file3'])\n\n        crop1 = crop_patch_by_spline(\n            img1, sample['all_coords1'], sample['disc_level'],\n            box_width=self.resize[0], box_height=self.resize[1]\n        )\n        crop2 = crop_patch_by_spline(\n            img2, sample['all_coords2'], sample['disc_level'],\n            box_width=self.resize[0], box_height=self.resize[1]\n        )\n        \n        stacked = np.stack([crop1, crop2, img3], axis=-1)  # [H, W, 3]\n\n        # apply albumentations nếu có\n        if self.transform:\n            augmented = self.transform(image=stacked)\n            stacked = augmented[\"image\"]\n\n        # chuyển thành [C, H, W]\n        stacked = np.transpose(stacked, (2, 0, 1))\n        stacked = stacked.reshape(3, 1, 1, 100, 150)\n\n        return torch.tensor(stacked, dtype=torch.float32), torch.tensor(sample['label'], dtype=torch.long)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:15.951388Z","iopub.execute_input":"2025-09-10T14:37:15.951656Z","iopub.status.idle":"2025-09-10T14:37:15.968565Z","shell.execute_reply.started":"2025-09-10T14:37:15.951626Z","shell.execute_reply":"2025-09-10T14:37:15.967757Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Augmentation\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\ntrain_transform = A.Compose([\n    A.RandomBrightnessContrast(p=0.5),\n    A.Blur(blur_limit=3, p=0.3),\n    A.GridDistortion(p=0.3),\n    A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.1, rotate_limit=15, p=0.5),\n    # A.CoarseDropout(max_holes=8, max_height=16, max_width=16, p=0.5),\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:15.969872Z","iopub.execute_input":"2025-09-10T14:37:15.970217Z","iopub.status.idle":"2025-09-10T14:37:15.980475Z","shell.execute_reply.started":"2025-09-10T14:37:15.970188Z","shell.execute_reply":"2025-09-10T14:37:15.979715Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dataset = SpineMultiViewDiscDataset(sagittal_t1_crop, sagittal_t2_crop, df_axis_cropped, transform=train_transform)\nprint(len(dataset))\nx_ = None\nfor i, (x, y) in enumerate(dataset):\n    if i == 10:\n        x_ = x\n        print(\"data:\", x.shape)\n        print(\"label:\", y)\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:15.981340Z","iopub.execute_input":"2025-09-10T14:37:15.981532Z","iopub.status.idle":"2025-09-10T14:37:18.896275Z","shell.execute_reply.started":"2025-09-10T14:37:15.981516Z","shell.execute_reply":"2025-09-10T14:37:18.895314Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(x_[0][0][0], cmap=\"gray\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:18.897342Z","iopub.execute_input":"2025-09-10T14:37:18.897639Z","iopub.status.idle":"2025-09-10T14:37:19.139053Z","shell.execute_reply.started":"2025-09-10T14:37:18.897615Z","shell.execute_reply":"2025-09-10T14:37:19.138115Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(x_[1][0][0], cmap=\"gray\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:19.143579Z","iopub.execute_input":"2025-09-10T14:37:19.143890Z","iopub.status.idle":"2025-09-10T14:37:19.306041Z","shell.execute_reply.started":"2025-09-10T14:37:19.143849Z","shell.execute_reply":"2025-09-10T14:37:19.305116Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(x_[2][0][0], cmap=\"gray\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T15:43:51.826326Z","iopub.execute_input":"2025-09-10T15:43:51.826691Z","iopub.status.idle":"2025-09-10T15:43:52.062132Z","shell.execute_reply.started":"2025-09-10T15:43:51.826664Z","shell.execute_reply":"2025-09-10T15:43:52.061299Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import random_split, DataLoader, WeightedRandomSampler\n\n# Dataset gốc\ndataset = SpineMultiViewDiscDataset(sagittal_t1_crop, sagittal_t2_crop, df_axis_cropped, transform=train_transform)\n\n# Tỉ lệ train/val, ví dụ 80/20\ntrain_size = int(0.8 * len(dataset))\nval_size = len(dataset) - train_size\ntrain_dataset, val_dataset = random_split(dataset, [train_size, val_size])\n\n# ---- Train loader với WeightedRandomSampler ----\nlabels_train = [train_dataset[i][1].item() for i in range(len(train_dataset))]\nsample_weights = [1.0 if l == 0 else 2.0 if l == 1 else 4.0 for l in labels_train]\nsample_weights = torch.tensor(sample_weights, dtype=torch.float)\n\nsampler = WeightedRandomSampler(\n    weights=sample_weights,\n    num_samples=len(sample_weights),\n    replacement=True\n)\n\ntrain_loader = DataLoader(train_dataset, batch_size=64, sampler=sampler)\n\n# ---- Val loader (ko sampler, chỉ shuffle=False) ----\nval_loader = DataLoader(val_dataset, batch_size=64, shuffle=False)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:37:19.307033Z","iopub.execute_input":"2025-09-10T14:37:19.307730Z","iopub.status.idle":"2025-09-10T14:38:03.272304Z","shell.execute_reply.started":"2025-09-10T14:37:19.307706Z","shell.execute_reply":"2025-09-10T14:38:03.271450Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Base Model","metadata":{}},{"cell_type":"markdown","source":"## 1. Encoder (resnet18 without FC layer)","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torchvision.models as models\nfrom torch import optim\nfrom torch.utils.data import DataLoader, WeightedRandomSampler\nfrom torch.optim.lr_scheduler import OneCycleLR","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:38:03.273467Z","iopub.execute_input":"2025-09-10T14:38:03.273797Z","iopub.status.idle":"2025-09-10T14:38:03.278293Z","shell.execute_reply.started":"2025-09-10T14:38:03.273770Z","shell.execute_reply":"2025-09-10T14:38:03.277427Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MultiModalEncoderSingleSlice(nn.Module):\n    def __init__(self, pretrained=True):\n        super().__init__()\n        base = models.resnet18(pretrained=pretrained)\n        in_features = base.fc.in_features\n        base.fc = nn.Identity()\n        self.backbone = base\n        self.d_model = in_features\n\n    def forward(self, x):\n        \"\"\"\n        x: [B, 3, 1, 1, H, W]\n        \"\"\"\n        B, V, N, C, H, W = x.shape\n        assert N == 1, \"Hiện tại chỉ support 1 lát mỗi view\"\n\n        x = x.view(B*V, C, H, W)   # [B*3,1,H,W]\n        x = x.repeat(1, 3, 1, 1)   # [B*3,3,H,W]\n\n        feat = self.backbone(x)    # [B*3,512]\n        feat = feat.view(B, V, -1) # [B,3,512]\n\n        return feat\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:38:03.279244Z","iopub.execute_input":"2025-09-10T14:38:03.279465Z","iopub.status.idle":"2025-09-10T14:38:03.289008Z","shell.execute_reply.started":"2025-09-10T14:38:03.279446Z","shell.execute_reply":"2025-09-10T14:38:03.288151Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\n\nx = torch.rand(2, 2, 1, 1, 150, 100)\nencoder = MultiModalEncoderSingleSlice(pretrained=True)  # pretrained=True nếu có internet\nfeat = encoder(x)\nprint(feat.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:38:03.290002Z","iopub.execute_input":"2025-09-10T14:38:03.290229Z","iopub.status.idle":"2025-09-10T14:38:03.580796Z","shell.execute_reply.started":"2025-09-10T14:38:03.290211Z","shell.execute_reply":"2025-09-10T14:38:03.579885Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Transformer Encoder","metadata":{}},{"cell_type":"code","source":"import torch.nn as nn\n\nclass ViewTransformer(nn.Module):\n    def __init__(self, d_model=512, nhead=8, num_layers=2, dim_feedforward=2048, dropout=0.1):\n        super().__init__()\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=d_model,\n            nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout,\n            batch_first=True   # để input là [B, seq, d]\n        )\n        self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)\n\n    def forward(self, x):\n        \"\"\"\n        x: [B, 3, 512]   # 3 views = sequence length 3\n        return: [B, 3, 512]\n        \"\"\"\n        out = self.transformer(x)\n        return out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:38:03.582254Z","iopub.execute_input":"2025-09-10T14:38:03.582577Z","iopub.status.idle":"2025-09-10T14:38:03.588397Z","shell.execute_reply.started":"2025-09-10T14:38:03.582547Z","shell.execute_reply":"2025-09-10T14:38:03.587409Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"view_transformer = ViewTransformer()\nout = view_transformer(feat)\nprint(out.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:38:03.589557Z","iopub.execute_input":"2025-09-10T14:38:03.590513Z","iopub.status.idle":"2025-09-10T14:38:03.637548Z","shell.execute_reply.started":"2025-09-10T14:38:03.590484Z","shell.execute_reply":"2025-09-10T14:38:03.636668Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Center - Spinal Canal Stenosis condition","metadata":{}},{"cell_type":"markdown","source":"## 1. Model","metadata":{}},{"cell_type":"code","source":"class MultiModalPipeline_center(nn.Module):\n    def __init__(self, pretrained=True, num_classes=3):\n        super().__init__()\n        self.encoder = MultiModalEncoderSingleSlice(pretrained=pretrained)  # [B,3,512]\n        self.transformer = ViewTransformer(d_model=512, nhead=8, num_layers=2)\n        self.fc = nn.Linear(512, num_classes)\n\n    def forward(self, x):\n        feat = self.encoder(x)              # [B,3,512]\n        feat_trans = self.transformer(feat) # [B,3,512]\n        # pool across views (mean pooling)\n        pooled = feat_trans.mean(dim=1)     # [B,512]\n        out = self.fc(pooled)               # [B,num_classes]\n        return out, feat_trans\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:38:03.638930Z","iopub.execute_input":"2025-09-10T14:38:03.639262Z","iopub.status.idle":"2025-09-10T14:38:03.644793Z","shell.execute_reply.started":"2025-09-10T14:38:03.639233Z","shell.execute_reply":"2025-09-10T14:38:03.643965Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x = torch.randn(2, 3, 1, 1, 100, 150)  # batch=2\nmodel = MultiModalPipeline_center(pretrained=True, num_classes=3)\n\nlogits, feat_trans = model(x)\nprint(\"Logits:\", logits.shape)       # [2,2]\nprint(\"Transformer output:\", feat_trans.shape)  # [2,3,512]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:38:03.645972Z","iopub.execute_input":"2025-09-10T14:38:03.646287Z","iopub.status.idle":"2025-09-10T14:38:03.997148Z","shell.execute_reply.started":"2025-09-10T14:38:03.646259Z","shell.execute_reply":"2025-09-10T14:38:03.996256Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Train model","metadata":{}},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = MultiModalPipeline_center(pretrained=True, num_classes=3).to(device)\n\n# định nghĩa weight cho từng class\nclass_weights = torch.tensor([1.0, 2.0, 4.0], dtype=torch.float32).to(device)\n\ncriterion = nn.CrossEntropyLoss(weight=class_weights)\n# ===== Optimizer =====\noptimizer = optim.AdamW(model.parameters(), lr=2.5e-4)\n\n# ===== Scheduler: OneCycleLR =====\nepochs = 30\nsteps_per_epoch = len(train_loader)\ntotal_steps = steps_per_epoch * epochs\nwarmup_steps = int(total_steps * 0.3)   # 3/10 = 30% warmup\n\nscheduler = torch.optim.lr_scheduler.OneCycleLR(\n    optimizer,\n    max_lr=2.5e-5,\n    steps_per_epoch=len(train_loader),\n    epochs=epochs\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:38:03.998240Z","iopub.execute_input":"2025-09-10T14:38:03.998502Z","iopub.status.idle":"2025-09-10T14:38:04.282212Z","shell.execute_reply.started":"2025-09-10T14:38:03.998481Z","shell.execute_reply":"2025-09-10T14:38:04.281499Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\nimport torch\nimport matplotlib.pyplot as plt\n\nnum_epochs = 30\nbest_acc = 0.0\nbest_model_state = None  # lưu state_dict tốt nhất\nsave_path = \"best_model.pth\"\n\n# ---- List để lưu val acc từng epoch ----\nval_acc_history = []\n\nfor epoch in range(num_epochs):\n    # ---- Training ----\n    model.train()\n    total_loss = 0\n    for imgs, labels in tqdm(train_loader, desc=f\"Epoch {epoch+1}/{num_epochs} [Train]\"):\n        imgs, labels = imgs.to(device), labels.to(device)\n\n        optimizer.zero_grad()\n        outputs = model(imgs)[0]\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        scheduler.step()\n        total_loss += loss.item()\n\n    avg_loss = total_loss / len(train_loader)\n    print(f\"Epoch {epoch+1}, Train Loss: {avg_loss:.4f}\")\n\n    # ---- Validation ----\n    model.eval()\n    correct, total = 0, 0\n    with torch.no_grad():\n        for imgs, labels in tqdm(val_loader, desc=f\"Epoch {epoch+1}/{num_epochs} [Val]\"):\n            imgs, labels = imgs.to(device), labels.to(device)\n            outputs = model(imgs)[0]\n            preds = outputs.argmax(dim=1)\n            correct += (preds == labels).sum().item()\n            total += labels.size(0)\n\n    acc = correct / total\n    val_acc_history.append(acc)  # lưu acc cho epoch này\n    print(f\"Epoch {epoch+1}, Validation Accuracy: {acc:.4f}\")\n\n    # ---- Lưu model tốt nhất ----\n    if acc > best_acc:\n        best_acc = acc\n        best_model_state = model.state_dict()\n        torch.save({\n            \"epoch\": epoch + 1,\n            \"model_state_dict\": best_model_state,\n            \"optimizer_state_dict\": optimizer.state_dict(),\n            \"best_acc\": best_acc,\n        }, save_path)\n        print(f\"Saved best model with val acc: {best_acc:.4f}\")\n\n# --- Load lại model tốt nhất để eval ---\nmodel.load_state_dict(best_model_state)\nmodel.eval()\nprint(f\"Best validation accuracy: {best_acc:.4f}\")\n\n# --- Plot val accuracy qua các epoch ---\nplt.figure()\nplt.plot(range(1, num_epochs+1), val_acc_history, marker='o')\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Validation Accuracy\")\nplt.title(\"Validation Accuracy per Epoch\")\nplt.grid(True)\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T14:38:04.283682Z","iopub.execute_input":"2025-09-10T14:38:04.284152Z","iopub.status.idle":"2025-09-10T15:07:32.727093Z","shell.execute_reply.started":"2025-09-10T14:38:04.284124Z","shell.execute_reply":"2025-09-10T15:07:32.726355Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Evaluate","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\nfrom sklearn.metrics import confusion_matrix\nfrom sklearn.metrics import classification_report, accuracy_score\n\n# ---- Validation + lưu nhãn để vẽ confusion matrix ----\nmodel.eval()\nall_preds = []\nall_labels = []\n\nwith torch.no_grad():\n    for imgs, labels in tqdm(val_loader, desc=\"Validation\"):\n        imgs, labels = imgs.to(device), labels.to(device)\n        outputs = model(imgs)[0]\n        preds = outputs.argmax(dim=1)\n\n        all_preds.extend(preds.cpu().numpy())\n        all_labels.extend(labels.cpu().numpy())\n\n# Accuracy\ncorrect = sum([p == t for p, t in zip(all_preds, all_labels)])\ntotal = len(all_labels)\nacc = correct / total\nprint(f\"Validation Accuracy: {acc:.4f}\")\n\n# Metric precision, recall, f1-score\n# Accuracy\nacc = accuracy_score(all_labels, all_preds)\nprint(f\"Accuracy: {acc:.4f}\")\n\n# Precision, Recall, F1-score cho từng class\nreport = classification_report(all_labels, all_preds, digits=4)\nprint(\"Classification Report:\\n\", report)\n\n# ---- Confusion Matrix ----\ncm = confusion_matrix(all_labels, all_preds)\nplt.figure(figsize=(8, 6))\nsns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\")\nplt.title(\"Confusion Matrix\")\nplt.xlabel(\"Predicted\")\nplt.ylabel(\"True\")\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T15:07:32.728189Z","iopub.execute_input":"2025-09-10T15:07:32.728434Z","iopub.status.idle":"2025-09-10T15:07:43.869817Z","shell.execute_reply.started":"2025-09-10T15:07:32.728413Z","shell.execute_reply":"2025-09-10T15:07:43.868921Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Side - Multilabel","metadata":{}},{"cell_type":"markdown","source":"## 1. Class dataset","metadata":{}},{"cell_type":"code","source":"class SpineMultiViewDiscDataset(Dataset):\n    def __init__(self, df1, df2, df3, resize=(150, 100), base_dir=TRAIN_DIR, base_dir_axial=BASE_DIR, transform=None, iloc=0, conditions=['scs']):\n        self.samples = []\n        self.resize = resize\n        self.base_dir = base_dir\n        self.base_dir_axial = base_dir_axial\n        self.transform = transform\n        self.iloc = iloc\n        self.conditions = conditions if isinstance(conditions, list) else [conditions]\n\n        label_map = {'N': 0, 'M': 1, 'S': 2}\n\n        # Lấy danh sách study_id\n        study_ids = df1['study_id'].unique()\n\n        for sid in study_ids:\n            rows1 = df1[df1['study_id'] == sid]\n            rows2 = df2[df2['study_id'] == sid]\n            rows3 = df3[df3['study_id'] == sid]\n            if len(rows1) == 0 or len(rows2) == 0 or len(rows3) != 5:\n                continue\n\n            row1 = rows1.iloc[self.iloc]\n            row2 = rows2.iloc[0]\n\n            coords1 = row1['coords']\n            coords2 = row2['coords']\n\n            # mỗi condition sẽ có 1 list label riêng\n            labels_dict = {cond: row1[cond] for cond in self.conditions}\n\n            for i in range(5):  # 5 disc levels\n                row3 = rows3[rows3[\"pred_level\"] == i + 1].iloc[0]\n\n                # tạo vector nhãn cho tất cả conditions\n                labels_vec = []\n                for cond in self.conditions:\n                    cond_labels = labels_dict[cond]\n                    label_char = cond_labels[0][i] if isinstance(cond_labels, list) else cond_labels[i]\n                    label = label_map.get(label_char)\n                    labels_vec.append(label)\n\n                sample = {\n                    'file1': f\"{self.base_dir}/{row1['file_path']}\",\n                    'file2': f\"{self.base_dir}/{row2['file_path']}\",\n                    'file3': f\"{self.base_dir_axial}/{row3['study_id']}___{row3['series_id']}___{row3['instance_number']}.png\",\n                    'all_coords1': coords1,\n                    'all_coords2': coords2,\n                    'labels': labels_vec,\n                    'disc_level': i\n                }\n                self.samples.append(sample)\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n\n        img1 = load_dcm_image(sample['file1'])\n        img2 = load_dcm_image(sample['file2'])\n        img3 = load_png_image(sample['file3'])\n\n        crop1 = crop_patch_by_spline(img1, sample['all_coords1'], sample['disc_level'],\n                                     box_width=self.resize[0], box_height=self.resize[1])\n        crop2 = crop_patch_by_spline(img2, sample['all_coords2'], sample['disc_level'],\n                                     box_width=self.resize[0], box_height=self.resize[1])\n        \n        stacked = np.stack([crop1, crop2, img3], axis=-1)  # [H, W, 3]\n\n        if self.transform:\n            augmented = self.transform(image=stacked)\n            stacked = augmented[\"image\"]\n\n        stacked = np.transpose(stacked, (2, 0, 1))\n        stacked = stacked.reshape(3, 1, 1, 100, 150)\n\n        labels = torch.tensor(sample['labels'], dtype=torch.long)  # vector nhãn\n\n        return torch.tensor(stacked, dtype=torch.float32), labels","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T15:07:43.870936Z","iopub.execute_input":"2025-09-10T15:07:43.871197Z","iopub.status.idle":"2025-09-10T15:07:43.883550Z","shell.execute_reply.started":"2025-09-10T15:07:43.871176Z","shell.execute_reply":"2025-09-10T15:07:43.882680Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import random_split, DataLoader, WeightedRandomSampler\n\ndataset_L = SpineMultiViewDiscDataset(sagittal_t1_crop, sagittal_t2_crop, df_axis_cropped, transform=train_transform, iloc=0, conditions=[\"l_nfn\", \"l_ss\"])\n# Tỉ lệ train/val, ví dụ 80/20 \ntrain_size = int(0.8 * len(dataset_L)) \nval_size = len(dataset_L) - train_size \ntrain_dataset_L, val_dataset_L = random_split(dataset_L, [train_size, val_size])\n\n# ---- Train loader ----\ntrain_loader_L = DataLoader(\n    train_dataset_L,\n    batch_size=64,\n    shuffle=True,       # bỏ sampler, dùng shuffle\n    num_workers=4,\n    pin_memory=True\n)\n\n# ---- Validation loader ----\nval_loader_L = DataLoader(\n    val_dataset_L,\n    batch_size=64,\n    shuffle=False,\n    num_workers=4,\n    pin_memory=True\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T15:07:43.884556Z","iopub.execute_input":"2025-09-10T15:07:43.884872Z","iopub.status.idle":"2025-09-10T15:07:46.587128Z","shell.execute_reply.started":"2025-09-10T15:07:43.884852Z","shell.execute_reply":"2025-09-10T15:07:46.586172Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Side Model","metadata":{}},{"cell_type":"code","source":"class MultiModalPipeline(nn.Module):\n    def __init__(self, pretrained=True, num_classes=[3]):\n        \"\"\"\n        num_classes: list[int], mỗi phần tử là số class cho một condition\n                     ví dụ [3, 3] nghĩa là 2 conditions, mỗi cái có 3 class\n        \"\"\"\n        super().__init__()\n        self.encoder = MultiModalEncoderSingleSlice(pretrained=pretrained)  # [B,3,512]\n        self.transformer = ViewTransformer(d_model=512, nhead=8, num_layers=2)\n\n        # tạo các head cho từng condition\n        self.heads = nn.ModuleList([nn.Linear(512, n_cls) for n_cls in num_classes])\n\n    def forward(self, x):\n        feat = self.encoder(x)              # [B,3,512]\n        feat_trans = self.transformer(feat) # [B,3,512]\n\n        # pool across views (mean pooling)\n        pooled = feat_trans.mean(dim=1)     # [B,512]\n\n        # mỗi condition 1 head riêng\n        outputs = [head(pooled) for head in self.heads]  # list of [B, num_classes]\n\n        return outputs, feat_trans","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Train and Val function","metadata":{}},{"cell_type":"code","source":"import torch\nfrom tqdm import tqdm\n\ndef train(model, train_loader, val_loader, criterion, optimizer, device, num_epochs=1, save_path=\"best_model.pth\"):\n    \"\"\"\n    Train a multi-task model, save the best model, return it and store val acc history per condition.\n    \"\"\"\n    best_acc = 0.0\n    best_model_state = None\n    val_acc_history = []  # lưu acc từng condition mỗi epoch\n\n    for epoch in range(num_epochs):\n        # ---- Training ----\n        model.train()\n        total_loss = 0\n        for imgs, labels in tqdm(train_loader, desc=f\"Epoch {epoch+1}/{num_epochs} [Train]\"):\n            imgs, labels = imgs.to(device), labels.to(device)\n\n            optimizer.zero_grad()\n            outputs, _ = model(imgs)\n\n            # Multi-task loss\n            loss = 0\n            for i, logits in enumerate(outputs):\n                loss += criterion(logits, labels[:, i])\n            loss = loss / len(outputs)\n\n            loss.backward()\n            optimizer.step()\n            total_loss += loss.item()\n\n        avg_loss = total_loss / len(train_loader)\n        print(f\"Epoch {epoch+1}, Train Loss: {avg_loss:.4f}\")\n\n        # ---- Validation ----\n        model.eval()\n        correct = [0] * len(outputs)\n        total = 0\n        with torch.no_grad():\n            for imgs, labels in tqdm(val_loader, desc=f\"Epoch {epoch+1}/{num_epochs} [Val]\"):\n                imgs, labels = imgs.to(device), labels.to(device)\n                outputs, _ = model(imgs)\n\n                for i, logits in enumerate(outputs):\n                    preds = logits.argmax(dim=1)\n                    correct[i] += (preds == labels[:, i]).sum().item()\n\n                total += labels.size(0)\n\n        accs = [c / total for c in correct]\n        val_acc_history.append(accs)  # lưu acc từng condition\n        avg_acc = sum(accs) / len(accs)\n        acc_str = \" | \".join([f\"Cond{i+1}: {a:.4f}\" for i, a in enumerate(accs)])\n        print(f\"Epoch {epoch+1}, Validation Accuracies: {acc_str}\")\n\n        # ---- Lưu model tốt nhất ----\n        if avg_acc > best_acc:\n            best_acc = avg_acc\n            best_model_state = model.state_dict()\n            torch.save({\n                \"epoch\": epoch + 1,\n                \"model_state_dict\": best_model_state,\n                \"optimizer_state_dict\": optimizer.state_dict(),\n                \"best_acc\": best_acc,\n                \"best_accs_per_condition\": accs,  # lưu luôn acc từng condition tốt nhất\n            }, save_path)\n            print(f\"Saved best model with avg acc: {best_acc:.4f}\")\n\n    # --- Trả về model tốt nhất và lịch sử val acc ---\n    model.load_state_dict(best_model_state)\n    model.eval()\n    return model, val_acc_history","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T15:07:46.588157Z","iopub.execute_input":"2025-09-10T15:07:46.588427Z","iopub.status.idle":"2025-09-10T15:07:46.599953Z","shell.execute_reply.started":"2025-09-10T15:07:46.588405Z","shell.execute_reply":"2025-09-10T15:07:46.599077Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom sklearn.metrics import confusion_matrix, classification_report\nfrom tqdm import tqdm\nimport torch\n\ndef validate_and_plot_cm(model, val_loader, device):\n    \"\"\"\n    Validate a multi-condition model, print accuracy, precision, recall, F1, \n    and plot confusion matrices for each condition.\n\n    Args:\n        model: PyTorch model, should return (outputs, aux) where outputs is a list of logits per condition.\n        val_loader: DataLoader for validation data.\n        device: torch.device.\n    \"\"\"\n    model.eval()\n    correct = None\n    total = 0\n\n    # Lưu nhãn để vẽ confusion matrix và tính metrics\n    all_preds, all_labels = None, None\n\n    with torch.no_grad():\n        for imgs, labels in tqdm(val_loader, desc=\"Validation\"):\n            imgs, labels = imgs.to(device), labels.to(device)\n            outputs, _ = model(imgs)\n\n            if correct is None:\n                correct = [0] * len(outputs)\n                all_preds = [[] for _ in range(len(outputs))]\n                all_labels = [[] for _ in range(len(outputs))]\n\n            for i, logits in enumerate(outputs):\n                preds = logits.argmax(dim=1)\n                correct[i] += (preds == labels[:, i]).sum().item()\n                all_preds[i].extend(preds.cpu().numpy())\n                all_labels[i].extend(labels[:, i].cpu().numpy())\n\n            total += labels.size(0)\n\n    # Accuracy\n    accs = [c / total for c in correct]\n    for i, a in enumerate(accs):\n        print(f\"Condition {i+1} Accuracy: {a:.4f}\")\n\n    # Vẽ confusion matrix và tính metrics cho từng condition\n    for i in range(len(outputs)):\n        cm = confusion_matrix(all_labels[i], all_preds[i])\n        plt.figure(figsize=(6, 5))\n        sns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\")\n        plt.title(f\"Confusion Matrix - Condition {i+1}\")\n        plt.xlabel(\"Predicted\")\n        plt.ylabel(\"True\")\n        plt.show()\n\n        # Classification report: precision, recall, f1-score\n        print(f\"Condition {i+1} Metrics:\")\n        report = classification_report(all_labels[i], all_preds[i], digits=4)\n        print(report)\n\nimport matplotlib.pyplot as plt\n\ndef plot_val_acc(val_acc_history):\n    \"\"\"\n    Plot validation accuracy per condition over epochs.\n\n    Args:\n        val_acc_history: list of list, mỗi phần tử là acc từng condition của 1 epoch\n    \"\"\"\n    num_conditions = len(val_acc_history[0])\n    epochs = range(1, len(val_acc_history) + 1)\n\n    plt.figure(figsize=(8, 6))\n    for i in range(num_conditions):\n        cond_acc = [epoch_acc[i] for epoch_acc in val_acc_history]\n        plt.plot(epochs, cond_acc, marker='o', label=f'Condition {i+1}')\n\n    plt.title(\"Validation Accuracy per Condition\")\n    plt.xlabel(\"Epoch\")\n    plt.ylabel(\"Accuracy\")\n    plt.ylim(0, 1)\n    plt.xticks(epochs)\n    plt.grid(True)\n    plt.legend()\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T15:07:46.601564Z","iopub.execute_input":"2025-09-10T15:07:46.601919Z","iopub.status.idle":"2025-09-10T15:07:46.617180Z","shell.execute_reply.started":"2025-09-10T15:07:46.601888Z","shell.execute_reply":"2025-09-10T15:07:46.616526Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Left Neural Foraminal Narrowing, Left Subarticular Stenosis","metadata":{}},{"cell_type":"code","source":"# ===== Device =====\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel_L = MultiModalPipeline(pretrained=True, num_classes=[3, 3]).to(device)\n\n# ===== Loss (có trọng số cho class imbalance) =====\nclass_weights = torch.tensor([1.0, 2.0, 4.0], dtype=torch.float32).to(device)\ncriterion = nn.CrossEntropyLoss(weight=class_weights)\n\n# ===== Optimizer =====\noptimizer = optim.AdamW(model_L.parameters(), lr=2.5e-4)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T15:07:46.633243Z","iopub.execute_input":"2025-09-10T15:07:46.633487Z","iopub.status.idle":"2025-09-10T15:07:46.914330Z","shell.execute_reply.started":"2025-09-10T15:07:46.633468Z","shell.execute_reply":"2025-09-10T15:07:46.913650Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_L, val_acc_history_L = train(model_L, train_loader_L, val_loader_L, criterion, optimizer, device, num_epochs=30)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T15:07:46.915330Z","iopub.execute_input":"2025-09-10T15:07:46.915575Z","iopub.status.idle":"2025-09-10T15:19:27.274206Z","shell.execute_reply.started":"2025-09-10T15:07:46.915555Z","shell.execute_reply":"2025-09-10T15:19:27.273190Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Evaluate","metadata":{}},{"cell_type":"code","source":"validate_and_plot_cm(model_L, val_loader_L, device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T15:19:27.275859Z","iopub.execute_input":"2025-09-10T15:19:27.276163Z","iopub.status.idle":"2025-09-10T15:19:33.582789Z","shell.execute_reply.started":"2025-09-10T15:19:27.276138Z","shell.execute_reply":"2025-09-10T15:19:33.581868Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_val_acc(val_acc_history_L)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T15:19:33.583926Z","iopub.execute_input":"2025-09-10T15:19:33.584231Z","iopub.status.idle":"2025-09-10T15:19:33.874321Z","shell.execute_reply.started":"2025-09-10T15:19:33.584208Z","shell.execute_reply":"2025-09-10T15:19:33.873462Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Right Neural Foraminal Narrowing, Right Subarticular Stenosis","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import random_split, DataLoader, WeightedRandomSampler\n\ndataset_R = SpineMultiViewDiscDataset(sagittal_t1_crop, sagittal_t2_crop, df_axis_cropped, transform=train_transform, iloc=1, conditions=[\"r_nfn\", \"r_ss\"])\n# Tỉ lệ train/val, ví dụ 80/20 \ntrain_size = int(0.8 * len(dataset_R)) \nval_size = len(dataset) - train_size \ntrain_dataset_R, val_dataset_R = random_split(dataset_R, [train_size, val_size])\n\n# ---- Train loader ----\ntrain_loader_R = DataLoader(\n    train_dataset_R,\n    batch_size=64,\n    shuffle=True,       # bỏ sampler, dùng shuffle\n    num_workers=4,\n    pin_memory=True\n)\n\n# ---- Validation loader ----\nval_loader_R = DataLoader(\n    val_dataset_R,\n    batch_size=64,\n    shuffle=False,\n    num_workers=4,\n    pin_memory=True\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T15:19:33.875347Z","iopub.execute_input":"2025-09-10T15:19:33.875632Z","iopub.status.idle":"2025-09-10T15:19:36.547111Z","shell.execute_reply.started":"2025-09-10T15:19:33.875610Z","shell.execute_reply":"2025-09-10T15:19:36.546366Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===== Device =====\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel_R = MultiModalPipeline(pretrained=True, num_classes=[3, 3]).to(device)\n\n# ===== Loss (có trọng số cho class imbalance) =====\nclass_weights = torch.tensor([1.0, 2.0, 4.0], dtype=torch.float32).to(device)\ncriterion = nn.CrossEntropyLoss(weight=class_weights)\n\n# ===== Optimizer =====\noptimizer = optim.AdamW(model_R.parameters(), lr=2.5e-4)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T15:19:36.548193Z","iopub.execute_input":"2025-09-10T15:19:36.548460Z","iopub.status.idle":"2025-09-10T15:19:36.847796Z","shell.execute_reply.started":"2025-09-10T15:19:36.548440Z","shell.execute_reply":"2025-09-10T15:19:36.847043Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_R, val_acc_history_R = train(model_R, train_loader_R, val_loader_R, criterion, optimizer, device, num_epochs=30)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T15:19:36.848960Z","iopub.execute_input":"2025-09-10T15:19:36.849235Z","iopub.status.idle":"2025-09-10T15:31:49.702809Z","shell.execute_reply.started":"2025-09-10T15:19:36.849214Z","shell.execute_reply":"2025-09-10T15:31:49.701667Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Evaluate","metadata":{}},{"cell_type":"code","source":"validate_and_plot_cm(model_R, val_loader_R, device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T15:31:49.704326Z","iopub.execute_input":"2025-09-10T15:31:49.704602Z","iopub.status.idle":"2025-09-10T15:31:55.294770Z","shell.execute_reply.started":"2025-09-10T15:31:49.704563Z","shell.execute_reply":"2025-09-10T15:31:55.293830Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_val_acc(val_acc_history_R)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-10T15:31:55.296268Z","iopub.execute_input":"2025-09-10T15:31:55.296643Z","iopub.status.idle":"2025-09-10T15:31:55.637133Z","shell.execute_reply.started":"2025-09-10T15:31:55.296608Z","shell.execute_reply":"2025-09-10T15:31:55.636244Z"}},"outputs":[],"execution_count":null}]}