{"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":8650305,"sourceType":"datasetVersion","datasetId":5181481}],"dockerImageVersionId":30733,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"print(\"\\n... IMPORTS STARTING ...\\n\")\nprint(\"\\n\\tVERSION INFORMATION\")\n\n# Competition Specific Import\nimport pydicom  \nfrom pydicom import dcmread\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut, apply_color_lut, apply_modality_lut\n\nimport pandas as pd; pd.options.mode.chained_assignment = None; pd.set_option('display.max_columns', None);\nimport numpy as np; print(f\"\\t\\t– NUMPY VERSION: {np.__version__}\");\nimport sklearn; print(f\"\\t\\t– SKLEARN VERSION: {sklearn.__version__}\");\nimport cv2; print(f\"\\t\\t– CV2 VERSION: {cv2.__version__}\");\n\n# Built-In Imports (mostly don't worry about these)\nfrom typing import Iterable, Any, Callable, Generator\nfrom kaggle_datasets import KaggleDatasets\nfrom dataclasses import dataclass\nfrom collections import Counter\nfrom datetime import datetime\nfrom zipfile import ZipFile\nfrom glob import glob\nimport subprocess\nimport warnings\nimport requests\nimport textwrap\nimport hashlib\nimport imageio\nimport IPython\nimport urllib\nimport zipfile\nimport pickle\nimport random\nimport shutil\nimport string\nimport json\nimport copy\nimport math\nimport time\nimport gzip\nimport ast\nimport sys\nimport io\nimport gc\nimport re\nimport os\n\n# Visualization Imports (overkill)\nfrom IPython.core.display import HTML, Markdown\nimport matplotlib.pyplot as plt\nfrom matplotlib import animation, rc; rc('animation', html='jshtml')\nfrom tqdm.notebook import tqdm; tqdm.pandas();\nimport plotly.express as px\nimport seaborn as sns\nfrom PIL import Image, ImageEnhance; Image.MAX_IMAGE_PIXELS = 5_000_000_000;\nimport matplotlib; print(f\"\\t\\t– MATPLOTLIB VERSION: {matplotlib.__version__}\");\nfrom colorama import Fore, Style, init; init()\nimport plotly\nimport PIL\n\ndef hex_to_rgb(hex_color: str) -> tuple:\n    \"\"\"Convert hex color to RGB tuple.\n\n    Args:\n        hex_color (str): The hex color string, starting with '#'.\n\n    Returns:\n        tuple: A tuple of RGB values.\n    \"\"\"\n    hex_color = hex_color.lstrip('#')\n    return tuple(int(hex_color[i:i+2], 16) for i in (0, 2, 4))\n\ndef clr_print(text: str, color: str = \"#B9508A\", bold: bool = True) -> None:\n    \"\"\"Print the given text with the specified color and bold formatting.\n\n    Args:\n        text (str): The text to format.\n        color (str): The hex color code to apply. Defaults to \"#752F55\".\n        bold (bool): Whether to apply bold formatting. Defaults to True.\n    \"\"\"\n    _text = text.replace('\\n', '<br>')\n    rgb = hex_to_rgb(color)\n    color_style = f\"color: rgb({rgb[0]}, {rgb[1]}, {rgb[2]});\"\n    bold_style = \"font-weight: bold;\" if bold else \"\"\n    style = f\"{color_style} {bold_style}\"\n    display(HTML(f\"<span style='{style}'>{_text}</span>\"))\n\ndef seed_it_all(seed=7):\n    \"\"\" Attempt to be Reproducible \"\"\"\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    random.seed(seed)\n    np.random.seed(seed)\n    # tf.random.set_seed(seed)\n    \nseed_it_all()\n\n# Define the hex colors\nnb_hex_colors = [\"#231942\", \"#5E548E\", \"#9F86C0\", \"#BE95C4\", \"#E0B1CB\", \"#B9508A\", \"#752F55\"]\n\n# Create a Seaborn color palette\nnb_palette = sns.color_palette(nb_hex_colors)\n\n# Is this notebook being run on the backend for scoring re-submission\nIS_DEBUG = False if os.getenv('KAGGLE_IS_COMPETITION_RERUN') else True\nprint(f\"IS DEBUG: {IS_DEBUG}\")\n\n# Plot the palette\nclr_print(\"\\n... NOTEBOOK COLOUR PALETTE ...\")\nsns.palplot(nb_palette)\nplt.show()\n\nprint(\"\\n\\n... IMPORTS COMPLETE ...\\n\")\n\n_POSSIBLE_DICOM_ATTRS = [\n    \"BitsAllocated\", \"BitsStored\", \"Columns\", \"ContentDate\", \"ContentTime\", \"FrameOfReferenceUID\", \n    \"HighBit\", \"ImageOrientationPatient\", \"ImagePositionPatient\", \"InstanceNumber\", \n    \"PatientID\", \"PatientPosition\", \"PhotometricInterpretation\", \"PixelData\", \n    \"PixelRepresentation\", \"PixelSpacing\", \"RescaleIntercept\", \"RescaleSlope\", \"RescaleType\", \n    \"Rows\", \"SOPInstanceUID\", \"SamplesPerPixel\", \"SeriesDescription\", \"SeriesInstanceUID\", \n    \"SliceLocation\", \"SliceThickness\", \"SpacingBetweenSlices\", \"StudyInstanceUID\", \n    \"WindowCenter\", \"WindowWidth\"\n]\n\n\ndef flatten_l_o_l(nested_list):\n    \"\"\" Flatten a list of lists into a single list.\n\n    Args:\n        nested_list (Iterable): \n            – A list of lists (or iterables) to be flattened.\n\n    Returns:\n        A flattened list containing all items from the input list of lists.\n    \"\"\"\n    return [item for sublist in nested_list for item in sublist]\n\n\ndef print_ln(symbol=\"-\", line_len=110, newline_before=False, newline_after=False):\n    \"\"\" Print a horizontal line of a specified length and symbol.\n\n    Args:\n        symbol (str, optional): \n            – The symbol to use for the horizontal line\n        line_len (int, optional): \n            – The length of the horizontal line in characters\n        newline_before (bool, optional): \n            – Whether to print a newline character before the line\n        newline_after (bool, optional): \n            – Whether to print a newline character after the line\n            \n    Returns:\n        None; A divider with pre/post new-lines (optional) is printed\n    \"\"\"\n    if newline_before: print();\n    print(symbol * line_len)\n    if newline_after: print();\n        \n        \ndef display_hr(newline_before=False, newline_after=False):\n    \"\"\" Renders a HTML <hr>\n\n    Args:\n        newline_before (bool, optional): \n            – Whether to print a newline character before the line\n        newline_after (bool, optional): \n            – Whether to print a newline character after the line\n            \n    Returns:\n        None; A divider with pre/post new-lines (optional) is printed\n    \"\"\"\n    if newline_before: print();\n    display(HTML(\"<hr>\"))\n    if newline_after: print();\n\n\ndef wrap_text(text, width=88):\n    \"\"\"Wrap text to a specified width.\n\n    Args:\n        text (str): \n            - The text to wrap.\n        width (int): \n            - The maximum width of a line. Default is 88.\n\n    Returns:\n        str: The wrapped text.\n    \"\"\"\n    return textwrap.fill(text, width)\n\n\ndef wrap_text_by_paragraphs(text, width=88):\n    \"\"\"Wrap text by paragraphs to a specified width.\n\n    Args:\n        text (str): \n            - The text containing multiple paragraphs to wrap.\n        width (int): \n            - The maximum width of a line. Default is 88.\n\n    Returns:\n        str: The wrapped text with preserved paragraph separation.\n    \"\"\"\n    paragraphs = text.split('\\n')  # Assuming paragraphs are separated by newlines\n    wrapped_paragraphs = [textwrap.fill(paragraph, width) for paragraph in paragraphs]\n    return '\\n\\n'.join(wrapped_paragraphs)\n\ndef extract_all_dcm_data(dcm_path: str, save_to_dir: str = \"/kaggle/working/pngs/train\", save_to_png: bool = True) -> dict:\n    \"\"\"Extract all DICOM data and optionally save the image to PNG.\n    \n    Args:\n        dcm_path (str): Path to the DICOM file.\n        save_to_dir (str): Directory to save the PNG images.\n        save_to_png (bool): Whether or not to save the image as PNG.\n    \n    Returns:\n        Dicom attributes\n    \"\"\"\n    dicom = pydicom.read_file(dcm_path)    \n    dicom_attr_dict = {\n        attr_key:dicom.get(attr_key)\n        for attr_key in _POSSIBLE_DICOM_ATTRS\n        if attr_key!=\"PixelData\"\n    }\n    \n    if save_to_png:\n        # Get save path info and create if not existing\n        file_path_ending = \"/\".join(dcm_path.rsplit(\"/\", 3)[1:]).replace(\".dcm\", \".png\")\n        save_path = os.path.join(save_to_dir, file_path_ending)\n        os.makedirs(save_path.rsplit(\"/\", 1)[0], exist_ok=True)\n        \n        # Save the image\n        img = Image.fromarray(dicom_array_to_image(dicom, dicom.pixel_array))\n        img.save(save_path)\n        dicom_attr_dict[\"PNGPath\"] = save_path\n    \n    return dicom_attr_dict\n\n\ndef create_dicom_df(all_dcm_paths: list[str], is_train: bool = True, save_to_png: bool = True) -> pd.DataFrame:\n    \"\"\"Create a DataFrame with DICOM data and optionally save images as PNGs.\n    \n    Args:\n        all_dcm_paths (list[str]): List of paths to DICOM files.\n        is_train (bool): Whether the data is training data or not.\n        save_to_dir (str): Directory to save the PNG images.\n        save_to_png (bool): Whether or not to save the image as PNG.\n    \n    Returns:\n        pd.DataFrame: DataFrame containing DICOM attributes and paths.\n    \"\"\"\n    # Set PNG directory based on training or test data\n    _png_dir = \"/kaggle/working/pngs/train\" if is_train else \"/kaggle/working/pngs/test\"\n    \n    # Create the initial DataFrame with DICOM paths\n    dicom_df = pd.DataFrame({\"dcm_path\": all_dcm_paths})\n    \n    # Extract study_id, series_id, and instance_number from DICOM paths\n    dicom_df[[\"study_id\", \"series_id\", \"instance_number\"]] = pd.DataFrame(\n        dicom_df.dcm_path.apply(lambda x: [x.replace(\".dcm\", \"\") for x in x.rsplit(\"/\", 3)[1:]]).tolist()\n    ).astype(\"int\")\n    \n    # Extract DICOM attributes and optionally save images as PNGs\n    _new_df_cols = [x for x in _POSSIBLE_DICOM_ATTRS if x != \"PixelData\"]\n    if save_to_png:\n        _new_df_cols += [\"PNGPath\",]\n        \n    # Create the new dataframe using parallel processing...\n    # I did this on Paperspace with 32 CPUs (vs the 2 on Kaggle)\n    dicom_df[_new_df_cols] = pd.DataFrame(dicom_df[\"dcm_path\"].progress_apply(lambda x: extract_all_dcm_data(x, save_to_dir=_png_dir)).tolist())\n    \n    return dicom_df\n\ndef dicom_path_to_image(path, try_lut=True, fix_monochrome=True):\n    \"\"\" Convert dicom file to numpy array \n    \n    Args:\n        path (str): \n            Path to the dicom file to be converted\n        try_lut (bool): \n            Whether or not VOI LUT is available.\n            VOI LUT (if available by DICOM device) is used to transform raw DICOM data to \"human-friendly\" view\n        fix_monochrome (bool): \n            Whether or not to apply monochrome fix\n        \n    Returns:\n        Numpy array of the respective dicom file \n        \n    \"\"\"\n    # (1)  Use the pydicom library to read the dicom file and get image array\n    dicom = pydicom.read_file(path)\n    arr = dicom.pixel_array\n    \n    # (2)  Some DICOM datasets store their output image pixel values in a lookup table (LUT)\n    #        - The values in Pixel Data are the index to a corresponding LUT entry. \n    #        - When a dataset’s (0028,0004) Photometric Interpretation value is PALETTE COLOR then we should ...\n    #          use the apply_color_lut() function to apply a palette color LUT to the pixel data to produce an RGB image.\n    if try_lut and dicom.PhotometricInterpretation==\"PALETTE COLOR\":\n        clr_print(\"\\n\\n... Applying COLOR LUT ...\\n\\n\")\n        arr = apply_color_lut(arr, ds)\n    \n    \n    # (3) The DICOM Modality LUT module (similar to the Color one) converts raw pixel data values to a specific (possibly unitless) physical quantity.\n    #     Examples are quantities such as Hounsfield units for CT scan . \n    #     The apply_modality_lut() function can be used with an input array of raw values and a dataset containing a Modality LUT module to return the converted values. \n    #     When a dicom dataset requires multiple grayscale transformations, the Modality LUT transformation is always applied first.\n    if try_lut:\n        hu = apply_modality_lut(arr, dicom)\n        if (arr!=hu).any(): \n            clr_print(\"\\n\\n... Applying MODALITY LUT ...\\n\\n\")\n    \n    # (4) The DICOM VOI LUT module applies a VOI or windowing operation to input values. \n    # The apply_voi_lut() function can be used with an input array and a dataset containing a VOI LUT module to return values with applied VOI LUT or windowing. \n    # When a dicom dataset contains multiple VOI or windowing views then a particular view can be returned by using the index keyword parameter. \n    # In this case the index 0 will be used.\n    # When a dataset requires multiple grayscale transformations, then it’s assumed that the modality LUT or rescale operation has already been applied.\n    if try_lut:\n        arr = apply_voi_lut(hu, dicom, index=0)\n        if (arr!=hu).any(): clr_print(\"\\n\\n... Applying VOI LUT ...\\n\\n\")\n        \n    # The XRAY may look inverted\n    #   - If we want to fix this we can\n    if fix_monochrome and dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        clr_print(\"\\n\\n... Applying MONOCHROME FIX ...\\n\\n\")\n        arr = np.amax(arr) - arr\n    \n    # Normalize the image array and return\n    lower, upper = np.percentile(x, (1, 99))\n    arr = np.clip(arr, lower, upper)\n    arr = arr - np.min(arr)\n    arr = arr / np.max(arr)\n    arr = (arr * 255).astype(np.uint8)\n    return arr\n\n\ndef dicom_array_to_image(dicom, arr, try_lut=True, fix_monochrome=True):\n    \"\"\"Convert dicom array to numpy image array.\n    \n    Args:\n        dicom (pydicom.dataset.FileDataset): \n            DICOM object containing metadata\n        dicom_array (numpy.ndarray): \n            Numpy array containing DICOM pixel data\n        try_lut (bool): \n            Whether or not VOI LUT is available.\n            VOI LUT (if available by DICOM device) is used to transform raw DICOM data to \"human-friendly\" view\n        fix_monochrome (bool): \n            Whether or not to apply monochrome fix\n        \n    Returns:\n        Numpy array of the respective dicom file \n    \"\"\"\n    if try_lut and dicom.PhotometricInterpretation == \"PALETTE COLOR\":\n        arr = apply_color_lut(arr, dicom)\n    \n    if try_lut:\n        arr = apply_modality_lut(arr, dicom)\n    \n    if try_lut:\n        arr = apply_voi_lut(arr, dicom, index=0)\n        \n    if fix_monochrome and dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        arr = np.amax(arr) - arr\n    \n    lower, upper = np.percentile(arr, (1, 99))\n    arr = np.clip(arr, lower, upper)\n    arr = arr - np.min(arr)\n    arr = arr / np.max(arr)\n    arr = (arr * 255).astype(np.uint8)\n    return arr\n\ndef get_study_labels(df: pd.DataFrame, study_id: int | str, one_hot_encode_labels: bool = False, sparse_encode_labels: bool = False,) -> dict[str, str | int | list[float]]:\n    \"\"\"Get the labels for a given study from a pandas dataframe (train).\n\n    Args:\n        df (pd.DataFrame): \n            DataFrame containing study data.\n        study_id (int | str): \n            ID of the study to retrieve labels for.\n        one_hot_encode_labels (bool, optional): \n            Whether to one-hot encode string labels.\n        sparse_encode_labels (bool, optional): \n            Whether to encode string labels as integers.\n    \n    Raises:\n        ValueError: \n            If study_id is not found within the provided dataframe.\n    \n    Returns:\n        dict[str, str | int | list[float]]: \n            Dictionary with columns as keys and labels (either strings, integers, or one-hot encoded lists) as values.\n    \"\"\"\n    # Extract the row corresponding to the given study_id\n    study_row = df[df['study_id'] == study_id]\n    \n    # Ensure the study_id exists in the dataframe\n    if study_row.empty:\n        raise ValueError(f\"Study ID {study_id} not found in the dataframe.\")\n\n    # Convert the single-row DataFrame to a dictionary\n    labels_dict = study_row.iloc[0].to_dict()\n\n    # Remove 'study_id' from the dictionary\n    labels_dict.pop('study_id')\n\n    # Encode labels if requested (priority given to one_hot)\n    if one_hot_encode_labels:\n        labels_dict = {key: [1.0 if i == str2int_severity[value] else 0.0 for i in range(len(LABEL_STRS))] for key, value in labels_dict.items()}\n    elif sparse_encode_labels:\n        labels_dict = {key: str2int_severity[value] for key, value in labels_dict.items()}\n    return labels_dict\n\n# ROOT PATHS\nWORKING_DIR = \"/kaggle/working\"\nINPUT_DIR = \"/kaggle/input\"\nCOMPETITION_DIR = os.path.join(INPUT_DIR, \"rsna-2024-lumbar-spine-degenerative-classification\")\nTRAIN_DCM_IMAGE_DIR = os.path.join(COMPETITION_DIR, \"train_images\")\nTEST_DCM_IMAGE_DIR = os.path.join(COMPETITION_DIR, \"test_images\")\nADDITIONAL_INPUT_DIR = \"/kaggle/input/rsna-lsdc-files/rsna_lsdc_files\"\nTRAIN_DICOM_ATTR_CSV = os.path.join(ADDITIONAL_INPUT_DIR, \"train_dicom.csv\")\nTEST_DICOM_ATTR_CSV = os.path.join(ADDITIONAL_INPUT_DIR, \"train_dicom.csv\")\nTRAIN_PNG_IMAGE_DIR = os.path.join(ADDITIONAL_INPUT_DIR, \"pngs\", \"train\")\nTRAIN_PNG_IMAGE_DIR = os.path.join(ADDITIONAL_INPUT_DIR, \"pngs\", \"test\")\n\n\n# COMPETITION FILE PATHS\nSS_CSV_PATH = os.path.join(COMPETITION_DIR, \"sample_submission.csv\")\nTRAIN_CSV_PATH = os.path.join(COMPETITION_DIR, \"train.csv\")\nTRAIN_SERIES_DESC_CSV_PATH  = os.path.join(COMPETITION_DIR, \"train_series_descriptions.csv\")\nTRAIN_LABEL_COORDINATES_CSV_PATH  = os.path.join(COMPETITION_DIR, \"train_label_coordinates.csv\")\nTEST_SERIES_DESC_CSV_PATH  = os.path.join(COMPETITION_DIR, \"test_series_descriptions.csv\")\n\n# DEFINE COMPETITION DATAFRAMES\nclr_print(\"\\n\\n... SAMPLE SUBMISSION DATAFRAME ...\\n\\n\")\nss_df = pd.read_csv(SS_CSV_PATH)\ndisplay(ss_df)\n\nclr_print(\"\\n\\n... TRAIN DATAFRAME ...\\n\\n\")\ntrain_df = pd.read_csv(TRAIN_CSV_PATH)\ndisplay(train_df)\n\nclr_print(\"\\n\\n... TRAIN SERIES DESCRIPTIONS ...\\n\\n\")\ntrain_series_desc_df = pd.read_csv(TRAIN_SERIES_DESC_CSV_PATH)\ndisplay(train_series_desc_df)\n\nclr_print(\"\\n\\n... TEST SERIES DESCRIPTIONS ...\\n\\n\")\ntest_series_desc_df = pd.read_csv(TEST_SERIES_DESC_CSV_PATH)\ndisplay(test_series_desc_df)\n\nclr_print(\"\\n\\n... TRAIN LABEL COORDINATES ...\\n\\n\")\ntrain_label_coords_df = pd.read_csv(TRAIN_LABEL_COORDINATES_CSV_PATH)\ndisplay(train_label_coords_df)\n\n# DEFINE CONSTANTS AND USEFUL GLOBALS\nSERIES_TYPES = [\"Axial T2\", \"Sagittal T1\", \"Sagittal T2/STIR\"]\nLABEL_STRS = ['Normal/Mild', 'Moderate', 'Severe']\nstr2int_severity = {lbl:i for i, lbl in enumerate(LABEL_STRS)}\nint2str_severity = {v:k for k,v in str2int_severity.items()}\nLABEL_INTS = [str2int_severity[lbl] for lbl in LABEL_STRS]\n\nDICOM_NAMING_MAP = {\n    \"PNGPath\": \"png_path\", \"Columns\":\"img_width\", \"Rows\":\"img_height\", \n    \"SliceLocation\": \"slice_location\", \"SliceThickness\": \"slice_thickness\", \n    \"PixelSpacing\": \"pixel_spacing\", \"SpacingBetweenSlices\": \"slice_spacing\",\n    \"WindowCenter\": \"window_center\", \"WindowWidth\": \"window_width\", \n}\n\n# Redundant, already handled in loading/png-creation, or not enough variability to retain\nDICOM_COLS_TO_DROP = [\n    \"StudyInstanceUID\", \"SeriesInstanceUID\", \"InstanceNumber\",  # Redundant\n    \"BitsAllocated\", \"BitsStored\", \"HighBit\", \"PixelRepresentation\",  # These determine how pixel values are stored and their range. The current loading/saving function should handle this already (I THINK!) --> https://pydicom.github.io/pydicom/stable/old/image_data_handlers.html\n    \"ContentDate\", \"ContentTime\",  # These fields are generally not relevant unless the model specifically requires temporal context.\n    \"FrameOfReferenceUID\", \"SOPInstanceUID\",  # Typically used for ensuring data integrity and traceability in clinical settings but not directly for image processing (also often 1-to-1 with series_id/instance_id)\n    \"PhotometricInterpretation\",  # Handled by our loading function,\n    \"PatientID\",  # Always 1-to-1 with study-id -- so it is superfluous\n    \"RescaleIntercept\",  # Only 1 value - Always 0.0 when present.\n    \"PhotometricInterpretation\",  # Always MONOCHROME2\n]\nif not os.path.isfile(TRAIN_DICOM_ATTR_CSV):\n    # CREATE A DICOM DATAFRAME... AND OPTIONALLY SAVE THE PIXEL ARRAYS AS PNGS\n    ALL_TRAIN_DCM_PATHS = glob(os.path.join(TRAIN_DCM_IMAGE_DIR, \"**\", \"*.dcm\"), recursive=True)\n    train_dicom_df = create_dicom_df(ALL_TRAIN_DCM_PATHS, save_to_png=False)\nelse:\n    train_dicom_df = pd.read_csv(TRAIN_DICOM_ATTR_CSV)\n    train_dicom_df[\"dcm_path\"] = train_dicom_df[\"dcm_path\"].str.replace(WORKING_DIR, ADDITIONAL_INPUT_DIR)\n    ALL_TRAIN_DCM_PATHS = train_dicom_df[\"dcm_path\"].tolist()\n    train_dicom_df[\"PNGPath\"] = train_dicom_df[\"PNGPath\"].str.replace(WORKING_DIR, ADDITIONAL_INPUT_DIR)\ntrain_dicom_df = pd.merge(train_series_desc_df, train_dicom_df, on=[\"study_id\", \"series_id\"])\ntrain_dicom_df = train_dicom_df.rename(columns=DICOM_NAMING_MAP).drop(columns=DICOM_COLS_TO_DROP)\ntrain_dicom_df = train_dicom_df[[x for x in train_dicom_df.columns if \"_\" in x]+[x for x in train_dicom_df.columns if \"_\" not in x]]\ntrain_dicom_df = pd.merge(train_dicom_df, train_label_coords_df, how=\"left\").sort_values(by=[\"study_id\", \"series_description\", \"series_id\", \"instance_number\"]).reset_index(drop=True)\n\nALL_TEST_DCM_PATHS = glob(os.path.join(TEST_DCM_IMAGE_DIR, \"**\", \"*.dcm\"), recursive=True)\ntest_dicom_df = create_dicom_df(ALL_TEST_DCM_PATHS, save_to_png=True)\ntest_dicom_df = pd.merge(test_series_desc_df, test_dicom_df, on=[\"study_id\", \"series_id\"])\ntest_dicom_df = test_dicom_df.rename(columns=DICOM_NAMING_MAP).drop(columns=DICOM_COLS_TO_DROP)\ntest_dicom_df = test_dicom_df[[x for x in test_dicom_df.columns if \"_\" in x]+[x for x in test_dicom_df.columns if \"_\" not in x]]\n\nclr_print(\"\\n\\n... TRAIN DICOM DF ...\\n\\n\")\ndisplay(train_dicom_df)\n\nclr_print(\"\\n\\n... TEST DICOM DF ...\\n\\n\")\ndisplay(test_dicom_df)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-06-23T21:27:31.355355Z","iopub.execute_input":"2024-06-23T21:27:31.355723Z","iopub.status.idle":"2024-06-23T21:27:41.490360Z","shell.execute_reply.started":"2024-06-23T21:27:31.355693Z","shell.execute_reply":"2024-06-23T21:27:41.489419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def determine_thresholds(image: np.ndarray) -> tuple[int, int]:\n    \"\"\"Determine upper and lower thresholds for image segmentation using a modified Otsu's method.\n\n    This function implements a two-stage Otsu's thresholding to determine optimal\n    lower and upper thresholds for segmenting vertebrae in MRI images.\n\n    Args:\n        image (np.ndarray): Input grayscale image.\n\n    Returns:\n        tuple[int, int]: Lower and upper threshold values.\n    \"\"\"\n    # Ensure the image is in uint8 format\n    if image.dtype != np.uint8:\n        image = (image / image.max() * 255).astype(np.uint8)\n\n    # Compute histogram\n    hist = cv2.calcHist([image], [0], None, [256], [0, 256])\n    hist_norm = hist.ravel() / hist.sum()\n    \n    # Compute cumulative sums\n    Q = hist_norm.cumsum()\n    \n    # Compute cumulative mean\n    bins = np.arange(256)\n    fn_min = np.inf\n    thresh = -1\n    \n    for i in range(1, 255):  # Avoid edge cases at 0 and 255\n        # Probabilities for the two classes\n        p1, p2 = np.hsplit(hist_norm, [i])\n        q1, q2 = Q[i], Q[255] - Q[i]\n        \n        if q1 < 1e-6 or q2 < 1e-6:  # Avoid divide by zero\n            continue\n        \n        b1, b2 = np.hsplit(bins, [i])\n\n        # Compute means and variances\n        m1, m2 = np.sum(p1 * b1) / q1, np.sum(p2 * b2) / q2\n        v1, v2 = np.sum(((b1 - m1) ** 2) * p1) / q1, np.sum(((b2 - m2) ** 2) * p2) / q2\n\n        # Compute the minimization function\n        fn = v1 * q1 + v2 * q2\n        if fn < fn_min:\n            fn_min = fn\n            thresh = i\n\n    # Ensure we have a valid threshold\n    if thresh == -1:\n        thresh = np.mean(image)\n\n    # Compute lower threshold\n    lower_thresh = max(0, int(0.825 * thresh))\n    \n    # Compute upper threshold using Otsu's method on the upper half of the histogram\n    upper_half = image[image > thresh]\n    if upper_half.size > 0:\n        upper_thresh = cv2.threshold(upper_half, 0, 255, cv2.THRESH_BINARY + cv2.THRESH_OTSU)[0]\n    else:\n        upper_thresh = np.max(image)\n    \n    # Ensure upper threshold is higher than lower threshold\n    upper_thresh = min(255, max(upper_thresh, lower_thresh + 1)*1.125)\n    \n    return int(lower_thresh), int(upper_thresh)\n\n\ndef calculate_rectangle_similarity(\n    contour: np.ndarray,\n    penalize_tall_rectangles: bool = True,\n    vertical_penalty_ratio: float = 1.25\n) -> float:\n    \"\"\"\n    Calculate the similarity between a contour and its minimum area rectangle.\n    \n    Args:\n        contour (np.ndarray): The contour to analyze.\n        penalize_tall_rectangles (bool): If True, penalize rectangles taller than they are wide.\n        vertical_penalty_ratio (float): The height-to-width ratio threshold above which to apply the penalty.\n    \n    Returns:\n        float: A value between 0 and 1, where 1 indicates a perfect rectangle.\n               The value may be halved if penalize_tall_rectangles is True and \n               the rectangle's height exceeds width * vertical_penalty_ratio.\n    \"\"\"\n    rot_rect = cv2.minAreaRect(contour)\n    rect = cv2.boundingRect(contour)\n    box = cv2.boxPoints(rot_rect)\n    rect_area = cv2.contourArea(box)\n    contour_area = cv2.contourArea(contour)\n\n    if rect_area > 0:\n        similarity = contour_area / rect_area\n        \n        if penalize_tall_rectangles:\n            # Get width and height from the rectangle\n            cx, cy, width, height = rect\n            \n            # If height is greater than width * vertical_penalty_ratio, halve the similarity\n            if height > (width * vertical_penalty_ratio):\n                similarity *= 0.5\n        \n        return similarity\n    \n    return 0\n\ndef preprocess_image(\n    image: np.ndarray,\n    threshold_low: int,\n    threshold_high: int,\n    kernel_size: tuple[int, int],\n    open_iterations: int,\n    close_iterations: int,\n    open_close_repeats: int,\n    gaussian_blur_kernel: tuple[int, int] = (3,3),\n    gaussian_blur_sigma: float = 0\n) -> np.ndarray:\n    \"\"\"\n    Preprocess the input image for vertebrae detection.\n\n    Args:\n        image (np.ndarray): Input grayscale image.\n        threshold_low (int): Lower threshold for binary segmentation.\n        threshold_high (int): Upper threshold for binary segmentation.\n        kernel_size (Tuple[int, int]): Size of the kernel for morphological operations.\n        open_iterations (int): Number of iterations for opening operation.\n        close_iterations (int): Number of iterations for closing operation.\n        open_close_repeats (int): Number of times to repeat the open-close cycle.\n        gaussian_blur_kernel (Optional[Tuple[int, int]]): Kernel size for Gaussian blur. If None, no blur is applied.\n        gaussian_blur_sigma (float): Sigma for Gaussian blur. If 0, it's computed automatically.\n\n    Returns:\n        np.ndarray: Preprocessed binary image.\n    \"\"\"\n    # Apply Gaussian blur if kernel size is provided\n    if gaussian_blur_kernel is not None:\n        image = cv2.GaussianBlur(image, gaussian_blur_kernel, gaussian_blur_sigma)\n    \n    # Apply threshold\n    _, binary = cv2.threshold(image, threshold_low, threshold_high, cv2.THRESH_BINARY)\n    \n    # Create a kernel for morphological operations\n    kernel = np.ones(kernel_size, np.uint8)\n    \n    for _ in range(open_close_repeats):\n        # Apply opening operation to remove small noise\n        binary = cv2.erode(binary, kernel, iterations=open_iterations)\n        \n        # Apply closing operation to fill small gaps\n        binary = cv2.dilate(binary, kernel, iterations=close_iterations)\n    \n    return binary\n\ndef detect_vertebrae(\n    binary_image: np.ndarray,\n    min_area: int = 1000, \n    max_area: int = 10000, \n    min_similarity: float = 0.6,\n) -> list[np.ndarray]:\n    \"\"\"\n    Detect vertebrae in the preprocessed binary image.\n\n    Args:\n        binary_image (np.ndarray): Preprocessed binary image.\n        min_area (int): Minimum contour area to consider.\n        max_area (int): Maximum contour area to consider.\n        min_similarity (float): Minimum rectangle similarity to consider.\n\n    Returns:\n        list[np.ndarray]: List of detected vertebrae contours, sorted by similarity and position.\n    \"\"\"\n    # Find contours\n    contours, _ = cv2.findContours(binary_image.astype(np.uint8), cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    \n    # Filter contours based on area and rectangle similarity\n    filtered_contours = [\n        c for c in contours \n        if ((min_area < cv2.contourArea(c) < max_area) and (calculate_rectangle_similarity(c) > min_similarity))\n    ]\n    \n    # Sort contours by rectangle similarity (highest to lowest) and then by y-coordinate (top to bottom)\n    sorted_contours = sorted(\n        filtered_contours, \n        key=lambda c: (-calculate_rectangle_similarity(c), cv2.boundingRect(c)[1])\n    )\n    \n    return sorted_contours, [calculate_rectangle_similarity(c) for c in sorted_contours]\n\ndef visualize_contours(image: np.ndarray, contours: list, n_contours: int = 10) -> np.ndarray:\n    \"\"\"\n    Visualize detected vertebrae contours on the original image.\n\n    Args:\n        image (np.ndarray): Original image.\n        contours (list): List of detected vertebrae contours.\n        n_contours (int): Number of top contours to visualize.\n\n    Returns:\n        np.ndarray: Image with visualized contours.\n    \"\"\"\n    # Convert to BGR for colored drawing if it's not already\n    if len(image.shape) == 2:\n        vis_image = cv2.cvtColor(image, cv2.COLOR_GRAY2BGR)\n    else:\n        vis_image = image.copy()\n    \n    colors = [(0, 255, 0), (255, 0, 0), (0, 0, 255), (255, 255, 0), (0, 255, 255)]\n    \n    for i, contour in enumerate(contours[:n_contours]):\n        color = colors[i % len(colors)]\n        \n        # Draw the contour\n        cv2.drawContours(vis_image, [contour], 0, color, 2)\n        \n        # Calculate rectangle similarity\n        similarity = calculate_rectangle_similarity(contour)\n        \n        # Get bounding rectangle for text placement\n        x, y, w, h = cv2.boundingRect(contour)\n        \n        # Draw similarity score\n        cv2.putText(vis_image, f'{similarity:.2f}', (x, y-10), \n                    cv2.FONT_HERSHEY_SIMPLEX, 0.9, color, 2)\n        \n        # Draw rotated rectangle\n        rect = cv2.minAreaRect(contour)\n        box = cv2.boxPoints(rect)\n        box = np.intp(box)\n        cv2.drawContours(vis_image, [box], 0, color, 2)\n    \n    return vis_image\n\ndef process_and_visualize(\n    image_path: str,\n    threshold_low: int | None = None, \n    threshold_high: int | None = None, \n    kernel_size: tuple[int, int] = (2,2), \n    open_iterations: int = 3, \n    close_iterations: int = 1,\n    open_close_repeats: int = 1,\n    min_area: int = 500, \n    max_area: int = 10000, \n    min_similarity: float = 0.6,\n    n_contours: int = 10,\n    brute_force: bool = True,\n) -> None:\n    \"\"\"Process the input image and visualize the detected vertebrae.\n\n    Args:\n        image_path (str): Path to the input image file.\n        threshold_low (int): Lower threshold for binary segmentation.\n        threshold_high (int): Upper threshold for binary segmentation.\n        kernel_size (tuple[int, int]): Size of the kernel for morphological operations.\n        open_iterations (int): Number of iterations for opening operation.\n        close_iterations (int): Number of iterations for closing operation.\n        open_close_repeats (int): Number of times to repeat the open-close cycle.\n        min_area (int): Minimum contour area to consider.\n        max_area (int): Maximum contour area to consider.\n        min_similarity (float): Minimum rectangle similarity to consider.\n        n_contours (int): Number of top contours to visualize.\n    \"\"\"\n    # Read the image\n    img = cv2.imread(image_path, cv2.IMREAD_GRAYSCALE)\n\n    if not brute_force and (threshold_low is None or threshold_high is None):\n        _thresh_l, _thresh_h = determine_thresholds(img)\n        threshold_low = threshold_low or _thresh_l\n        threshold_high = threshold_high or _thresh_h\n        print(\"THRESH: \", threshold_low, threshold_high)\n    \n    if brute_force:\n        max_sim = 0\n        best_threshold_low=1\n        threshold_low=1\n        threshold_high=200\n        while threshold_low<125:       \n            # Preprocess the image\n            binary = preprocess_image(\n                img, threshold_low, threshold_high, kernel_size, \n                open_iterations, close_iterations, open_close_repeats\n            )\n\n            # Detect vertebrae\n            vertebrae_contours, list_of_sims = detect_vertebrae(binary, min_area, max_area, min_similarity)\n            if sum(list_of_sims[:5])>max_sim:\n                best_threshold_low = threshold_low\n                max_sim = sum(list_of_sims[:5])\n            threshold_low+=1\n        threshold_low=best_threshold_low\n        \n        \n    # Preprocess the image\n    binary = preprocess_image(\n        img, threshold_low, threshold_high, kernel_size, \n        open_iterations, close_iterations, open_close_repeats\n    )\n\n    # Detect vertebrae\n    vertebrae_contours, _ = detect_vertebrae(binary, min_area, max_area, min_similarity)\n    \n    # Visualize contours\n    vis_image = visualize_contours(img, vertebrae_contours, n_contours=n_contours)\n    \n    # Display results\n    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(20, 10))\n    ax1.imshow(binary, cmap='bone')\n    ax1.set_title('Preprocessed Binary Image')\n    ax1.axis('off')\n    ax2.imshow(cv2.cvtColor(vis_image, cv2.COLOR_BGR2RGB))\n    ax2.set_title('Detected Vertebrae Contours')\n    ax2.axis('off')\n    plt.tight_layout()\n    plt.show()\n\nDEMO_ROW = train_dicom_df[train_dicom_df.series_description.str.contains(\"Sagittal\")].iloc[50]\nDEMO_IMG = DEMO_ROW.png_path\ndisplay(DEMO_ROW.to_frame().T)\nprocess_and_visualize(DEMO_IMG, brute_force=True)","metadata":{"execution":{"iopub.status.busy":"2024-06-23T23:18:15.963152Z","iopub.execute_input":"2024-06-23T23:18:15.963495Z","iopub.status.idle":"2024-06-23T23:18:17.080719Z","shell.execute_reply.started":"2024-06-23T23:18:15.963466Z","shell.execute_reply":"2024-06-23T23:18:17.079808Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndetermine_thresholds(cv2.imread(DEMO_IMG, cv2.IMREAD_GRAYSCALE))\n# process_and_visualize(DEMO_IMG)","metadata":{"execution":{"iopub.status.busy":"2024-06-23T22:25:11.601794Z","iopub.execute_input":"2024-06-23T22:25:11.602847Z","iopub.status.idle":"2024-06-23T22:25:11.652080Z","shell.execute_reply.started":"2024-06-23T22:25:11.602801Z","shell.execute_reply":"2024-06-23T22:25:11.651011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i, (series_id, series_df) in enumerate(train_dicom_df[train_dicom_df.series_description == \"Sagittal T1\"].groupby(\"series_id\")):\n    if series_id in [10996, 3619813]:\n        continue\n    if i>15:\n        break\n    print(series_id)\n    DEMO_IMG = series_df[series_df.instance_number==int(series_df.instance_number.median())][\"png_path\"].values[0]\n    process_and_visualize(DEMO_IMG)","metadata":{"execution":{"iopub.status.busy":"2024-06-23T23:18:44.336148Z","iopub.execute_input":"2024-06-23T23:18:44.336518Z","iopub.status.idle":"2024-06-23T23:18:58.834207Z","shell.execute_reply.started":"2024-06-23T23:18:44.336487Z","shell.execute_reply":"2024-06-23T23:18:58.833237Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}