{"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":"none","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":6124,"sourceType":"modelInstanceVersion","modelInstanceId":4599,"modelId":2797},{"sourceId":6125,"sourceType":"modelInstanceVersion","modelInstanceId":4596,"modelId":2797},{"sourceId":6127,"sourceType":"modelInstanceVersion","modelInstanceId":4598,"modelId":2797}],"dockerImageVersionId":30732,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Testing if MONAI can work to process these images\n\nDraft version -- not working yet\n\n\n# Setup","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install monai\n","metadata":{"execution":{"iopub.status.busy":"2024-09-15T23:56:16.212014Z","iopub.execute_input":"2024-09-15T23:56:16.212462Z","iopub.status.idle":"2024-09-15T23:56:33.681062Z","shell.execute_reply.started":"2024-09-15T23:56:16.212430Z","shell.execute_reply":"2024-09-15T23:56:33.679662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport time\n\n\nos.environ[\"KERAS_BACKEND\"] = \"jax\"  # @param [\"tensorflow\", \"jax\", \"torch\"]\n\n#from tensorflow import data as tf_data\n#import tensorflow_datasets as tfds\nimport keras\n#import keras_cv\n#from keras_cv import bounding_box\n#from keras_cv import visualization\nimport tqdm\nimport pandas as pd\nimport numpy as np\nimport os\n#import pydicom\nimport tensorflow as tf\n#import tensorflow_io as tfio\nimport monai\n\nimport warnings\nwarnings.filterwarnings('ignore')","metadata":{"execution":{"iopub.status.busy":"2024-09-15T23:56:33.683383Z","iopub.execute_input":"2024-09-15T23:56:33.683787Z","iopub.status.idle":"2024-09-15T23:57:43.237806Z","shell.execute_reply.started":"2024-09-15T23:56:33.683753Z","shell.execute_reply":"2024-09-15T23:57:43.236323Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BASE_DIR = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/'\nTRAIN_DIR = BASE_DIR+'train_images/'\nTEST_DIR = BASE_DIR+'test_images/'\n\nPRETRAINED = 'efficientnetv2_s_imagenet'\n\nSPLIT_RATIO = .2\nBATCH_SIZE = 32\nEPOCH = 8\n\n#IMG_SIZE = [320,320]\nIMG_SIZE = [240, 240]","metadata":{"execution":{"iopub.status.busy":"2024-09-15T23:57:43.239659Z","iopub.execute_input":"2024-09-15T23:57:43.241133Z","iopub.status.idle":"2024-09-15T23:57:43.247894Z","shell.execute_reply.started":"2024-09-15T23:57:43.241084Z","shell.execute_reply":"2024-09-15T23:57:43.246630Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data processing","metadata":{}},{"cell_type":"code","source":"studies = os.listdir(TRAIN_DIR)\n\n# Print the first 10 studies\nprint(\"List of the first 10 studies:\")\nfor study in studies[:10]:\n    print(study)\n\n# Splitting studies into train and validation sets\ntrain_studies = studies[:int(len(studies)*(1-SPLIT_RATIO))]\nval_studies = studies[int(len(studies)*(1-SPLIT_RATIO)):]\n\n# Print the number of studies in each set\nprint(\"Number of studies in train set:\", len(train_studies))\nprint(\"Number of studies in validation set:\", len(val_studies))","metadata":{"execution":{"iopub.status.busy":"2024-09-15T23:57:43.251204Z","iopub.execute_input":"2024-09-15T23:57:43.252062Z","iopub.status.idle":"2024-09-15T23:57:43.374747Z","shell.execute_reply.started":"2024-09-15T23:57:43.252020Z","shell.execute_reply":"2024-09-15T23:57:43.373510Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\n\n# Read the series descriptions CSV file\ndesc = pd.read_csv(BASE_DIR+'train_series_descriptions.csv')\n\n# Convert study_id and series_id columns to string data type\ndesc.study_id = desc.study_id.astype(str)\ndesc.series_id = desc.series_id.astype(str)\n\n# Read the labels CSV file\nlabels = pd.read_csv(BASE_DIR+'train.csv')\n\n# Convert study_id column to string data type\nlabels.study_id = labels.study_id.astype(str)\n\n# Get unique conditions from the labels dataframe\nconditions = np.unique(labels.columns[1:])\n\n# Create a list of class names for each condition and severity level\nclasses = []\nfor c in conditions:\n    classes.append(c+'_normal')\n    classes.append(c+'_moderate')\n    classes.append(c+'_severe')\n\n# Create a mapping from class names to indices\nclasses_map = {classes[i]:i for i in range(len(classes))}\n\n# Create a reverse mapping from indices to class names\nclass_mapping = {i:classes[i] for i in range(len(classes))}\n\n# Determine the total number of classes\nN_CLASSES = len(class_mapping)\n\n# Print the classes and the number of classes\nprint(\"Classes:\", classes)\nprint(\"Number of classes:\", N_CLASSES)","metadata":{"execution":{"iopub.status.busy":"2024-09-15T23:57:43.376208Z","iopub.execute_input":"2024-09-15T23:57:43.376568Z","iopub.status.idle":"2024-09-15T23:57:43.451006Z","shell.execute_reply.started":"2024-09-15T23:57:43.376538Z","shell.execute_reply":"2024-09-15T23:57:43.449591Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Creates a output file where there is a 1 for each condition in the list above\ndef prepare_labels(X):\n    out = np.zeros(N_CLASSES)\n    cols = X.index[1:]\n    X = X.values[1:]\n    \n    for x in cols[X=='Normal/Mild']: out[classes_map[x+'_normal']] = 1\n    for x in cols[X=='Moderate']: out[classes_map[x+'_moderate']] = 1\n    for x in cols[X=='Severe']: out[classes_map[x+'_severe']] = 1\n        \n    return out\n\n","metadata":{"execution":{"iopub.status.busy":"2024-09-15T23:57:43.452472Z","iopub.execute_input":"2024-09-15T23:57:43.452881Z","iopub.status.idle":"2024-09-15T23:57:43.460524Z","shell.execute_reply.started":"2024-09-15T23:57:43.452850Z","shell.execute_reply":"2024-09-15T23:57:43.459287Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_dicom_images(image_dir):\n    \"\"\"Loads DICOM images from a directory.\n    Args:  image_dir: Path to the directory containing DICOM images.\n    Returns:  A list of loaded image arrays.\n    \"\"\"\n    MAX_SLICES = 10                 # The maximum number of slices to load in for each DICOM folde\n    # Get a list of all DICOM files in the directory\n    dicom_files = [os.path.join(image_dir, f\"{i+1}.dcm\") for i in range(MAX_SLICES)]  # Adjust the maximum number of images as needed\n    dicom_files = [f for f in dicom_files if os.path.exists(f)]  # Filter existing files\n\n    # original code without resize\n    #ds = monai.data.Dataset(data=dicom_files, transform=monai.transforms.LoadImage(image_only=True))\n    #images = [np.squeeze(img) for img in ds]  # Convert to numpy arrays and squeeze\n    #return images\n\n    transforms = [monai.transforms.LoadImage(image_only=True)]\n    transforms.append(monai.transforms.Resize(spatial_size=IMG_SIZE))\n    ds = monai.data.Dataset(data=dicom_files, transform=monai.transforms.Compose(transforms))\n    images = [np.squeeze(img) for img in ds]  # Convert to numpy arrays and squeeze\n    return images\n\n# Example usage:\n\n#image_dir = TRAIN_DIR+\"/10728036/2399638375\"\n#images = load_dicom_images(image_dir)\n\n#print(images)","metadata":{"execution":{"iopub.status.busy":"2024-09-16T00:07:18.190276Z","iopub.execute_input":"2024-09-16T00:07:18.190981Z","iopub.status.idle":"2024-09-16T00:07:18.202866Z","shell.execute_reply.started":"2024-09-16T00:07:18.190881Z","shell.execute_reply":"2024-09-16T00:07:18.201046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Example usage:\n\nimage_dir = TRAIN_DIR+\"/10728036/2399638375\"\nimages = load_dicom_images(image_dir)\nprint(\"DICOM image loaded\")\n#print(images)","metadata":{"execution":{"iopub.status.busy":"2024-09-16T00:10:18.548884Z","iopub.execute_input":"2024-09-16T00:10:18.549413Z","iopub.status.idle":"2024-09-16T00:10:27.018033Z","shell.execute_reply.started":"2024-09-16T00:10:18.549376Z","shell.execute_reply":"2024-09-16T00:10:27.016709Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# New Code on Sept 15\n","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pydicom\nimport tensorflow as tf\nfrom tensorflow.keras import layers, models","metadata":{"execution":{"iopub.status.busy":"2024-09-16T00:12:02.843070Z","iopub.execute_input":"2024-09-16T00:12:02.843682Z","iopub.status.idle":"2024-09-16T00:12:02.892439Z","shell.execute_reply.started":"2024-09-16T00:12:02.843635Z","shell.execute_reply":"2024-09-16T00:12:02.891277Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_dicom_images(input_folder, num_slices=3):\n    \"\"\"\n    Load multiple DICOM images, stack them into a volume, and prepare them for CNN input.\n    \n    Parameters:\n    input_folder (str): Path to the folder containing DICOM files.\n    num_slices (int): Number of slices to sample from each DICOM series.\n\n    Returns:\n    np.array: Stacked DICOM images for CNN input.\n    \"\"\"\n    dicom_images = []\n    \n    for root, _, files in os.walk(input_folder):\n        dicom_files = [f for f in files if f.endswith('.dcm')]\n        slices = []\n        for file in dicom_files[:num_slices]:  # Sample `num_slices` from the series\n            dicom_path = os.path.join(root, file)\n            dicom_data = pydicom.dcmread(dicom_path)\n            slices.append(dicom_data.pixel_array)\n        \n        print(\"slices = \",slices.__len__)\n\n        if len(slices) > 0:\n            # Stack slices along the depth axis to create a 3D volume\n            dicom_volume = np.stack(slices, axis=-1)\n            dicom_images.append(dicom_volume)\n    \n    return np.array(dicom_images)\n\n\n","metadata":{"execution":{"iopub.status.busy":"2024-09-16T00:23:47.827313Z","iopub.execute_input":"2024-09-16T00:23:47.828641Z","iopub.status.idle":"2024-09-16T00:23:47.838403Z","shell.execute_reply.started":"2024-09-16T00:23:47.828588Z","shell.execute_reply":"2024-09-16T00:23:47.836923Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def build_multi_input_cnn(input_shape, num_inputs=2):\n    \"\"\"\n    Build a CNN with multiple input layers for DICOM image analysis.\n\n    Parameters:\n    input_shape (tuple): Shape of the input image.\n    num_inputs (int): Number of separate input layers.\n\n    Returns:\n    tensorflow.keras.Model: CNN model with multiple input layers.\n    \"\"\"\n    # Create a list of input layers\n    inputs = [layers.Input(shape=input_shape) for _ in range(num_inputs)]\n    \n    # Create a shared CNN block for each input\n    cnn_blocks = []\n    for inp in inputs:\n        x = layers.Conv2D(32, kernel_size=(3, 3), activation='relu')(inp)\n        x = layers.MaxPooling2D(pool_size=(2, 2))(x)\n        x = layers.Conv2D(64, kernel_size=(3, 3), activation='relu')(x)\n        x = layers.MaxPooling2D(pool_size=(2, 2))(x)\n        x = layers.Flatten()(x)\n        cnn_blocks.append(x)\n    \n    # Combine the CNN blocks from each input layer\n    combined = layers.Concatenate()(cnn_blocks)\n    x = layers.Dense(128, activation='relu')(combined)\n    x = layers.Dense(64, activation='relu')(x)\n    output = layers.Dense(1, activation='sigmoid')(x)\n    \n    # Create the model\n    model = models.Model(inputs=inputs, outputs=output)\n    model.compile(optimizer='adam', loss='binary_crossentropy', metrics=['accuracy'])\n    \n    return model\n\n","metadata":{"execution":{"iopub.status.busy":"2024-09-16T00:14:42.887109Z","iopub.execute_input":"2024-09-16T00:14:42.887652Z","iopub.status.idle":"2024-09-16T00:14:42.901274Z","shell.execute_reply.started":"2024-09-16T00:14:42.887611Z","shell.execute_reply":"2024-09-16T00:14:42.900019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"input_folder_healthy = TRAIN_DIR\nhealthy_images = load_dicom_images(input_folder_healthy, num_slices=3)","metadata":{"execution":{"iopub.status.busy":"2024-09-16T00:23:52.474633Z","iopub.execute_input":"2024-09-16T00:23:52.476205Z","iopub.status.idle":"2024-09-16T00:26:21.652142Z","shell.execute_reply.started":"2024-09-16T00:23:52.476161Z","shell.execute_reply":"2024-09-16T00:26:21.649702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# Load DICOM images (you can specify multiple folders for healthy and degenerative images)\ninput_folder_healthy = 'path_to_healthy_dicom'\ninput_folder_degenerative = 'path_to_degenerative_dicom'\n\nhealthy_images = load_dicom_images(input_folder_healthy, num_slices=3)\ndegenerative_images = load_dicom_images(input_folder_degenerative, num_slices=3)\n\n# Prepare input data (e.g., X_train, y_train for healthy and degenerative)\nX_train = np.concatenate([healthy_images, degenerative_images], axis=0)\ny_train = np.concatenate([np.zeros(len(healthy_images)), np.ones(len(degenerative_images))], axis=0)\n\n# Build a CNN with two input layers for demonstration (could be more)\ncnn_model = build_multi_input_cnn(input_shape=(512, 512, 3), num_inputs=2)\n\n# Train the model (example)\ncnn_model.fit([X_train, X_train], y_train, batch_size=8, epochs=10)\n\n# Save the model\ncnn_model.save('multi_input_cnn_model.h5')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import time\n\n\ndef preprocess_images(images):\n    # possibly add some preprocessing here\n    # TODO:  resize images here\n    return images\n            \n            \ndef process_patients(num_patients=50):\n    start_time = time.time()\n    # Get a list of patient IDs\n    patient_ids = os.listdir(TRAIN_DIR)\n\n    all_images = []\n    all_labels = []  # Assuming you have corresponding labels\n\n    for i, patient_id in enumerate(patient_ids[:num_patients]):\n        patient_dir = os.path.join(TRAIN_DIR, patient_id)\n        series_dirs = [os.path.join(patient_dir, series_id) \n                       for series_id in os.listdir(patient_dir) if os.path.isdir(os.path.join(patient_dir, series_id))]\n\n        for series_dir in series_dirs:\n            images = load_dicom_images(series_dir)\n            # Process images for the series (e.g., stack, preprocess, etc.)\n            processed_images = preprocess_images(images)  # Add preprocessing function\n            all_images.append(processed_images)\n            \n            # Append corresponding labels here\n            study_dir = TRAIN_DIR+patient_id+'/'\n            sub_sample = labels[labels.study_id==patient_id].fillna('Normal/Mild')\n            label = prepare_labels(sub_sample.iloc[0])\n            all_labels.append(label)  # Assuming you have a way to get labels\n        # Display progress since this will take a long time...\n        if (i+1) % 3 == 0:\n            elapsed_time = time.time() - start_time\n            print(f\"Processed {i+1}/{num_patients} patients in {elapsed_time:.2f} seconds\")\n            start_time = time.time()\n    return all_images, all_labels\n\n#### TODO process all the patients in the future\nprint (\"Only training of first 50 patients --- fix this in the future\")\nall_images, all_labels = process_patients(10)\n\n\n","metadata":{"execution":{"iopub.status.busy":"2024-09-15T23:44:52.893828Z","iopub.execute_input":"2024-09-15T23:44:52.894268Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# This code is from https://towardsdatascience.com/simple-3d-mri-classification-ranked-bronze-on-kaggle-87edfdef018a\n\ndef load_dicom(path_file: str) -> Optional[np.ndarray]:\n    dicom = pydicom.dcmread(path_file)\n    # TODO: adjust spacing in particular dimension according DICOM meta\n    try:\n        img = apply_voi_lut(dicom.pixel_array, dicom).astype(np.float32)\n    except RuntimeError as err:\n        print(err)\n        return None\n    return img\n\n\ndef load_volume(path_volume: str, percentile: Optional[float] = 0.01) -> Tensor:\n    path_slices = glob.glob(os.path.join(path_volume, '*.dcm'))\n    path_slices = sorted(path_slices, key=parse_name_index)\n    vol = []\n    for p_slice in path_slices:\n        img = load_dicom(p_slice)\n        if img is None:\n            continue\n        vol.append(img.T)\n    volume = torch.tensor(vol, dtype=torch.float32)\n    if percentile is not None:\n        # get extreme values\n        p_low = np.quantile(volume, percentile) if percentile else volume.min()\n        p_high = np.quantile(volume, 1 - percentile) if percentile else volume.max()\n        # normalize\n        volume = (volume - p_low) / (p_high - p_low)\n    return volume.T\n","metadata":{"execution":{"iopub.status.busy":"2024-09-16T00:07:25.124498Z","iopub.execute_input":"2024-09-16T00:07:25.125072Z","iopub.status.idle":"2024-09-16T00:07:25.221313Z","shell.execute_reply.started":"2024-09-16T00:07:25.125026Z","shell.execute_reply":"2024-09-16T00:07:25.219488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print (\"all_labels type = \", type(all_labels))\nprint(\"Length of all_labels:\", len(all_labels))\nprint(\"First element :\", all_labels[0])\n#print (\"all_labels = \", all_labels)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print (\"all_images type = \", type(all_images))\nprint(\"Length of all_images:\", len(all_images))\nprint(\"First element :\", all_images[0])\n#print (\"all_images = \", all_images)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"    #From https://github.com/Project-MONAI/tutorials/blob/main/3d_classification/torch/densenet_training_array.py\n        \n    # create a training data loader\n    train_ds = ImageDataset(image_files=images[:10], labels=labels[:10], transform=train_transforms)\n    train_loader = DataLoader(train_ds, batch_size=2, shuffle=True, num_workers=2, pin_memory=torch.cuda.is_available())\n\n    # create a validation data loader\n    val_ds = ImageDataset(image_files=images[-10:], labels=labels[-10:], transform=val_transforms)\n    val_loader = DataLoader(val_ds, batch_size=2, num_workers=2, pin_memory=torch.cuda.is_available())","metadata":{"execution":{"iopub.status.busy":"2024-09-15T23:44:47.769322Z","iopub.execute_input":"2024-09-15T23:44:47.770061Z","iopub.status.idle":"2024-09-15T23:44:47.836588Z","shell.execute_reply.started":"2024-09-15T23:44:47.770020Z","shell.execute_reply":"2024-09-15T23:44:47.835111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef process_patient(patient_id):\n      patient_dir = os.path.join(TRAIN_DIR, patient_id)\n      series_dirs = [os.path.join(patient_dir, series_id) for series_id in os.listdir(patient_dir) if os.path.isdir(os.path.join(patient_dir, series_id))]\n\n      for series_dir in series_dirs:\n            images = load_dicom_images(series_dir)\n            # Process images for the series (e.g., stack, preprocess, etc.)\n            # ... your image processing logic here\n\n\n# Get a list of patient IDs\npatient_ids = os.listdir(TRAIN_DIR)\n# Need for progress display below\ntotal_patients = len(patient_ids)\n\n\n# Only process the first 20 patients for now\n#### TODO process all the patients in the future\npatients_to_train = 50\nprint (\"Only training of first 50 patients --- fix this in the future\")\n\n# Process each patient\n#### TODO process all the patients in the future\nfor i, patient_id in enumerate(patient_ids[:patients_to_train]):\n#for i, patient_id in enumerate(patient_ids):\n    process_patient(patient_id)\n        \n    if (i+1) % 10 == 0:\n        elapsed_time = time.time() - start_time\n        print(f\"Processed {i+1}/{total_patients} patients in {elapsed_time:.2f} seconds\")\n        start_time = time.time()","metadata":{"execution":{"iopub.status.busy":"2024-08-03T20:49:32.481481Z","iopub.status.idle":"2024-08-03T20:49:32.482051Z","shell.execute_reply.started":"2024-08-03T20:49:32.481826Z","shell.execute_reply":"2024-08-03T20:49:32.481848Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import nibabel as nib\nimport pydicom\n\ndef dicom_to_nifti(dicom_folder, output_nifti_file):\n    # Read the DICOM files from the folder\n    dicom_files = [os.path.join(dicom_folder, f) for f in os.listdir(dicom_folder) if f.endswith('.dcm')]\n    \n    # Sort files by slice location (if the DICOM headers have SliceLocation)\n    #dicom_files.sort(key=lambda f: pydicom.dcmread(f).SliceLocation)\n\n    # Read the pixel arrays from the DICOM files\n    slices = [pydicom.dcmread(f) for f in dicom_files]\n    pixel_arrays = [s.pixel_array for s in slices]\n\n    # Convert the pixel arrays into a 3D numpy array\n    volume_3d = np.stack(pixel_arrays, axis=-1)\n    \n    # Get the affine transformation from DICOM (you may need to adjust this based on DICOM metadata)\n    # Here we're assuming identity affine, but more sophisticated transformations may be needed\n    affine = np.eye(4)\n    \n    # Create a NIfTI image\n    nifti_image = nib.Nifti1Image(volume_3d, affine)\n\n    # Save as a NIfTI file\n    nib.save(nifti_image, output_nifti_file)\n    print(f\"Conversion complete. NIfTI file saved to {output_nifti_file}\")\n\n\n","metadata":{"execution":{"iopub.status.busy":"2024-09-16T00:00:32.701322Z","iopub.execute_input":"2024-09-16T00:00:32.701893Z","iopub.status.idle":"2024-09-16T00:00:32.711983Z","shell.execute_reply.started":"2024-09-16T00:00:32.701853Z","shell.execute_reply":"2024-09-16T00:00:32.710656Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Example usage:\ndicom_folder = TRAIN_DIR+\"/10728036/2399638375\"\noutput_nifti_file = '2399638375.nii'\ndicom_to_nifti(dicom_folder, output_nifti_file)\n","metadata":{"execution":{"iopub.status.busy":"2024-09-16T00:00:38.237310Z","iopub.execute_input":"2024-09-16T00:00:38.238292Z","iopub.status.idle":"2024-09-16T00:00:38.807312Z","shell.execute_reply.started":"2024-09-16T00:00:38.238249Z","shell.execute_reply":"2024-09-16T00:00:38.805608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#### this version is not working\n# This version uses the SimpleITK libary\nimport SimpleITK as sitk\n\ndef prepare_data_SimpleITK(studies, df):\n    labels = []\n    image_data = []  # Changed to store processed data\n            \n    for study_id in studies:\n        study_dir = TRAIN_DIR + study_id + '/'\n        sub_sample = df[df.study_id == study_id].fillna('Normal/Mild')\n        label = prepare_labels(sub_sample.iloc[0])\n        print(\"study_dir = \", study_dir)\n\n        # Get a list of all series subdirectories\n        series_dirs = [os.path.join(study_dir, d) for d in os.listdir(study_dir) if os.path.isdir(os.path.join(study_dir, d))]\n\n        for series_dir in series_dirs:\n            # Load the entire series as a 3D volume\n            reader = sitk.ImageSeriesReader()\n            dicom_names = reader.GetGDCMSeriesFileNames(series_dir)\n\n            # Check if any DICOM files were found\n            if not dicom_names:\n                print(f\"WARNING: No DICOM files found in {series_dir}\")\n                continue\n\n            reader.SetFileNames(dicom_names)\n            image = reader.Execute()\n\n            # Convert to numpy array and preprocess\n            image_np = sitk.GetArrayFromImage(image)\n            processed_image = load_image(image_np)\n\n            labels.append(label)\n            image_data.append(processed_image)\n\n    # ... rest of the code (consider returning a dictionary for labels and data)\n    return labels, image_data\n","metadata":{"execution":{"iopub.status.busy":"2024-08-03T18:58:42.836800Z","iopub.status.idle":"2024-08-03T18:58:42.837174Z","shell.execute_reply.started":"2024-08-03T18:58:42.836999Z","shell.execute_reply":"2024-08-03T18:58:42.837014Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Original Prepaer Data\ndef prepare_data_orig(studies, df):\n    labels = []\n    image_paths = []\n    \n    for study_id in studies:\n        study_dir = TRAIN_DIR+study_id+'/'\n        sub_sample = df[df.study_id==study_id].fillna('Normal/Mild')\n        label = prepare_labels(sub_sample.iloc[0])\n        \n        for series_id in os.listdir(study_dir):\n            series_dir = study_dir+series_id+'/'\n            for z in os.listdir(series_dir):\n                z = z.split('.')[0]\n                #sub_desc = desc.where(desc.study_id==study_id).where(desc.series_id==series_id).dropna()\n                path = series_dir+z+'.dcm'\n                labels.append(label)\n                image_paths.append(path)\n    \n    data = tf.data.Dataset.from_tensor_slices((np.array(image_paths),np.array(labels, dtype='float32')))\n    return data\n\n# TEG New Prepare data to load 10 image slices in\n# Modify the prepare_data function to ensure that for each series, only the paths for the first 10 .dcm files are added to \n#     the image_paths list. You can use slicing or a loop to control this.\n# prepare_data now groups the first 10 DICOM image paths per series into lists and associates them with their corresponding label.\n\n\n\ndef prepare_data(studies, df):\n    label_list = []\n    stacked_image_list = []  # This will now store lists of 10 image paths for each series\n    \n    for study_id in studies:\n        study_dir = f\"{TRAIN_DIR}{study_id}/\"\n        sub_sample = df[df.study_id == study_id].fillna('Normal/Mild')\n        label = prepare_labels(sub_sample.iloc[0])\n        \n        for series_id in os.listdir(study_dir):\n            series_dir = f\"{study_dir}{series_id}/\"\n            # Assuming files are named sequentially as '1.dcm', '2.dcm', ..., '10.dcm'\n            paths = [f\"{series_dir}{z+1}.dcm\" for z in range(10)]\n            label_list.append(label)\n            stacked_image_list.append(paths)  # Store the list of 10 paths as a single entry\n    \n    #print(\"label_list = \", label_list)\n    return stacked_image_list, label_list\n    #data = tf.data.Dataset.from_tensor_slices((np.array(stacked_image_list), np.array(label_list, dtype='float32')))\n    #return data\n\n","metadata":{"execution":{"iopub.status.busy":"2024-07-22T16:06:42.99662Z","iopub.execute_input":"2024-07-22T16:06:42.997153Z","iopub.status.idle":"2024-07-22T16:06:43.018792Z","shell.execute_reply.started":"2024-07-22T16:06:42.997111Z","shell.execute_reply":"2024-07-22T16:06:43.017175Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#train_data = prepare_data(train_studies, labels)\n#val_data = prepare_data(val_studies, labels)\n\ntrain_stacked_image_list, train_label_list  = prepare_data(train_studies, labels)\nval_stacked_image_list, val_label_list  = prepare_data(val_studies, labels)","metadata":{"execution":{"iopub.status.busy":"2024-07-22T16:22:17.372092Z","iopub.execute_input":"2024-07-22T16:22:17.372567Z","iopub.status.idle":"2024-07-22T16:22:23.190175Z","shell.execute_reply.started":"2024-07-22T16:22:17.372532Z","shell.execute_reply":"2024-07-22T16:22:23.189018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(type(train_label_list))\nprint(\"Size of label_list:\", len(train_label_list))\n\nprint(\"First 5 elements:\")\nfor i in range(5):\n    print(train_label_list[i])\n\nprint(type(train_stacked_image_list))\nprint(\"Size of label_list:\", len(train_stacked_image_list))\n\nprint(\"First 5 elements:\")\nfor i in range(5):\n    print(train_stacked_image_list[i])","metadata":{"execution":{"iopub.status.busy":"2024-07-22T16:22:36.770255Z","iopub.execute_input":"2024-07-22T16:22:36.770748Z","iopub.status.idle":"2024-07-22T16:22:36.786418Z","shell.execute_reply.started":"2024-07-22T16:22:36.770711Z","shell.execute_reply":"2024-07-22T16:22:36.785021Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Original Load image function\ndef load_image_orig(image_path):\n    raw_image = tf.io.read_file(image_path)\n    sp = tf.strings.split(tf.gather(tf.strings.split(image_path, 'images/'), 1), '/')\n    N = tf.size(sp)\n    LEN = tf.strings.length(tf.gather(sp, 0))+tf.strings.length(tf.gather(sp, 2))\n    \n    # Add missing file metadata to avoid warnnigs flooding\n    if   LEN==12: raw_image = tf.strings.regex_replace(raw_image, pattern=b'DICM\\x02\\x00\\x01\\x00', rewrite=b'DICM\\x02\\x00\\x00\\x00UL\\x04\\x00\\x92\\x00\\x00\\x00\\x02\\x00\\x01\\x00')\n    elif LEN==13: raw_image = tf.strings.regex_replace(raw_image, pattern=b'DICM\\x02\\x00\\x01\\x00', rewrite=b'DICM\\x02\\x00\\x00\\x00UL\\x04\\x00\\x92\\x00\\x00\\x00\\x02\\x00\\x01\\x00')\n    elif LEN==14: raw_image = tf.strings.regex_replace(raw_image, pattern=b'DICM\\x02\\x00\\x01\\x00', rewrite=b'DICM\\x02\\x00\\x00\\x00UL\\x04\\x00\\x94\\x00\\x00\\x00\\x02\\x00\\x01\\x00')\n    elif LEN==15: raw_image = tf.strings.regex_replace(raw_image, pattern=b'DICM\\x02\\x00\\x01\\x00', rewrite=b'DICM\\x02\\x00\\x00\\x00UL\\x04\\x00\\x94\\x00\\x00\\x00\\x02\\x00\\x01\\x00')\n    elif LEN==16: raw_image = tf.strings.regex_replace(raw_image, pattern=b'DICM\\x02\\x00\\x01\\x00', rewrite=b'DICM\\x02\\x00\\x00\\x00UL\\x04\\x00\\x96\\x00\\x00\\x00\\x02\\x00\\x01\\x00')\n    elif LEN==17: raw_image = tf.strings.regex_replace(raw_image, pattern=b'DICM\\x02\\x00\\x01\\x00', rewrite=b'DICM\\x02\\x00\\x00\\x00UL\\x04\\x00\\x96\\x00\\x00\\x00\\x02\\x00\\x01\\x00')\n    elif LEN==18: raw_image = tf.strings.regex_replace(raw_image, pattern=b'DICM\\x02\\x00\\x01\\x00', rewrite=b'DICM\\x02\\x00\\x00\\x00UL\\x04\\x00\\x98\\x00\\x00\\x00\\x02\\x00\\x01\\x00')\n    \n    img = tfio.image.decode_dicom_image(raw_image, scale='auto', dtype=tf.float32)\n    m, M=tf.math.reduce_min(img), tf.math.reduce_max(img)\n    img = (tf.image.grayscale_to_rgb(img)-m)/(M-m)\n    img = tf.image.resize(img, IMG_SIZE)[0]\n    return img\n\n# TEG New load image function that loads 10 stacked images\n#  Modify the load_image function to handle multiple DICOM images and stack them along the channel axis.\n# load_image is modified to accept a list of paths, load each image, preprocess it, and then stack them into a single tensor with multiple channels.\n\n@tf.function\ndef load_image(image_paths):\n    def read_and_preprocess(path):\n        raw_image = tf.io.read_file(path)\n        \n        sp = tf.strings.split(tf.gather(tf.strings.split(path, 'images/'), 1), '/')\n        N = tf.size(sp)\n        LEN = tf.strings.length(tf.gather(sp, 0))+tf.strings.length(tf.gather(sp, 2))\n\n        # Add missing file metadata to avoid warnnigs flooding\n        if   LEN==12: raw_image = tf.strings.regex_replace(raw_image, pattern=b'DICM\\x02\\x00\\x01\\x00', rewrite=b'DICM\\x02\\x00\\x00\\x00UL\\x04\\x00\\x92\\x00\\x00\\x00\\x02\\x00\\x01\\x00')\n        elif LEN==13: raw_image = tf.strings.regex_replace(raw_image, pattern=b'DICM\\x02\\x00\\x01\\x00', rewrite=b'DICM\\x02\\x00\\x00\\x00UL\\x04\\x00\\x92\\x00\\x00\\x00\\x02\\x00\\x01\\x00')\n        elif LEN==14: raw_image = tf.strings.regex_replace(raw_image, pattern=b'DICM\\x02\\x00\\x01\\x00', rewrite=b'DICM\\x02\\x00\\x00\\x00UL\\x04\\x00\\x94\\x00\\x00\\x00\\x02\\x00\\x01\\x00')\n        elif LEN==15: raw_image = tf.strings.regex_replace(raw_image, pattern=b'DICM\\x02\\x00\\x01\\x00', rewrite=b'DICM\\x02\\x00\\x00\\x00UL\\x04\\x00\\x94\\x00\\x00\\x00\\x02\\x00\\x01\\x00')\n        elif LEN==16: raw_image = tf.strings.regex_replace(raw_image, pattern=b'DICM\\x02\\x00\\x01\\x00', rewrite=b'DICM\\x02\\x00\\x00\\x00UL\\x04\\x00\\x96\\x00\\x00\\x00\\x02\\x00\\x01\\x00')\n        elif LEN==17: raw_image = tf.strings.regex_replace(raw_image, pattern=b'DICM\\x02\\x00\\x01\\x00', rewrite=b'DICM\\x02\\x00\\x00\\x00UL\\x04\\x00\\x96\\x00\\x00\\x00\\x02\\x00\\x01\\x00')\n        elif LEN==18: raw_image = tf.strings.regex_replace(raw_image, pattern=b'DICM\\x02\\x00\\x01\\x00', rewrite=b'DICM\\x02\\x00\\x00\\x00UL\\x04\\x00\\x98\\x00\\x00\\x00\\x02\\x00\\x01\\x00')\n        \n        img = tfio.image.decode_dicom_image(raw_image, scale='auto', dtype=tf.float32)\n        img = tf.image.resize(img, IMG_SIZE)[0]  # Assuming IMG_SIZE is defined\n        #print(img.shape)\n\n        return img\n\n    # Use flat_map to apply load_image to each element in image_paths\n    #return tf.data.Dataset.from_tensor_slices(image_paths).flat_map(read_and_preprocess)\n\n    images = [read_and_preprocess(path) for path in image_paths]\n    stacked_image = tf.stack(images, axis=-1)  # Stack along the channel dimension\n    return stacked_image\n\ndef load_image3(image_paths):\n    def read_and_preprocess(path):\n        raw_image = tf.io.read_file(path)\n        img = tfio.image.decode_dicom_image(raw_image, scale='auto', dtype=tf.float32)\n        img = tf.image.resize(img, IMG_SIZE)[0]  # Ensure all images are the same size\n        print(img.shape)\n        return img\n\n    images = tf.map_fn(read_and_preprocess, image_paths, dtype=tf.float32)\n    # Ensure all images have the same number of channels\n    # Assuming all images have been resized and have consistent channel size set during preprocessing\n    return tf.stack(images, axis=-1)  # Stack along the channel dimension\n\n# Original Load Dataset\ndef load_dataset_orig(image_path, labels):\n    image = load_image(image_path)\n    return {\"images\": tf.cast(image, tf.float32), \"labels\": tf.cast(labels, tf.float32)}\n\ndef load_dataset(image_paths, labels):\n    image = tf.py_function(func=load_image, inp=[image_paths], Tout=tf.float32)\n    image.set_shape((IMG_SIZE[0], IMG_SIZE[1], 10))  # Assuming all images now will have 10 channels\n    return {\"images\": image, \"labels\": labels}\n\ndef prepare_datasets(data, batch_size):\n    return data.map(load_dataset, num_parallel_calls=tf.data.AUTOTUNE).batch(batch_size, drop_remainder=True)\n\n\n\n","metadata":{"execution":{"iopub.status.busy":"2024-07-22T16:17:43.300741Z","iopub.execute_input":"2024-07-22T16:17:43.301629Z","iopub.status.idle":"2024-07-22T16:17:43.332532Z","shell.execute_reply.started":"2024-07-22T16:17:43.301586Z","shell.execute_reply":"2024-07-22T16:17:43.331104Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_pre_stacked_images = []\nval_pre_stacked_images = []\n\n#for image_paths in train_stacked_image_list:\n#    # Call load_image to process the group of slices\n#    stacked_image = load_image(image_paths)\n#    train_pre_stacked_images.append(stacked_image)\n    \nfor image_paths in val_stacked_image_list:\n    # Call load_image to process the group of slices\n    stacked_image = load_image(image_paths)\n    val_pre_stacked_images.append(stacked_image)","metadata":{"execution":{"iopub.status.busy":"2024-07-22T16:24:09.873028Z","iopub.execute_input":"2024-07-22T16:24:09.874463Z","iopub.status.idle":"2024-07-22T16:41:52.03793Z","shell.execute_reply.started":"2024-07-22T16:24:09.874419Z","shell.execute_reply":"2024-07-22T16:41:52.035935Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(type(val_pre_stacked_images))\nprint(\"Size of label_list:\", len(val_pre_stacked_images))\n\nprint(\"First 5 elements:\")\nfor i in range(5):\n    print(val_pre_stacked_images[i])","metadata":{"execution":{"iopub.status.busy":"2024-07-22T16:47:46.452681Z","iopub.execute_input":"2024-07-22T16:47:46.453932Z","iopub.status.idle":"2024-07-22T16:47:46.485728Z","shell.execute_reply.started":"2024-07-22T16:47:46.453856Z","shell.execute_reply":"2024-07-22T16:47:46.484019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# original code\n#train_ds = train_data.map(load_dataset, num_parallel_calls=tf.data.AUTOTUNE)\n#train_ds = train_ds.ragged_batch(BATCH_SIZE, drop_remainder=True)\n\n#val_ds = val_data.map(load_dataset, num_parallel_calls=tf.data.AUTOTUNE)\n#val_ds = val_ds.ragged_batch(BATCH_SIZE, drop_remainder=True)\n\n# TEG updated code\n# When mapping the dataset, you'll need to update the lambda function to handle lists of image paths:\n\n#train_ds = train_data.map(lambda x, y: load_dataset(x, y), num_parallel_calls=tf.data.AUTOTUNE)\n#train_ds = train_ds.batch(BATCH_SIZE, drop_remainder=True)\n\ntrain_ds =  train_data.map(load_dataset, num_parallel_calls=tf.data.AUTOTUNE).batch(BATCH_SIZE, drop_remainder=True)\n\n\nval_ds = val_data.map(lambda x, y: load_dataset(x, y), num_parallel_calls=tf.data.AUTOTUNE)\nval_ds = val_ds.batch(BATCH_SIZE, drop_remainder=True)\n\n# Application of prepare_datasets function\n#train_ds = prepare_datasets(train_data, BATCH_SIZE)\n#val_ds = prepare_datasets(val_data, BATCH_SIZE)\n\n# The final model input will consist of single tensors with 10 channels each, representing the first 10 image slices of each series. \n# This approach assumes that each series directory contains at least 10 .dcm files and that they are named sequentially \n# from 1.dcm to 10.dcm. Adjustments may be necessary if these assumptions do not hold.","metadata":{"execution":{"iopub.status.busy":"2024-07-22T01:38:06.192739Z","iopub.execute_input":"2024-07-22T01:38:06.193184Z","iopub.status.idle":"2024-07-22T01:38:06.278798Z","shell.execute_reply.started":"2024-07-22T01:38:06.193152Z","shell.execute_reply":"2024-07-22T01:38:06.277644Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(labels)\nprint(type(labels))\n\nprint(image_paths)\nprint(type(image_paths))","metadata":{"execution":{"iopub.status.busy":"2024-07-22T01:39:40.896071Z","iopub.execute_input":"2024-07-22T01:39:40.896492Z","iopub.status.idle":"2024-07-22T01:39:40.904903Z","shell.execute_reply.started":"2024-07-22T01:39:40.896461Z","shell.execute_reply":"2024-07-22T01:39:40.903609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create the dataset\ntrain_ds = tf.data.Dataset.from_tensor_slices((image_paths, labels))\nprint(traind_ds)\n\n# Map the load_image function to the image paths\ntrain_ds = train_ds.map(lambda image_paths, label: (load_image(image_paths), label))\nprint(traind_ds)\n\n# Batch the data\nbatch_size = 32\ntrain_ds = train_ds.batch(batch_size)\nprint(traind_ds)\n","metadata":{"execution":{"iopub.status.busy":"2024-07-22T01:17:49.817895Z","iopub.execute_input":"2024-07-22T01:17:49.818359Z","iopub.status.idle":"2024-07-22T01:17:49.901634Z","shell.execute_reply.started":"2024-07-22T01:17:49.818326Z","shell.execute_reply":"2024-07-22T01:17:49.900153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define batch size (adjust as needed)\nbatch_size = 32\n\n# Map the load_image function to each element in train_data\ntrain_ds = train_data.map(lambda image_paths, labels: (load_image(image_paths), labels))\n\n# Combine images and labels into tuples\n#train_ds = train_ds.zip((lambda x: x[1]))  # Access labels from the second element\n\n# Batch the data\ntrain_ds = train_ds.batch(batch_size)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print (train_ds)\nprint (type(train_ds))\n","metadata":{"execution":{"iopub.status.busy":"2024-07-22T01:38:43.875886Z","iopub.execute_input":"2024-07-22T01:38:43.876286Z","iopub.status.idle":"2024-07-22T01:38:43.881788Z","shell.execute_reply.started":"2024-07-22T01:38:43.876259Z","shell.execute_reply":"2024-07-22T01:38:43.88062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Iterate through the dataset and print the first 3 elements\nfor images, labels in train_ds.take(3):\n    # Print some information about the images and labels\n    #print(\"Image batch shape:\", images.shape)\n    print(\"Sample image values (first 3 pixels):\")\n    #print(images[0, :3, :3, :3])  # Print the first 3x3x3 values from the first image\n    #print(\"Label batch shape:\", labels.shape)\n    print(\"Sample labels (first 3):\")\n    #print(labels[:3])  # Print the first 3 labels\n    break  # Stop after printing 3 elements\n\n# Remember that datasets are stateful, so iterating through it might consume elements","metadata":{"execution":{"iopub.status.busy":"2024-07-22T01:38:48.44901Z","iopub.execute_input":"2024-07-22T01:38:48.450181Z","iopub.status.idle":"2024-07-22T01:38:48.997259Z","shell.execute_reply.started":"2024-07-22T01:38:48.450131Z","shell.execute_reply":"2024-07-22T01:38:48.993247Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dict_to_tuple(inputs):\n    return inputs[\"images\"], inputs[\"labels\"]\n\ntrain_ds_tuple = train_ds.map(dict_to_tuple, num_parallel_calls=tf.data.AUTOTUNE)\ntrain_ds_tuple = train_ds_tuple.prefetch(tf.data.AUTOTUNE)\n\nval_ds_tuple = val_ds.map(dict_to_tuple, num_parallel_calls=tf.data.AUTOTUNE)\nval_ds_tuple = val_ds_tuple.prefetch(tf.data.AUTOTUNE)","metadata":{"execution":{"iopub.status.busy":"2024-07-22T00:44:52.820139Z","iopub.execute_input":"2024-07-22T00:44:52.820616Z","iopub.status.idle":"2024-07-22T00:44:52.855295Z","shell.execute_reply.started":"2024-07-22T00:44:52.820581Z","shell.execute_reply":"2024-07-22T00:44:52.853799Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Iterate through the dataset and print the first 3 elements\nfor images, labels in train_ds_tuple.take(3):\n    # Print some information about the images and labels\n    print(\"Image batch shape:\", images.shape)\n    print(\"Sample image values (first 3 pixels):\")\n    print(images[0, :3, :3, :3])  # Print the first 3x3x3 values from the first image\n    print(\"Label batch shape:\", labels.shape)\n    print(\"Sample labels (first 3):\")\n    print(labels[:3])  # Print the first 3 labels\n    break  # Stop after printing 3 elements\n\n# Remember that datasets are stateful, so iterating through it might consume elements\n","metadata":{"execution":{"iopub.status.busy":"2024-07-22T00:44:53.842509Z","iopub.execute_input":"2024-07-22T00:44:53.842952Z","iopub.status.idle":"2024-07-22T00:44:54.118369Z","shell.execute_reply.started":"2024-07-22T00:44:53.842917Z","shell.execute_reply":"2024-07-22T00:44:54.108538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model ","metadata":{}},{"cell_type":"code","source":"import math\nimport matplotlib.pyplot as plt\n\n# Re-used the lr_scheduler from https://www.kaggle.com/code/awsaf49/birdclef24-kerascv-starter-train\ndef get_lr_callback(batch_size=8, mode='cos', epochs=10, plot=False):\n    lr_start, lr_max, lr_min = 5e-5, 8e-6 * batch_size, 1e-5\n    lr_ramp_ep, lr_sus_ep, lr_decay = 6, 0, 0.75\n\n    def lrfn(epoch):  # Learning rate update function\n        if epoch < lr_ramp_ep: lr = (lr_max - lr_start) / lr_ramp_ep * epoch + lr_start\n        elif epoch < lr_ramp_ep + lr_sus_ep: lr = lr_max\n        elif mode == 'exp': lr = (lr_max - lr_min) * lr_decay**(epoch - lr_ramp_ep - lr_sus_ep) + lr_min\n        elif mode == 'step': lr = lr_max * lr_decay**((epoch - lr_ramp_ep - lr_sus_ep) // 2)\n        elif mode == 'cos':\n            decay_total_epochs, decay_epoch_index = epochs - lr_ramp_ep - lr_sus_ep + 3, epoch - lr_ramp_ep - lr_sus_ep\n            phase = math.pi * decay_epoch_index / decay_total_epochs\n            lr = (lr_max - lr_min) * 0.5 * (1 + math.cos(phase)) + lr_min\n        return lr\n\n    if plot:  # Plot lr curve if plot is True\n        plt.figure(figsize=(10, 5))\n        plt.plot(np.arange(epochs), [lrfn(epoch) for epoch in np.arange(epochs)], marker='o')\n        plt.xlabel('epoch'); plt.ylabel('lr')\n        plt.title('LR Scheduler')\n        plt.show()\n\n    return keras.callbacks.LearningRateScheduler(lrfn, verbose=False)  # Create lr callback","metadata":{"execution":{"iopub.status.busy":"2024-07-22T00:30:55.948729Z","iopub.execute_input":"2024-07-22T00:30:55.950146Z","iopub.status.idle":"2024-07-22T00:30:55.962225Z","shell.execute_reply.started":"2024-07-22T00:30:55.950092Z","shell.execute_reply":"2024-07-22T00:30:55.960889Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"backbone = keras_cv.models.EfficientNetV2Backbone.from_preset(PRETRAINED)\nmodel_orig = keras.Sequential(\n    [\n        keras.layers.Input(shape=(None, None, 3)),\n        backbone,\n        keras.layers.GlobalMaxPooling2D(),\n        keras.layers.Dropout(rate=0.3),\n        keras.layers.Dense(N_CLASSES, activation=\"sigmoid\"),\n    ]\n)\n\n# Define the model\nmodel = keras.Sequential([\n    # Adjust the input shape to accept 10 channels\n    keras.layers.Input(shape=(None, None, 10)),  # Specify desired input dimensions, e.g., (224, 224, 10) for specific size\n\n    # 1x1 Convolution to reduce channel dimension from 10 to 3\n    keras.layers.Conv2D(3, (1, 1), padding='same', activation='relu'),\n\n    # EfficientNet Backbone\n    backbone,\n\n    # Following layers remain unchanged\n    keras.layers.GlobalMaxPooling2D(),\n    keras.layers.Dropout(rate=0.3),\n    keras.layers.Dense(10, activation=\"sigmoid\"),\n])\n\n\nmodel.compile(optimizer=\"adam\",\n              loss=keras.losses.BinaryCrossentropy(),\n              metrics=[keras.metrics.AUC(name='auc')],\n             )\n\nmodel.summary()","metadata":{"execution":{"iopub.status.busy":"2024-07-22T00:31:00.527241Z","iopub.execute_input":"2024-07-22T00:31:00.527697Z","iopub.status.idle":"2024-07-22T00:31:23.255737Z","shell.execute_reply.started":"2024-07-22T00:31:00.527653Z","shell.execute_reply":"2024-07-22T00:31:23.254448Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lr_cb = get_lr_callback(BATCH_SIZE, epochs=EPOCH, plot=True)","metadata":{"execution":{"iopub.status.busy":"2024-07-22T00:31:35.671733Z","iopub.execute_input":"2024-07-22T00:31:35.672552Z","iopub.status.idle":"2024-07-22T00:31:35.929747Z","shell.execute_reply.started":"2024-07-22T00:31:35.672503Z","shell.execute_reply":"2024-07-22T00:31:35.928489Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ckpt_cb = keras.callbacks.ModelCheckpoint(\"best_model.weights.h5\",\n                                         monitor='val_auc',\n                                         save_best_only=True,\n                                         save_weights_only=True,\n                                         mode='max')","metadata":{"execution":{"iopub.status.busy":"2024-07-22T00:31:39.783699Z","iopub.execute_input":"2024-07-22T00:31:39.784633Z","iopub.status.idle":"2024-07-22T00:31:39.789737Z","shell.execute_reply.started":"2024-07-22T00:31:39.784592Z","shell.execute_reply":"2024-07-22T00:31:39.78855Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import warnings\n\n","metadata":{"execution":{"iopub.status.busy":"2024-07-22T00:31:23.264491Z","iopub.execute_input":"2024-07-22T00:31:23.264913Z","iopub.status.idle":"2024-07-22T00:31:23.277905Z","shell.execute_reply.started":"2024-07-22T00:31:23.264881Z","shell.execute_reply":"2024-07-22T00:31:23.276757Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = model.fit(\n    train_ds, \n    validation_data=val_ds, \n    epochs=EPOCH,\n    callbacks=[lr_cb, ckpt_cb], \n    verbose=1\n)","metadata":{"execution":{"iopub.status.busy":"2024-07-22T00:45:23.71099Z","iopub.execute_input":"2024-07-22T00:45:23.711466Z","iopub.status.idle":"2024-07-22T00:45:23.938483Z","shell.execute_reply.started":"2024-07-22T00:45:23.711434Z","shell.execute_reply":"2024-07-22T00:45:23.926954Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nauc = history.history['auc']\nval_auc = history.history['val_auc']\nloss = history.history['loss']\nval_loss = history.history['val_loss']\n\nepochs = range(len(loss))\nplt.plot(epochs, auc, 'r', label='Training auc')\nplt.plot(epochs, val_auc, 'b', label='Validation auc')\nplt.plot(epochs, loss, 'r', label='Training loss')\nplt.plot(epochs, val_loss, 'b', label='Validation loss')\nplt.legend(loc=0)\nplt.figure()\n\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-07-15T22:55:40.765387Z","iopub.status.idle":"2024-07-15T22:55:40.767321Z","shell.execute_reply.started":"2024-07-15T22:55:40.767012Z","shell.execute_reply":"2024-07-15T22:55:40.767043Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}