{"metadata":{"colab":{"provenance":[],"gpuType":"T4"},"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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"},"accelerator":"GPU","kaggle":{"accelerator":"none","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":3951115,"sourceType":"datasetVersion","datasetId":1027206},{"sourceId":9197435,"sourceType":"datasetVersion","datasetId":5560485},{"sourceId":9220371,"sourceType":"datasetVersion","datasetId":5576017},{"sourceId":9220475,"sourceType":"datasetVersion","datasetId":5575996},{"sourceId":9467045,"sourceType":"datasetVersion","datasetId":5756451},{"sourceId":9471621,"sourceType":"datasetVersion","datasetId":5760035},{"sourceId":193044752,"sourceType":"kernelVersion"},{"sourceId":193398994,"sourceType":"kernelVersion"},{"sourceId":193487897,"sourceType":"kernelVersion"},{"sourceId":119225,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":100262,"modelId":124431}],"dockerImageVersionId":30747,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\n\ndf = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_series_descriptions.csv')\ndf.head()\n","metadata":{"id":"Q3hWy2qQ8G91","outputId":"e9c4f529-8476-425c-8250-4cac2f644506","execution":{"iopub.status.busy":"2024-10-02T23:49:56.740329Z","iopub.execute_input":"2024-10-02T23:49:56.740725Z","iopub.status.idle":"2024-10-02T23:49:57.962611Z","shell.execute_reply.started":"2024-10-02T23:49:56.740692Z","shell.execute_reply":"2024-10-02T23:49:57.961239Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sag2 = df[df['series_description'] == 'Sagittal T2/STIR']","metadata":{"id":"r_xJrrnW-Pju","execution":{"iopub.status.busy":"2024-10-02T23:50:05.345345Z","iopub.execute_input":"2024-10-02T23:50:05.345740Z","iopub.status.idle":"2024-10-02T23:50:05.355676Z","shell.execute_reply.started":"2024-10-02T23:50:05.345707Z","shell.execute_reply":"2024-10-02T23:50:05.354331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\n\nstudy_base_path = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images\"\n\nall_study_ids = os.listdir(study_base_path)\n\nall_study_ids = [int(study_id) for study_id in all_study_ids if study_id.isdigit()]\n\nall_study_ids","metadata":{"execution":{"iopub.status.busy":"2024-10-02T23:50:53.109903Z","iopub.execute_input":"2024-10-02T23:50:53.110621Z","iopub.status.idle":"2024-10-02T23:50:53.124114Z","shell.execute_reply.started":"2024-10-02T23:50:53.110583Z","shell.execute_reply":"2024-10-02T23:50:53.123093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\n\n# Assuming sag2 is a dataframe containing the columns 'study_id' and 'series_id'\nbase_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images'\n# base_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images'\n\n# Initialize a dictionary to store paths for each study_id\nsagittal_series_paths = {}\n\nfor study_id, group in sag2.groupby('study_id'):\n    # Get the list of series_id for this study_id\n    series_ids = group['series_id'].tolist()\n\n    # Create the full paths for each series_id under the study_id\n    paths = [os.path.join(base_path, str(study_id), str(\n        series_id)) for series_id in series_ids]\n\n    # Store the paths in the dictionary\n    sagittal_series_paths[study_id] = paths\n","metadata":{"execution":{"iopub.status.busy":"2024-10-02T23:50:58.951689Z","iopub.execute_input":"2024-10-02T23:50:58.952072Z","iopub.status.idle":"2024-10-02T23:50:58.964417Z","shell.execute_reply.started":"2024-10-02T23:50:58.952027Z","shell.execute_reply":"2024-10-02T23:50:58.963152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sagittal_series_paths.keys()\n\nstudy_ids_without_sag2 = [study_id for study_id in sagittal_series_paths.keys() if study_id not in all_study_ids]\n\nstudy_ids_without_sag2","metadata":{"execution":{"iopub.status.busy":"2024-10-02T23:51:04.081505Z","iopub.execute_input":"2024-10-02T23:51:04.081918Z","iopub.status.idle":"2024-10-02T23:51:04.089988Z","shell.execute_reply.started":"2024-10-02T23:51:04.081880Z","shell.execute_reply":"2024-10-02T23:51:04.088583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\n!pip install -q segmentation_models_pytorch\n!pip install -q natsort\n'''","metadata":{"execution":{"iopub.status.busy":"2024-09-26T12:05:08.037482Z","iopub.execute_input":"2024-09-26T12:05:08.037859Z","iopub.status.idle":"2024-09-26T12:05:08.044703Z","shell.execute_reply.started":"2024-09-26T12:05:08.037829Z","shell.execute_reply":"2024-09-26T12:05:08.043744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n!pip install ../input/segmentation-models-pytorch-wheels-install-once/efficientnet_pytorch-0.7.1-py3-none-any.whl\n!pip install ../input/segmentation-models-pytorch-wheels-install-once/munch-4.0.0-py2.py3-none-any.whl\n!pip install ../input/segmentation-models-pytorch-wheels-install-once/pretrainedmodels-0.7.4-py3-none-any.whl\n!pip install ../input/segmentation-models-pytorch-wheels-install-once/timm-0.9.2-py3-none-any.whl\n!pip install ../input/segmentation-models-pytorch-wheels-install-once/segmentation_models_pytorch-0.3.3-py3-none-any.whl\n","metadata":{"execution":{"iopub.status.busy":"2024-10-02T23:51:17.399789Z","iopub.execute_input":"2024-10-02T23:51:17.400230Z","iopub.status.idle":"2024-10-02T23:54:13.583872Z","shell.execute_reply.started":"2024-10-02T23:51:17.400196Z","shell.execute_reply":"2024-10-02T23:54:13.582410Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir -p /root/.cache/torch/hub/checkpoints/\n!cp ../input/pytorch-pretrained-resnet34-model/resnet34-333f7ec4.pth /root/.cache/torch/hub/checkpoints/resnet34-333f7ec4.pth","metadata":{"execution":{"iopub.status.busy":"2024-10-02T23:54:13.586880Z","iopub.execute_input":"2024-10-02T23:54:13.587380Z","iopub.status.idle":"2024-10-02T23:54:17.844963Z","shell.execute_reply.started":"2024-10-02T23:54:13.587333Z","shell.execute_reply":"2024-10-02T23:54:17.843470Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install /kaggle/input/natsort/natsort/natsort-8.4.0-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2024-10-02T23:54:17.846719Z","iopub.execute_input":"2024-10-02T23:54:17.847122Z","iopub.status.idle":"2024-10-02T23:54:52.469250Z","shell.execute_reply.started":"2024-10-02T23:54:17.847085Z","shell.execute_reply":"2024-10-02T23:54:52.467955Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport pydicom\nimport numpy as np\nimport torch\nfrom segmentation_models_pytorch import Unet\nfrom natsort import natsorted\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom PIL import Image\nfrom skimage.measure import regionprops\n\n# Preprocessing transformation (based on your training setup)\ntransforms_valid = A.Compose([\n    A.Resize(256, 256),\n    A.Normalize(mean=[0.485], std=[0.229]),\n    ToTensorV2()\n])\n\ndef load_middle_slice(study_id, sagittal_series_paths):\n    # Get the folder containing slices for this study_id\n    folder_path = sagittal_series_paths[study_id]\n\n    # List all .dcm files and sort them\n    slice_files = natsorted([f for f in os.listdir(folder_path[0]) if f.endswith('.dcm')])\n\n    # Find the middle slice\n    middle_index = len(slice_files) // 2\n    middle_slice_path = os.path.join(folder_path[0], slice_files[middle_index])\n\n    # Load and preprocess the DICOM image\n    dicom_image = pydicom.dcmread(middle_slice_path).pixel_array\n\n    if (dicom_image<0).any():\n        normalized_image = (dicom_image - dicom_image.min()) / (dicom_image.max() - dicom_image.min())\n        dicom_image=(normalized_image * 255).astype(np.uint8)\n\n    image = Image.fromarray(dicom_image).convert('L')\n    image = np.asarray(image)\n\n    if (image > 1).any():  # Normalize if pixel values are between 0-255\n        image = image / 255.0\n\n    transformed = transforms_valid(image=image)\n    image_tensor = transformed[\"image\"].unsqueeze(0)  # Add batch dimension\n\n    return image_tensor, dicom_image, middle_slice_path  # Returning dicom_image for visualization if needed\n\ndef predict_segmentation(model, image_tensor):\n    model.eval()\n    with torch.no_grad():\n        prediction = model(image_tensor)\n        predicted_mask = torch.argmax(prediction, dim=1).squeeze(0).cpu().numpy()\n    return predicted_mask\n\ndef calculate_centroids(segmentation_mask):\n    centroids = {}\n    for vertebrae_class in range(1, 6):  # Classes 1 to 5 correspond to L5 to L1\n        vertebrae_mask = (segmentation_mask == vertebrae_class).astype(np.uint8)\n        regions = regionprops(vertebrae_mask)\n        if regions:\n            centroid = regions[0].centroid  # Take the first region's centroid\n            centroids[vertebrae_class] = centroid  # Store as {class_label: (y, x)}\n    return centroids\n\n'''\nmodel = Unet('resnet34', classes=6, in_channels=1)  # Assuming you have a pre-trained model\nmodel.load_state_dict(torch.load('/kaggle/input/lumbar-vertebrae-segmentation-using-spider-datset/simple_unet.pth', map_location='cpu'))\n\nstudy_id = 4646740 #44036939    #1217004843#296314829           # Change to your study id\n\nimage_tensor, dicom_image, slice_path = load_middle_slice(study_id, sagittal_series_paths)\npredicted_mask = predict_segmentation(model, image_tensor)\ncentroids = calculate_centroids(predicted_mask)\n\nprint(f\"Centroids for study {study_id}: {centroids}\")\n'''","metadata":{"id":"KzUNUYF8-s2K","outputId":"97dab300-9e45-49e7-9d5e-3c32afcb0f5a","execution":{"iopub.status.busy":"2024-10-02T23:56:30.834442Z","iopub.execute_input":"2024-10-02T23:56:30.835273Z","iopub.status.idle":"2024-10-02T23:56:40.131339Z","shell.execute_reply.started":"2024-10-02T23:56:30.835237Z","shell.execute_reply":"2024-10-02T23:56:40.130066Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\ndef plot_segmentation_with_centroids(dicom_image, predicted_mask, centroids):\n    original_height, original_width = dicom_image.shape\n    resized_height, resized_width = 256, 256  # The mask size after resizing\n\n    # Scaling factors to map the 256x256 mask coordinates back to the original DICOM size\n    x_scale = original_width / resized_width\n    y_scale = original_height / resized_height\n\n    fig, axes = plt.subplots(1, 2, figsize=(12, 6))\n\n    # Plot the original DICOM image\n    axes[0].imshow(dicom_image, cmap='gray')\n    axes[0].set_title('Original DICOM Image')\n    axes[0].axis('off')\n\n    # Plot the predicted segmentation mask\n    axes[1].imshow(predicted_mask, cmap='nipy_spectral')\n    axes[1].set_title('Predicted Segmentation Mask (256x256)')\n    axes[1].axis('off')\n\n    # Overlay scaled centroids on both images\n    for class_id, centroid in centroids.items():\n        y, x = centroid\n        # Rescale centroids to original DICOM size\n        y_original = y * y_scale\n        x_original = x * x_scale\n\n        # Overlay on the DICOM image\n        axes[0].plot(x_original, y_original, 'ro')\n        axes[0].text(x_original + 2, y_original, f'L{6-class_id}', color='white', fontsize=8)\n\n        # Overlay on the resized segmentation mask\n        axes[1].plot(x, y, 'ro')\n        axes[1].text(x + 2, y, f'L{6-class_id}', color='white', fontsize=8)\n\n    plt.tight_layout()\n    plt.show()\n'''\n# Example usage:\nplot_segmentation_with_centroids(dicom_image, predicted_mask, centroids)\n'''\n","metadata":{"id":"SkjYRyxz_R5G","outputId":"90610b66-0312-40df-a54d-1b973dee74f9","execution":{"iopub.status.busy":"2024-10-02T23:56:40.133408Z","iopub.execute_input":"2024-10-02T23:56:40.133990Z","iopub.status.idle":"2024-10-02T23:56:40.150590Z","shell.execute_reply.started":"2024-10-02T23:56:40.133957Z","shell.execute_reply":"2024-10-02T23:56:40.149310Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nfrom numpy.polynomial import Polynomial\nfrom scipy.interpolate import interp1d\n\ndef fit_polynomial_curve(centroids):\n    # Extract Y and X coordinates from the centroids dictionary\n    y_coords = np.array([centroid[0] for centroid in centroids.values()])  # Y-coordinates\n    x_coords = np.array([centroid[1] for centroid in centroids.values()])  # X-coordinates\n\n    # Fit a polynomial (e.g., degree 2 or 3)\n    polynomial = np.polyfit(y_coords, x_coords, deg=2)  # Degree 2 polynomial\n    poly_func = np.poly1d(polynomial)\n\n    return poly_func\n\n\ndef predict_disc_centroids(centroids, poly_func):\n    y_coords = np.array([centroid[0] for centroid in centroids.values()])\n\n    # Calculate midpoints between consecutive vertebrae for the intervertebral discs\n    disc_y_coords = [(y_coords[i] + y_coords[i + 1]) / 2 for i in range(len(y_coords) - 1)]\n\n    # Corrected naming\n    disc_centroids = {f'L{6-(i+2)}/L{6-(i+1)}': (y, poly_func(y)) for i, y in enumerate(disc_y_coords)}\n\n    # Handle L5/S1 approximation\n    l4_l5_y = disc_centroids['L4/L5'][0]\n    l5_y = centroids[1][0]  # L5 vertebra Y-coordinate\n    l5_s1_y = (l5_y + (l5_y-l4_l5_y))\n    l5_s1_x = poly_func(l5_s1_y)\n\n    # Add the L5/S1 disc centroid\n    disc_centroids['L5/S1'] = (l5_s1_y, l5_s1_x)\n\n    return disc_centroids\n'''\n# Example usage:\npoly_func = fit_polynomial_curve(centroids)\ndisc_centroids = predict_disc_centroids(centroids, poly_func)\n\nprint(f\"Predicted disc centroids for study {study_id}: {disc_centroids}\")\n'''","metadata":{"id":"-SfCgGiMAAN6","outputId":"4f458287-319c-49df-9102-ec6e3da37fc7","execution":{"iopub.status.busy":"2024-10-02T23:56:40.152006Z","iopub.execute_input":"2024-10-02T23:56:40.152383Z","iopub.status.idle":"2024-10-02T23:56:40.173129Z","shell.execute_reply.started":"2024-10-02T23:56:40.152353Z","shell.execute_reply":"2024-10-02T23:56:40.171787Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_all_centroids(dicom_image, predicted_mask, vertebrae_centroids, disc_centroids):\n    original_height, original_width = dicom_image.shape\n    resized_height, resized_width = 256, 256  # The mask size after resizing\n\n    # Scaling factors to map the 256x256 mask coordinates back to the original DICOM size\n    x_scale = original_width / resized_width\n    y_scale = original_height / resized_height\n\n    fig, axes = plt.subplots(1, 2, figsize=(12, 6))\n\n    # Plot the original DICOM image\n    axes[0].imshow(dicom_image, cmap='gray')\n    axes[0].set_title('Original DICOM Image')\n    axes[0].axis('off')\n\n    # Plot the predicted segmentation mask\n    axes[1].imshow(predicted_mask, cmap='nipy_spectral')\n    axes[1].set_title('Predicted Segmentation Mask (256x256)')\n    axes[1].axis('off')\n\n    # Overlay vertebrae centroids\n    for class_id, centroid in vertebrae_centroids.items():\n        y, x = centroid\n        y_original = y * y_scale\n        x_original = x * x_scale\n        axes[0].plot(x_original, y_original, 'ro')\n        axes[0].text(x_original + 2, y_original, f'L{6-class_id}', color='white', fontsize=8)\n        axes[1].plot(x, y, 'ro')\n        axes[1].text(x + 2, y, f'L{6-class_id}', color='white', fontsize=8)\n\n    # Overlay disc centroids\n    for disc_label, (y_disc, x_disc) in disc_centroids.items():\n        y_disc_original = y_disc * y_scale\n        x_disc_original = x_disc * x_scale\n        axes[0].plot(x_disc_original, y_disc_original, 'bo')  # Blue dots for disc centroids\n        axes[0].text(x_disc_original + 2, y_disc_original, disc_label, color='cyan', fontsize=8)\n        axes[1].plot(x_disc, y_disc, 'bo')\n        axes[1].text(x_disc + 2, y_disc, disc_label, color='cyan', fontsize=8)\n\n    plt.tight_layout()\n    plt.show()\n'''\n# Example usage:\nplot_all_centroids(dicom_image, predicted_mask, centroids, disc_centroids)\n'''","metadata":{"id":"OCWNqwZRASwf","outputId":"57243c4a-debf-4a40-bd3b-32cafe2a04ca","execution":{"iopub.status.busy":"2024-10-02T23:56:40.175846Z","iopub.execute_input":"2024-10-02T23:56:40.176588Z","iopub.status.idle":"2024-10-02T23:56:40.195954Z","shell.execute_reply.started":"2024-10-02T23:56:40.176554Z","shell.execute_reply":"2024-10-02T23:56:40.194610Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def extract_disc_crops(sagittal_series_paths, study_id, disc_centroids, crop_size=(64, 64)):\n    folder_path = sagittal_series_paths[study_id][0]\n    slice_files = natsorted([f for f in os.listdir(folder_path) if f.endswith('.dcm')])\n\n    disc_crops = {}\n    half_crop_height = crop_size[0] // 2\n    half_crop_width = crop_size[1] // 2\n\n    # Loop through each disc level and extract crops across all slices\n    for disc_label, (y_centroid, x_centroid) in disc_centroids.items():\n        crops = []\n        for slice_file in slice_files:\n            slice_path = os.path.join(folder_path, slice_file)\n            dicom_image = pydicom.dcmread(slice_path).pixel_array\n\n            # Get the actual dimensions of the current slice\n            original_height, original_width = dicom_image.shape\n\n            # Calculate the scaling factors based on the current slice's shape\n            x_scale = original_width / 256\n            y_scale = original_height / 256\n\n            # Scale the centroids to the actual size of this slice\n            y_centroid_original = y_centroid * y_scale\n            x_centroid_original = x_centroid * x_scale\n\n            # Calculate the bounding box coordinates using the scaled centroids\n            top_left_y = int(max(0, y_centroid_original - half_crop_height))\n            bottom_right_y = int(min(original_height, y_centroid_original + half_crop_height))\n            top_left_x = int(max(0, x_centroid_original - half_crop_width))\n            bottom_right_x = int(min(original_width, x_centroid_original + half_crop_width))\n\n            # Extract the crop\n            crop = dicom_image[top_left_y:bottom_right_y, top_left_x:bottom_right_x]\n            crops.append(crop)\n\n        # Store the crops for this disc level as a 3D array (slice_count × width × height)\n        disc_crops[disc_label] = np.stack(crops, axis=0)\n\n    return disc_crops\n'''\n# Example usage:\ndisc_crops = extract_disc_crops(sagittal_series_paths, study_id, disc_centroids, crop_size=(100, 250))\n\nimport matplotlib.pyplot as plt\n\n# Assuming disc_crops is a dictionary with disc labels as keys and lists of images as values\nsample_disc_label = 'L5/S1'\nsample_crop = disc_crops[sample_disc_label]\n\n# Calculate the number of rows based on the number of images and columns (4 in this case)\nnum_images = len(sample_crop)\ncols = 4\nrows = (num_images // cols) + (1 if num_images % cols > 0 else 0)\n\n# Create the subplots\nfig, ax = plt.subplots(rows, cols, figsize=(40, 40))\n\n# Flatten the axes array for easy iteration (in case of multiple rows)\nax = ax.flatten()\n\nfor i, slice_image in enumerate(sample_crop):\n    ax[i].imshow(slice_image, cmap='gray')\n    ax[i].axis('off')\n\n# Hide any unused subplots\nfor i in range(num_images, rows * cols):\n    ax[i].axis('off')\n\nplt.tight_layout()\nplt.show()\n\n'''","metadata":{"id":"WwuuDEQeAW4-","outputId":"d3f69593-7a67-4c67-ea0a-b0372fcc024a","execution":{"iopub.status.busy":"2024-10-02T23:56:40.197563Z","iopub.execute_input":"2024-10-02T23:56:40.197931Z","iopub.status.idle":"2024-10-02T23:56:40.217378Z","shell.execute_reply.started":"2024-10-02T23:56:40.197901Z","shell.execute_reply":"2024-10-02T23:56:40.216221Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport torch\n\n# Configuration\ncrop_size = (84, 160)\noutput_dataset_path = \"./sagittalT2\"\nos.makedirs(output_dataset_path, exist_ok=True)\n\n# List to store the study_ids that encounter errors\nskipped_studies = []\nif len(study_ids_without_sag2) > 0:\n    skipped_studies.append(study_ids_without_sag2)\n\ndef create_dataset(sagittal_series_paths):\n    dataset = {}\n\n    for study_id in sagittal_series_paths:\n        try:\n\n\n            # Processing steps for each study_id\n            image_tensor, dicom_image, middle_slice_path = load_middle_slice(study_id, sagittal_series_paths)\n            predicted_mask = predict_segmentation(model, image_tensor)\n            centroids = calculate_centroids(predicted_mask)\n            poly_func = fit_polynomial_curve(centroids)\n            disc_centroids = predict_disc_centroids(centroids, poly_func)\n            disc_crops = extract_disc_crops(sagittal_series_paths, study_id, disc_centroids, crop_size)\n            # Create a directory for each study_id\n            os.makedirs(os.path.join(output_dataset_path, str(study_id)), exist_ok=True)\n            # Save cropped disc slices\n            for disc_label, crops in disc_crops.items():\n\n                save_path = os.path.join(output_dataset_path, str(study_id), f\"{disc_label.replace('/', '_')}.npy\")\n                np.save(save_path, crops)  # Save the crops as numpy arrays\n\n                if study_id not in dataset:\n                    dataset[study_id] = {}\n                dataset[study_id][disc_label] = save_path  # Store the path in the dataset structure\n\n        except Exception as e:\n            print(f\"Study {study_id} was skipped due to error: {e}\")\n            skipped_studies.append(study_id)  # Add the skipped study_id to the list\n\n    return pd.DataFrame(dataset).T  # Convert the dataset to a pandas DataFrame for easy analysis\n\nmodel = Unet('resnet34', classes=6, in_channels=1)  # Assuming you have a pre-trained model\nmodel.load_state_dict(torch.load('/kaggle/input/lumbar-vertebrae-segmentation-using-spider-datset/simple_unet.pth', map_location='cpu'))\nmodel.eval()\n\n# Generate the dataset\ndataset_df = create_dataset(sagittal_series_paths)\n\n# Save the dataset information\ndataset_df.to_csv(os.path.join(output_dataset_path, 'dataset_metadata.csv'), index=True)\n'''\n# Save the skipped study_ids to a text file or print them for review\nif skipped_studies:\n    print(f\"The following studies were skipped due to errors: {skipped_studies}\")\n    with open(os.path.join(output_dataset_path, 'skipped_studies.txt'), 'w') as f:\n        for study_id in skipped_studies:\n            f.write(f\"{study_id}\\n\")\n'''\nprint(\"Pipeline completed successfully!\")","metadata":{"id":"gY_78ctbAeFD","outputId":"dd03c067-c755-4f21-96f8-4f0803264ee9","execution":{"iopub.status.busy":"2024-10-02T23:57:02.160422Z","iopub.execute_input":"2024-10-02T23:57:02.161100Z","iopub.status.idle":"2024-10-02T23:57:10.895755Z","shell.execute_reply.started":"2024-10-02T23:57:02.161068Z","shell.execute_reply":"2024-10-02T23:57:10.894713Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"skipped_studies","metadata":{"execution":{"iopub.status.busy":"2024-10-03T00:00:33.880868Z","iopub.execute_input":"2024-10-03T00:00:33.881720Z","iopub.status.idle":"2024-10-03T00:00:33.888330Z","shell.execute_reply.started":"2024-10-03T00:00:33.881685Z","shell.execute_reply":"2024-10-03T00:00:33.887249Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ax=df[df['series_description']=='Axial T2']\n\nimport os\n\n# Assuming sag2 is a dataframe containing the columns 'study_id' and 'series_id'\nbase_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images'\n#base_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images'\n\n# Initialize a dictionary to store paths for each study_id\naxial_series_paths = {}\n\nfor study_id, group in ax.groupby('study_id'):\n    # Get the list of series_id for this study_id\n    series_ids = group['series_id'].tolist()\n    #print(group['study_id'])\n\n    # Create the full paths for each series_id under the study_id\n    paths = [os.path.join(base_path, str(study_id), str(series_id)) for series_id in series_ids]\n\n    # Store the paths in the dictionary\n    axial_series_paths[study_id] = paths\n\nstudy_ids_without_ax = [study_id for study_id in axial_series_paths.keys() if study_id not in all_study_ids]\n\nif len(study_ids_without_ax) > 0:\n    skipped_studies.append(study_ids_without_ax)","metadata":{"id":"2BXFpxZdB-bn","execution":{"iopub.status.busy":"2024-10-03T00:02:26.076719Z","iopub.execute_input":"2024-10-03T00:02:26.077448Z","iopub.status.idle":"2024-10-03T00:02:26.091602Z","shell.execute_reply.started":"2024-10-03T00:02:26.077416Z","shell.execute_reply":"2024-10-03T00:02:26.090538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"skipped_studies\n\nskipped_studies = list(dict.fromkeys(skipped_studies))\n\nlen(skipped_studies)","metadata":{"execution":{"iopub.status.busy":"2024-09-26T12:05:35.504225Z","iopub.execute_input":"2024-09-26T12:05:35.504656Z","iopub.status.idle":"2024-09-26T12:05:35.512167Z","shell.execute_reply.started":"2024-09-26T12:05:35.504615Z","shell.execute_reply":"2024-09-26T12:05:35.511097Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport glob\nimport pydicom\nimport numpy as np\nimport cv2\n\ndef load_and_sort_axial_slices(axial_series_paths, study_id, target_shape=(256, 256)):\n    series_folders = axial_series_paths[study_id]\n\n    dicom_files = []\n    for folder in series_folders:\n        dicom_files.extend(glob.glob(os.path.join(folder, \"*.dcm\")))\n\n    dicoms = [pydicom.dcmread(f) for f in dicom_files]\n    z_positions = [float(d.ImagePositionPatient[2]) for d in dicoms]\n\n    sorted_indices = np.argsort(-np.array(z_positions))\n    sorted_dicoms = [dicoms[i] for i in sorted_indices]\n\n    # Resize each pixel array to the target shape\n    pixel_arrays = []\n    for d in sorted_dicoms:\n        pixel_array = d.pixel_array\n        resized_pixel_array = cv2.resize(pixel_array, target_shape, interpolation=cv2.INTER_LINEAR)\n        pixel_arrays.append(resized_pixel_array)\n\n    # Stack the resized arrays\n    pixel_arrays = np.stack(pixel_arrays)\n    sorted_positions = np.array([d.ImagePositionPatient for d in sorted_dicoms])\n\n    return {\n        \"array\": pixel_arrays,\n        \"positions\": sorted_positions,\n        \"pixel_spacing\": np.array(sorted_dicoms[0].PixelSpacing, dtype=np.float32)\n    }\n'''\n# Example usage:\nstudy_id=4646740  #44036939\naxial_slices = load_and_sort_axial_slices(axial_series_paths, study_id, target_shape=(256, 256))\n# Visualizing the sorted Z-positions\nprint(\"Sorted Z-positions:\", axial_slices[\"positions\"][:, 2])\n'''","metadata":{"id":"MyalFs-SCOhh","outputId":"24f0f81b-c331-4bc6-d5ba-4609e2f7b8e7","execution":{"iopub.status.busy":"2024-09-26T12:05:37.632139Z","iopub.execute_input":"2024-09-26T12:05:37.632568Z","iopub.status.idle":"2024-09-26T12:05:37.647688Z","shell.execute_reply.started":"2024-09-26T12:05:37.632534Z","shell.execute_reply":"2024-09-26T12:05:37.646480Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def sag_middle(file_path):\n\n    slice_files = natsorted([f for f in os.listdir(file_path) if f.endswith('.dcm')])\n\n    # Find the middle slice\n    middle_index = len(slice_files) // 2\n    middle_slice_path = os.path.join(file_path, slice_files[middle_index])\n    dcmfile=pydicom.dcmread(middle_slice_path)\n\n    return {\n        'array':dcmfile.pixel_array,\n        'positions':dcmfile.ImagePositionPatient,\n        \"pixel_spacing\":dcmfile.PixelSpacing\n    }\n\ndef map_sagittal_y_to_axial_slices(sagittal_slice, study_id, axial_slices):\n\n    top_left_z_coordinate = sagittal_slice[\"positions\"][2]\n    pixel_spacing_y = sagittal_slice[\"pixel_spacing\"][1]  # Pixel spacing in the Y direction (corresponds to z-axis in world space)\n\n    # Create the Y-coordinate in pixel space for the sagittal image\n    sag_y_axis_to_pixel_space = [top_left_z_coordinate]\n    for _ in range(sagittal_slice[\"array\"].shape[1] - 1):  # For each pixel row in the sagittal image\n        sag_y_axis_to_pixel_space.append(sag_y_axis_to_pixel_space[-1] - pixel_spacing_y)\n\n    # Map sagittal Y-coordinates to corresponding axial slices\n    sag_y_coord_to_axial_slice = {}\n    for ax_slice, ax_position in zip(axial_slices[\"array\"], axial_slices[\"positions\"]):\n        # Find the closest match between the axial slice's Z-coordinate and the sagittal Y-coordinates\n        diffs = np.abs(np.asarray(sag_y_axis_to_pixel_space) - ax_position[2])\n        closest_y_coord = np.argmin(diffs)\n        sag_y_coord_to_axial_slice[closest_y_coord] = ax_slice\n\n    return sag_y_coord_to_axial_slice\n'''\n# Example usage:\naxial_slices = load_and_sort_axial_slices(axial_series_paths, study_id,target_shape=(256, 256))\nsagittal_slice=sag_middle(sagittal_series_paths[study_id][0])\nsag_y_coord_to_axial_mapping = map_sagittal_y_to_axial_slices(sagittal_slice, study_id, axial_slices)\n\nprint(\"Sagittal Y to Axial Slice Mapping:\", sag_y_coord_to_axial_mapping.keys())\n'''","metadata":{"id":"ak7J8mQ_CVZz","outputId":"d5c6d9fd-2d28-4d92-bb4b-81426b02da08","execution":{"iopub.status.busy":"2024-09-26T12:05:39.853567Z","iopub.execute_input":"2024-09-26T12:05:39.853993Z","iopub.status.idle":"2024-09-26T12:05:39.868906Z","shell.execute_reply.started":"2024-09-26T12:05:39.853958Z","shell.execute_reply":"2024-09-26T12:05:39.867711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\ndef verify_mapping(sagittal_slice, sag_y_coord_to_axial_mapping):\n    # Plot the sagittal image\n    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 6))\n\n    ax1.imshow(sagittal_slice[\"array\"], cmap='gray')\n    ax1.set_title(\"Sagittal View with Axial Mapping\")\n\n    # Overlay horizontal lines indicating the mapped Y-coordinates\n    for y_coord in sag_y_coord_to_axial_mapping.keys():\n        ax1.axhline(y=y_coord, color='red', linestyle='--', linewidth=1)\n\n    ax1.set_xlim([0, sagittal_slice[\"array\"].shape[1]])\n    ax1.set_ylim([sagittal_slice[\"array\"].shape[0], 0])  # Invert the y-axis for correct orientation\n\n    # Plot a sample axial slice\n    sample_axial_slice = list(sag_y_coord_to_axial_mapping.values())[0]  # Just for visualization\n    ax2.imshow(sample_axial_slice, cmap='gray')\n    ax2.set_title(\"Sample Axial Slice\")\n\n    plt.tight_layout()\n    plt.show()\n'''\n# Example usage:\nverify_mapping(sagittal_slice, sag_y_coord_to_axial_mapping)\n'''","metadata":{"id":"yNU2vKJGCd2g","outputId":"f79ccbf3-bfac-44f2-aa5c-181ff8948552","execution":{"iopub.status.busy":"2024-09-26T12:05:42.146922Z","iopub.execute_input":"2024-09-26T12:05:42.147350Z","iopub.status.idle":"2024-09-26T12:05:42.160125Z","shell.execute_reply.started":"2024-09-26T12:05:42.147315Z","shell.execute_reply":"2024-09-26T12:05:42.158885Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\nmodel = Unet('resnet34', classes=6, in_channels=1)  # Assuming you have a pre-trained model\nmodel.load_state_dict(torch.load('/kaggle/input/lumbar-vertebrae-segmentation-using-spider-datset/simple_unet.pth', map_location='cpu'))\nmodel.eval()\n\nimage_tensor, dicom_image, middle_slice_path = load_middle_slice(study_id, sagittal_series_paths)\npredicted_mask = predict_segmentation(model, image_tensor)\ncentroids = calculate_centroids(predicted_mask)\npoly_func = fit_polynomial_curve(centroids)\ndisc_centroids = predict_disc_centroids(centroids, poly_func)\n'''","metadata":{"id":"0hKRrERcCi8p","execution":{"iopub.status.busy":"2024-09-26T11:58:35.981175Z","iopub.execute_input":"2024-09-26T11:58:35.982146Z","iopub.status.idle":"2024-09-26T11:58:35.996135Z","shell.execute_reply.started":"2024-09-26T11:58:35.982101Z","shell.execute_reply":"2024-09-26T11:58:35.994843Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def extract_axial_slices_for_discs(sag_y_coord_to_axial_mapping, disc_centroids, original_sagittal_shape, num_slices=3):\n    disc_level_slices = {}\n\n    # Calculate scaling factors based on the original sagittal image shape\n    original_width, original_height = original_sagittal_shape\n    y_scale = original_height / 256  # Since the centroids were calculated on a 256x256 image\n\n    for disc_label, (y_centroid, _) in disc_centroids.items():\n        # Rescale the Y-centroid to the original size\n        y_centroid_original = y_centroid * y_scale\n\n        # Find the closest Y-coordinate in the mapping\n        closest_y_coord = min(sag_y_coord_to_axial_mapping.keys(), key=lambda y: abs(y - y_centroid_original))\n\n        # Get the index of the closest Y-coordinate\n        closest_index = list(sag_y_coord_to_axial_mapping.keys()).index(closest_y_coord)\n\n        # Extract slices around the closest index (3 above and 3 below)\n        start_index = max(0, closest_index - num_slices)\n        end_index = min(len(sag_y_coord_to_axial_mapping), closest_index + num_slices)\n        print(disc_label,list(sag_y_coord_to_axial_mapping.keys())[closest_index])\n        print(list(sag_y_coord_to_axial_mapping.keys())[start_index:end_index])\n        selected_slices = [sag_y_coord_to_axial_mapping[y] for y in list(sag_y_coord_to_axial_mapping.keys())[start_index:end_index]]\n\n        # Store the selected slices for this disc level\n        disc_level_slices[disc_label] = np.stack(selected_slices, axis=0)  # Stack into a 3D array\n\n    return disc_level_slices\n'''\n# Example usage:\noriginal_sagittal_shape = dicom_image.shape  # Replace this with the actual original shape\ndisc_level_slices = extract_axial_slices_for_discs(sag_y_coord_to_axial_mapping, disc_centroids, original_sagittal_shape)\n\n# Visualize the axial slices for a specific disc level (e.g., L4/L5)\nimport matplotlib.pyplot as plt\n\nsample_disc_label = 'L1/L2'\nsample_slices = disc_level_slices[sample_disc_label]\n\n# Plot each slice in the 3D array\nfig, axes = plt.subplots(1, len(sample_slices), figsize=(20, 5))\nfor i, slice_image in enumerate(sample_slices):\n    axes[i].imshow(slice_image, cmap='gray')\n    axes[i].axis('off')\n    axes[i].set_title(f\"Slice {i + 1}\")\nplt.tight_layout()\nplt.show()\n'''","metadata":{"id":"UnKNY-ZpC5YU","outputId":"334ccebf-01ce-4c96-c857-b71047817a9d","execution":{"iopub.status.busy":"2024-09-26T12:05:44.502817Z","iopub.execute_input":"2024-09-26T12:05:44.503300Z","iopub.status.idle":"2024-09-26T12:05:44.522191Z","shell.execute_reply.started":"2024-09-26T12:05:44.503262Z","shell.execute_reply":"2024-09-26T12:05:44.521019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport numpy as np\nimport torch\nfrom tqdm import tqdm\n\n# Configuration\noutput_dataset_path = \"./axialT2\"\nos.makedirs(output_dataset_path, exist_ok=True)\n\n# List of skipped study IDs\n#skipped_studies = [305101752, 1082591956, 1249090513, 1406734395, 1451886888, 1614488620, 2155667219, 3295980651, 3647150026]\n\ndef create_axial_dataset(axial_series_paths, sagittal_series_paths, skipped_studies):\n    study_ids = [sid for sid in axial_series_paths if int(sid) not in skipped_studies]\n\n    for study_id in tqdm(study_ids, desc=\"Processing studies\"):\n        try:\n            # Create a directory for each study_id\n            os.makedirs(os.path.join(output_dataset_path, str(study_id)), exist_ok=True)\n\n            # Load the sagittal slices and axial slices\n            sagittal_slice = sag_middle(sagittal_series_paths[study_id][0])\n            axial_slices = load_and_sort_axial_slices(axial_series_paths, study_id, target_shape=(256, 256))\n\n            # Map sagittal Y-coordinates to axial slices\n            sag_y_coord_to_axial_mapping = map_sagittal_y_to_axial_slices(sagittal_slice, study_id, axial_slices)\n\n            # Get disc centroids from the sagittal middle slice\n            image_tensor, dicom_image, middle_slice_path = load_middle_slice(study_id, sagittal_series_paths)\n            predicted_mask = predict_segmentation(model, image_tensor)\n            centroids = calculate_centroids(predicted_mask)\n            poly_func = fit_polynomial_curve(centroids)\n            disc_centroids = predict_disc_centroids(centroids, poly_func)\n\n            # Extract and save the axial slices for each disc level\n            disc_level_slices = extract_axial_slices_for_discs(sag_y_coord_to_axial_mapping, disc_centroids,dicom_image.shape)\n            for disc_label, slices in disc_level_slices.items():\n                save_path = os.path.join(output_dataset_path, str(study_id), f\"{disc_label.replace('/', '_')}.npy\")\n                np.save(save_path, slices)  # Save the slices as numpy arrays\n\n        except Exception as e:\n            print(f\"Study {study_id} was skipped due to error: {e}\")\n            skipped_studies.append(study_id)\n            continue\n    return skipped_studies\n# Load the pre-trained model\nmodel = Unet('resnet34', classes=6, in_channels=1)  # Assuming you have a pre-trained model\nmodel.load_state_dict(torch.load('/kaggle/input/lumbar-vertebrae-segmentation-using-spider-datset/simple_unet.pth', map_location='cpu'))\nmodel.eval()\n\n# Generate the axial dataset\nskipped_studies=create_axial_dataset(axial_series_paths, sagittal_series_paths, skipped_studies)\n\nprint(\"Axial dataset generation completed successfully!\")\n","metadata":{"id":"yr8m6p8xC_MC","outputId":"4d7c4cca-a879-4c48-e2e6-c71236bc6dd5","execution":{"iopub.status.busy":"2024-09-26T12:05:47.248844Z","iopub.execute_input":"2024-09-26T12:05:47.249285Z","iopub.status.idle":"2024-09-26T12:05:48.738257Z","shell.execute_reply.started":"2024-09-26T12:05:47.249252Z","shell.execute_reply":"2024-09-26T12:05:48.737197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"RESNeXT","metadata":{"id":"fLxrrPkhDazN"}},{"cell_type":"code","source":"o1 = \"./sagittalT2\"\no2 = \"./axialT2\"\n\nfor study_id in skipped_studies:\n    # Correct the function name to os.path.join\n    p = os.path.join(o1, str(study_id))\n    \n    # Create the directory if it does not exist\n    os.makedirs(p, exist_ok=True)\n    \n    # Initialize arrays with zeros\n    l1_l2 = np.zeros((17, 84, 160))\n    l2_l3 = np.zeros((17, 84, 160))\n    l3_l4 = np.zeros((17, 84, 160))\n    l4_l5 = np.zeros((17, 84, 160))\n    l5_s1 = np.zeros((17, 84, 160))\n    \n    # Correct method to save numpy arrays\n    np.save(os.path.join(p, \"L1_L2.npy\"), l1_l2)\n    np.save(os.path.join(p, \"L2_L3.npy\"), l2_l3)\n    np.save(os.path.join(p, \"L3_L4.npy\"), l3_l4)\n    np.save(os.path.join(p, \"L4_L5.npy\"), l4_l5)\n    np.save(os.path.join(p, \"L5_S1.npy\"), l5_s1)\n    \nfor study_id in skipped_studies:\n    # Correct the function name to os.path.join\n    p = os.path.join(o2, str(study_id))\n    \n    # Create the directory if it does not exist\n    os.makedirs(p, exist_ok=True)\n    \n    # Initialize arrays with zeros\n    # Initialize arrays with zeros\n    l1_l2 = np.zeros((6, 256, 256))\n    l2_l3 = np.zeros((6, 256, 256))\n    l3_l4 = np.zeros((6, 256, 256))\n    l4_l5 = np.zeros((6, 256, 256))\n    l5_s1 = np.zeros((6, 256, 256))\n    \n    # Correct method to save numpy arrays\n    np.save(os.path.join(p, \"L1_L2.npy\"), l1_l2)\n    np.save(os.path.join(p, \"L2_L3.npy\"), l2_l3)\n    np.save(os.path.join(p, \"L3_L4.npy\"), l3_l4)\n    np.save(os.path.join(p, \"L4_L5.npy\"), l4_l5)\n    np.save(os.path.join(p, \"L5_S1.npy\"), l5_s1)","metadata":{"execution":{"iopub.status.busy":"2024-09-26T12:05:50.569706Z","iopub.execute_input":"2024-09-26T12:05:50.570105Z","iopub.status.idle":"2024-09-26T12:05:50.584473Z","shell.execute_reply.started":"2024-09-26T12:05:50.570074Z","shell.execute_reply":"2024-09-26T12:05:50.583204Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub=pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/sample_submission.csv')\nsub.iloc[:5,:]","metadata":{"id":"wCVpqCgjEcpo","outputId":"77b74b52-257a-48fe-8be3-0b49a708beae","execution":{"iopub.status.busy":"2024-09-26T12:05:52.540196Z","iopub.execute_input":"2024-09-26T12:05:52.540594Z","iopub.status.idle":"2024-09-26T12:05:52.557683Z","shell.execute_reply.started":"2024-09-26T12:05:52.540565Z","shell.execute_reply":"2024-09-26T12:05:52.556346Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport torch\nimport pandas as pd\nimport numpy as np\nfrom torch.utils.data import Dataset, DataLoader\nimport cv2\n\n# Define the test dataset\nclass TestSpineDataset(Dataset):\n    def __init__(self, sagittal_path, axial_path, study_ids, levels=[\"L1_L2\", \"L2_L3\", \"L3_L4\", \"L4_L5\", \"L5_S1\"]):\n        self.sagittal_path = sagittal_path\n        self.axial_path = axial_path\n        self.study_ids = study_ids\n        self.levels = levels\n        self.sagittal_slices_required = 15\n        self.axial_slices_required = 6\n\n\n    def load_slices(self, folder_path, level, required_slices, is_sagittal=False, target_shape=(160,84)):\n        file_path = os.path.join(folder_path, f\"{level}.npy\")\n        if not os.path.exists(file_path):\n            return None  # Return None if the file is missing\n        slices = np.load(file_path)\n        slices = [self.normalize_slice(s) for s in slices]\n        slices = self.pad_or_crop_slices(list(slices), required_slices, is_sagittal)\n\n        # Resize slices to the target shape\n        resized_slices = []\n        if is_sagittal:\n            for s in slices:\n                s_resized = cv2.resize(s, target_shape, interpolation=cv2.INTER_LINEAR)\n                resized_slices.append(s_resized)\n\n            return np.stack(resized_slices, axis=0)\n        return np.stack(slices, axis=0)\n\n    def normalize_slice(self, slice_array):\n        slice_array = slice_array.astype(np.float32)\n        if slice_array.min() < 0 or slice_array.max() > 1:\n            slice_array = (slice_array - slice_array.min()) / (slice_array.max() - slice_array.min())\n        slice_array -= slice_array.mean()\n        return slice_array\n\n    def pad_or_crop_slices(self, slices, required_slices, is_sagittal):\n        if len(slices) < required_slices:\n            padding = [np.zeros_like(slices[0])] * (required_slices - len(slices))\n            slices.extend(padding)\n        elif len(slices) > required_slices and is_sagittal:\n            start_idx = (len(slices) - required_slices) // 2\n            slices = slices[start_idx:start_idx + required_slices]\n        return slices\n\n    def apply_transform(self, slices):\n        # Apply the transformation to each slice individuallynp.transpose(slice_array, (1, 2, 0))\n        transformed_slices = [torch.tensor(slice_array) for slice_array in slices]\n        return torch.stack(transformed_slices, axis=0)\n\n    def __len__(self):\n        return len(self.study_ids)\n\n    def __getitem__(self, idx):\n        study_id = str(self.study_ids[idx])\n        sagittal_slices = []\n        axial_slices = []\n\n        for level in self.levels:\n            sagittal_folder = os.path.join(self.sagittal_path, study_id)\n            sagittal_slice = self.load_slices(sagittal_folder, level, self.sagittal_slices_required, is_sagittal=True)\n            if sagittal_slice is None:\n                # If any level is missing, skip this study ID\n                return None\n\n            axial_folder = os.path.join(self.axial_path, study_id)\n            axial_slice = self.load_slices(axial_folder, level, self.axial_slices_required)\n            if axial_slice is None:\n                # If any level is missing, skip this study ID\n                return None\n\n            sagittal_slices.append(sagittal_slice)\n            axial_slices.append(axial_slice)\n\n        sagittal_slices = self.apply_transform(sagittal_slices)\n        axial_slices = self.apply_transform(axial_slices)\n\n        return sagittal_slices, axial_slices, study_id\n    \ndef collate_fn(batch):\n    # Remove any `None` entries from the batch\n    batch = [item for item in batch if item is not None]\n    if len(batch) == 0:\n        return []  # Return None if the entire batch is empty\n    return torch.utils.data.dataloader.default_collate(batch)\n\n# Loading the test dataset\nsagittal_path = \"./sagittalT2\"\naxial_path = \"./axialT2\"\n#study_ids = os.listdir('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images') # Replace with actual test study IDs\nstudy_ids=list(axial_series_paths.keys())\n\ntest_dataset = TestSpineDataset(sagittal_path, axial_path, study_ids)\ntest_loader = DataLoader(test_dataset, batch_size=1, shuffle=False, collate_fn=collate_fn)","metadata":{"id":"n96l_GAvDPrX","execution":{"iopub.status.busy":"2024-09-26T12:05:55.543112Z","iopub.execute_input":"2024-09-26T12:05:55.543528Z","iopub.status.idle":"2024-09-26T12:05:55.573668Z","shell.execute_reply.started":"2024-09-26T12:05:55.543492Z","shell.execute_reply":"2024-09-26T12:05:55.572435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# Function to plot axial slices\ndef plot_axial_slices(axial_slices):\n    num_levels = axial_slices.shape[0]\n    fig, axes = plt.subplots(1, num_levels, figsize=(15, 5))\n    for i in range(num_levels):\n\n        axes[i].imshow(axial_slices[i, 0, :, :], cmap='gray')  # Show the first slice of each level\n        axes[i].set_title(f\"Level {i+1}\")\n        axes[i].axis('off')\n    #print(axial_slices[i, 0, :, :].shape)\n    plt.show()\n\n# Loop through the test loader and print shapes\nfor sagittal_slices, axial_slices, study_id in test_loader:\n    print(f\"Study ID: {study_id}\")\n    print(f\"Sagittal Slices Shape: {sagittal_slices.shape}\")  # Expected: [5, 15, 84, 160]\n    print(f\"Axial Slices Shape: {axial_slices.shape}\")  # Expected: [5, 6, 256, 256]\n\n    # Plot axial slices for a study ID\n    plot_axial_slices(sagittal_slices.squeeze(0))  # Remove the batch dimension for visualization\n\n    # Break the loop after the first example for testing\n    break\n","metadata":{"id":"GoAF5ZbwHVW3","outputId":"2bfd36be-c85e-42dd-ecee-53db73988f6b","execution":{"iopub.status.busy":"2024-09-26T12:05:58.645892Z","iopub.execute_input":"2024-09-26T12:05:58.646820Z","iopub.status.idle":"2024-09-26T12:05:59.047410Z","shell.execute_reply.started":"2024-09-26T12:05:58.646786Z","shell.execute_reply":"2024-09-26T12:05:59.046202Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append('../input/timm-pytorch-image-models/pytorch-image-models-master')\nimport timm","metadata":{"execution":{"iopub.status.busy":"2024-10-03T00:48:40.669602Z","iopub.execute_input":"2024-10-03T00:48:40.670437Z","iopub.status.idle":"2024-10-03T00:48:40.677237Z","shell.execute_reply.started":"2024-10-03T00:48:40.670400Z","shell.execute_reply":"2024-10-03T00:48:40.675874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport timm\n\nclass SpineClassificationModel2D(nn.Module):\n    def __init__(self, num_tasks=5, num_classes=3, dropout_rate=0.2):\n        super(SpineClassificationModel2D, self).__init__()\n\n        # 2D ResNeXt-50 for Axial Input\n        self.axial_branch = timm.create_model('resnext50_32x4d', pretrained=False, num_classes=0)  # Removing final FC layer\n        self.axial_branch.conv1 = nn.Conv2d(6, 64, kernel_size=7, stride=2, padding=3, bias=False)  # Adjusting for 6 channels\n\n        # 2D ResNeXt-50 for Sagittal Input (previously was planned as 3D)\n        self.sagittal_branch = timm.create_model('resnext50_32x4d', pretrained=False, num_classes=0)\n        self.sagittal_branch.conv1 = nn.Conv2d(15, 64, kernel_size=7, stride=2, padding=3, bias=False)  # Adjusting for 15 channels\n\n        # Global Average Pooling Layers\n        self.axial_global_pool = nn.AdaptiveAvgPool2d((1, 1))\n        self.sagittal_global_pool = nn.AdaptiveAvgPool2d((1, 1))\n\n        # Fully Connected Layers for Each Task\n        self.fcs = nn.ModuleList([nn.Sequential(\n            nn.Dropout(dropout_rate),\n            nn.Linear(2048 + 2048, num_classes),  # 2048 features from each branch\n        ) for _ in range(num_tasks)])\n\n    def forward(self, axial_input, sagittal_input):\n        # Axial branch forward pass\n        axial_features = self.axial_branch.forward_features(axial_input)\n        axial_features = self.axial_global_pool(axial_features)\n        axial_features = torch.flatten(axial_features, 1)\n\n        # Sagittal branch forward pass\n        sagittal_features = self.sagittal_branch.forward_features(sagittal_input)\n        sagittal_features = self.sagittal_global_pool(sagittal_features)\n        sagittal_features = torch.flatten(sagittal_features, 1)\n\n        # Concatenate the axial and sagittal features\n        combined_features = torch.cat((axial_features, sagittal_features), dim=1)\n\n        # Task-specific classification heads (logits)\n        outputs = [fc(combined_features) for fc in self.fcs]\n        return outputs  # Logits for each task\n","metadata":{"id":"ujaU_7hRQBob","execution":{"iopub.status.busy":"2024-09-26T12:06:04.275782Z","iopub.execute_input":"2024-09-26T12:06:04.276170Z","iopub.status.idle":"2024-09-26T12:06:04.291174Z","shell.execute_reply.started":"2024-09-26T12:06:04.276140Z","shell.execute_reply":"2024-09-26T12:06:04.289864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\n\n# Paths to your saved models\nmodel_paths = [\n    \"/kaggle/input/best_model/pytorch/default/1/best_model.pth\"\n]\n\n# Load models\nmodels = []\nfor path in model_paths:\n    model = SpineClassificationModel2D()  # Replace with your model class\n    optimizer = torch.optim.Adam(model.parameters())\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    checkpoint = torch.load(path, map_location=device)\n    model.load_state_dict(checkpoint['model_state_dict'])\n    optimizer.load_state_dict(checkpoint['optimizer_state_dict'])\n    epoch = checkpoint['epoch']\n    best_loss = checkpoint['loss']\n    model.to(device)\n    model.eval()\n    model_info = {\n        'model': model,\n        'optimizer': optimizer,\n        'epoch': epoch,\n        'best_loss': best_loss\n    }\n\n    models.append(model_info)\n    \nfor model_info in models:\n    model = model_info['model']\n    optimizer = model_info['optimizer']\n    epoch = model_info['epoch']\n    best_loss = model_info['best_loss']\n    print(f\"Model loaded from epoch {epoch} with best loss {best_loss}\")","metadata":{"id":"6zydRtoDP3RA","execution":{"iopub.status.busy":"2024-09-26T12:06:07.183097Z","iopub.execute_input":"2024-09-26T12:06:07.183901Z","iopub.status.idle":"2024-09-26T12:06:08.631273Z","shell.execute_reply.started":"2024-09-26T12:06:07.183864Z","shell.execute_reply":"2024-09-26T12:06:08.630188Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport pandas as pd\nimport numpy as np\n\n# Initialize list to store results\nresults = []\n\n# Device configuration\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Function to get predictions from ensemble models for each level\ndef get_ensemble_predictions(models, sagittal_slice, axial_slice):\n    ensemble_predictions = []\n\n    with torch.no_grad():\n        for model in models:\n            output = model(axial_slice, sagittal_slice)\n            output = [torch.softmax(o, dim=1) for o in output]\n            ensemble_predictions.append(output)\n\n    # Average the predictions across models\n    final_predictions = [torch.mean(torch.stack(pred), dim=0) for pred in zip(*ensemble_predictions)]\n\n    return final_predictions\n\n# Process each study ID in the test loader\nfor sagittal_slices, axial_slices, study_id in test_loader:\n    sagittal_slices, axial_slices = sagittal_slices.to(device), axial_slices.to(device)\n\n    task_names = [\n        \"spinal_canal_stenosis\",\n        \"left_neural_foraminal_narrowing\",\n        \"right_neural_foraminal_narrowing\",\n        \"left_subarticular_stenosis\",\n        \"right_subarticular_stenosis\"\n    ]\n    levels = [\"l1_l2\", \"l2_l3\", \"l3_l4\", \"l4_l5\", \"l5_s1\"]\n\n    for task_idx, task in enumerate(task_names):\n        for level_idx, level in enumerate(levels):\n            sagittal_slice = sagittal_slices[:, level_idx, :, :, :].squeeze(1)\n            axial_slice = axial_slices[:, level_idx, :, :, :].squeeze(1)\n\n            final_predictions = get_ensemble_predictions([model_info['model'] for model_info in models], sagittal_slice, axial_slice)\n\n            # Get probabilities and handle NaN values\n            if not np.isnan(final_predictions[task_idx].cpu().numpy()).any():\n                prob_normal_mild, prob_moderate, prob_severe = final_predictions[task_idx].cpu().numpy()[0]\n            else:\n                prob_normal_mild, prob_moderate, prob_severe = 0.34, 0.33, 0.33  # 或者其他处理方式\n\n            row_id = f\"{study_id[0]}_{task}_{level}\"\n            results.append([row_id, prob_normal_mild, prob_moderate, prob_severe])\n\n# Convert results to a DataFrame and save as CSV\ndf = pd.DataFrame(results, columns=[\"row_id\", \"normal_mild\", \"moderate\", \"severe\"])\ndf.to_csv(\"submission.csv\", index=False)\n\nprint(\"Submission file saved as submission.csv\")\n","metadata":{"id":"BhsgfphdYSNO","execution":{"iopub.status.busy":"2024-09-26T12:06:10.968573Z","iopub.execute_input":"2024-09-26T12:06:10.969681Z","iopub.status.idle":"2024-09-26T12:06:17.031497Z","shell.execute_reply.started":"2024-09-26T12:06:10.969634Z","shell.execute_reply":"2024-09-26T12:06:17.030374Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df","metadata":{"execution":{"iopub.status.busy":"2024-09-26T12:06:17.033238Z","iopub.execute_input":"2024-09-26T12:06:17.033573Z","iopub.status.idle":"2024-09-26T12:06:17.050304Z","shell.execute_reply.started":"2024-09-26T12:06:17.033545Z","shell.execute_reply":"2024-09-26T12:06:17.049161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Debug","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}