{"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":8713513,"sourceType":"datasetVersion","datasetId":5227521},{"sourceId":9187072,"sourceType":"datasetVersion","datasetId":5504483}],"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":"2024-09-09T15:34:04.597600Z","iopub.execute_input":"2024-09-09T15:34:04.598494Z","iopub.status.idle":"2024-09-09T15:34:04.605888Z","shell.execute_reply.started":"2024-09-09T15:34:04.598450Z","shell.execute_reply":"2024-09-09T15:34:04.604964Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Lumbar Coordinate Dataset\n\n\nThis notebook shows how the [Lumbar Coordinate Dataset](https://www.kaggle.com/datasets/brendanartley/lumbar-coordinate-pretraining-dataset) can be used in the RSNA 2024 competition. First I will showcase improved coordinates for the competition data, and then I will showcase using external data to pretrain backbones.","metadata":{}},{"cell_type":"markdown","source":"## 笔记本的主要内容\n- 展示改进的坐标数据：笔记本的第一部分展示了如何利用Lumbar Coordinate Dataset改进RSNA竞赛中的坐标数据。这可能包括对已有的坐标进行校正或细化，以提高模型的准确性。\n\n- 使用外部数据进行预训练：笔记本的第二部分展示了如何使用外部数据（如Lumbar Coordinate Dataset）来对模型的骨干网络（backbone）进行预训练。预训练可以帮助模型更好地理解医学影像数据，提高在特定任务（如识别或分类）上的表现。","metadata":{"execution":{"iopub.status.busy":"2024-09-01T03:55:15.444423Z","iopub.execute_input":"2024-09-01T03:55:15.445145Z","iopub.status.idle":"2024-09-01T03:55:15.451517Z","shell.execute_reply.started":"2024-09-01T03:55:15.445097Z","shell.execute_reply":"2024-09-01T03:55:15.450355Z"}}},{"cell_type":"markdown","source":"## 1. Improved RSNA Coordinates\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":"## 作者展示了如何通过标注每个椎间盘的左侧来改进RSNA竞赛中的坐标数据。这种方法可以帮助我们确定椎间盘的旋转角度，从而在对图像进行裁剪时提高精度。\n\n- 关键点：\n标注左侧：在RSNA竞赛中，通常只提供椎间盘的中心坐标。然而，作者通过额外标注每个椎间盘的左侧（左侧的坐标）来进一步优化这些坐标信息。\n\n- 确定角度：通过知道椎间盘的中心和左侧的位置，可以计算出椎间盘的旋转角度。这种角度信息对于调整图像裁剪非常有帮助。例如，如果你知道椎间盘是倾斜的，那么你可以按照这个角度裁剪图像，以便更好地聚焦在目标区域。\n\n- 改进裁剪：有了旋转角度后，裁剪过程可以更加精确，确保图像的ROI（感兴趣区域）只包含椎间盘的相关部分，从而减少不必要的信息和噪声。这种方法可以有效提升后续模型的性能。\n\n","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport glob\nimport os\nimport pydicom\n\ndef convert_to_8bit(x):\n    # 获取x中第1和第99百分位数的值作为上下限\n    lower, upper = np.percentile(x, (1, 99))\n    # 将x限制在lower和upper之间\n    x = np.clip(x, lower, upper)\n    # 将x的最小值平移至0\n    x = x - np.min(x)\n    # 将x归一化至0-1之间\n    x = x / np.max(x)\n    # 将x扩展至0-255并转换为8位无符号整数（uint8）\n    return (x * 255).astype(\"uint8\")\n\ndef load_dicom_stack(dicom_folder, plane, reverse_sort=False):\n    # 获取DICOM文件夹中的所有.dcm文件路径\n    dicom_files = glob.glob(os.path.join(dicom_folder, \"*.dcm\"))\n    # 读取所有DICOM文件\n    dicoms = [pydicom.dcmread(f) for f in dicom_files]\n    \n    # 根据给定的切面选择对应的平面轴，0: 矢状面, 1: 冠状面, 2: 横断面\n    plane = {\"sagittal\": 0, \"coronal\": 1, \"axial\": 2}[plane.lower()]\n    \n    # 提取每个DICOM图像在所选平面的位置\n    positions = np.asarray([float(d.ImagePositionPatient[plane]) for d in dicoms])\n    # 如果reverse_sort=False，则增加的数组索引将从右到左，并且从尾侧到头侧，\n    # 因此我们在横断面时将reverse_sort设置为True，使得增加的数组索引为头尾方向（头侧->尾侧）\n    idx = np.argsort(-positions if reverse_sort else positions)\n    # 获取图像的患者位置数组并按照排序后的索引排列\n    ipp = np.asarray([d.ImagePositionPatient for d in dicoms]).astype(\"float\")[idx]\n    # 获取DICOM图像的像素数据并将其转换为float32类型，然后按排序后的索引排列\n    array = np.stack([d.pixel_array.astype(\"float32\") for d in dicoms])\n    array = array[idx]\n    # 将图像转换为8位，并返回图像数组、位置数组和像素间距信息\n    return {\"array\": convert_to_8bit(array), \"positions\": ipp, \"pixel_spacing\": np.asarray(dicoms[0].PixelSpacing).astype(\"float\")}\n\n# DICOM图像文件夹路径\nimage_dir = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/\"\n","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:04.607632Z","iopub.execute_input":"2024-09-09T15:34:04.607917Z","iopub.status.idle":"2024-09-09T15:34:04.621711Z","shell.execute_reply.started":"2024-09-09T15:34:04.607894Z","shell.execute_reply":"2024-09-09T15:34:04.620867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# 图像调整转换，使用Albumentations库\nresize_transform = A.Compose([\n    # 将图像的最长边调整为256像素，使用三次插值法\n    A.LongestMaxSize(max_size=256, interpolation=cv2.INTER_CUBIC, always_apply=True),\n    # 如果图像不足256x256像素，则进行填充，边界填充颜色为黑色\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    \"\"\"\n    计算两点之间连线的角度（度数）。\n    \n    参数：\n    x1, y1 -- 第一个点的坐标\n    x2, y2 -- 第二个点的坐标\n    返回：\n    连线的角度（度数）\n    \"\"\"\n    return math.degrees(math.atan2(-(y2 - y1), x2 - x1))\n\ndef plot_img(img, coords_temp):\n    \"\"\"\n    在图像上绘制关键点。\n    \n    参数：\n    img -- 图像数组\n    coords_temp -- 包含关键点坐标的DataFrame，包含'level', 'relative_x', 'relative_y'列\n    \"\"\"\n    # 创建绘图对象和坐标轴\n    fig, ax = plt.subplots()\n    # 显示灰度图像\n    ax.imshow(img, cmap='gray')\n    h, w = img.shape\n    \n    # 按level分组并获取关键点对\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    # 在图像上绘制每组关键点\n    for _, row in p.iterrows():\n        level = row['level']\n        # 将相对坐标转换为图像像素坐标\n        x = [coord[0] * w for coord in row[\"vals\"]]\n        y = [coord[1] * h for coord in row[\"vals\"]]\n        # 绘制关键点\n        ax.plot(x, y, marker='o')\n    # 关闭坐标轴\n    ax.axis('off')\n    # 显示图像\n    plt.show()\n\ndef plot_5_crops(img, coords_temp):\n    \"\"\"\n    根据关键点信息裁剪图像，并展示五个裁剪结果。\n    \n    参数：\n    img -- 原始图像数组\n    coords_temp -- 包含关键点坐标的DataFrame，包含'level', 'relative_x', 'relative_y'列\n    \"\"\"\n    # 创建绘图对象和网格布局\n    fig = plt.figure(figsize=(10, 10))\n    gs = gridspec.GridSpec(1, 5, width_ratios=[1]*5)\n    \n    # 按level分组并获取关键点对\n    p = coords_temp.groupby(\"level\").apply(\n        lambda g: list(zip(g['relative_x'], g['relative_y'])),\n        include_groups=False\n    ).reset_index(drop=False, name=\"vals\")\n    \n    # 遍历每组关键点并进行裁剪\n    for idx, (_, row) in enumerate(p.iterrows()):\n        # 复制原始图像\n        img_copy = img.copy()\n        h, w = img.shape\n\n        # 获取当前组的关键点，并按x坐标排序\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        # 计算连线的旋转角度\n        rotate_angle = angle_of_line(a[0], a[1], b[0], b[1])\n        # 创建旋转变换，仅旋转指定角度\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        # 根据旋转后的关键点裁剪图像\n        img_copy = crop_between_keypoints(img_copy, a, b)\n        # 调整裁剪后的图像大小\n        img_copy = resize_transform(image=img_copy)[\"image\"]\n        \n        # 在网格中绘制裁剪后的图像\n        ax = plt.subplot(gs[idx])\n        ax.imshow(img_copy, cmap='gray')\n        ax.set_title(level)\n        ax.axis('off')\n    # 显示所有裁剪图像\n    plt.show()\n\ndef crop_between_keypoints(img, keypoint1, keypoint2):\n    \"\"\"\n    根据两个关键点的位置裁剪图像的特定区域。\n    \n    参数：\n    img -- 原始图像数组\n    keypoint1 -- 第一个关键点的坐标 (x, y)\n    keypoint2 -- 第二个关键点的坐标 (x, y)\n    \n    返回：\n    裁剪后的图像数组\n    \"\"\"\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    # 计算包含两个关键点的边界框\n    left = int(min(x1, x2))\n    right = int(max(x1, x2))\n    top = int(min(y1, y2) - (h * 0.1))  # 上边界向上扩展10%的图像高度\n    bottom = int(max(y1, y2) + (h * 0.1))  # 下边界向下扩展10%的图像高度\n            \n    # 确保裁剪区域在图像范围内\n    top = max(top, 0)\n    bottom = min(bottom, h)\n    left = max(left, 0)\n    right = min(right, w)\n    \n    # 裁剪图像\n    return img[top:bottom, left:right]\n\n# DICOM图像文件夹路径\nimage_dir = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/\"\n","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:04.622944Z","iopub.execute_input":"2024-09-09T15:34:04.623231Z","iopub.status.idle":"2024-09-09T15:34:04.647007Z","shell.execute_reply.started":"2024-09-09T15:34:04.623185Z","shell.execute_reply":"2024-09-09T15:34:04.646078Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Feel free to change the seed, or increase N to see more samples.\n- 你可以随意改变种子，或者增加N来看到更多的样品。","metadata":{}},{"cell_type":"code","source":"SEED = 10  # 设置随机种子，以确保结果可复现\nN = 1  # 设置要加载的样本数量\n\n# 加载series_ids\n# 从CSV文件中加载系列描述数据\ndfd = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv\")\n# 过滤出系列描述为\"Sagittal T2/STIR\"的数据\ndfd = dfd[dfd.series_description == \"Sagittal T2/STIR\"]\n# 随机打乱数据并选取前N个样本\ndfd = dfd.sample(frac=1, random_state=SEED).head(N)\n","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:04.717853Z","iopub.execute_input":"2024-09-09T15:34:04.718254Z","iopub.status.idle":"2024-09-09T15:34:04.734410Z","shell.execute_reply.started":"2024-09-09T15:34:04.718228Z","shell.execute_reply":"2024-09-09T15:34:04.733289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dfd##只筛选出一个人的切片信息","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:04.736294Z","iopub.execute_input":"2024-09-09T15:34:04.736651Z","iopub.status.idle":"2024-09-09T15:34:04.746190Z","shell.execute_reply.started":"2024-09-09T15:34:04.736613Z","shell.execute_reply":"2024-09-09T15:34:04.745040Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# 针对比赛的数据自己制作的数据集合\ncoords = pd.read_csv(\"/kaggle/input/lumbar-coordinate-pretraining-dataset/coords_rsna_improved.csv\")\ncoords.head()\n#特点是，对每一块骨头都标注了两个点，确定出一个关节之间的具体坐标信息","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:04.747580Z","iopub.execute_input":"2024-09-09T15:34:04.747909Z","iopub.status.idle":"2024-09-09T15:34:04.840671Z","shell.execute_reply.started":"2024-09-09T15:34:04.747886Z","shell.execute_reply":"2024-09-09T15:34:04.839697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"coords.shape","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:04.843659Z","iopub.execute_input":"2024-09-09T15:34:04.844074Z","iopub.status.idle":"2024-09-09T15:34:04.850052Z","shell.execute_reply.started":"2024-09-09T15:34:04.844040Z","shell.execute_reply":"2024-09-09T15:34:04.849091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# debug\n### 下面这个表格处理得出了针对比赛训练数据的所有的坐标信息（\"Sagittal T2/STIR\"）/特定方向","metadata":{}},{"cell_type":"code","source":"# 对坐标数据进行排序并重置索引\n#换一下数据 就是片子/骨头/左右\ncoords = coords.sort_values([\"series_id\", \"level\", \"side\"]).reset_index(drop=True)\n# 选择我们需要的列\n\ncoords = coords[[\"series_id\", \"level\", \"side\", \"relative_x\", \"relative_y\"]]\ncoords.head()\n","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:04.851381Z","iopub.execute_input":"2024-09-09T15:34:04.851657Z","iopub.status.idle":"2024-09-09T15:34:04.886387Z","shell.execute_reply.started":"2024-09-09T15:34:04.851632Z","shell.execute_reply":"2024-09-09T15:34:04.885547Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# 单个样本进行测试%\ndicom_files = glob.glob(os.path.join(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/153831832/2054214528\", \"*.dcm\"))\ndicom_files#按照顺序的列表","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:04.887754Z","iopub.execute_input":"2024-09-09T15:34:04.888048Z","iopub.status.idle":"2024-09-09T15:34:04.895732Z","shell.execute_reply.started":"2024-09-09T15:34:04.888023Z","shell.execute_reply":"2024-09-09T15:34:04.894828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# int(len(dicom_files)/2)","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:04.897097Z","iopub.execute_input":"2024-09-09T15:34:04.897492Z","iopub.status.idle":"2024-09-09T15:34:04.901731Z","shell.execute_reply.started":"2024-09-09T15:34:04.897442Z","shell.execute_reply":"2024-09-09T15:34:04.900858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#f'/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/153831832/2054214528/{int(len(dicom_files)/2)}.dcm'","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:04.903027Z","iopub.execute_input":"2024-09-09T15:34:04.903709Z","iopub.status.idle":"2024-09-09T15:34:04.909634Z","shell.execute_reply.started":"2024-09-09T15:34:04.903684Z","shell.execute_reply":"2024-09-09T15:34:04.908841Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dicoms = [pydicom.dcmread(f) for f in dicom_files]\ndicoms[0]","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:04.910758Z","iopub.execute_input":"2024-09-09T15:34:04.911411Z","iopub.status.idle":"2024-09-09T15:34:04.947146Z","shell.execute_reply.started":"2024-09-09T15:34:04.911377Z","shell.execute_reply":"2024-09-09T15:34:04.946321Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plane = \"sagittal\"\n# 根据给定的切面选择对应的平面轴，0: 矢状面, 1: 冠状面, 2: 横断面\nplane = {\"sagittal\": 0, \"coronal\": 1, \"axial\": 2}[plane.lower()]\nplane","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:04.950265Z","iopub.execute_input":"2024-09-09T15:34:04.950541Z","iopub.status.idle":"2024-09-09T15:34:04.956172Z","shell.execute_reply.started":"2024-09-09T15:34:04.950519Z","shell.execute_reply":"2024-09-09T15:34:04.955271Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# 提取每个DICOM图像在所选平面的位置，此处指的是在X平面上的位置（应该是读取进来一堆并不一定是按照顺序读进来的）\npositions = np.asarray([float(d.ImagePositionPatient[plane]) for d in dicoms])\npositions","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:04.957276Z","iopub.execute_input":"2024-09-09T15:34:04.957611Z","iopub.status.idle":"2024-09-09T15:34:04.967713Z","shell.execute_reply.started":"2024-09-09T15:34:04.957584Z","shell.execute_reply":"2024-09-09T15:34:04.966800Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 如果reverse_sort=False，则增加的数组索引将从右到左，并且从尾侧到头侧，\n# 因此我们在横断面时将reverse_sort设置为True，使得增加的数组索引为头尾方向（头侧->尾侧）\nreverse_sort = 'False'\nidx = np.argsort(-positions if reverse_sort else positions)\nidx","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:04.968842Z","iopub.execute_input":"2024-09-09T15:34:04.970767Z","iopub.status.idle":"2024-09-09T15:34:04.977092Z","shell.execute_reply.started":"2024-09-09T15:34:04.970743Z","shell.execute_reply":"2024-09-09T15:34:04.976095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 获取图像的患者位置数组并按照排序后的索引排列\nipp = np.asarray([d.ImagePositionPatient for d in dicoms]).astype(\"float\")[idx]\nipp","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:04.978323Z","iopub.execute_input":"2024-09-09T15:34:04.979113Z","iopub.status.idle":"2024-09-09T15:34:04.990361Z","shell.execute_reply.started":"2024-09-09T15:34:04.979082Z","shell.execute_reply":"2024-09-09T15:34:04.989603Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\n# 获取DICOM图像的像素数据并将其转换为float32类型，然后按排序后的索引排列\narray = np.stack([d.pixel_array.astype(\"float32\") for d in dicoms])\narray = array[idx]\n# 将图像转换为8位，并返回图像数组、位置数组和像素间距信息\nprint({\"array\": convert_to_8bit(array), \"positions\": ipp, \"pixel_spacing\": np.asarray(dicoms[0].PixelSpacing).astype(\"float\")})","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:04.991478Z","iopub.execute_input":"2024-09-09T15:34:04.991767Z","iopub.status.idle":"2024-09-09T15:34:05.225179Z","shell.execute_reply.started":"2024-09-09T15:34:04.991730Z","shell.execute_reply":"2024-09-09T15:34:05.224165Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\n# 遍历样本并绘制图像和裁剪结果\nfor idx, row in dfd.iterrows():\n    try:\n        # 输出当前处理的study_id和series_id\n        print(\"-\" * 25, \" STUDY_ID: {}, SERIES_ID: {} \".format(row.study_id, row.series_id), \"-\" * 25)\n        \n        # 加载指定study_id和series_id对应的DICOM图像（最优质的通道方向）\n        sag_t2 = load_dicom_stack(os.path.join(image_dir, str(row.study_id), str(row.series_id)), plane=\"sagittal\")\n        # 获取图像数据及其对应的关键点坐标\n        img = sag_t2[\"array\"][len(sag_t2[\"array\"])//2]  # 使用中间层的图像进行展示（也就是说只选择了最中间的一个通道进行展示）\n        coords_temp = coords[coords[\"series_id\"] == row.series_id].copy()\n        \n        print('该切片对应的关键点的信息:',coords_temp)\n        # 绘制图像及其关键点\n        plot_img(img, coords_temp)\n        # 绘制图像的五个裁剪区域\n        plot_5_crops(img, coords_temp)\n        \n    except Exception as e:\n        # 如果在处理某个样本时出现异常，则输出异常信息并继续处理下一个样本\n        print(e)\n        pass","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:05.226475Z","iopub.execute_input":"2024-09-09T15:34:05.226794Z","iopub.status.idle":"2024-09-09T15:34:06.141456Z","shell.execute_reply.started":"2024-09-09T15:34:05.226768Z","shell.execute_reply":"2024-09-09T15:34:06.140485Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2. Pretraining\n\nNext, I show a simple pipeline to train a model to predict the x,y coordinates of the 5 lower lumbar vertabrae. The idea is that we first train our image model on a similar task so that the model better suited to tackle our main objective. \n\nThis data was put together in the first version so it does not include left side coordinates. \n\nFor more information on the data, see [here](https://www.kaggle.com/datasets/brendanartley/lumbar-coordinate-pretraining-dataset).\n\n<h1 align=\"left\">\n<img src=\"https://storage.googleapis.com/kaggle-datasets-images/5464745/9091594/db0b402668602e8a6eb772a162f47eb3/dataset-cover.png?t=2024-08-02-23-33-03\" alt=\"spine_img\" width=\"700\">\n</h1>\n","metadata":{}},{"cell_type":"markdown","source":"### 2. 预训练\n\n在这一部分，我将展示一个简单的流水线，用于训练模型来预测下腰椎的x,y坐标。这部分工作的主要思路是首先在一个相似的任务上训练图像模型，使得模型更好地适应并解决我们的主要目标。\n\n**预训练的目的**：  \n预训练通过在相关任务上进行初步训练，使模型能够在随后的主要任务上更好地表现。在这里，模型首先被训练用于预测下腰椎的x,y坐标，这样模型可以学到与脊椎相关的特征，从而在主要任务（比如诊断或分割任务）中表现得更好。\n\n**数据集说明**：  \n这次预训练使用的数据集是在最初版本中汇总的，它不包含左侧坐标的数据。该数据集专注于下腰椎区域的5个椎体的坐标信息。\n\n**数据集的更多信息**：  \n有关这个数据集的更多详细信息可以参考[这里](https://www.kaggle.com/datasets/brendanartley/lumbar-coordinate-pretraining-dataset)。\n\n### 图像展示\n\n为了更好地理解数据集，这里展示了一张数据集的封面图像，该图像展示了脊柱的相关区域。\n\n<img src=\"https://storage.googleapis.com/kaggle-datasets-images/5464745/9091594/db0b402668602e8a6eb772a162f47eb3/dataset-cover.png?t=2024-08-02-23-33-03\" alt=\"脊柱图像\" width=\"700\">\n\n通过这张图像，我们可以看到脊柱的解剖结构，特别是下腰椎区域，这也是我们要预测的关键区域。\n\n**总结**：  \n通过先在类似的任务上对模型进行预训练，模型可以在特征学习方面打下良好的基础，从而在处理更复杂的主要任务时表现得更好。在这个例子中，模型将通过预测脊柱的关键点坐标来进行预训练。","metadata":{}},{"cell_type":"code","source":"# 配置模型训练的参数\ncfg = SimpleNamespace(\n    img_dir=\"/kaggle/input/lumbar-coordinate-pretraining-dataset/data/\",  # 数据集的路径\n    device=torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\"),  # 设置设备为GPU（如果可用），否则使用CPU\n    n_frames=3,  # 处理的帧数，可能是视频数据的时间步长\n    epochs=15,  # 训练的轮数\n    lr=0.0005,  # 学习率，控制模型权重更新的步长\n    batch_size=16,  # 批量大小，每次训练时处理的样本数\n    backbone=\"resnet18\",  # 网络的主干模型，这里使用 ResNet18\n    seed=0,  # 随机种子，确保结果的可重复性\n)\n\n# 设置随机种子，以保证结果的可重复性\nset_seed(seed=cfg.seed)  # Makes results reproducable\n","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:06.142719Z","iopub.execute_input":"2024-09-09T15:34:06.143029Z","iopub.status.idle":"2024-09-09T15:34:06.149055Z","shell.execute_reply.started":"2024-09-09T15:34:06.143003Z","shell.execute_reply":"2024-09-09T15:34:06.148001Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(\"/kaggle/input/lumbar-coordinate-pretraining-dataset/coords_pretrain.csv\")  # 从CSV文件中加载数据\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:06.150369Z","iopub.execute_input":"2024-09-09T15:34:06.150659Z","iopub.status.idle":"2024-09-09T15:34:06.174839Z","shell.execute_reply.started":"2024-09-09T15:34:06.150637Z","shell.execute_reply":"2024-09-09T15:34:06.173954Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['source'].unique()","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:06.175876Z","iopub.execute_input":"2024-09-09T15:34:06.176147Z","iopub.status.idle":"2024-09-09T15:34:06.182787Z","shell.execute_reply.started":"2024-09-09T15:34:06.176124Z","shell.execute_reply":"2024-09-09T15:34:06.181822Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = df.sort_values([\"source\", \"filename\", \"level\"]).reset_index(drop=True)  # 按 \"source\"、\"filename\" 和 \"level\" 列排序，并重置索引\ndf.head(6)","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:06.183952Z","iopub.execute_input":"2024-09-09T15:34:06.184288Z","iopub.status.idle":"2024-09-09T15:34:06.202460Z","shell.execute_reply.started":"2024-09-09T15:34:06.184257Z","shell.execute_reply":"2024-09-09T15:34:06.201618Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df[\"filename\"] = df[\"filename\"].str.replace(\".jpg\", \".npy\")  # 将文件名中的 \".jpg\" 扩展名替换为 \".npy\"\ndf[\"series_id\"] = df[\"source\"] + \"_\" + df[\"filename\"].str.split(\".\").str[0]  # 创建一个新的 \"series_id\" 列，合并 \"source\" 和去掉扩展名的文件名部分\n\ndf.head(6)","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:06.203730Z","iopub.execute_input":"2024-09-09T15:34:06.204384Z","iopub.status.idle":"2024-09-09T15:34:06.233963Z","shell.execute_reply.started":"2024-09-09T15:34:06.204351Z","shell.execute_reply":"2024-09-09T15:34:06.232959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 打印每个来源的图像数量\nprint(\"----- 每个来源的图像数量 -----\")\ndisplay((df.source.value_counts() / 5).astype(int).reset_index())  # 统计每个来源的图像数量，并除以5后显示\n","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:06.235193Z","iopub.execute_input":"2024-09-09T15:34:06.235503Z","iopub.status.idle":"2024-09-09T15:34:06.247829Z","shell.execute_reply.started":"2024-09-09T15:34:06.235478Z","shell.execute_reply":"2024-09-09T15:34:06.246829Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.head()","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:06.249160Z","iopub.execute_input":"2024-09-09T15:34:06.249472Z","iopub.status.idle":"2024-09-09T15:34:06.261445Z","shell.execute_reply.started":"2024-09-09T15:34:06.249441Z","shell.execute_reply":"2024-09-09T15:34:06.260449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset\n\nHere we define the torch dataset that will be used during training.","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        # 将数据转换为字典\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  # 确保每个系列有5个数据点\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)  # 归一化到 [0, 1] 范围\n        return img\n        \n    def __getitem__(self, idx):\n        d = self.records[idx]  # 获取数据信息 records[i] = {\"series_id\": k, \"label\": np.array(v).flatten()}  # 将坐标数据展平并存储\n        label = d[\"label\"]  # 获取标签（一个长度为10的array）\n        source = d[\"series_id\"].split(\"_\")[0]  # 提取来源（总共就四种来源，确定是哪一种方向来的数据）\n        series_id = \"_\".join(d[\"series_id\"].split(\"_\")[1:])  # 提取系列ID\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    \n# 创建数据集实例\nds = PreTrainDataset(df, cfg)    \n\n# 打印单个样本的信息\nprint(\"---- 单个样本的形状 -----\")\nfor k, v in ds[0].items():\n    print(k, v.shape)  # 打印每个字段的形状\n","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:06.262937Z","iopub.execute_input":"2024-09-09T15:34:06.263387Z","iopub.status.idle":"2024-09-09T15:34:06.451437Z","shell.execute_reply.started":"2024-09-09T15:34:06.263355Z","shell.execute_reply":"2024-09-09T15:34:06.450499Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Utils\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.\n### 工具函数\n\n这里有几个辅助函数：\n\n1. **将数据移动到 GPU（如果启用）**：\n   这个函数的作用是将数据从 CPU 移动到 GPU，以利用 GPU 的加速计算能力。这个操作通常是在训练模型时进行的，可以显著提高训练速度。函数的实现通常会检查是否有可用的 GPU，如果有，则将数据移动到 GPU 上；否则，数据将保持在 CPU 上。\n\n2. **在训练期间可视化预测结果**：\n   这个函数用于在训练过程中可视化模型的预测结果。它通常会在训练的每个阶段显示或保存模型的输出图像与实际标签的对比。这有助于监控模型的训练进展，评估模型的性能，并发现潜在的问题。\n\n3. **加载权重时处理形状不匹配**：\n   这个函数的作用是加载模型的预训练权重，并处理权重形状不匹配的情况。在模型的结构发生变化（如调整层的数量或形状）时，预训练的权重可能与当前模型的结构不完全匹配。这个函数可以帮助调整和匹配权重的形状，以便顺利加载权重，避免因形状不匹配导致的错误。\n\n这些工具函数可以帮助简化和加速模型的训练过程，提高训练的效率和准确性。","metadata":{}},{"cell_type":"code","source":"def batch_to_device(batch, device, skip_keys=[]):\n    # 将批次数据移动到指定设备（如 GPU）\n    batch_dict = {}\n    for key in batch:\n        if key in skip_keys:\n            # 如果键在跳过列表中，则不移动到设备\n            batch_dict[key] = batch[key]\n        else:\n            # 否则将数据移动到指定设备\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    for idx in range(1):  # 这里只选择第一个样本进行可视化\n\n        # 选择数据\n        img = batch[\"img\"][idx, mid, :, :].cpu().numpy() * 255  # 获取图像数据并转换为 NumPy 数组\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        # 绘制坐标\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            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            # 在坐标附近添加文本标签\n            for i, (x, y) in enumerate(zip(coords[0::2], coords[1::2])):\n                if i < len(text_labels):  # 确保有足够的标签\n                    ax.text(x + 10, y, text_labels[i], color='white', fontsize=15, bbox=dict(facecolor='black', alpha=0.5))\n\n        plt.show()  # 显示图像\n        # plt.close(fig)  # 可选，关闭图像以释放内存\n    return\n\ndef load_weights_skip_mismatch(model, weights_path, device):\n    # 加载权重时处理形状不匹配\n    state_dict = torch.load(weights_path, map_location=device)  # 加载权重\n    model_dict = model.state_dict()  # 获取模型当前的状态字典\n    \n    params = {}\n    for (sdk, sfv), (mdk, mdv) in zip(state_dict.items(), model_dict.items()):\n        if sfv.size() == mdv.size():\n            # 如果权重的形状匹配，则添加到参数字典中\n            params[sdk] = sfv\n        else:\n            print(\"Skipping param: {}, {} != {}\".format(sdk, sfv.size(), mdv.size()))  # 打印形状不匹配的信息\n    \n    # 加载权重并忽略形状不匹配的参数\n    model.load_state_dict(params, strict=False)\n    print(\"Loaded weights from:\", weights_path)  # 打印加载权重的路径\n","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:06.452536Z","iopub.execute_input":"2024-09-09T15:34:06.452832Z","iopub.status.idle":"2024-09-09T15:34:06.468187Z","shell.execute_reply.started":"2024-09-09T15:34:06.452806Z","shell.execute_reply":"2024-09-09T15:34:06.467207Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model Training\n\nHere we train on all sources except for the spider dataset, which is used for validation.- \n\n- 在所有数据集上进行训练","metadata":{}},{"cell_type":"code","source":"df[\"source\"].unique()","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:06.469214Z","iopub.execute_input":"2024-09-09T15:34:06.469551Z","iopub.status.idle":"2024-09-09T15:34:06.481987Z","shell.execute_reply.started":"2024-09-09T15:34:06.469528Z","shell.execute_reply":"2024-09-09T15:34:06.481146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 数据框\ntrain_df = df[df[\"source\"] != \"spider\"]  # 从数据框中筛选出 source 不等于 \"spider\" 的样本作为训练数据\nval_df = df[df[\"source\"] == \"spider\"]  # 从数据框中筛选出 source 等于 \"spider\" 的样本作为验证数据\nprint(\"训练集大小: {}, 验证集大小: {}\".format(len(train_df) // 5, len(val_df) // 5))\n# 输出训练集和验证集的大小，这里将每个样本数据点的数量除以 5 是因为每个样本可能包含 5 个数据点\n\n# 数据集和数据加载器\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# 模型\nmodel = timm.create_model('resnet18', pretrained=True, num_classes=10)  # 创建 ResNet18 模型，使用预训练权重，并设置输出类别数为 10\nmodel = model.to(cfg.device)  # 将模型移动到指定设备（如 GPU 或 CPU）\n\n# 损失函数和优化器\ncriterion = nn.MSELoss()  # 设置损失函数为均方误差损失，用于回归任务\noptimizer = torch.optim.AdamW(model.parameters(), lr=cfg.lr)  # 设置优化器为 AdamW，学习率为配置中的值\n","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:06.483035Z","iopub.execute_input":"2024-09-09T15:34:06.483385Z","iopub.status.idle":"2024-09-09T15:34:07.019341Z","shell.execute_reply.started":"2024-09-09T15:34:06.483353Z","shell.execute_reply":"2024-09-09T15:34:07.018574Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for epoch in range(cfg.epochs + 1):\n    \n    # 训练循环\n    loss = torch.tensor([0.]).float().to(cfg.device)  # 初始化训练损失为 0\n    if epoch != 0:  # 第 0 轮跳过训练\n        model = model.train()  # 设置模型为训练模式\n        for batch in tqdm(train_dl):  # 遍历训练数据加载器中的批次\n            batch = batch_to_device(batch, cfg.device)  # 将批次数据移动到指定设备\n            \n            optimizer.zero_grad()  # 清零梯度\n            \n            x_out = model(batch[\"img\"].float())  # 前向传播，得到模型输出\n            x_out = torch.sigmoid(x_out)  # 使用 Sigmoid 函数将输出映射到 [0, 1] 范围\n            \n            loss = criterion(x_out, batch[\"label\"].float())  # 计算损失\n            loss.backward()  # 反向传播计算梯度\n            optimizer.step()  # 更新模型参数\n        \n    # 验证循环\n    val_loss = 0  # 初始化验证损失为 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)  # 使用 Sigmoid 函数将预测映射到 [0, 1] 范围\n            \n            val_loss += criterion(pred, batch[\"label\"].float()).item()  # 计算并累加验证损失\n        val_loss /= len(val_dl)  # 计算平均验证损失\n            \n    # 可视化\n    visualize_prediction(batch, pred, epoch)  # 可视化当前批次的预测结果\n            \n    print(f\"Epoch {epoch + 1}, Training Loss: {loss.item()}, Validation Loss: {val_loss}\")\n    # 打印当前轮次的训练损失和验证损失\n\nprint(\"训练完成...\")\n","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:34:07.020579Z","iopub.execute_input":"2024-09-09T15:34:07.020874Z","iopub.status.idle":"2024-09-09T15:35:20.704745Z","shell.execute_reply.started":"2024-09-09T15:34:07.020849Z","shell.execute_reply":"2024-09-09T15:35:20.703658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Save\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":{"execution":{"iopub.status.busy":"2024-09-09T15:35:20.711339Z","iopub.execute_input":"2024-09-09T15:35:20.711741Z","iopub.status.idle":"2024-09-09T15:35:20.788744Z","shell.execute_reply.started":"2024-09-09T15:35:20.711706Z","shell.execute_reply":"2024-09-09T15:35:20.787944Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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)\n\n#总结来说，这个函数的主要功能是从权重文件中加载权重，同时处理模型和权重之间可能存在的形状不匹配问题。\n#对于不匹配的权重，函数会忽略它们，并且不会影响模型的其他参数。\nload_weights_skip_mismatch(model, f, cfg.device)","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:35:20.789878Z","iopub.execute_input":"2024-09-09T15:35:20.790172Z","iopub.status.idle":"2024-09-09T15:35:21.143319Z","shell.execute_reply.started":"2024-09-09T15:35:20.790148Z","shell.execute_reply":"2024-09-09T15:35:21.142098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 保存的模型直接在训练的数据上进行推理查看坐标点\n","metadata":{}},{"cell_type":"markdown","source":"\n- 1对训练数据的表格整理一下可以做推理的格式\n- 2定义dataset\n- 3加载模型\n- 4对数据进行预测\n- 5将预测的坐标结果进行绘制\n\n","metadata":{}},{"cell_type":"markdown","source":"## 2加载训练数据","metadata":{}},{"cell_type":"code","source":"#############GPT########修改后\nimport cv2\nimport numpy as np\nimport torch\nimport glob\nimport os\n\nclass Param:\n    debug = True","metadata":{"execution":{"iopub.status.busy":"2024-09-09T16:48:54.843655Z","iopub.execute_input":"2024-09-09T16:48:54.844006Z","iopub.status.idle":"2024-09-09T16:48:54.848569Z","shell.execute_reply.started":"2024-09-09T16:48:54.843981Z","shell.execute_reply":"2024-09-09T16:48:54.847615Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from IPython.display import display\n\nRSNA_df =  pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv')\n# RSNA_df.head()\nRSNA_coor = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_label_coordinates.csv')\ndisplay(RSNA_coor.head())\n\n# 排序所有数据\nRSNA_coor = RSNA_coor.sort_values(['study_id', 'series_id', 'level'])\n# 显示前 10 行以确认\nRSNA_coor['path_id'] = RSNA_coor.apply(lambda row: f\"{row['study_id']}/{row['series_id']}\", axis=1)\nRSNA_coor = RSNA_coor.drop_duplicates(subset='study_id', keep='first')\nRSNA_coor = RSNA_coor.reset_index(drop=True)\nif Param.debug:\n    RSNA_coor = RSNA_coor.head(34)\nRSNA_coor.tail(2)","metadata":{"execution":{"iopub.status.busy":"2024-09-09T16:48:55.807658Z","iopub.execute_input":"2024-09-09T16:48:55.807998Z","iopub.status.idle":"2024-09-09T16:48:56.656367Z","shell.execute_reply.started":"2024-09-09T16:48:55.807971Z","shell.execute_reply":"2024-09-09T16:48:56.655454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"############=================================================================使用T1效果更好============================================================##\n# #\n# import cv2\n# import matplotlib.pyplot as plt\n\n# # 读取 PNG 图像\n# image = cv2.imread('/kaggle/input/lsdcgcs/cvt_png/100206310/Sagittal T1/005.png', cv2.IMREAD_COLOR)\n# # image = cv2.imread('/kaggle/input/lsdcgcs/cvt_png/100206310/Sagittal T2_STIR/005.png', cv2.IMREAD_COLOR)\n# # image = cv2.imread('/kaggle/input/lsdcgcs/cvt_png/100206310/Axial T2/012.png', cv2.IMREAD_COLOR)\n\n# # 将图像从 BGR 转换为 RGB\n# image_rgb = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\n# # 使用 matplotlib 可视化图像\n# plt.imshow(image_rgb)\n# plt.axis('off')  # 不显示坐标轴\n# plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-09-09T16:48:56.955675Z","iopub.execute_input":"2024-09-09T16:48:56.956034Z","iopub.status.idle":"2024-09-09T16:48:56.960789Z","shell.execute_reply.started":"2024-09-09T16:48:56.956005Z","shell.execute_reply":"2024-09-09T16:48:56.959798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3定义dataset\n# 模型预测只需要获得该具体图像对应的路径即可，不需要其他信息","metadata":{"execution":{"iopub.status.busy":"2024-09-08T14:30:10.051194Z","iopub.execute_input":"2024-09-08T14:30:10.051611Z","iopub.status.idle":"2024-09-08T14:30:10.057994Z","shell.execute_reply.started":"2024-09-08T14:30:10.051576Z","shell.execute_reply":"2024-09-08T14:30:10.056626Z"}}},{"cell_type":"code","source":"center = 1  # 示例变量\nformatted_center = f'{center:03}'  # 将中心数字格式化为3位整数\npath_sag = f'/kaggle/input/lsdcgcs/cvt_png/2/Sagittal T1/{formatted_center}.png'\n","metadata":{"execution":{"iopub.status.busy":"2024-09-09T16:52:54.659126Z","iopub.execute_input":"2024-09-09T16:52:54.659932Z","iopub.status.idle":"2024-09-09T16:52:54.665797Z","shell.execute_reply.started":"2024-09-09T16:52:54.659901Z","shell.execute_reply":"2024-09-09T16:52:54.664814Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\nclass PreTrainDataset(torch.utils.data.Dataset):\n    def __init__(self, df, cfg):\n        self.cfg = cfg  # 配置参数\n        self.df = df\n    \n    def load_img(self, image_path, target_size=(256, 256)):\n        image = cv2.imread(image_path)\n        if image is None:\n            print(f\"Error: Unable to load image at path: {image_path}\")\n            return None\n        \n        # Resize 图像到目标大小\n        image_resized = cv2.resize(image, target_size)\n        # 将图像从 BGR 转换为 RGB\n        image_rgb = cv2.cvtColor(image_resized, cv2.COLOR_BGR2RGB)\n        # 转换图像的维度顺序为 (channels, height, width)\n        image_transposed = np.transpose(image_rgb, (2, 0, 1))\n        # 将图像数据转换为浮点数\n        image_float = image_transposed.astype(np.float32)\n        # 归一化到 [0, 1] 范围\n        image_normalized = image_float / 255.0\n        \n        return image_normalized\n    \n    def __len__(self):\n        return len(self.df)  # 返回数据集的大小\n    \n    def __getitem__(self, idx):\n        # 获取人员ID\n        st_id = self.df['study_id'][idx]\n        \n        # 查找 sagittal 图像\n        num_sag = glob.glob(f'/kaggle/input/lsdcgcs/cvt_png/{st_id}/Sagittal T1/*.png')\n        \n        if len(num_sag) == 0:  # 如果没有找到图片\n            print(f'数据缺失: {st_id}')\n            return {'img': torch.zeros(3, 256, 256)}\n        # 获取中间的切片\n        center = len(num_sag) // 2\n        formatted_center = f'{center:03}'  # 将中心数字格式化为3位整数\n        path_sag = f'/kaggle/input/lsdcgcs/cvt_png/{st_id}/Sagittal T1/{formatted_center}.png'\n        \n        # 检查图像文件是否存在\n        if not os.path.exists(path_sag):\n            print(f\"文件不存在: {path_sag}\")\n            return {\n                'img': torch.zeros(3, 256, 256),  # 返回全0张量\n            }\n        \n        # 加载图像\n        img = self.load_img(path_sag)\n        \n        # 检查图像是否成功加载\n        if img is None:\n            print(f\"加载图像失败: {path_sag}\")\n            img = torch.zeros(3, 256, 256)  # 返回全0张量\n        \n        # 返回图片和标签\n        return {\n            'img': img, \n        }\n","metadata":{"execution":{"iopub.status.busy":"2024-09-09T16:53:24.933942Z","iopub.execute_input":"2024-09-09T16:53:24.934645Z","iopub.status.idle":"2024-09-09T16:53:24.946363Z","shell.execute_reply.started":"2024-09-09T16:53:24.934614Z","shell.execute_reply":"2024-09-09T16:53:24.945363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 加载模型","metadata":{}},{"cell_type":"code","source":"# 加载保存的权重\ndef load_weights_skip_mismatch(model, weights_path, device):\n    # 从权重文件中加载状态字典\n    state_dict = torch.load(weights_path, map_location=device)\n    # 获取模型当前的状态字典\n    model_state_dict = model.state_dict()\n    # 只保留在模型中匹配的权重\n    new_state_dict = {k: v for k, v in state_dict.items() if k in model_state_dict}\n    # 更新模型的状态字典\n    model_state_dict.update(new_state_dict)\n    # 将更新后的状态字典加载到模型中\n    model.load_state_dict(model_state_dict)\n    \n# 重新定义模型结构\nmodel = timm.create_model('resnet18', pretrained=True, num_classes=10)\nmodel = model.to(cfg.device)  # 将模型移动到指定设备\n# 示例使用：\nweights_path = \"/kaggle/working/resnet18_0.pt\"\nload_weights_skip_mismatch(model, weights_path, cfg.device)  # 加载权重到模型\nprint(\"模型的最后一层:\")\nlast_layer = list(model.children())[-1]  # 获取最后一层\nprint(last_layer)","metadata":{"execution":{"iopub.status.busy":"2024-09-09T16:53:26.791926Z","iopub.execute_input":"2024-09-09T16:53:26.792594Z","iopub.status.idle":"2024-09-09T16:53:27.179825Z","shell.execute_reply.started":"2024-09-09T16:53:26.792552Z","shell.execute_reply":"2024-09-09T16:53:27.178925Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 进行模型推理\n- 准备测试数据\n- 测试数据的要过滤重复的行\n- 生成dataset\n- 生成dataloader","metadata":{}},{"cell_type":"code","source":"# 创建数据集实例\ndf = RSNA_coor  # 你的 DataFrame\nprint(df.head(5))\ntrain_ds = PreTrainDataset(df, cfg)  # 创建训练数据集对象，传入配置参数\n\n# 打印单个样本的信息\nprint(\"---- 单个样本的形状 -----\")\nfor k, v in ds[0].items():\n    if isinstance(v, np.ndarray):\n        print(k, v.shape)  # 打印每个字段的形状\n    else:\n        print(k, v)\ntrain_dl = torch.utils.data.DataLoader(train_ds, batch_size=cfg.batch_size, shuffle=True, drop_last=False)  # 创建训练数据加载器，设置批量大小，打乱数据，并丢弃最后一个不完整的批次","metadata":{"execution":{"iopub.status.busy":"2024-09-09T16:53:28.094635Z","iopub.execute_input":"2024-09-09T16:53:28.095432Z","iopub.status.idle":"2024-09-09T16:53:28.112895Z","shell.execute_reply.started":"2024-09-09T16:53:28.095397Z","shell.execute_reply":"2024-09-09T16:53:28.112017Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict(model, data_loader, device):\n    model.eval()  # 将模型设置为评估模式\n    all_preds = []\n    \n    with torch.no_grad():  # 禁用梯度计算\n        for batch in tqdm(data_loader):  # 遍历 DataLoader\n            # 将数据移动到指定设备（CPU或GPU）\n            inputs = batch['img'].to(device).float()  # 假设图像数据位于 'img' 键\n            preds = model(inputs)  # 模型预测\n            preds = torch.sigmoid(preds)  # 假设模型输出 logits，需要用 sigmoid 进行转换\n            all_preds.append(preds.cpu().numpy())  # 收集预测结果\n\n    all_preds = np.concatenate(all_preds, axis=0)  # 合并所有批次的预测\n    return all_preds\n\n# 使用示例\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nall_predictions = predict(model, train_dl, device)\n","metadata":{"execution":{"iopub.status.busy":"2024-09-09T16:53:28.711743Z","iopub.execute_input":"2024-09-09T16:53:28.712589Z","iopub.status.idle":"2024-09-09T16:53:29.190108Z","shell.execute_reply.started":"2024-09-09T16:53:28.712549Z","shell.execute_reply":"2024-09-09T16:53:29.189132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nprint('预测结果的维度:',all_predictions.shape)\nprint('##########')\nprint(all_predictions)","metadata":{"execution":{"iopub.status.busy":"2024-09-09T16:53:33.592533Z","iopub.execute_input":"2024-09-09T16:53:33.592911Z","iopub.status.idle":"2024-09-09T16:53:33.601018Z","shell.execute_reply.started":"2024-09-09T16:53:33.592879Z","shell.execute_reply":"2024-09-09T16:53:33.600100Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 对第一张图片以及预测结果进行可视化\n\n### 写一个绘制的函数\n- 输入的参数是图像的路径，一个列表\n- 该列表包括10个长度的浮点数，代表5个坐标点，就是相对的坐标点   ，原来是这样生成的 \ndef load_coords(self, df):\n        # 将数据转换为字典\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  # 确保每个系列有5个数据点\n        return records\n- 然后将5个坐标点标注在图像中","metadata":{}},{"cell_type":"code","source":"import cv2\nimport matplotlib.pyplot as plt\n\ndef plot_image_with_coords(image_path, coords):\n    \"\"\"\n    绘制图像并在图像上标注坐标点\n\n    Parameters:\n    - image_path: 图像的路径\n    - coords: 一个包含10个浮点数的列表，代表5个坐标点的相对位置（归一化）\n    \"\"\"\n    # 加载图像\n    image = cv2.imread(image_path)\n    if image is None:\n        print(f\"Error: Unable to load image at path: {image_path}\")\n        return\n    \n    # 将图像从 BGR 转换为 RGB\n    image_rgb = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    \n    # 获取图像的宽度和高度\n    height, width, _ = image.shape\n    \n    # 提取坐标点\n    x_coords = coords[::2]  # 提取 x 坐标（归一化）\n    y_coords = coords[1::2]  # 提取 y 坐标（归一化）\n    \n    # 将归一化坐标转换为实际坐标\n    x_coords_actual = [x * width for x in x_coords]\n    y_coords_actual = [y * height for y in y_coords]\n    \n    # 创建图像显示\n    plt.figure(figsize=(5,5))\n    plt.imshow(image_rgb)\n    plt.axis('off')  # 关闭坐标轴\n\n    # 在图像上标注坐标点\n    for x, y in zip(x_coords_actual, y_coords_actual):\n        plt.plot(x, y, 'ro')  # 红色点标记坐标\n        plt.text(x, y, f'({x:.2f}, {y:.2f})', fontsize=12, color='white', ha='right')\n    \n    plt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-09-09T17:02:16.215863Z","iopub.execute_input":"2024-09-09T17:02:16.216223Z","iopub.status.idle":"2024-09-09T17:02:16.225448Z","shell.execute_reply.started":"2024-09-09T17:02:16.216194Z","shell.execute_reply":"2024-09-09T17:02:16.224357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display(RSNA_coor.head(2))","metadata":{"execution":{"iopub.status.busy":"2024-09-09T17:02:20.036406Z","iopub.execute_input":"2024-09-09T17:02:20.037287Z","iopub.status.idle":"2024-09-09T17:02:20.048202Z","shell.execute_reply.started":"2024-09-09T17:02:20.037245Z","shell.execute_reply":"2024-09-09T17:02:20.047208Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"idx_plt = 3\n# 获取第一行的study_id\nstid = RSNA_coor.study_id[idx_plt]\nprint(stid)\n# 获取对应的路径\nnum_sag = glob.glob(f'/kaggle/input/lsdcgcs/cvt_png/{stid}/Sagittal T1/*.png')\ncenter = len(num_sag) // 2\nformatted_center = f'{center:03}'  # 将中心数字格式化为3位整数\npath_sag = f'/kaggle/input/lsdcgcs/cvt_png/{stid}/Sagittal T1/{formatted_center}.png'\n\n# 对应的坐标数据\nCORD = all_predictions[idx_plt]\nprint('坐标数据',CORD)\n#调用函数\n########=================绘制预测结果\nplot_image_with_coords(path_sag,CORD)","metadata":{"execution":{"iopub.status.busy":"2024-09-09T17:02:49.688128Z","iopub.execute_input":"2024-09-09T17:02:49.688770Z","iopub.status.idle":"2024-09-09T17:02:49.902719Z","shell.execute_reply.started":"2024-09-09T17:02:49.688739Z","shell.execute_reply":"2024-09-09T17:02:49.901795Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 接下来需要绘制标准结果","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # 图像可视化\n# import cv2\n# import numpy as np\n# import matplotlib.pyplot as plt\n\n# def load_and_visualize_image(image_path):\n#     # 使用 cv2 加载图像\n#     image = cv2.imread(image_path)\n    \n#     # 检查图像是否加载成功\n#     if image is None:\n#         print(f\"Error: Unable to load image at path: {image_path}\")\n#         return\n    \n#     # 打印图像的原始形状\n#     print(f\"Original image shape: {image.shape}\")\n    \n#     # 将图像从 BGR 转换为 RGB\n#     image_rgb = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    \n#     # 转换图像的维度顺序为 (channels, height, width)\n#     image_transposed = np.transpose(image_rgb, (2, 0, 1))\n    \n#     # 打印转换后的形状\n#     print(f\"Transposed image shape: {image_transposed.shape}\")\n    \n#     # 使用 matplotlib 显示图像\n#     plt.imshow(image_rgb)\n#     plt.axis('off')  # 不显示坐标轴\n#     plt.show()\n\n# # 示例图像路径\n# image_path = '/kaggle/input/lsdcgcs/cvt_png/100206310/Sagittal T1/005.png'\n# load_and_visualize_image(image_path)\n","metadata":{"execution":{"iopub.status.busy":"2024-09-09T15:35:23.926623Z","iopub.status.idle":"2024-09-09T15:35:23.926936Z","shell.execute_reply.started":"2024-09-09T15:35:23.926780Z","shell.execute_reply":"2024-09-09T15:35:23.926793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 2根据重新打标的训练数据坐标点进行训练（进行的是双点的）","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}