{"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":"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":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2025-09-06T16:19:55.715919Z","iopub.execute_input":"2025-09-06T16:19:55.716293Z","iopub.status.idle":"2025-09-06T16:19:55.722978Z","shell.execute_reply.started":"2025-09-06T16:19:55.71626Z","shell.execute_reply":"2025-09-06T16:19:55.722032Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Lumbar Coordinate Dataset","metadata":{}},{"cell_type":"markdown","source":"## 1. Improved RSNA Coordinates","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-09-06T16:19:55.724478Z","iopub.execute_input":"2025-09-06T16:19:55.724722Z","iopub.status.idle":"2025-09-06T16:19:55.739491Z","shell.execute_reply.started":"2025-09-06T16:19:55.724702Z","shell.execute_reply":"2025-09-06T16:19:55.738684Z"},"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\ndef angle_of_line(x1, y1, x2, y2):\n    return math.degrees(math.atan2(-(y2-y1), x2-x1))\n\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\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    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-09-06T16:19:55.740788Z","iopub.execute_input":"2025-09-06T16:19:55.741177Z","iopub.status.idle":"2025-09-06T16:19:55.755212Z","shell.execute_reply.started":"2025-09-06T16:19:55.741148Z","shell.execute_reply":"2025-09-06T16:19:55.754349Z"},"trusted":true},"outputs":[],"execution_count":null},{"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\"]\ndfd= dfd.sample(frac=1, random_state=SEED).head(N)\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)\ncoords= coords[[\"series_id\", \"level\", \"side\", \"relative_x\", \"relative_y\"]]\n\n# Plot samples\nfor idx, row in dfd.iterrows():\n    try:\n        print(\"-\"*25, \" STUDY_ID: {}, SERIES_ID: {} \".format(row.study_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        # 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","metadata":{"execution":{"iopub.status.busy":"2025-09-06T16:19:55.756332Z","iopub.execute_input":"2025-09-06T16:19:55.757108Z","iopub.status.idle":"2025-09-06T16:19:57.156043Z","shell.execute_reply.started":"2025-09-06T16:19:55.75708Z","shell.execute_reply":"2025-09-06T16:19:57.15508Z"},"trusted":true},"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=200,\n    lr=0.0005,\n    batch_size=16,\n    backbone=\"resnet18\",\n    seed= 0,\n)\nset_seed(seed=cfg.seed) # Makes results reproducable","metadata":{"execution":{"iopub.status.busy":"2025-09-06T16:19:57.158224Z","iopub.execute_input":"2025-09-06T16:19:57.158507Z","iopub.status.idle":"2025-09-06T16:19:57.163577Z","shell.execute_reply.started":"2025-09-06T16:19:57.158485Z","shell.execute_reply":"2025-09-06T16:19:57.162702Z"},"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-09-06T16:19:57.16453Z","iopub.execute_input":"2025-09-06T16:19:57.164763Z","iopub.status.idle":"2025-09-06T16:19:57.207694Z","shell.execute_reply.started":"2025-09-06T16:19:57.164744Z","shell.execute_reply":"2025-09-06T16:19:57.206759Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"class PreTrainDataset(torch.utils.data.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        # Convert to dict\n        d = df.groupby(\"series_id\")[[\"relative_x\", \"relative_y\"]].apply(lambda x: list(x.itertuples(index=False, name=None)))\n        records= {}\n        for i, (k,v) in enumerate(d.items()):\n            records[i]= {\"series_id\": k, \"label\": np.array(v).flatten()}\n            assert len(v) == 5\n            \n        return records\n    \n    def pad_image(self, img):\n        n= img.shape[-1]\n        if n >= self.cfg.n_frames:\n            start_idx = (n - self.cfg.n_frames) // 2\n            return img[:, :, start_idx:start_idx + self.cfg.n_frames]\n        else:\n            pad_left = (self.cfg.n_frames - n) // 2\n            pad_right = self.cfg.n_frames - n - pad_left\n            return np.pad(img, ((0,0), (0,0), (pad_left, pad_right)), 'constant', constant_values=0)\n    \n    def load_img(self, source, series_id):\n        fname= os.path.join(self.cfg.img_dir, \"processed_{}/{}.npy\".format(source, series_id))\n        img= np.load(fname).astype(np.float32)\n        img= self.pad_image(img)\n        img= np.transpose(img, (2, 0, 1))\n        img= (img / 255.0)\n        return img\n        \n        \n    def __getitem__(self, idx):\n        d= self.records[idx]\n        label= d[\"label\"]\n        source= d[\"series_id\"].split(\"_\")[0]\n        series_id= \"_\".join(d[\"series_id\"].split(\"_\")[1:])     \n                \n        img= self.load_img(source, series_id)\n        return {\n            'img': img, \n            'label': label,\n            }\n    \n    def __len__(self,):\n        return len(self.records)\n    \nds= PreTrainDataset(df, cfg)    \n\n# Plot a Single Sample\nprint(\"---- Sample Shapes -----\")\nfor k, v in ds[0].items():\n    print(k, v.shape)","metadata":{"execution":{"iopub.status.busy":"2025-09-06T16:19:57.209047Z","iopub.execute_input":"2025-09-06T16:19:57.209447Z","iopub.status.idle":"2025-09-06T16:19:57.40255Z","shell.execute_reply.started":"2025-09-06T16:19:57.209413Z","shell.execute_reply":"2025-09-06T16:19:57.401611Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Utils","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\ndef visualize_prediction(batch, pred, epoch):\n    \n    mid= cfg.n_frames//2\n    \n    # Plot\n    for idx in range(1):\n    \n        # Select Data\n        img= batch[\"img\"][idx, mid, :, :].cpu().numpy()*255\n        cs_true= batch[\"label\"][idx, ...].cpu().numpy()*256\n        cs= pred[idx, ...].cpu().numpy()*256\n                \n        coords_list = [(\"TRUE\", \"lightblue\", cs_true), (\"PRED\", \"orange\", cs)]\n        text_labels = [str(x) for x in range(1,6)]\n        \n        # Plot coords\n        fig, axes = plt.subplots(1, len(coords_list), figsize=(10,4))\n        fig.suptitle(\"EPOCH: {}\".format(epoch))\n        for ax, (title, color, coords) in zip(axes, coords_list):\n\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            # Add text labels near the coordinates\n            for i, (x, y) in enumerate(zip(coords[0::2], coords[1::2])):\n                if i < len(text_labels):  # Ensure there are enough labels\n                    ax.text(x + 10, y, text_labels[i], color='white', fontsize=15, bbox=dict(facecolor='black', alpha=0.5))\n\n\n        fig.suptitle(\"EPOCH: {}\".format(epoch))\n        plt.show()\n#         plt.close(fig)\n    return\n\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":{"execution":{"iopub.status.busy":"2025-09-06T16:54:35.89047Z","iopub.execute_input":"2025-09-06T16:54:35.891207Z","iopub.status.idle":"2025-09-06T16:54:35.903946Z","shell.execute_reply.started":"2025-09-06T16:54:35.891178Z","shell.execute_reply":"2025-09-06T16:54:35.902943Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Training","metadata":{}},{"cell_type":"code","source":"# Dataframes\ntrain_df= df[df[\"source\"] != \"spider\"]\nval_df= df[df[\"source\"] == \"spider\"]\nprint(\"TRAIN_SIZE: {}, VAL_SIZE: {}\".format(len(train_df)//5, len(val_df)//5))\n\n# Datasets + Dataloaders\ntrain_ds= PreTrainDataset(train_df, cfg)\nval_ds= PreTrainDataset(val_df, cfg)\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)\n\n# Model\nmodel = timm.create_model('resnet18', pretrained=True, num_classes=10)\nmodel = model.to(cfg.device)\n\n# Loss / Optim\ncriterion = nn.MSELoss()\noptimizer = torch.optim.AdamW(model.parameters(), lr=cfg.lr)","metadata":{"execution":{"iopub.status.busy":"2025-09-06T16:19:57.416615Z","iopub.execute_input":"2025-09-06T16:19:57.416879Z","iopub.status.idle":"2025-09-06T16:19:57.89644Z","shell.execute_reply.started":"2025-09-06T16:19:57.416858Z","shell.execute_reply":"2025-09-06T16:19:57.895681Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_model = model\nbest_val_loss = float(\"inf\")\nfor 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\n            pred = model(batch[\"img\"].float())\n            pred = torch.sigmoid(pred)\n            \n            val_loss += criterion(pred, batch[\"label\"].float()).item()\n        val_loss /= len(val_dl)\n        if val_loss < best_val_loss:\n            best_model = model\n            best_val_loss = val_loss\n            \n    # Viz\n    # visualize_prediction(batch, pred, epoch)           \n            \n    print(f\"Epoch {epoch+1}, Training Loss: {loss.item()}, Validation Loss: {val_loss}, Best val loss: {best_val_loss}\")\nprint(\"Training complete...\")\n","metadata":{"execution":{"iopub.status.busy":"2025-09-06T16:47:05.548069Z","iopub.execute_input":"2025-09-06T16:47:05.548773Z","iopub.status.idle":"2025-09-06T16:47:05.553124Z","shell.execute_reply.started":"2025-09-06T16:47:05.548746Z","shell.execute_reply":"2025-09-06T16:47:05.552172Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Save\n","metadata":{}},{"cell_type":"code","source":"f= \"{}_{}.pt\".format(cfg.backbone, cfg.seed)\ntorch.save(best_model.state_dict(), f)\nprint(\"Saved weights: {}\".format(f))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-06T16:32:47.3166Z","iopub.execute_input":"2025-09-06T16:32:47.317299Z","iopub.status.idle":"2025-09-06T16:32:47.428664Z","shell.execute_reply.started":"2025-09-06T16:32:47.317272Z","shell.execute_reply":"2025-09-06T16:32:47.427672Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load backbone for RSNA 2024 task\nmodel = timm.create_model('resnet18', pretrained=True, num_classes=10)\nmodel = model.to(cfg.device)\nload_weights_skip_mismatch(model, f, cfg.device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-06T16:53:29.375333Z","iopub.execute_input":"2025-09-06T16:53:29.3757Z","iopub.status.idle":"2025-09-06T16:53:29.756125Z","shell.execute_reply.started":"2025-09-06T16:53:29.375674Z","shell.execute_reply":"2025-09-06T16:53:29.755169Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Prediction","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, (256, 256),\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-09-06T17:05:57.484468Z","iopub.execute_input":"2025-09-06T17:05:57.485202Z","iopub.status.idle":"2025-09-06T17:05:57.492317Z","shell.execute_reply.started":"2025-09-06T17:05:57.485176Z","shell.execute_reply":"2025-09-06T17:05:57.491453Z"}},"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/train_images/100206310/2092806862/1.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(5):\n    x = pred[i*2] * img.shape[1]  # scale x\n    y = pred[i*2+1] * img.shape[0]  # scale y\n    print(x, 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-09-06T17:19:27.301877Z","iopub.execute_input":"2025-09-06T17:19:27.302255Z","iopub.status.idle":"2025-09-06T17:19:27.541809Z","shell.execute_reply.started":"2025-09-06T17:19:27.302229Z","shell.execute_reply":"2025-09-06T17:19:27.540632Z"}},"outputs":[],"execution_count":null}]}