{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"from typing import Tuple, Union\nimport datetime\nimport os\nimport sys\nimport numpy as np\nimport pandas as pd\nfrom pandas.core.groupby.generic import DataFrameGroupBy\nimport pydicom\nimport SimpleITK as sitk\nfrom SimpleITK import Image\nfrom sklearn.model_selection import train_test_split\nimport tensorflow as tf","metadata":{"execution":{"iopub.status.busy":"2023-09-15T09:39:21.425518Z","iopub.execute_input":"2023-09-15T09:39:21.426292Z","iopub.status.idle":"2023-09-15T09:39:21.431926Z","shell.execute_reply.started":"2023-09-15T09:39:21.426255Z","shell.execute_reply":"2023-09-15T09:39:21.430801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Helpers and Preprocessing >>>\n\ndef set_parameters(\n        max_slices: int, root: str, seed: int, target_size: int, target_spacing: float, window: int,\n        working: str) -> dict:\n    \"\"\"\n    Set uniform parameters for the functions in the scope of the project.\n    :param max_slices: Maximum number of slices in CT, for the purpose of padding/cropping the image during inference,\n    :param root: Path to the project directory, where 'train_images' and 'test_images' directories are located,\n    :param seed: Seed for setting the random state,\n    :param target_size: Target image size along x- and y-axis, pixels,\n    :param target_spacing: Target voxel size in channel-first format, mm, tuple of floats,\n    :param window: Window size, images (slices),\n    :param working: Path to the working directory.\n    :return: Dictionary of uniform parameters.\n    \"\"\"\n\n    # Get dictionary of uniform parameters\n    inputs = {\n        'image_labels': os.path.join(root, 'image_level_labels.csv'),\n        'images_train': os.path.join(root, 'train_images'),\n        'images_test': os.path.join(root, 'test_images'),\n        'max_slices': max_slices,\n        'max_windows': int(max_slices/window),\n        'patient_labels': os.path.join(root, 'train.csv'),\n        'root': root,\n        'seed': seed,\n        'segmentations': os.path.join(root, 'segmentations'),\n        'series_meta_train': os.path.join(root, 'train_series_meta.csv'),\n        'series_meta_test': os.path.join(root, 'test_series_meta.csv'),\n        'target_size': target_size,\n        'target_spacing': target_spacing,\n        'window': window,\n        'working': working}\n\n    return inputs\n\n\ndef get_patient(\n        patient_id: str, path: str, series_meta: str, **kwargs) -> pd.DataFrame:\n    \"\"\"\n    Get series metadata for given 'patient_id' and 'series_id'.\n    :param patient_id: Patient ID,\n    :param path: Path to images directory,\n    :param series_meta: Path to series metadata.\n    :return: Series metadata.\n    \"\"\"\n\n    # Iterate over series for given patient_id\n    meta = []\n    series = os.listdir(os.path.join(path, patient_id))\n    for series_id in series:\n\n        # Get list of images for the given 'patient_id' and 'series_id'\n        images = os.listdir(os.path.join(path, patient_id, series_id))\n\n        # Make metadata Dataframe\n        num_instances = len(images)\n        data = {\n            'patient_id': [int(patient_id)] * num_instances,\n            'series_id': [int(series_id)] * num_instances,\n            'instance_number': [int(x.split('.')[0]) for x in images],\n            'path': [os.path.join(path, patient_id, series_id, x) for x in images]}\n        df = pd.DataFrame(data=data).sort_values(by='instance_number').reset_index(drop=True)\n\n        # Get series-specific DICOM attributes\n        image = pydicom.read_file(df.iloc[0]['path'])\n        df['kvp'] = image.KVP\n        df['bits_allocated'] = image.BitsAllocated\n        df['bits_stored'] = image.BitsStored\n        df['pixel_representation'] = image.PixelRepresentation\n        df['intercept'] = image.RescaleIntercept\n        df['slope'] = image.RescaleSlope\n        df['patient_position'] = image.PatientPosition\n        df['orientation'] = df.apply(lambda x: tuple([float(x) for x in image.ImageOrientationPatient]), axis=1)\n        df['spacing'] = df.apply(lambda x: tuple([float(x) for x in image.PixelSpacing]), axis=1)\n        df['size'] = df.apply(lambda x: (image.Rows, image.Columns, num_instances), axis=1)\n\n        # Get origin\n        position_first = tuple([float(x) for x in image.ImagePositionPatient])\n        position_last = tuple([float(x) for x in pydicom.read_file(df.iloc[-1]['path']).ImagePositionPatient])\n        if position_first[-1] < position_first[-1]:\n            df['origin'] = df.apply(lambda x: position_first, axis=1)\n        else:\n            df['origin'] = df.apply(lambda x: position_last, axis=1)\n            df.sort_values(by='instance_number', ascending=False, inplace=True)\n            df.reset_index(drop=True, inplace=True)\n            df['instance_index'] = list(range(num_instances))\n\n        # Get spacing\n        if num_instances == 1:\n            spacing_z_axis = image.SliceThickness\n        else:\n            spacing_z_axis = abs(position_first[-1] - position_last[-1]) / (num_instances - 1)\n        df['spacing'] = df['spacing'].apply(lambda x: (x[0], x[1], spacing_z_axis))\n\n        # Append series metadata to 'meta'\n        meta.append(df)\n\n    # Concatenate series and add 'num_series' attribute\n    num_series = len(meta)\n    meta = meta[0] if len(meta) == 1 else pd.concat(meta, axis=0)\n    meta['num_series'] = num_series\n\n    # Add 'aortic_hu' attribute\n    cols = ['series_id', 'aortic_hu']\n    meta = meta.merge(pd.read_csv(series_meta, usecols=cols), on='series_id', how='left')\n\n    return meta\n\n\ndef get_metadata(\n        inference: bool, image_labels: str, images_train: str, images_test: str, patient_labels: str,\n        segmentations: str, series_meta_train: str, series_meta_test: str, **kwargs) -> pd.DataFrame:\n    \"\"\"\n    Marge and process image-, series-, and patient-level metadata and labels.\n    :param inference: Get metadata for train set (if False) or Test set (if True),\n    :param image_labels: Uniform parameter,\n    :param images_train: Uniform parameter,\n    :param images_test: Uniform parameter,\n    :param patient_labels: Uniform parameter,\n    :param segmentations: Uniform parameter,\n    :param series_meta_train: Uniform parameter,\n    :param series_meta_test: Uniform parameter,\n    :return: pd.DataFrame.\n    \"\"\"\n\n    # Get image metadata\n    meta = []\n    path = images_train if not inference else images_test\n    series_meta = series_meta_train if not inference else series_meta_test\n    patients = os.listdir(path)\n    for patient_id in patients:\n        df = get_patient(patient_id=patient_id, path=path, series_meta=series_meta, **parameters)\n        meta.append(df)\n    meta = pd.concat(meta, axis=0)\n\n    # Add extra data if not inference mode\n    if not inference:\n\n        # Add path to NIfTI segmentation files\n        subset = [int(x.split('.')[0]) for x in os.listdir(segmentations)]\n        meta['segmentation'] = meta['series_id'].apply(\n            lambda x: os.path.join(segmentations, '{}.nii'.format(x)) if x in subset else None)\n\n        # Get patient-level labels\n        patient_labels = pd.read_csv(patient_labels, index_col='patient_id')\n        meta = meta.merge(right=patient_labels, left_on='patient_id', right_index=True, how='left')\n\n        # Add image-level labels >>>\n\n        # Helper function\n        def get_image_labels(x):\n            \"\"\"\n            Convert one-how encoded injury labels to ordinal: 1 - bowel, 2 - extravasation, 3 - both.\n            :param x: Iterator.\n            :return: Ordinal encoding.\n            \"\"\"\n            if x['bowel'] == 1 and x['extravasation'] == 1:\n                return 3\n            else:\n                if x['bowel'] == 1:\n                    return 1\n                else:\n                    return 2\n\n        # Read image labels, convert injury labels to one-hot\n        image_labels = pd.read_csv(image_labels)\n        image_labels = pd.get_dummies(data=image_labels, columns=['injury_name'], dtype=int, prefix='', prefix_sep='')\n\n        cols = {'Active_Extravasation': 'extravasation', 'Bowel': 'bowel'}\n        image_labels.rename(columns=cols, inplace=True)\n\n        # Sum image duplicates (multiple injuries)\n        image_labels = image_labels.groupby(by=['series_id', 'instance_number']).sum().reset_index().astype(int)\n\n        # Convert to ordinal\n        image_labels['image_label'] = image_labels.apply(lambda x: get_image_labels(x), axis=1)\n\n        # Merge with metadata and fill NaN with zero\n        image_labels = image_labels[['series_id', 'instance_number', 'image_label']]\n        cols = ['series_id', 'instance_number']\n        meta = meta.merge(right=image_labels, left_on=cols, right_on=cols, how='left')\n        meta['image_label'] = meta['image_label'].fillna(0)\n\n    return meta\n\n\ndef get_direction(\n        orientation: tuple) -> tuple:\n    \"\"\"\n    Convert 'ImageOrientationPatient' DICOM attribute to SimpleITK 'direction'.\n    :param orientation: ImageOrientationPatient, tuple of floats.\n    :return: SimpleITK 'direction'.\n    \"\"\"\n\n    row_cosines = np.array(orientation[:3])\n    column_cosines = np.array(orientation[3:])\n    normal_vector = np.cross(row_cosines, column_cosines)\n    direction = np.array([row_cosines, column_cosines, normal_vector]).T\n    direction = tuple(direction.flatten())\n\n    return direction\n\n\ndef dicom_to_sitk(\n        series: pd.DataFrame, bits_allocated: int, bits_stored: int, pixel_representation: int, slope: float,\n        intercept: int, spacing: tuple, origin: tuple, orientation: tuple, **kwargs) -> Image:\n    \"\"\"\n    Convert DICOM files to SimpleITK. Rescale image given slope and intercept.\n    Set spacing, origin and direction.\n    :param series: Filtered metadata for a given 'series_id',\n    :param bits_allocated: (0028, 0100) Bits Allocated DICOM attribute,\n    :param bits_stored: (0028, 0100) Bits Allocated DICOM attribute,\n    :param pixel_representation: (0028, 0101) Bits Stored DICOM attribute,\n    :param slope: (0028, 1053) Rescale Slope DICOM attribute,\n    :param intercept: (0028, 1052) Rescale Intercept DICOM attribute,\n    :param spacing: (0028, 0030) Pixel Spacing DICOM attribute amended with spacing along z-axis,\n    :param origin: analogue of SimpleITK 'origin',\n    :param orientation: (0020, 0037) Image Orientation (Patient) DICOM attribute,\n    :return: Rescaled SimpleITK image.\n    \"\"\"\n\n    # Stack DICOM files to 3D numpy array\n    image = series.sort_values(by='instance_index')['path'].tolist()\n    image = [pydicom.read_file(file).pixel_array for file in image]\n    image = np.stack(image)\n\n    # Bit shift for images where PixelRepresentation == 1 and BitsStored < BitsAllocated\n    if pixel_representation == 1 and bits_stored < bits_allocated:\n        bit_shift = int(bits_allocated - bits_stored)\n        image = (image << bit_shift) >> bit_shift\n\n    # Rescale image\n    image = image * slope + intercept\n\n    # Convert image to SimpleITK and set spacing, origin and direction\n    image = sitk.GetImageFromArray(image)\n    image.SetSpacing(spacing)\n    image.SetOrigin(origin)\n    image.SetDirection(get_direction(orientation))\n\n    return image\n\n\ndef resample_sitk(\n        image: Image, labels: bool, image_level: bool, reference: Union[Image, None], target_size: int,\n        target_spacing: float, window: int, **kwargs) -> Image:\n    \"\"\"\n    Convert image to SimpleITK and resample to target size and spacing.\n    :param image: SimpleITK image,\n    :param labels: Type of input: image (False) or organ labels (True),\n    :param image_level: Image- (extravasation and bowel injury) or pixel-level (organ level) labels,\n    :param reference: reference SimpleITK image or None,\n    :param target_size: Uniform parameter,\n    :param target_spacing: Uniform parameter.\n    :param window: Uniform parameter.\n    :return: SimpleITK image.\n    \"\"\"\n\n    # Get default pixel value\n    default_pixel_value = -1 if labels else np.min(sitk.GetArrayFromImage(image)).astype(float)\n\n    # Create instance of ResampleImageFilter\n    resampler = sitk.ResampleImageFilter()\n    resampler.SetReferenceImage(reference if reference is not None else image)\n    resampler.SetDefaultPixelValue(default_pixel_value)\n    resampler.SetInterpolator(sitk.sitkNearestNeighbor if labels else sitk.sitkLinear)\n\n    # Compute output size and spacing, if no reference image provided >>>\n    if reference is None:\n\n        # Get input size, spacing and dimensions\n        size = np.array(image.GetSize())\n        spacing = np.array(image.GetSpacing())\n        dims = size * spacing\n\n        # Set target_spacing\n        target_spacing = target_spacing if size[2] <= 256 else dims[2] / 256\n\n        # Compute output size and spacing\n        output_size_z = int(((dims[2] / target_spacing) // window + 1) * window)\n        size = (target_size, target_size, output_size_z) if not image_level else (1, 1, output_size_z)\n        output_spacing_xy = np.max(dims[:2]) / target_size\n        spacing = (output_spacing_xy, output_spacing_xy, target_spacing)\n\n        # Assign size and spacing to resampler\n        resampler.SetOutputSpacing(spacing)\n        resampler.SetSize(size)\n\n    # Set target size for image-level labels\n    elif image_level:\n        size = (1, 1, reference.GetSize()[2])\n        resampler.SetSize(size)\n\n    # Resample image\n    image = sitk.Cast(image, sitk.sitkInt8) if labels else image\n    image = resampler.Execute(image)\n\n    return image\n","metadata":{"execution":{"iopub.status.busy":"2023-09-15T09:39:21.635488Z","iopub.execute_input":"2023-09-15T09:39:21.636257Z","iopub.status.idle":"2023-09-15T09:39:21.689402Z","shell.execute_reply.started":"2023-09-15T09:39:21.63622Z","shell.execute_reply":"2023-09-15T09:39:21.688347Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Make and read datasets >>>\n\ndef validation_split(patient_labels: str, **kwargs) -> Tuple[list, list]:\n    \"\"\"\n    Make an injury-stratified split of patients.\n    :param patient_labels: Uniform parameter,\n    :return: train/validation split of patients.\n    \"\"\"\n\n    # Drop columns with healthy organ labels\n    df = pd.read_csv(patient_labels, index_col='patient_id')\n    cols = [x for x in df.columns.tolist() if 'healthy' in x.split('_')]\n    df = df.drop(columns=cols)\n\n    # Add 'healthy' column (inverse of 'any_injury')\n    df['healthy'] = 1 - df['any_injury']\n    df = df.drop(columns='any_injury')\n\n    # Convert one-hot to dense\n    df = df * df.columns\n    df['Group'] = df.apply(lambda x: tuple([x[i] for i in df.columns if x[i] != '']), axis=1)\n    groups = dict(zip(df['Group'].unique(), list(range(len(df['Group'].unique())))))\n    df['Group'] = df['Group'].apply(lambda x: groups[x])\n\n    # Merge categories with total count <= 3\n    counts = df.groupby(by='Group').size()\n    counts.name = 'Count'\n    df = df.merge(right=counts, left_on='Group', right_index=True, how='left')[['Group', 'Count']]\n    df['Group'] = df.apply(lambda x: x['Group'] if x['Count'] > 3 else -1, axis=1)\n\n    # Make stratified split of patients\n    train, validation = train_test_split(df, stratify=df['Group'], test_size=0.2)\n\n    return train.index.tolist(), validation.index.tolist()\n\n\ndef serialize_patient(patient_id: int, patients: DataFrameGroupBy, metadata: DataFrameGroupBy, **kwargs) -> tuple:\n    \"\"\"\n    Get equally-sampled images, organ labels, slice- and patient-level injury labels, padding indicator,\n    and instance attributes (number of slices, number of images per series, true segmentation indicator).\n    :param metadata: Metadata, grouped by SeriesID,\n    :param patient_id: PatientID,\n    :param patients: Metadata, grouped by SeriesID and PatientID ,\n    :return: Example Proto.\n    \"\"\"\n\n    # Helper functions >>>\n\n    def resample_image_level(series: DataFrameGroupBy, attributes: dict, reference: Union[Image, None]) -> Image:\n        \"\"\"\n        Get and resample image-level injury labels (bowel, extravasation).\n        :param series: DataFrameGroupBy object,\n        :param attributes: Attributes of a given series.\n        :param reference: Refrence Image.\n        :return: SimpleITK image.\n        \"\"\"\n        labels = series['image_label'].to_numpy()\n        labels = np.expand_dims(np.expand_dims(labels, axis=-1), axis=-1)\n        labels = sitk.GetImageFromArray(labels)\n        labels.SetSpacing(attributes['spacing'])\n        labels.SetOrigin(attributes['origin'])\n        labels.SetDirection(get_direction(attributes['orientation']))\n        labels = resample_sitk(labels, labels=True, image_level=True, reference=reference, **parameters)\n\n        return labels\n\n    def image_labels_to_tensor(labels: Image) -> tuple:\n        \"\"\"\n        Convert image labels to one-how slise-level bowel injury and extravasation labels, and padding labels.\n        :param labels: image-level labels, SimpleITK image.\n        :return: tf.Tensor\n        \"\"\"\n\n        labels = tf.constant(sitk.GetArrayFromImage(labels), dtype=tf.int8)\n        labels = tf.add(labels, tf.constant(1, dtype=tf.int8))\n        labels = tf.squeeze(tf.one_hot(labels, depth=5))  # padding, healthy, bowel, extravasation, both\n        bowel = tf.reduce_max(tf.stack([labels[:, 2], labels[:, 4]], axis=-1), axis=-1)\n        extravasation = tf.reduce_max(tf.stack([labels[:, 3], labels[:, 4]], axis=-1), axis=-1)\n        padding = tf.cast(labels[:, 0], dtype=tf.float16)\n        labels = tf.cast(tf.stack([bowel, extravasation], axis=-1), dtype=tf.float16)\n\n        return labels, padding\n\n    def organ_labels_to_tensor(labels: Image) -> tf.Tensor:\n        \"\"\"\n        Convert 3D organ labels from 3D dense [slices, height, width] to 1D one-hot [slices, classes].\n        :param labels: SimpleITK image.\n        :return: tf.Tensor, tf.float16\n        \"\"\"\n\n        labels = tf.constant(sitk.GetArrayFromImage(labels), dtype=tf.int8)\n        labels = tf.add(labels, tf.constant(1, dtype=tf.int8))\n        labels = tf.one_hot(labels, depth=7)\n        labels = tf.cast(tf.reduce_max(labels, axis=[1, 2]), dtype=tf.float16)[:, 2:]\n\n        return labels\n\n    # Get patient pd.DataFrame\n    patient = patients.get_group(patient_id)\n\n    # One series per PatientID:\n    if len(patient) == 1:\n\n        # Get series and attributes\n        attributes1 = patient.iloc[0].to_dict()\n        series1 = metadata.get_group(attributes1['series_id'])\n\n        # Resample image to target size and target spacing\n        image1 = dicom_to_sitk(series=series1, **attributes1, **parameters)\n        image1 = resample_sitk(image1, labels=False, image_level=False, reference=None, **parameters)\n        image2 = None\n\n        # Resample image-level injury labels\n        image_labels = resample_image_level(series=series1, attributes=attributes1, reference=None)\n        image_labels, padding = image_labels_to_tensor(image_labels)\n\n        # Resample organ labels\n        organ_labels = attributes1['segmentation']\n        if organ_labels is not None:\n            organ_labels = sitk.ReadImage(organ_labels)\n            organ_labels = resample_sitk(organ_labels, labels=True, image_level=False, reference=None, **parameters)\n            organ_labels = organ_labels_to_tensor(organ_labels)\n        else:\n            organ_labels = tf.constant(-1, dtype=tf.float16)\n\n    # Two series per PatientID:\n    else:\n\n        # Get series and attributes\n        attributes1, attributes2 = patient.iloc[0].to_dict(), patient.iloc[1].to_dict()\n        series1, series2 = metadata.get_group(attributes1['series_id']), metadata.get_group(attributes2['series_id'])\n\n        # Read images\n        image1 = dicom_to_sitk(series=series1, **attributes1, **parameters)\n        image2 = dicom_to_sitk(series=series2, **attributes2, **parameters)\n\n        # Get input size, spacing and dimensions\n        size1, size2 = np.array(image1.GetSize()), np.array(image2.GetSize())\n        spacing1, spacing2 = np.array(image1.GetSpacing()), np.array(image2.GetSpacing())\n        dims1, dims2 = size1 * spacing1, size2 * spacing2\n\n        # Resample images. Reference image is the image with larger dimensions along z-axis.\n        if dims1[2] > dims2[2]:\n            image1 = resample_sitk(image1, labels=False, image_level=False, reference=None, **parameters)\n            image2 = resample_sitk(image2, labels=False, image_level=False, reference=image1, **parameters)\n        else:\n            image2 = resample_sitk(image2, labels=False, image_level=False, reference=None, **parameters)\n            image1 = resample_sitk(image1, labels=False, image_level=False, reference=image2, **parameters)\n\n        # Resample and merge image-level injury labels >>>\n        image_labels1 = resample_image_level(series=series1, attributes=attributes1, reference=image1)\n        image_labels2 = resample_image_level(series=series2, attributes=attributes2, reference=image2)\n        image_labels1, padding1 = image_labels_to_tensor(image_labels1)\n        image_labels2, padding2 = image_labels_to_tensor(image_labels2)\n        image_labels = tf.stack([image_labels1, image_labels2], axis=-1)\n        image_labels = tf.reduce_max(image_labels, axis=-1)\n        padding = tf.stack([padding1, padding2], axis=-1)\n        padding = tf.reduce_min(padding, axis=-1)  # Images may overlap\n\n        # Resample and merge organ labels\n        organ_labels1, organ_labels2 = attributes1['segmentation'], attributes2['segmentation']\n        if organ_labels1 is not None and organ_labels2 is not None:\n            organ_labels1, organ_labels2 = sitk.ReadImage(organ_labels1), sitk.ReadImage(organ_labels2)\n            organ_labels1 = resample_sitk(organ_labels1, labels=True, image_level=False, reference=image1, **parameters)\n            organ_labels2 = resample_sitk(organ_labels2, labels=True, image_level=False, reference=image2, **parameters)\n            organ_labels1 = organ_labels_to_tensor(organ_labels1)\n            organ_labels2 = organ_labels_to_tensor(organ_labels2)\n            organ_labels = tf.stack([organ_labels1, organ_labels2], axis=-1)\n            organ_labels = tf.reduce_max(organ_labels, axis=-1)\n        else:\n            organ_labels = tf.constant(-1, dtype=tf.float16)\n\n    # Convert images to tf.Tensors and expand dimensions\n    images = [image1, image2]\n    for i in range(2):\n        image = images[i]\n        if image is not None:\n            image = tf.constant(sitk.GetArrayFromImage(image), dtype=tf.float16)\n            image = tf.expand_dims(image, axis=-1)\n            images[i] = image\n        else:\n            images[i] = tf.constant(-1, dtype=tf.float16)\n    image1, image2 = images\n\n    # Get patient labels\n    cols = ['bowel_healthy', 'bowel_injury',\n            'extravasation_healthy', 'extravasation_injury',\n            'kidney_healthy', 'kidney_low', 'kidney_high',\n            'liver_healthy', 'liver_low', 'liver_high',\n            'spleen_healthy', 'spleen_low', 'spleen_high']\n    patient_labels = tf.constant(patient.iloc[0][cols].to_numpy(), dtype=tf.float16)\n\n    # Create a Feature dictionary\n    segmentation = 1 if attributes1['segmentation'] is not None else 0\n    features = {'patient_id': tf.constant(patient_id, dtype=tf.int32),\n                'image1': image1,\n                'image2': image2,\n                'organ_labels': organ_labels,\n                'image_labels': image_labels,\n                'padding': padding,\n                'patient_labels': patient_labels,\n                'num_slices': tf.constant(image1.shape[0], dtype=tf.int16),\n                'num_series': tf.constant(attributes1['num_series'], dtype=tf.int16),\n                'segmentation': tf.constant(segmentation, dtype=tf.int16)}\n    for key in features.keys():\n        features[key] = tf.train.Feature(\n            bytes_list=tf.train.BytesList(value=[tf.io.serialize_tensor(features[key]).numpy()]))\n\n    # Create an Example protocol buffer, serialize to a string and write to the TFRecord file\n    example_proto = tf.train.Example(features=tf.train.Features(feature=features))\n\n    return example_proto\n\n\ndef make_dataset(\n        seed: int, working: str, **kwargs):\n    \"\"\"\n    Convert patients to TFRecord train and validation datasets.\n    :param pretrain: Include only patients with segmentations (if True), or predict missing organ labels,\n    :param seed: Uniform parameter,\n    :param working: Uniform parameter.\n    :return: None.\n    \"\"\"\n\n    # Set random seed\n    np.random.seed(seed)\n\n    # Get metadata\n    metadata = get_metadata(inference=False, **parameters)\n\n    # Split patients into train and validation set >>>\n    train, validation = validation_split(metadata=metadata, **parameters)\n    subsets = {'train': train, 'validation': validation}\n\n    # Group metadata by series_id and by patient_id\n    metadata = metadata.sort_values(by='instance_index').groupby(by='series_id')\n    patients = metadata.agg(lambda x: list(set(x))[0]).sort_values(by='aortic_hu').reset_index().groupby(\n        by='patient_id')\n\n    # Set starting values of time and counter\n    start = datetime.datetime.now()\n    counter = 0\n    count = len(train) + len(validation)\n\n    # Iterate over series in train and validation sets >>>\n    for key in subsets.keys():\n\n        # Create a TFRecord writer\n        path = os.path.join(working, 'abdominal_{}.tfrecord'.format(key))\n        writer = tf.io.TFRecordWriter(path)\n\n        # Iterate over patients and write to TFRecord dataset\n        subset = subsets[key]\n        np.random.shuffle(subset)\n        for patient_id in subset[:1]:\n            data = serialize_patient(patient_id=patient_id, patients=patients, metadata=metadata, **parameters)\n            writer.write(data.SerializeToString())\n\n            # Print progress report\n            counter += 1\n            progress = '[' + '.' * int(counter / 50) + ' ' * int((count - counter) / 50) + ']'\n            sys.stdout.write('\\rProcessed {}/{} patients: {} elapsed time: {}'.format(\n                counter, count, progress, datetime.datetime.now() - start))\n            sys.stdout.flush()\n\n        # Close TFRecord writer\n        writer.close()\n    print('\\r')\n\n    return None\n\n\ndef read_dataset(working: str, **kwargs) -> tuple:\n    \"\"\"\n    Read TFRecord train and validation datasets.\n    :param working: Uniform parameter.\n    :return: Train and Validation tf.data.Datasets.\n    \"\"\"\n\n    # Helper functions to parse and deserialize TFR dataset >>>\n\n    def parse_example(x):\n        \"\"\"\n        Parse TFR dataset.\n        :param x: Instance of dataset.\n        :return: Parsed instance.\n        \"\"\"\n\n        features = ['patient_id',\n                    'image1', 'image2',\n                    'image_labels',\n                    'organ_labels',\n                    'padding',\n                    'patient_labels',\n                    'num_slices',\n                    'num_series',\n                    'segmentation']\n        features = {key: tf.io.FixedLenFeature([], tf.string) for key in features}\n        example = tf.io.parse_single_example(serialized=x, features=features)\n        return example\n\n    def deserialize(x):\n        \"\"\"\n        Deserialize tensors.\n        :param x: Instance of dataset.\n        :return: Dictionary of deserialized tensors.\n        \"\"\"\n\n        deserialized = {}\n        for feature in ['image1', 'image2', 'image_labels', 'organ_labels', 'padding', 'patient_labels']:\n            deserialized[feature] = tf.io.parse_tensor(x[feature], out_type=tf.float16)\n        for feature in ['patient_id']:\n            deserialized[feature] = tf.io.parse_tensor(x[feature], out_type=tf.int32)\n        for feature in ['num_slices', 'num_series', 'segmentation']:\n            deserialized[feature] = tf.io.parse_tensor(x[feature], out_type=tf.int16)\n\n        return deserialized\n\n    # Read train and validation subsets\n    subsets = {'train': None, 'validation': None}\n    for key in subsets.keys():\n\n        dataset = tf.data.TFRecordDataset(os.path.join(working, 'abdominal_{}.tfrecord'.format(key)))\n        dataset = dataset.map(parse_example, num_parallel_calls=tf.data.AUTOTUNE)\n        dataset = dataset.map(deserialize, num_parallel_calls=tf.data.AUTOTUNE)\n        dataset = dataset.prefetch(buffer_size=tf.data.AUTOTUNE)\n        subsets[key] = dataset\n\n    return subsets['train'], subsets['validation']\n","metadata":{"execution":{"iopub.status.busy":"2023-09-15T09:39:21.691575Z","iopub.execute_input":"2023-09-15T09:39:21.692243Z","iopub.status.idle":"2023-09-15T09:39:21.758338Z","shell.execute_reply.started":"2023-09-15T09:39:21.692205Z","shell.execute_reply":"2023-09-15T09:39:21.757331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Set parameters\nparameters = set_parameters(\n    root='/kaggle/input/rsna-2023-abdominal-trauma-detection',\n    working='/kaggle/working/',\n    target_size=480, target_spacing=3, max_slices=256, seed=42, window=64)","metadata":{"execution":{"iopub.status.busy":"2023-09-15T09:39:21.759726Z","iopub.execute_input":"2023-09-15T09:39:21.760405Z","iopub.status.idle":"2023-09-15T09:39:21.777629Z","shell.execute_reply.started":"2023-09-15T09:39:21.760372Z","shell.execute_reply":"2023-09-15T09:39:21.776452Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Make train and validation TFRecord datasets\nmake_dataset(**parameters)","metadata":{"execution":{"iopub.status.busy":"2023-09-15T09:39:21.780197Z","iopub.execute_input":"2023-09-15T09:39:21.781428Z","iopub.status.idle":"2023-09-15T09:50:28.256226Z","shell.execute_reply.started":"2023-09-15T09:39:21.781389Z","shell.execute_reply":"2023-09-15T09:50:28.253538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Read train and validation TFRecord datasets\ntrain, validation = read_dataset(**parameters)","metadata":{"execution":{"iopub.status.busy":"2023-09-15T09:50:28.2605Z","iopub.execute_input":"2023-09-15T09:50:28.261091Z","iopub.status.idle":"2023-09-15T09:50:28.701983Z","shell.execute_reply.started":"2023-09-15T09:50:28.261041Z","shell.execute_reply":"2023-09-15T09:50:28.700551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Test the reader function\nfor i in train.take(1):\n    for key in i.keys():\n        print(key, i[key].shape)","metadata":{"execution":{"iopub.status.busy":"2023-09-15T09:53:48.987644Z","iopub.execute_input":"2023-09-15T09:53:48.988127Z","iopub.status.idle":"2023-09-15T09:53:49.455184Z","shell.execute_reply.started":"2023-09-15T09:53:48.988094Z","shell.execute_reply":"2023-09-15T09:53:49.453701Z"},"trusted":true},"execution_count":null,"outputs":[]}]}