{"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}],"dockerImageVersionId":30747,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 1. Import thư viện cần thiết","metadata":{}},{"cell_type":"code","source":"import os\nimport glob\nimport random\nimport math\n\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport matplotlib.gridspec as gridspec\nfrom tqdm import tqdm\nfrom types import SimpleNamespace\nimport albumentations as A\nimport cv2\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport timm\nimport pydicom\n\ndef set_seed(seed=1234):\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = False\n    torch.backends.cudnn.benchmark = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-22T01:45:38.263058Z","iopub.execute_input":"2025-08-22T01:45:38.263908Z","iopub.status.idle":"2025-08-22T01:45:45.201758Z","shell.execute_reply.started":"2025-08-22T01:45:38.263867Z","shell.execute_reply":"2025-08-22T01:45:45.200709Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 2. Xử lý dữ liệu\n\n\nWe improve the competition coordinates by labelling the left side of each disc. This gives us the angle of orientation which can be used to improve cropping of each disc.\n\nThanks to [Ian Pan](https://www.kaggle.com/vaillant) for sharing some helper functions for loading the dicom data. See his great notebook [here](https://www.kaggle.com/code/vaillant/cross-reference-images-in-different-mri-planes?scriptVersionId=182551992&cellId=2).","metadata":{}},{"cell_type":"markdown","source":"## a. Dùng bộ dữ liệu đã được gán nhãn 2 phía của mỗi đĩa đệm","metadata":{}},{"cell_type":"code","source":"def convert_to_8bit(x):\n    lower, upper = np.percentile(x, (1, 99))\n    x = np.clip(x, lower, upper)\n    x = x - np.min(x)\n    x = x / np.max(x) \n    return (x * 255).astype(\"uint8\")\n\n\ndef load_dicom_stack(dicom_folder, plane, reverse_sort=False):\n    dicom_files = glob.glob(os.path.join(dicom_folder, \"*.dcm\"))\n    dicoms = [pydicom.dcmread(f) for f in dicom_files]\n    plane = {\"sagittal\": 0, \"coronal\": 1, \"axial\": 2}[plane.lower()]\n    positions = np.asarray([float(d.ImagePositionPatient[plane]) for d in dicoms])\n    # if reverse_sort=False, then increasing array index will be from RIGHT->LEFT and CAUDAL->CRANIAL\n    # thus we do reverse_sort=True for axial so increasing array index is craniocaudal\n    idx = np.argsort(-positions if reverse_sort else positions)\n    ipp = np.asarray([d.ImagePositionPatient for d in dicoms]).astype(\"float\")[idx]\n    array = np.stack([d.pixel_array.astype(\"float32\") for d in dicoms])\n    array = array[idx]\n    return {\"array\": convert_to_8bit(array), \"positions\": ipp, \"pixel_spacing\": np.asarray(dicoms[0].PixelSpacing).astype(\"float\")}\n\nimage_dir = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/\"","metadata":{"execution":{"iopub.status.busy":"2025-08-22T01:45:53.051568Z","iopub.execute_input":"2025-08-22T01:45:53.052376Z","iopub.status.idle":"2025-08-22T01:45:53.059122Z","shell.execute_reply.started":"2025-08-22T01:45:53.052348Z","shell.execute_reply":"2025-08-22T01:45:53.058135Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"resize_transform= A.Compose([\n    A.LongestMaxSize(max_size=256, interpolation=cv2.INTER_CUBIC, always_apply=True),\n    A.PadIfNeeded(min_height=256, min_width=256, border_mode=cv2.BORDER_CONSTANT, value=(0, 0, 0), always_apply=True),\n])\n\n# Hàm tính góc giữa 2 điểm so với trục hoành\ndef angle_of_line(x1, y1, x2, y2):\n    return math.degrees(math.atan2(-(y2-y1), x2-x1))\n\n# Hàm vẽ \ndef plot_img(img, coords_temp):\n    # Plot img\n    fig, ax = plt.subplots()\n    ax.imshow(img, cmap='gray')\n    h, w = img.shape\n    \n    # Kepoints as pairs\n    p= coords_temp.groupby(\"level\") \\\n                  .apply(lambda g: list(zip(g['relative_x'], g['relative_y'])), include_groups=False) \\\n                  .reset_index(drop=False, name=\"vals\")\n    \n    # Plot keypoints\n    for _, row in p.iterrows():\n        level = row['level']\n        x = [_[0]*w for _ in row[\"vals\"]]\n        y = [_[1]*h for _ in row[\"vals\"]]\n        ax.plot(x, y, marker='o')\n    ax.axis('off')\n    plt.show()\n\n# Hàm cắt 5 đĩa đệm ra khỏi ảnh gốc\ndef plot_5_crops(img, coords_temp):\n    # Create a figure and axis for the grid\n    fig = plt.figure(figsize=(10, 10))\n    gs = gridspec.GridSpec(1, 5, width_ratios=[1]*5)\n    \n    # Plot the crops\n    p= coords_temp.groupby(\"level\").apply(lambda g: list(zip(g['relative_x'], g['relative_y'])), include_groups=False).reset_index(drop=False, name=\"vals\")\n    print(\"P: \", p)\n    for idx, (_, row) in enumerate(p.iterrows()):\n        # Copy of img\n        img_copy= img.copy()\n        h, w = img.shape\n\n        # Extract Keypoints\n        level = row['level']\n        vals = sorted(row[\"vals\"], key=lambda x: x[0])\n        a,b= vals\n        a= (a[0]*w, a[1]*h)\n        b= (b[0]*w, b[1]*h)\n\n        # Rotate\n        rotate_angle= angle_of_line(a[0], a[1], b[0], b[1])\n        transform = A.Compose([\n            A.Rotate(limit=(-rotate_angle, -rotate_angle), p=1.0),\n        ], keypoint_params= A.KeypointParams(format='xy', remove_invisible=False),\n        )\n\n        t= transform(image=img_copy, keypoints=[a,b])\n        img_copy= t[\"image\"]\n        a,b= t[\"keypoints\"]\n        \n        # Crop + Resize\n        img_copy= crop_between_keypoints(img_copy, a, b)\n        img_copy= resize_transform(image=img_copy)[\"image\"]\n        \n        # Plot\n        ax = plt.subplot(gs[idx])\n        ax.imshow(img_copy, cmap='gray')\n        ax.set_title(level)\n        ax.axis('off')\n    plt.show()\n    \n    \ndef crop_between_keypoints(img, keypoint1, keypoint2):\n    h, w = img.shape\n    x1, y1 = int(keypoint1[0]), int(keypoint1[1])\n    x2, y2 = int(keypoint2[0]), int(keypoint2[1])\n    \n    # Calculate bounding box around the keypoints\n    left = int(min(x1, x2))\n    right = int(max(x1, x2))\n    top = int(min(y1, y2) - (h * 0.1))\n    bottom = int(max(y1, y2) + (h * 0.1))\n            \n    # Crop the image\n    return img[top:bottom, left:right]","metadata":{"execution":{"iopub.status.busy":"2025-08-22T01:45:56.988858Z","iopub.execute_input":"2025-08-22T01:45:56.989301Z","iopub.status.idle":"2025-08-22T01:45:57.002573Z","shell.execute_reply.started":"2025-08-22T01:45:56.989272Z","shell.execute_reply":"2025-08-22T01:45:57.001753Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Trực quan hóa bộ dữ liệu phụ","metadata":{}},{"cell_type":"code","source":"SEED= 10\nN= 2\n\n# Load series_ids\ndfd= pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv\")\ndfd= dfd[dfd.series_description == \"Sagittal T2/STIR\"]\n# dfd= dfd.sample(frac=1, random_state=SEED).head(N) # chỉ dùng hàng này khi plot\nprint(dfd)\n\n# Load coords\ncoords= pd.read_csv(\"/kaggle/input/lumbar-coordinate-pretraining-dataset/coords_rsna_improved.csv\")\ncoords= coords.sort_values([\"series_id\", \"level\", \"side\"]).reset_index(drop=True)\n\n# DataFrame để chứa tất cả coords_temp\nall_coords = []\n\n# Plot samples\nfor idx, row in dfd.iterrows():\n    try:\n        # print(\"-\"*25, f\" STUDY_ID: {row.study_id}, SERIES_ID: {row.series_id} \", \"-\"*25)\n        # sag_t2 = load_dicom_stack(os.path.join(image_dir, str(row.study_id), str(row.series_id)), plane=\"sagittal\")\n        \n        # # Img + Coords\n        # img= sag_t2[\"array\"][len(sag_t2[\"array\"])//2]\n        # coords_temp= coords[coords[\"series_id\"] == row.series_id].copy()\n        \n        # Thêm study_id để tracking\n        coords_temp[\"study_id\"] = row.study_id\n        all_coords.append(coords_temp)\n        \n        # # Plot\n        # plot_img(img, coords_temp)\n        # plot_5_crops(img, coords_temp)\n        \n    except Exception as e:\n        print(e)\n        pass\n\n# Gộp lại thành 1 dataframe\nall_coords_df = pd.concat(all_coords, ignore_index=True)\nprint(all_coords_df.head())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-22T01:59:57.726536Z","iopub.execute_input":"2025-08-22T01:59:57.726881Z","iopub.status.idle":"2025-08-22T01:59:59.654462Z","shell.execute_reply.started":"2025-08-22T01:59:57.726856Z","shell.execute_reply":"2025-08-22T01:59:59.653579Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Bộ dữ liệu sau khi xử lý\nall_coords_df.head(10)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-21T17:06:49.427022Z","iopub.execute_input":"2025-08-21T17:06:49.427521Z","iopub.status.idle":"2025-08-21T17:06:49.447004Z","shell.execute_reply.started":"2025-08-21T17:06:49.427478Z","shell.execute_reply":"2025-08-21T17:06:49.445628Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Pretraining","metadata":{}},{"cell_type":"code","source":"# Config\ncfg= SimpleNamespace(\n    img_dir= \"/kaggle/input/lumbar-coordinate-pretraining-dataset/data/\",\n    device= torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\"),\n    n_frames=3,\n    epochs=30,\n    lr=0.0005,\n    batch_size=16,\n    backbone=\"resnet18\",\n    seed= 0,\n    image_size=256\n)\nset_seed(seed=cfg.seed) # Makes results reproducable","metadata":{"execution":{"iopub.status.busy":"2025-08-21T17:32:09.381739Z","iopub.execute_input":"2025-08-21T17:32:09.382132Z","iopub.status.idle":"2025-08-21T17:32:09.388434Z","shell.execute_reply.started":"2025-08-21T17:32:09.382103Z","shell.execute_reply":"2025-08-21T17:32:09.387244Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load metadata\ndf= pd.read_csv(\"/kaggle/input/lumbar-coordinate-pretraining-dataset/coords_pretrain.csv\")\ndf= df.sort_values([\"source\", \"filename\", \"level\"]).reset_index(drop=True)\ndf[\"filename\"] = df[\"filename\"].str.replace(\".jpg\", \".npy\")\ndf[\"series_id\"] = df[\"source\"] + \"_\" + df[\"filename\"].str.split(\".\").str[0]\n\nprint(\"----- IMGS per source -----\")\ndisplay((df.source.value_counts()/5).astype(int).reset_index())","metadata":{"execution":{"iopub.status.busy":"2025-08-21T17:06:49.462462Z","iopub.execute_input":"2025-08-21T17:06:49.462793Z","iopub.status.idle":"2025-08-21T17:06:49.510084Z","shell.execute_reply.started":"2025-08-21T17:06:49.462765Z","shell.execute_reply":"2025-08-21T17:06:49.508799Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Lấy 2 điểm thay vì một điểm\n# Load coords\ncoords= all_coords_df.copy()\ncoords= coords.sort_values([\"series_id\", \"level\", \"side\"]).reset_index(drop=True)\n# coords= coords[[\"series_id\", \"level\", \"side\", \"relative_x\", \"relative_y\"]]\ncoords_new = coords.groupby([\"level\", \"series_id\", \"study_id\"]).apply(lambda g: list(zip(g['relative_x'], g['relative_y'])), include_groups=False).reset_index(drop=False, name=\"vals\")\n\ncoords_new[\"instance_vals\"] = coords.groupby(\n    [\"level\", \"series_id\", \"study_id\"]\n).apply(\n    lambda g: list(g['instance_number']),\n    include_groups=False\n).reset_index(drop=True)\ncoords_new= coords_new.sort_values([\"series_id\", \"level\"]).reset_index(drop=True)\n\nlen(coords_new)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-21T17:06:49.511642Z","iopub.execute_input":"2025-08-21T17:06:49.511948Z","iopub.status.idle":"2025-08-21T17:06:50.818912Z","shell.execute_reply.started":"2025-08-21T17:06:49.511922Z","shell.execute_reply":"2025-08-21T17:06:50.817791Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Xóa luôn những hàng có vals chỉ có 1 phần tử\ncoords_new = coords_new[coords_new[\"vals\"].apply(lambda x: len(x) != 1)].reset_index(drop=True)\ncoords_new.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-21T17:06:50.820584Z","iopub.execute_input":"2025-08-21T17:06:50.821030Z","iopub.status.idle":"2025-08-21T17:06:50.846906Z","shell.execute_reply.started":"2025-08-21T17:06:50.820993Z","shell.execute_reply":"2025-08-21T17:06:50.845554Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"coords_new[\"relative_x_l\"] = coords_new[\"vals\"].apply(lambda x: x[0][0])\ncoords_new[\"relative_y_l\"] = coords_new[\"vals\"].apply(lambda x: x[0][1])\ncoords_new[\"relative_x_r\"] = coords_new[\"vals\"].apply(lambda x: x[1][0])\ncoords_new[\"relative_y_r\"] = coords_new[\"vals\"].apply(lambda x: x[1][1])\nlen(coords_new)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-21T17:06:50.848665Z","iopub.execute_input":"2025-08-21T17:06:50.849403Z","iopub.status.idle":"2025-08-21T17:06:50.884984Z","shell.execute_reply.started":"2025-08-21T17:06:50.849374Z","shell.execute_reply":"2025-08-21T17:06:50.883932Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\n# Tạo thành một cột label bao gồm 20 phần tử tương ứng với 10 điểm của một ảnh\ndef agg_func(g):\n    # ghép thành list 1 chiều\n    labels = g.apply(lambda row: [row[\"relative_x_l\"], row[\"relative_y_l\"],\n                                  row[\"relative_x_r\"], row[\"relative_y_r\"]], axis=1).sum()\n    return pd.Series({\n        \"label\": labels,\n        \"instance_val\": g.iloc[0][\"instance_vals\"][1],   # lấy phần tử thứ 2 của hàng đầu tiên\n        # \"series_description\": g.iloc[0][\"series_description\"]\n    })\n\ngrouped = coords_new.groupby([\"study_id\", \"series_id\"]).apply(agg_func).reset_index()\ngrouped.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-21T17:06:50.889077Z","iopub.execute_input":"2025-08-21T17:06:50.889583Z","iopub.status.idle":"2025-08-21T17:06:52.747013Z","shell.execute_reply.started":"2025-08-21T17:06:50.889552Z","shell.execute_reply":"2025-08-21T17:06:52.745922Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 3. Chuẩn bị dữ liệu\n\nHere we define the torch dataset that will be used during training.","metadata":{}},{"cell_type":"code","source":"# Split data\nfrom sklearn.model_selection import train_test_split\n\n# Chia dữ liệu\ntrain_df, val_df = train_test_split(grouped, test_size=0.2, random_state=42)\nprint(len(train_df))\nprint(len(val_df))\ntrain_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-21T17:06:52.748475Z","iopub.execute_input":"2025-08-21T17:06:52.748803Z","iopub.status.idle":"2025-08-21T17:06:52.766407Z","shell.execute_reply.started":"2025-08-21T17:06:52.748775Z","shell.execute_reply":"2025-08-21T17:06:52.765246Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"grouped[\"instance_val\"].value_counts()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-21T17:06:52.767936Z","iopub.execute_input":"2025-08-21T17:06:52.768334Z","iopub.status.idle":"2025-08-21T17:06:52.781996Z","shell.execute_reply.started":"2025-08-21T17:06:52.768304Z","shell.execute_reply":"2025-08-21T17:06:52.780897Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pydicom\nimport cv2\nimport torch\nfrom torch.utils.data import Dataset\n\n# Class dataset để trả ra dict gồm img và label\nclass PreTrainDataset(Dataset):\n    def __init__(self, df, cfg):\n        self.cfg = cfg\n        self.records = self.load_coords(df)\n\n    def load_coords(self, df):\n        \"\"\"\n        Gom nhóm theo series_id, lưu lại series_id, study_id và 20 điểm (label).\n        \"\"\"\n        records = {}\n        df = df.reset_index(drop=True)   # reset về 0..N-1\n        for i, row in df.iterrows():\n            records[i] = {\n                \"series_id\": row[\"series_id\"],\n                \"study_id\": row[\"study_id\"],\n                \"instances\": row[\"instance_val\"],   # chỉ 1 số nguyên\n                \"label\": np.array(row[\"label\"], dtype=np.float32)\n            }\n        return records\n\n    def pad_and_resize(self, img):\n        \"\"\"\n        Chuẩn hóa ảnh về (1, img_size, img_size).\n        - Pad thành hình vuông\n        - Resize về kích thước chuẩn\n        \"\"\"\n        h, w = img.shape\n        target = max(h, w)\n\n        # Pad thành hình vuông\n        pad_h = (target - h) // 2\n        pad_w = (target - w) // 2\n        img = np.pad(img,\n                     ((pad_h, target - h - pad_h),\n                      (pad_w, target - w - pad_w)),\n                     mode=\"constant\", constant_values=0)\n\n        # Resize\n        img = cv2.resize(img, (self.cfg.image_size, self.cfg.image_size),\n                         interpolation=cv2.INTER_LINEAR)\n\n        # Thêm channel\n        img = np.expand_dims(img, axis=0)  # (1, size, size)\n        return img\n\n    def load_img(self, study_id, series_id, instance):\n        dcm_path = f\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/{study_id}/{series_id}/{instance}.dcm\"\n        dcm = pydicom.dcmread(dcm_path)\n        img = dcm.pixel_array.astype(np.float32)\n\n        img = self.pad_and_resize(img)   # (1, img_size, img_size)\n        img = img / 255.0\n        return img\n\n    def __getitem__(self, idx):\n        d = self.records[idx]\n        img = self.load_img(str(d[\"study_id\"]), str(d[\"series_id\"]), str(d[\"instances\"]))\n        label = d[\"label\"]\n        img = np.repeat(img, 3, axis=0)   # (1,H,W) -> (3,H,W)\n        return {\n            \"img\": torch.tensor(img, dtype=torch.float32),      # (1, size, size)\n            \"label\": torch.tensor(label, dtype=torch.float32),  # (20,)\n        }\n\n    def __len__(self):\n        return len(self.records)\n\n\ntrain_ds = PreTrainDataset(train_df, cfg)\nval_ds = PreTrainDataset(val_df, cfg)\n\n# Test 1 sample\nprint(\"---- Sample Shapes -----\")\nsample = train_ds[0]\nfor k, v in sample.items():\n    print(k, v.shape, v.dtype)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-21T17:06:52.783698Z","iopub.execute_input":"2025-08-21T17:06:52.784093Z","iopub.status.idle":"2025-08-21T17:06:52.942422Z","shell.execute_reply.started":"2025-08-21T17:06:52.784064Z","shell.execute_reply":"2025-08-21T17:06:52.941237Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Các hàm bổ sung (chuyển dữ liệu lên GPU, vẽ ảnh trong lúc train, ...)\n\n\nHere we have a couple helpers functions. \n\nThe first moves data to the GPU if enabled, the second visualizes predictions during training, and the third loads weights when dealing with mismatched shapes.","metadata":{}},{"cell_type":"code","source":"def batch_to_device(batch, device, skip_keys=[]):\n    batch_dict= {}\n    for key in batch:\n        if key in skip_keys:\n             batch_dict[key]= batch[key]\n        else:    \n            batch_dict[key]= batch[key].to(device)\n    return batch_dict\n\n# Hàm vẽ ảnh trong lúc huấn luyện\ndef visualize_prediction(batch, pred, epoch, cfg):\n    \"\"\"\n    Hiển thị ảnh + tọa độ thật và dự đoán\n    batch: dict từ DataLoader\n    pred: tensor (B, 20) -> 10 điểm (x,y)\n    epoch: số epoch\n    cfg: config chứa n_frames\n    \"\"\"\n    mid = cfg.n_frames // 2   # lấy frame giữa\n    B = batch[\"img\"].shape[0]\n\n    for idx in range(min(1, B)):\n        # ảnh: (n_frames, H, W), lấy frame giữa\n        img = batch[\"img\"][idx, mid].cpu().numpy()  # (H, W)\n        H, W = img.shape\n\n        # toạ độ thật và dự đoán (20 giá trị = 10 cặp (x,y))\n        cs_true = batch[\"label\"][idx].cpu().numpy() * H\n        cs_pred = pred[idx].detach().cpu().numpy() * H\n\n        coords_list = [\n            (\"TRUE\", \"lightblue\", cs_true),\n            (\"PRED\", \"orange\", cs_pred)\n        ]\n        text_labels = [str(x) for x in range(1, 11)]  # 10 điểm\n\n        # Vẽ\n        fig, axes = plt.subplots(1, len(coords_list), figsize=(12, 5))\n        if len(coords_list) == 1:\n            axes = [axes]\n        \n        for ax, (title, color, coords) in zip(axes, coords_list):\n            ax.imshow(img, cmap=\"gray\")\n            ax.scatter(coords[0::2], coords[1::2], c=color, s=50)\n            ax.axis(\"off\")\n            ax.set_title(title)\n\n            # Thêm nhãn số cho từng điểm\n            for i, (x, y) in enumerate(zip(coords[0::2], coords[1::2])):\n                ax.text(\n                    x + 5, y, text_labels[i],\n                    color=\"white\", fontsize=10,\n                    bbox=dict(facecolor=\"black\", alpha=0.5)\n                )\n\n        fig.suptitle(f\"EPOCH: {epoch}\")\n        plt.show()\n\n    return\n    \n# utils\ndef load_weights_skip_mismatch(model, weights_path, device):\n    # Load Weights\n    state_dict = torch.load(weights_path, map_location=device)\n    model_dict = model.state_dict()\n    \n    # Iter models\n    params = {}\n    for (sdk, sfv), (mdk, mdv) in zip(state_dict.items(), model_dict.items()):\n        if sfv.size() == mdv.size():\n            params[sdk] = sfv\n        else:\n            print(\"Skipping param: {}, {} != {}\".format(sdk, sfv.size(), mdv.size()))\n    \n    # Reload + Skip\n    model.load_state_dict(params, strict=False)\n    print(\"Loaded weights from:\", weights_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-21T17:06:52.944046Z","iopub.execute_input":"2025-08-21T17:06:52.944402Z","iopub.status.idle":"2025-08-21T17:06:52.960503Z","shell.execute_reply.started":"2025-08-21T17:06:52.944374Z","shell.execute_reply":"2025-08-21T17:06:52.959272Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Huấn luyện mô hình\n\n- Finetuning resnet18 để có 20 đầu ra tương ứng với 10 điểm","metadata":{}},{"cell_type":"code","source":"# Datasets + Dataloaders\nprint(len(train_ds))\nprint(len(val_ds))\ntrain_dl = torch.utils.data.DataLoader(train_ds, batch_size=cfg.batch_size, shuffle=True, drop_last=True)\nval_dl = torch.utils.data.DataLoader(val_ds, batch_size=cfg.batch_size, shuffle=False)\nprint(len(train_dl))\nprint(len(val_dl))\n# Model\nmodel = timm.create_model('resnet18', pretrained=True, num_classes=20)\nmodel = model.to(cfg.device)\nprint(model.fc)\n\n# Loss / Optim\ncriterion = nn.MSELoss()\noptimizer = torch.optim.AdamW(model.parameters(), lr=cfg.lr)","metadata":{"execution":{"iopub.status.busy":"2025-08-21T17:18:38.970329Z","iopub.execute_input":"2025-08-21T17:18:38.971367Z","iopub.status.idle":"2025-08-21T17:18:39.391336Z","shell.execute_reply.started":"2025-08-21T17:18:38.971329Z","shell.execute_reply":"2025-08-21T17:18:39.390261Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for epoch in range(cfg.epochs+1):\n    \n    # Train Loop\n    loss= torch.tensor([0.]).float().to(cfg.device)\n    if epoch != 0:\n        model= model.train()\n        for batch in tqdm(train_dl):\n            batch = batch_to_device(batch, cfg.device)\n            optimizer.zero_grad()\n\n            x_out = model(batch[\"img\"].float())\n            x_out = torch.sigmoid(x_out)\n\n            loss = criterion(x_out, batch[\"label\"].float())\n            loss.backward()\n            optimizer.step()\n        \n    # Validation Loop\n    val_loss = 0\n    with torch.no_grad():\n        model = model.eval()\n        for batch in tqdm(val_dl):\n            batch = batch_to_device(batch, cfg.device)\n            pred = model(batch[\"img\"].float())\n            pred = torch.sigmoid(pred)\n            val_loss += criterion(pred, batch[\"label\"].float()).item()\n        val_loss /= len(val_dl)\n            \n    # Viz\n    visualize_prediction(batch, pred, epoch, cfg)           \n            \n    print(f\"Epoch {epoch+1}, Training Loss: {loss.item()}, Validation Loss: {val_loss}\")\nprint(\"Training complete...\")","metadata":{"execution":{"iopub.status.busy":"2025-08-21T17:32:29.916270Z","iopub.execute_input":"2025-08-21T17:32:29.916642Z","iopub.status.idle":"2025-08-21T17:46:44.822498Z","shell.execute_reply.started":"2025-08-21T17:32:29.916616Z","shell.execute_reply":"2025-08-21T17:46:44.821520Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 4. Dự đoán với dữ liệu mới","metadata":{}},{"cell_type":"code","source":"import pydicom\nimport cv2\nimport numpy as np\nimport torch\nfrom types import SimpleNamespace\n\n# Tiền xử lý để có dạng giống lúc huấn luyện\ndef preprocess_new_image(dicom_path, cfg):\n    \"\"\"\n    Tiền xử lý một ảnh DICOM mới để chuẩn bị cho việc dự đoán.\n\n    Args:\n        dicom_path (str): Đường dẫn tới file DICOM.\n        cfg (SimpleNamespace): Đối tượng cấu hình chứa image_size.\n\n    Returns:\n        torch.Tensor: Tensor ảnh đã được xử lý (3, img_size, img_size).\n    \"\"\"\n\n    # 1. Đọc và chuẩn hóa ảnh\n    dcm = pydicom.dcmread(dicom_path)\n    img = dcm.pixel_array.astype(np.float32)\n\n    # Đảm bảo ảnh 1 kênh là đầu vào cho các bước tiếp theo\n    if len(img.shape) == 3:\n        img = img[:, :, 0]\n\n    # 2. Pad thành hình vuông\n    h, w = img.shape\n    target = max(h, w)\n    pad_h = (target - h) // 2\n    pad_w = (target - w) // 2\n    img = np.pad(img,\n                 ((pad_h, target - h - pad_h),\n                  (pad_w, target - w - pad_w)),\n                 mode=\"constant\", constant_values=0)\n\n    # 3. Resize về kích thước chuẩn\n    img_resized = cv2.resize(img, (cfg.image_size, cfg.image_size),\n                             interpolation=cv2.INTER_LINEAR)\n\n    # 4. Chuẩn hóa pixel về khoảng [0, 1]\n    img_normalized = img_resized / 255.0\n\n    # 5. Thêm kênh và lặp lại để tạo thành 3 kênh\n    img_1_channel = np.expand_dims(img_normalized, axis=0)  # (1, size, size)\n    img_3_channels = np.repeat(img_1_channel, 3, axis=0)  # (3, size, size)\n\n    # 6. Chuyển đổi sang tensor PyTorch\n    return torch.tensor(img_3_channels, dtype=torch.float32)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-21T17:46:58.919322Z","iopub.execute_input":"2025-08-21T17:46:58.920248Z","iopub.status.idle":"2025-08-21T17:46:58.928280Z","shell.execute_reply.started":"2025-08-21T17:46:58.920211Z","shell.execute_reply":"2025-08-21T17:46:58.927248Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 1. Xác định thiết bị\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Sử dụng thiết bị: {device}\")\n\n# 2. Chuyển mô hình sang thiết bị\nmodel.to(device)\nmodel.eval()\n\n# 3. Đường dẫn tới file DICOM mới của bạn\nnew_image_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images/44036939/3844393089/10.dcm'\n\n# 4. Tiền xử lý ảnh\ninput_tensor = preprocess_new_image(new_image_path, cfg)\nprint(input_tensor.shape)\n# 5. Thêm chiều batch (batch size = 1) và CHUYỂN TENSOR SANG CÙNG THIẾT BỊ\ninput_tensor = input_tensor.unsqueeze(0).to(device)\n\n# 6. Đưa vào mô hình để dự đoán\nwith torch.no_grad():\n    prediction = model(input_tensor)  # prediction trên GPU\n    print(prediction.shape)\n    # pred = prediction[0].cpu().numpy()  # Chuyển batch đầu tiên về CPU rồi mới .numpy()\n    pred = torch.sigmoid(prediction)[0].cpu().numpy()\n\nprint(pred)\n# --- Plot kết quả ---\ndcm = pydicom.dcmread(new_image_path)\nimg = dcm.pixel_array\nprint(img.shape)\nplt.imshow(img, cmap=\"gray\")\n\n\nfor i in range(10):\n    x = pred[i*2] * img.shape[1]  # scale x\n    y = pred[i*2+1] * img.shape[0]  # scale y\n    plt.scatter(x, y, c=\"orange\", s=60)\n    plt.text(x+5, y, f\"P{i+1}\", color=\"white\", fontsize=10, bbox=dict(facecolor=\"black\", alpha=0.5))\nplt.axis(\"off\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-21T18:36:00.279431Z","iopub.execute_input":"2025-08-21T18:36:00.279799Z","iopub.status.idle":"2025-08-21T18:36:00.702124Z","shell.execute_reply.started":"2025-08-21T18:36:00.279771Z","shell.execute_reply":"2025-08-21T18:36:00.701070Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Các hàm để plot ảnh và cắt ảnh các đốt dựa trên các điểm đã dự đoán","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport albumentations as A\nimport math\n\n# Giữ nguyên hàm angle\ndef angle_of_line(x1, y1, x2, y2):\n    return math.degrees(math.atan2(-(y2-y1), x2-x1))\n\n# Crop giữa 2 điểm\ndef crop_between_keypoints(img, keypoint1, keypoint2):\n    h, w = img.shape\n    x1, y1 = int(keypoint1[0]), int(keypoint1[1])\n    x2, y2 = int(keypoint2[0]), int(keypoint2[1])\n    \n    # Calculate bounding box\n    left = max(int(min(x1, x2)), 0)\n    right = min(int(max(x1, x2)), w)\n    top = max(int(min(y1, y2) - 0.1*h), 0)\n    bottom = min(int(max(y1, y2) + 0.1*h), h)\n    \n    return img[top:bottom, left:right]\n\n# Resize transform (giữ nguyên)\nresize_transform= A.Compose([\n    A.LongestMaxSize(max_size=640, interpolation=1, always_apply=True), # cv2.INTER_CUBIC\n    A.PadIfNeeded(min_height=640, min_width=640, border_mode=0, value=(0,0,0), always_apply=True),\n])\n\n# Vẽ keypoints trên ảnh gốc\ndef plot_img_from_coords(img, coords):\n    \"\"\"\n    img: numpy 2D (H, W)\n    coords: numpy 1D (20,) normalized [0,1] -> 10 điểm (x,y)\n    \"\"\"\n    H, W = img.shape\n    fig, ax = plt.subplots()\n    ax.imshow(img, cmap='gray')\n    \n    xs = coords[0::2] * W\n    ys = coords[1::2] * H\n    ax.scatter(xs, ys, c=\"orange\", s=50)\n    \n    # Add text labels\n    ax.set_title(\"Kết quả dự đoán vị trí 2 phía các đĩa đệm bằng ResNet18\")\n    for i, (x, y) in enumerate(zip(xs, ys)):\n        ax.text(x+5, y, str(i+1), color=\"white\", fontsize=10,\n                bbox=dict(facecolor=\"black\", alpha=0.5))\n    ax.axis('off')\n    plt.show()\n\n# Tạo 5 crops dựa trên 5 cặp điểm (mỗi đốt 2 điểm)\ndef plot_5_crops_from_coords(img, coords):\n    \"\"\"\n    img: numpy 2D (H, W)\n    coords: numpy 1D (20,) normalized [0,1] -> 10 điểm (x,y)\n    \"\"\"\n    import matplotlib.gridspec as gridspec\n    \n    H, W = img.shape\n    fig = plt.figure(figsize=(10,10))\n    gs = gridspec.GridSpec(1, 5, width_ratios=[1]*5)\n    \n    # Chia coords thành 5 cặp điểm (mỗi cặp cho 1 crop)\n    pairs = [(coords[i*4:i*4+2], coords[i*4+2:i*4+4]) for i in range(5)]\n    map_dict = {1: \"L1/L2\", \n                2: \"L2/L3\",\n                3: \"L3/L4\", \n                4: \"L4/L5\",\n                5: \"L5/S1\"}\n    fig.suptitle(\"Các vùng ảnh đĩa đệm đã được cắt dựa trên tọa độ dự đoán\", fontsize=16, y=0.65)\n    for idx, (a_norm, b_norm) in enumerate(pairs):\n        # Convert normalized -> pixel\n        a = (a_norm[0]*W, a_norm[1]*H)\n        b = (b_norm[0]*W, b_norm[1]*H)\n        \n        img_copy = img.copy()\n        \n        # Rotate ảnh theo đường nối 2 điểm\n        rotate_angle = angle_of_line(a[0], a[1], b[0], b[1])\n        transform = A.Compose([\n            A.Rotate(limit=(-rotate_angle, -rotate_angle), p=1.0),\n        ], keypoint_params=A.KeypointParams(format='xy', remove_invisible=False))\n        t = transform(image=img_copy, keypoints=[a,b])\n        img_copy = t[\"image\"]\n        a, b = t[\"keypoints\"]\n        \n        # Crop + Resize\n        img_crop = crop_between_keypoints(img_copy, a, b)\n        img_crop = resize_transform(image=img_crop)[\"image\"]\n        \n        # Plot\n        ax = plt.subplot(gs[idx])\n        ax.imshow(img_crop, cmap='gray')\n        ax.set_title(f\"{map_dict[idx+1]}\")\n        ax.axis('off')\n    \n    plt.show()\n\ndef crop_between_keypoints(img, keypoint1, keypoint2):\n    h, w = img.shape\n    x1, y1 = int(keypoint1[0]), int(keypoint1[1])\n    x2, y2 = int(keypoint2[0]), int(keypoint2[1])\n    \n    # Calculate bounding box around the keypoints\n    left = int(min(x1, x2))\n    right = int(max(x1, x2))\n    top = int(min(y1, y2) - (h * 0.1))\n    bottom = int(max(y1, y2) + (h * 0.1))\n            \n    # Crop the image\n    return img[top:bottom, left:right]\nplot_img_from_coords(img, pred)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-21T18:36:04.826418Z","iopub.execute_input":"2025-08-21T18:36:04.827436Z","iopub.status.idle":"2025-08-21T18:36:05.085351Z","shell.execute_reply.started":"2025-08-21T18:36:04.827385Z","shell.execute_reply":"2025-08-21T18:36:05.084289Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_5_crops_from_coords(img, pred)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-21T18:36:11.550536Z","iopub.execute_input":"2025-08-21T18:36:11.550914Z","iopub.status.idle":"2025-08-21T18:36:12.068778Z","shell.execute_reply.started":"2025-08-21T18:36:11.550885Z","shell.execute_reply":"2025-08-21T18:36:12.067621Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Save model\n\nNext, we save the backbone weights.","metadata":{}},{"cell_type":"code","source":"f= \"{}_{}.pt\".format(cfg.backbone, cfg.seed)\ntorch.save(model.state_dict(), f)\nprint(\"Saved weights: {}\".format(f))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-21T17:09:40.767646Z","iopub.execute_input":"2025-08-21T17:09:40.767960Z","iopub.status.idle":"2025-08-21T17:09:40.856376Z","shell.execute_reply.started":"2025-08-21T17:09:40.767933Z","shell.execute_reply":"2025-08-21T17:09:40.855265Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Finally, the weights can be loaded for a new task (eg. this competition).","metadata":{}},{"cell_type":"code","source":"# Load backbone for RSNA 2024 task\nmodel = timm.create_model('resnet18', pretrained=True, num_classes=75)\nmodel = model.to(cfg.device)\nload_weights_skip_mismatch(model, f, cfg.device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-21T17:09:40.857806Z","iopub.execute_input":"2025-08-21T17:09:40.858148Z","iopub.status.idle":"2025-08-21T17:09:41.236991Z","shell.execute_reply.started":"2025-08-21T17:09:40.858117Z","shell.execute_reply":"2025-08-21T17:09:41.235757Z"}},"outputs":[],"execution_count":null}]}