{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.16","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpu1vmV38","dataSources":[{"sourceId":20270,"databundleVersionId":1222630,"sourceType":"competition"},{"sourceId":9825415,"sourceType":"datasetVersion","datasetId":6025253},{"sourceId":11957481,"sourceType":"datasetVersion","datasetId":6591488}],"dockerImageVersionId":30841,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"DEVICE = \"TPU\"\n\nif DEVICE == \"TPU\":\n    !pip install -q pydicom\n\nRUN_TRAINING = False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T09:12:14.674769Z","iopub.execute_input":"2025-05-27T09:12:14.675014Z","iopub.status.idle":"2025-05-27T09:12:18.107563Z","shell.execute_reply.started":"2025-05-27T09:12:14.674987Z","shell.execute_reply":"2025-05-27T09:12:18.106373Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# imports\nimport os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport tensorflow as tf\nimport pydicom # needed to load .dcm images\nfrom keras.applications.densenet import DenseNet121\nfrom tensorflow.keras.applications import EfficientNetB0\nfrom tensorflow.keras import layers, models, Sequential\nfrom tensorflow.keras.layers import Input, Dense, Activation, Flatten, Conv2D, Layer\nfrom tensorflow.keras.layers import (\n    RandomFlip, RandomRotation, RandomZoom, RandomTranslation,\n    RandomBrightness, RandomContrast, GaussianNoise, Resizing, Lambda\n)\nfrom tensorflow.keras.models import Model\nfrom tensorflow.keras.callbacks import EarlyStopping, ModelCheckpoint, ReduceLROnPlateau\nfrom sklearn.metrics import (\n    confusion_matrix, precision_score, recall_score,\n    f1_score, roc_auc_score, roc_curve, classification_report, precision_recall_curve,\n    auc, precision_recall_fscore_support\n)\nfrom sklearn.model_selection import KFold, StratifiedGroupKFold\nfrom tensorflow.keras import regularizers\nimport math\nfrom tensorflow.keras import backend as keras_backend\nimport cv2\nfrom tensorflow.keras.models import load_model\nfrom tensorflow.keras import backend as K\nimport json\nimport re\nfrom enum import Enum\n\nSEED = 42\nnp.random.seed(SEED)\nnum_additional_features = 0\nAUTO = tf.data.experimental.AUTOTUNE\n\n# Define normalization constants\nMEAN = tf.constant([0.485, 0.456, 0.406], shape=(1, 1, 3), dtype=tf.float32)  # Pretrained mean\nSTD = tf.constant([0.229, 0.224, 0.225], shape=(1, 1, 3), dtype=tf.float32)   # Pretrained std\n\nIMAGE_RESIZE = 224","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-27T09:12:18.108557Z","iopub.execute_input":"2025-05-27T09:12:18.108950Z","iopub.status.idle":"2025-05-27T09:12:25.706194Z","shell.execute_reply.started":"2025-05-27T09:12:18.108922Z","shell.execute_reply":"2025-05-27T09:12:25.704847Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if DEVICE == \"TPU\":\n    print(\"connecting to TPU...\")\n    try:\n        tpu = tf.distribute.cluster_resolver.TPUClusterResolver(tpu='local')\n        print('Running on TPU ', tpu.master())\n    except ValueError:\n        print(\"Could not connect to TPU\")\n        tpu = None\n\n    if tpu:\n        try:\n            print(\"initializing  TPU ...\")\n            tf.config.experimental_connect_to_cluster(tpu)\n            tf.tpu.experimental.initialize_tpu_system(tpu)\n            strategy = tf.distribute.TPUStrategy(tpu)\n            print(\"TPU initialized\")\n        except _:\n            print(\"failed to initialize TPU\")\n    else:\n        DEVICE = \"GPU\"\n\nif DEVICE != \"TPU\":\n    print(\"Using default strategy for CPU and single GPU\")\n    strategy = tf.distribute.get_strategy()\n\nif DEVICE == \"GPU\":\n    print(\"Num GPUs Available: \", len(tf.config.experimental.list_physical_devices('GPU')))\n    \n\nAUTO     = tf.data.experimental.AUTOTUNE\nREPLICAS = strategy.num_replicas_in_sync\nprint(f'REPLICAS: {REPLICAS}')\n\nFOLDS = 3\n# BATCH_SIZES = [128]*FOLDS\nBATCH_SIZES = [128]*FOLDS\nEPOCHS = [15]*FOLDS","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T09:12:25.707092Z","iopub.execute_input":"2025-05-27T09:12:25.707495Z","iopub.status.idle":"2025-05-27T09:12:31.384990Z","shell.execute_reply.started":"2025-05-27T09:12:25.707467Z","shell.execute_reply":"2025-05-27T09:12:31.384014Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define your TFRecord files\n#INPUT_DIR = '/kaggle/input/siim-isic-melanoma-classification/tfrecords/'\n# INPUT_DIR = '/kaggle/input/isic-2020-melanoma-images-and-metadata/enhanced_dataset/'\nINPUT_DIR = '/kaggle/input/isic-2020-melanoma-images-and-metadata/isic2020-tfrecords-with-metadata/'\ninput_tfrec_train_pattern = INPUT_DIR + 'train*.tfrec'\ninput_tfrec_train = tf.io.gfile.glob(input_tfrec_train_pattern)\nprint(f\"Number of training TFRecord files found: {len(input_tfrec_train)}\")\n\ninput_tfrec_test_pattern = INPUT_DIR + 'test*.tfrec'\ninput_tfrec_test = tf.io.gfile.glob(input_tfrec_test_pattern)\nprint(f\"Number of training TFRecord files found: {len(input_tfrec_test)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T09:12:31.385930Z","iopub.execute_input":"2025-05-27T09:12:31.386177Z","iopub.status.idle":"2025-05-27T09:12:31.410022Z","shell.execute_reply.started":"2025-05-27T09:12:31.386152Z","shell.execute_reply":"2025-05-27T09:12:31.409080Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def inspect_tfrecord(file_path, num_samples=3):\n    raw_dataset = tf.data.TFRecordDataset(file_path)\n    for raw_record in raw_dataset.take(num_samples):\n        try:\n            example = tf.train.Example()\n            example.ParseFromString(raw_record.numpy())\n            print(example)\n        except Exception as e:\n            print(f\"Error parsing record: {e}\")\n            continue\n\n#inspect_tfrecord(INPUT_DIR + 'train00-2071-enhanced.tfrec', num_samples=3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T09:12:31.410968Z","iopub.execute_input":"2025-05-27T09:12:31.411198Z","iopub.status.idle":"2025-05-27T09:12:31.416201Z","shell.execute_reply.started":"2025-05-27T09:12:31.411175Z","shell.execute_reply":"2025-05-27T09:12:31.415063Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def extract_target(example):\n    tfrec_format = {\n        'target': tf.io.FixedLenFeature([], tf.int64),\n    }\n    parsed = tf.io.parse_single_example(example, tfrec_format)\n    return parsed['target']\n\n\ndef compute_class_ratio_from_tfrecords(tfrecord_files, max_samples=5000):\n    dataset = tf.data.TFRecordDataset(tfrecord_files, num_parallel_reads=tf.data.AUTOTUNE)\n    dataset = dataset.map(extract_target, num_parallel_calls=tf.data.AUTOTUNE)\n    dataset = dataset.take(max_samples)  # limit for speed\n\n    num_pos = 0\n    num_total = 0\n\n    for label in dataset:\n        label_val = label.numpy()\n        num_total += 1\n        if label_val == 1:\n            num_pos += 1\n\n    if num_total == 0:\n        print(\"Warning: No samples found for oversampling calculation.\")\n        return 0.0, 0, 1\n\n    pos_ratio = num_pos / num_total\n\n    # Compute safe oversampling factor\n    if pos_ratio == 0:\n        repeat_factor = 1  # Can't oversample if no positives\n    elif pos_ratio >= 0.5:\n        repeat_factor = 1  # No need to oversample if positives dominate\n    else:\n        repeat_factor = int((1 - pos_ratio) / pos_ratio)\n\n    return pos_ratio, num_total, repeat_factor\n\n\ndef compute_class_counts_from_tfrecords(tfrecord_files):\n    dataset = tf.data.TFRecordDataset(tfrecord_files, num_parallel_reads=tf.data.AUTOTUNE)\n    dataset = dataset.map(extract_target, num_parallel_calls=tf.data.AUTOTUNE)\n\n    num_pos = 0\n    num_neg = 0\n    num_total = 0\n\n    for label in dataset:\n        val = label.numpy()\n        if val == 1:\n            num_pos += 1\n        else:\n            num_neg += 1\n        num_total += 1\n\n    return num_pos, num_neg, num_total\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T09:12:31.417036Z","iopub.execute_input":"2025-05-27T09:12:31.417262Z","iopub.status.idle":"2025-05-27T09:12:31.426792Z","shell.execute_reply.started":"2025-05-27T09:12:31.417239Z","shell.execute_reply":"2025-05-27T09:12:31.425830Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Compute Global Age Statistics (Mean and Max Age) from TFRecords\ndef extract_age_approx(example):\n    tfrec_format = {\n        'age_approx': tf.io.FixedLenFeature([], tf.int64),\n    }\n    parsed_example = tf.io.parse_single_example(example, tfrec_format)\n    age = tf.cast(parsed_example['age_approx'], tf.float32)\n    return age\n\ndef create_age_dataset(tfrecord_files):\n    age_dataset = tf.data.TFRecordDataset(tfrecord_files, num_parallel_reads=AUTO)\n    age_dataset = age_dataset.map(extract_age_approx, num_parallel_calls=AUTO)\n    return age_dataset\n\ndef compute_age_statistics(age_dataset):\n    # Initialize accumulators\n    sum_age = tf.constant(0.0, dtype=tf.float32)\n    count = tf.constant(0, dtype=tf.int64)\n    max_age = tf.constant(0.0, dtype=tf.float32)\n    \n    def reducer(accum, age):\n        sum_age, count, max_age = accum\n        # Only consider non-missing values (assuming missing values are represented as NaN)\n        condition = tf.logical_not(tf.math.is_nan(age))\n        sum_age += tf.where(condition, age, 0.0)\n        count += tf.cast(condition, tf.int64)\n        max_age = tf.maximum(max_age, tf.where(condition, age, 0.0))\n        return sum_age, count, max_age\n    \n    sum_age, count, max_age = age_dataset.reduce(\n        (sum_age, count, max_age),\n        reducer\n    )\n    \n    # Avoid division by zero\n    mean_age = tf.cond(\n        tf.equal(count, 0),\n        lambda: tf.constant(0.0, dtype=tf.float32),\n        lambda: sum_age / tf.cast(count, tf.float32)\n    )\n    \n    mean_age = mean_age.numpy()\n    max_age = max_age.numpy()\n    \n    print(f\"Global Mean Age: {mean_age}\")\n    print(f\"Maximum Age: {max_age}\")\n    \n    return mean_age, max_age\n\n\n# Create the age dataset\nage_dataset = create_age_dataset(input_tfrec_train)\n# Compute statistics\nglobal_mean_age, global_max_age = compute_age_statistics(age_dataset)\n\nnum_additional_features += 1 # age is one numerical feature\nprint(f'Number of features after age computation: {num_additional_features}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T09:12:31.427870Z","iopub.execute_input":"2025-05-27T09:12:31.428104Z","iopub.status.idle":"2025-05-27T09:12:33.037028Z","shell.execute_reply.started":"2025-05-27T09:12:31.428081Z","shell.execute_reply":"2025-05-27T09:12:33.035464Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def extract_sex(example):\n    tfrec_format = {\n        'sex': tf.io.FixedLenFeature([], tf.string),\n    }\n    parsed_example = tf.io.parse_single_example(example, tfrec_format)\n    sex = parsed_example['sex']\n    # Decode bytes to string\n    sex = sex.numpy().decode('utf-8')\n    return sex\n\nsex_set = set()\n\nfor tfrecord in input_tfrec_train:\n    # Create a TFRecordDataset\n    dataset = tf.data.TFRecordDataset(tfrecord, num_parallel_reads=tf.data.AUTOTUNE)\n    \n    # Iterate through each serialized example in the TFRecord\n    for raw_record in dataset:\n        try:\n            sex = extract_sex(raw_record)\n            # Replace empty strings or specific missing indicators with 'unknown'\n            if not sex or sex.lower() == 'nan':\n                sex = 'unknown'\n            sex_set.add(sex)\n        except Exception as e:\n            print(f\"Error parsing record: {e}\")\n            continue\n\n# Add 'unknown' category to handle missing values\nsex_set.add('unknown')\n\n# Convert the set to a sorted list\nsex_categories = sorted(sex_set)\nprint(f\"Unique 'sex' Categories: {sex_categories}\")\nnum_sex_categories = len(sex_categories)\nprint(f\"Number of sex categories: {num_sex_categories}\")\n\nsex_to_index = {sex: idx for idx, sex in enumerate(sex_categories)}\nprint(f\"'sex' Mapping: {sex_to_index}\")\n\nnum_additional_features +=  num_sex_categories\nprint(f'Number of features after sex computation: {num_additional_features}')\n\ndef create_sex_lookup_table(sex_to_index):\n    # Create TensorFlow tensors for keys and values\n    keys = tf.constant(list(sex_to_index.keys()))\n    values = tf.constant(list(sex_to_index.values()), dtype=tf.int64)\n\n    # Set the default value (e.g., index for 'unknown')\n    default_value = sex_to_index.get('unknown', 2)\n    \n    # Create a KeyValueTensorInitializer\n    initializer = tf.lookup.KeyValueTensorInitializer(keys, values)\n    \n    # Use StaticHashTable which accepts default_value\n    table = tf.lookup.StaticHashTable(initializer, default_value)\n    return table\n\n# Create the lookup table\nsex_lookup_table = create_sex_lookup_table(sex_to_index)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T09:12:33.037893Z","iopub.execute_input":"2025-05-27T09:12:33.038169Z","iopub.status.idle":"2025-05-27T09:12:53.424138Z","shell.execute_reply.started":"2025-05-27T09:12:33.038141Z","shell.execute_reply":"2025-05-27T09:12:53.422557Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def extract_anatom_site(example):\n    tfrec_format = {\n        'anatom_site_general_challenge': tf.io.FixedLenFeature([], tf.string),\n    }\n    parsed_example = tf.io.parse_single_example(example, tfrec_format)\n    anatom_site = parsed_example['anatom_site_general_challenge']\n    # Decode bytes to string\n    anatom_site = anatom_site.numpy().decode('utf-8')\n    return anatom_site\n\n# Initialize an empty set to store unique 'anatom_site_general_challenge' categories\nanatom_site_set = set()\n\n# Iterate through each TFRecord file\nfor tfrecord in input_tfrec_train:\n    # Create a TFRecordDataset\n    dataset = tf.data.TFRecordDataset(tfrecord, num_parallel_reads=tf.data.AUTOTUNE)\n    \n    # Iterate through each serialized example in the TFRecord\n    for raw_record in dataset:\n        try:\n            anatom_site = extract_anatom_site(raw_record)\n            # Replace empty strings or specific missing indicators with 'unknown'\n            if not anatom_site or anatom_site.lower() == 'nan':\n                anatom_site = 'unknown'\n            anatom_site_set.add(anatom_site)\n        except Exception as e:\n            print(f\"Error parsing record: {e}\")\n            continue\n\n# Add 'unknown' category to handle missing values\nanatom_site_set.add('unknown')\n\n# Convert the set to a sorted list\nanatom_site_categories = sorted(anatom_site_set)\nprint(f\"Unique 'anatom_site_general_challenge' Categories: {anatom_site_categories}\")\nnum_anatom_sites = len(anatom_site_categories)\nprint(f\"Number of anatom sites: {num_anatom_sites}\")\n\n# Create a mapping from 'anatom_site_general_challenge' categories to indices\nanatom_site_to_index = {site: idx for idx, site in enumerate(anatom_site_categories)}\nprint(f\"'anatom_site_general_challenge' Mapping: {anatom_site_to_index}\")\n\nnum_additional_features += num_anatom_sites\nprint(f\"Total number of additional tabular features = {num_additional_features}\")\n\ndef create_anatom_site_lookup_table(anatom_site_to_index):\n    # Create TensorFlow tensors for keys and values\n    keys = tf.constant(list(anatom_site_to_index.keys()))\n    values = tf.constant(list(anatom_site_to_index.values()), dtype=tf.int64)\n\n    # Initialize the lookup table with a default value (index of 'unknown')\n    default_value = anatom_site_to_index.get('unknown', 5)  # Default to 'unknown' if not found\n    initializer = tf.lookup.KeyValueTensorInitializer(keys, values)\n    table = tf.lookup.StaticHashTable(initializer, default_value)\n    return table\n\n# Create the lookup table\nanatom_site_lookup_table = create_anatom_site_lookup_table(anatom_site_to_index)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T09:12:53.425045Z","iopub.execute_input":"2025-05-27T09:12:53.425316Z","iopub.status.idle":"2025-05-27T09:13:13.622790Z","shell.execute_reply.started":"2025-05-27T09:12:53.425289Z","shell.execute_reply":"2025-05-27T09:13:13.621391Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def check_missing_data(input_df):\n    print(\"=== Checking for NaNs ===\")\n    print(input_df.isna().sum())\n    \n    print(\"\\n=== Checking numeric columns for inf ===\")\n    numeric_cols = input_df.select_dtypes(include=[np.number]).columns\n    for col in numeric_cols:\n        if np.isinf(input_df[col]).any():\n            print(f\"Column '{col}' has infinite values!\")\n    \n    print(\"\\n=== Checking 'target' column uniqueness ===\")\n    print(\"Unique target labels:\", input_df[\"target\"].unique())\n    \n    print(\"\\n=== Some numeric bounds ===\")\n    for col in numeric_cols:\n        cmin = input_df[col].min()\n        cmax = input_df[col].max()\n        print(f\"{col} -> min: {cmin}, max: {cmax}\")\n    \n    print(\"Mixed precision policy:\", tf.keras.mixed_precision.global_policy()) # must be \"float32\"\n\n\ndef validate_training_dfs(train_df, validation_df):\n\n    positives_samples = train_df['target'].value_counts()\n    print(f\"Training set class distribution (neg:pos): {1-positives_samples}:{positives_samples}\")\n    positives_samples = validation_df['target'].value_counts()\n    print(f\"Validation set class distribution (neg:pos): {1-positives_samples}:{positives_samples}\")\n\n    # Verify no 'patient id' leakage\n    train_patient_ids = set(train_df['patient_id'])\n    val_patient_ids = set(validation_df['patient_id'])\n    leakage_check1 = train_patient_ids.intersection(val_patient_ids)\n    assert not leakage_check1, f\"Leakage check ::: 'Patient ID' leakage detected for patients: {leakage_check1}\"\n\n    # Check duplicates within validation set\n    train_images = set(train_df['image_name'])\n    val_images = set(validation_df['image_name'])\n    leakage_check2 = train_images.intersection(val_images)\n    assert not leakage_check2, f\"Leakage check ::: 'Image' leakage detected for images: {leakage_check2}\"\n\n    check_missing_data(train_df)\n    check_missing_data(validation_df)\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T09:13:13.623740Z","iopub.execute_input":"2025-05-27T09:13:13.624012Z","iopub.status.idle":"2025-05-27T09:13:13.632941Z","shell.execute_reply.started":"2025-05-27T09:13:13.623985Z","shell.execute_reply":"2025-05-27T09:13:13.631686Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def microscope_crop_tf(img, radius_offset=0):\n    # Assume square image [H, W, 3]\n    shape = tf.shape(img)\n    height = tf.cast(shape[0], tf.float32)\n    width = tf.cast(shape[1], tf.float32)\n    center_x = width / 2\n    center_y = height / 2\n    radius = (tf.minimum(width, height) / 2) - radius_offset\n\n    # Create meshgrid\n    y = tf.range(0.0, height)\n    x = tf.range(0.0, width)\n    Y, X = tf.meshgrid(y, x, indexing='ij')\n\n    dist_from_center = tf.sqrt(tf.square(X - center_x) + tf.square(Y - center_y))\n    circular_mask = tf.cast(dist_from_center <= radius, tf.float32)  # shape (H, W)\n\n    # Expand to 3 channels\n    circular_mask = tf.expand_dims(circular_mask, axis=-1)\n    circular_mask = tf.tile(circular_mask, [1, 1, 3])  # shape (H, W, 3)\n\n    return img * circular_mask\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T09:13:13.633961Z","iopub.execute_input":"2025-05-27T09:13:13.634195Z","iopub.status.idle":"2025-05-27T09:13:13.648190Z","shell.execute_reply.started":"2025-05-27T09:13:13.634170Z","shell.execute_reply":"2025-05-27T09:13:13.647519Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def label_counter(tfrecords):\n    tfrec_format = {'target': tf.io.FixedLenFeature([], tf.int64)}\n    counter = {0: 0, 1: 0}\n\n    for tfrec in tfrecords:\n        raw_ds = tf.data.TFRecordDataset(tfrec)\n        parsed = raw_ds.map(lambda x: tf.io.parse_single_example(x, tfrec_format))\n\n        for example in parsed:\n            label = example['target'].numpy()\n            counter[label] += 1\n\n    return counter\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T09:13:13.649046Z","iopub.execute_input":"2025-05-27T09:13:13.649265Z","iopub.status.idle":"2025-05-27T09:13:13.660039Z","shell.execute_reply.started":"2025-05-27T09:13:13.649242Z","shell.execute_reply":"2025-05-27T09:13:13.659249Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def parse_tfrec(example, training, num_sex_categories, num_anatom_sites, mean, std,\n                global_mean_age, global_max_age, sex_lookup_table, anatom_site_lookup_table, \n                has_label=True):\n    tfrec_format = {\n        'target'                       : tf.io.FixedLenFeature([], tf.int64),\n        'image'                        : tf.io.FixedLenFeature([], tf.string),\n        'image_name'                   : tf.io.FixedLenFeature([], tf.string),\n        'sex'                          : tf.io.FixedLenFeature([], tf.string),\n        'age_approx'                   : tf.io.FixedLenFeature([], tf.int64),\n        'patient_id'                   : tf.io.FixedLenFeature([], tf.string),\n        'anatom_site_general_challenge': tf.io.FixedLenFeature([], tf.string),\n    }\n    \n    if not has_label:\n        tfrec_format.pop('target')\n    \n    parsed_example = tf.io.parse_single_example(example, tfrec_format)\n\n    ####################\n    ###### IMAGE #######\n    ####################\n    # Decode image\n    image = tf.image.decode_jpeg(parsed_example['image'], channels=3)\n    image = tf.image.resize(image, [IMAGE_RESIZE, IMAGE_RESIZE])\n    \n    # Apply data augmentation if Training set\n    # Val and Test sets are training=False\n    if training:\n        image = data_augmentation(image)\n    \n    # # Apply microscope crop\n    # # image = microscope_crop(image, radius_offset=5)\n    image = microscope_crop_tf(image, radius_offset=5)\n    \n    # Normalize image\n    image = tf.cast(image, tf.float32) / 255.0  # Scale pixel values to [0, 1]\n    image = (image - mean) / std\n\n    ####################\n    ### TABULAR DATA ###\n    ####################\n    \n    # Handle 'age_approx': cast to float32, impute missing values and normalize\n    age_approx = tf.cast(parsed_example['age_approx'], tf.float32)\n    # Check for NaN\n    age_approx = tf.where(tf.math.is_nan(age_approx), \n                          tf.constant(global_mean_age, dtype=tf.float32), \n                          age_approx)\n    # Normalize by max_age\n    age_approx = age_approx / global_max_age\n    age_approx = tf.expand_dims(age_approx, axis=-1)  # Shape: (1,)\n\n    # Handle 'sex': impute missing with 'unknown' and map to index\n    sex = parsed_example['sex']\n    # Replace empty strings or 'nan' with 'unknown'\n    sex = tf.cond(\n        tf.logical_or(tf.equal(sex, ''), tf.equal(tf.strings.lower(sex), 'nan')),\n        lambda: tf.constant('unknown'),\n        lambda: sex\n    )\n    # Lookup the 'sex' index\n    sex_index = sex_lookup_table.lookup(sex)\n    sex_encoded = tf.one_hot(sex_index, depth=num_sex_categories, dtype=tf.float32)\n\n    # Handle 'anatom_site_general_challenge': impute missing with 'unknown' and map to index\n    anatom_site = parsed_example['anatom_site_general_challenge']\n    # Replace empty strings or 'nan' with 'unknown'\n    anatom_site = tf.cond(\n        tf.logical_or(tf.equal(anatom_site, ''), tf.equal(tf.strings.lower(anatom_site), 'nan')),\n        lambda: tf.constant('unknown'),\n        lambda: anatom_site\n    )\n    # Lookup the 'anatom_site_general_challenge' index\n    anatom_site_index = anatom_site_lookup_table.lookup(anatom_site)\n    # One-hot encode 'anatom_site_general_challenge'\n    anatom_site_encoded = tf.one_hot(tf.cast(anatom_site_index, tf.int32), depth=num_anatom_sites)\n    \n    \n    # Concatenate additional features\n    additional_features = tf.concat([sex_encoded, age_approx, anatom_site_encoded], axis=-1)\n    \n    if has_label:\n        label = tf.cast(parsed_example['target'], tf.float32)\n        return (image, additional_features), label\n    else:\n        return {\n            'image_input': image,\n            'tabular_input': additional_features,\n            'image_name': parsed_example['image_name']\n        }\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T09:13:13.661166Z","iopub.execute_input":"2025-05-27T09:13:13.661402Z","iopub.status.idle":"2025-05-27T09:13:13.673245Z","shell.execute_reply.started":"2025-05-27T09:13:13.661381Z","shell.execute_reply":"2025-05-27T09:13:13.672376Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data_augmentation = Sequential([\n    RandomFlip('horizontal_and_vertical'),\n    RandomRotation(0.05),          # ±18 degrees\n    RandomZoom(0.01),               # ±1%\n    RandomTranslation(0.05, 0.05),   # ±5%\n    RandomBrightness(0.15),\n    RandomContrast(0.15),\n    GaussianNoise(0.01)\n])\n\nclass DatasetType(Enum):\n    TRAINING = 'training'\n    VALIDATION = 'validation'\n    TEST = 'test'\n\ndef get_dataset(tfrecords, df_type, batch_size=128,\n                num_sex_categories=None, num_anatom_sites=None,\n                mean=MEAN, std=STD, global_mean_age=0.0, global_max_age=1.0,\n                sex_lookup_table=None, anatom_site_lookup_table=None,\n                has_label=True):\n\n    dataset = tf.data.TFRecordDataset(tfrecords, num_parallel_reads=AUTO)\n    dataset = dataset.cache()\n\n    \n    if df_type == DatasetType.TRAINING:\n        \n        dataset = dataset.repeat()\n    \n        dataset = dataset.shuffle(8192)\n        opt = tf.data.Options()\n        opt.experimental_deterministic = False\n        dataset = dataset.with_options(opt)\n        \n        dataset = dataset.map(\n            lambda x: parse_tfrec(\n                x,\n                training=True,\n                num_sex_categories=num_sex_categories, \n                num_anatom_sites=num_anatom_sites, \n                mean=mean, \n                std=std,\n                global_mean_age=global_mean_age, \n                global_max_age=global_max_age, \n                sex_lookup_table=sex_lookup_table, \n                anatom_site_lookup_table=anatom_site_lookup_table, \n                has_label=has_label\n            ),\n            num_parallel_calls=tf.data.AUTOTUNE\n        )       \n        \n        dataset = dataset.batch(batch_size * REPLICAS)\n\n\n    elif df_type == DatasetType.VALIDATION or df_type == DatasetType.TEST:\n        \n        dataset = dataset.map(\n            lambda x: parse_tfrec(\n                x,\n                training=False,\n                num_sex_categories=num_sex_categories, \n                num_anatom_sites=num_anatom_sites, \n                mean=mean, \n                std=std,\n                global_mean_age=global_mean_age, \n                global_max_age=global_max_age, \n                sex_lookup_table=sex_lookup_table, \n                anatom_site_lookup_table=anatom_site_lookup_table, \n                has_label=has_label\n            ),\n            num_parallel_calls=tf.data.AUTOTUNE\n        )\n        \n        dataset = dataset.batch(batch_size * REPLICAS)\n        \n    else:\n        raise ValueError('Wrong dataset type. Abort.')\n    \n    return dataset.prefetch(AUTO)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T09:13:13.674308Z","iopub.execute_input":"2025-05-27T09:13:13.674529Z","iopub.status.idle":"2025-05-27T09:13:13.704128Z","shell.execute_reply.started":"2025-05-27T09:13:13.674509Z","shell.execute_reply":"2025-05-27T09:13:13.703187Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class BaseModel(Enum):\n    DENSENET121 = 'DenseNet121'\n    EFFICIENTNETB0 = 'EfficientNetB0'\n\ndef build_model(fold_no, model_name, num_additional_features):\n    # Define input layers\n    image_input = layers.Input(shape=(224, 224, 3), name='image_input')\n    tabular_input = layers.Input(shape=(num_additional_features,), name='tabular_input')\n\n    # Base model\n    if model_name == BaseModel.EFFICIENTNETB0:\n        base_model = EfficientNetB0(weights='imagenet', include_top=False)\n        print(f\"Fold {fold_no} ::: Base model: {BaseModel.EFFICIENTNETB0.value}.\")\n    elif model_name == BaseModel.DENSENET121:\n        base_model = DenseNet121(weights='imagenet', include_top=False)\n        print(f\"Fold {fold_no} ::: Base model: {BaseModel.DENSENET121.value}.\")\n    else:\n        base_model = DenseNet121(weights='imagenet', include_top=False)\n        print(f\"Fold {fold_no} ::: Base model: default.\")\n    \n    # Train entire base model\n    for layer in base_model.layers:\n        layer.trainable = True\n\n    # Image processing\n    x1 = base_model(image_input)\n    x1 = layers.GlobalAveragePooling2D()(x1)\n    x1 = layers.BatchNormalization()(x1)\n\n    # Process tabular features\n    x2 = layers.Dense(32, activation='relu', kernel_regularizer=regularizers.l2(1e-4))(tabular_input)\n    x2 = layers.BatchNormalization()(x2)\n    # x2 = layers.Dropout(0.05)(x2)\n    x2 = layers.Dense(16, activation='relu', kernel_regularizer=regularizers.l2(1e-4))(x2)\n    x2 = layers.BatchNormalization()(x2)\n    # x2 = layers.Dropout(0.05)(x2)\n\n    # Combine image and tabular features\n    combined = layers.concatenate([x1, x2])\n    combined = layers.BatchNormalization()(combined)\n\n    # Add final dense layers\n    x = layers.Dense(8, activation='relu', kernel_regularizer=regularizers.l2(1e-4))(combined)\n    x = layers.BatchNormalization()(x)\n    # x = layers.Dropout(0.1)(x)\n    output = layers.Dense(1, activation='sigmoid')(x)  # Binary classification\n\n    # Create the model\n    model = models.Model(inputs=[image_input, tabular_input], outputs=output)\n\n    opt = tf.keras.optimizers.Adam(learning_rate=1e-4)\n    loss = tf.keras.losses.BinaryCrossentropy(label_smoothing=0.05) \n    # loss = binary_focal_loss(alpha=0.5, gamma=2.25)\n    model.compile(\n        optimizer=opt,\n        loss=loss,\n        metrics=[\n            tf.keras.metrics.AUC(name='AUC'),\n            tf.keras.metrics.Precision(name='Precision'),\n            tf.keras.metrics.Recall(name='Recall')]\n    )\n    \n    return model\n\ndef get_lr_callback(batch_size):\n    lr_start   = 0.000005\n    lr_max     = 0.00000125 * batch_size * REPLICAS\n    lr_min     = 0.000001\n    lr_ramp_ep = 5\n    lr_sus_ep  = 0\n    lr_decay   = 0.8\n   \n    def lrfn(epoch):\n        if epoch < lr_ramp_ep:\n            lr = (lr_max - lr_start) / lr_ramp_ep * epoch + lr_start\n            \n        elif epoch < lr_ramp_ep + lr_sus_ep:\n            lr = lr_max\n            \n        else:\n            lr = (lr_max - lr_min) * lr_decay**(epoch - lr_ramp_ep - lr_sus_ep) + lr_min\n            \n        return lr\n\n    lr_callback = tf.keras.callbacks.LearningRateScheduler(lrfn, verbose=False)\n    return lr_callback\n\ndef get_callbacks(fold_no):\n    early_stopping = EarlyStopping(\n        monitor='val_AUC', patience=5, mode='max', restore_best_weights=True\n    )\n\n    checkpoint = ModelCheckpoint(\n        filepath=f'MelanomaModel_fold_{fold_no}_AUC_{{val_AUC:.5f}}.keras',\n        monitor='val_AUC',\n        mode='max',\n        save_best_only=True,\n        save_weights_only=False,  # saves full model (set True for just weights)\n        verbose=1\n    )\n\n    return [early_stopping, get_lr_callback(BATCH_SIZES[fold_no]), checkpoint]\n    # return [get_lr_callback(BATCH_SIZES[fold_no]), checkpoint]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T09:13:13.705036Z","iopub.execute_input":"2025-05-27T09:13:13.705238Z","iopub.status.idle":"2025-05-27T09:13:13.718385Z","shell.execute_reply.started":"2025-05-27T09:13:13.705218Z","shell.execute_reply":"2025-05-27T09:13:13.717149Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#input_df.head()\n\n\ndef count_data_items(filenames):\n    n = []\n    for filename in filenames:\n        match = re.search(r\"-(\\d+)-enhanced\\.tfrec\", filename)\n        if match:\n            n.append(int(match.group(1)))\n        else:\n            print(f\"⚠️ Warning: Could not extract count from {filename}\")\n    return np.sum(n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T09:13:13.719077Z","iopub.execute_input":"2025-05-27T09:13:13.719302Z","iopub.status.idle":"2025-05-27T09:13:13.729409Z","shell.execute_reply.started":"2025-05-27T09:13:13.719279Z","shell.execute_reply":"2025-05-27T09:13:13.728361Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def check_for_problems(df):\n    print(\"=== Checking for NaNs ===\")\n    print(df.isna().sum())\n    \n    print(\"\\n=== Checking numeric columns for inf ===\")\n    numeric_cols = df.select_dtypes(include=[np.number]).columns\n    for col in numeric_cols:\n        if np.isinf(df[col]).any():\n            print(f\"Column '{col}' has infinite values!\")\n    \n    print(\"\\n=== Checking 'target' column uniqueness ===\")\n    print(\"Unique target labels:\", df[\"target\"].unique())\n    \n    print(\"\\n=== Some numeric bounds ===\")\n    for col in numeric_cols:\n        cmin = df[col].min()\n        cmax = df[col].max()\n        print(f\"{col} -> min: {cmin}, max: {cmax}\")\n    \n    print(\"Mixed precision policy:\", tf.keras.mixed_precision.global_policy()) # must be \"float32\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T09:13:13.730183Z","iopub.execute_input":"2025-05-27T09:13:13.730405Z","iopub.status.idle":"2025-05-27T09:13:13.737396Z","shell.execute_reply.started":"2025-05-27T09:13:13.730382Z","shell.execute_reply":"2025-05-27T09:13:13.736539Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def display_images(dataset, mean, std, num_batches=1, images_per_batch=12):\n    for (images, additional_features), labels in dataset.take(num_batches):\n        num_images = min(images_per_batch, images.shape[0])\n\n        # Calculate grid size dynamically\n        cols = math.ceil(math.sqrt(num_images))\n        rows = math.ceil(num_images / cols)\n\n        plt.figure(figsize=(cols * 3, rows * 3))\n\n        for i in range(num_images):\n            plt.subplot(rows, cols, i + 1)\n            img = images[i].numpy() * std.numpy() + mean.numpy()\n            img = np.clip(img, 0.0, 1.0)\n            plt.imshow(img)\n            # Display label as \"Positive\" or \"Negative\"\n            label_text = \"Positive\" if labels[i].numpy() == 1 else \"Negative\"\n            plt.title(label_text)\n            plt.axis(\"off\")\n\n        # Hide any remaining subplots if the grid has more slots than images\n        total_slots = rows * cols\n        if total_slots > num_images:\n            for i in range(num_images, total_slots):\n                plt.subplot(rows, cols, i + 1)\n                plt.axis(\"off\")\n\n        plt.tight_layout()\n        plt.show()\n\n\ndef display_balanced_images(dataset, mean, std, num_batches=1, images_per_batch=12):\n    import random  # for shuffling\n\n    for (images, additional_features), labels in dataset.take(num_batches):\n        # Convert tensors to numpy arrays\n        images_np = images.numpy()\n        labels_np = labels.numpy()\n\n        # Separate positives and negatives\n        positives = [(img, 1) for img, label in zip(images_np, labels_np) if label == 1]\n        negatives = [(img, 0) for img, label in zip(images_np, labels_np) if label == 0]\n\n        # Shuffle for variety\n        random.shuffle(positives)\n        random.shuffle(negatives)\n\n        # Choose half from each class\n        half = images_per_batch // 2\n        selected = positives[:half] + negatives[:half]\n\n        # In case we don't have enough of one class\n        if len(selected) < images_per_batch:\n            additional = negatives[half:] + positives[half:]\n            selected += additional[:images_per_batch - len(selected)]\n\n        # Final safety check\n        selected = selected[:images_per_batch]\n\n        # Calculate grid size dynamically\n        cols = math.ceil(math.sqrt(images_per_batch))\n        rows = math.ceil(images_per_batch / cols)\n\n        plt.figure(figsize=(cols * 3, rows * 3))\n\n        for i, (img, label) in enumerate(selected):\n            plt.subplot(rows, cols, i + 1)\n            # Unnormalize\n            img = img * std.numpy() + mean.numpy()\n            img = np.clip(img, 0.0, 1.0)\n            plt.imshow(img)\n            plt.title(\"Positive\" if label == 1 else \"Negative\")\n            plt.axis(\"off\")\n\n        plt.tight_layout()\n        plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T09:13:13.738316Z","iopub.execute_input":"2025-05-27T09:13:13.738507Z","iopub.status.idle":"2025-05-27T09:13:13.750918Z","shell.execute_reply.started":"2025-05-27T09:13:13.738487Z","shell.execute_reply":"2025-05-27T09:13:13.750078Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if RUN_TRAINING:\n    all_fold_ids = np.arange(len(input_tfrec_train))\n    skf = KFold(n_splits=FOLDS,shuffle=True,random_state=SEED)\n    \n    # fold_no = 1\n    results = []\n    optimal_thresholds = []\n    \n    def count_examples(tfrecords):\n        count = 0\n        for _ in tf.data.TFRecordDataset(tfrecords):\n            count += 1\n        return count\n    \n    for fold_no, (train_index, val_index) in enumerate(skf.split(all_fold_ids)):\n        print(f\"::: Fold {fold_no} START :::\")\n        \n        # Split the data\n        # For each index in train_index, build the corresponding file pattern and expand it.\n        train_tfrecords = []\n        for idx in train_index:\n            pattern = f\"{INPUT_DIR}train{idx:02d}*.tfrec\"\n            # Expand the pattern to get the actual file names\n            files = tf.io.gfile.glob(pattern)\n            train_tfrecords.extend(files)\n        # Similarly for validation\n        validation_tfrecords = []\n        for idx in val_index:\n            pattern = f\"{INPUT_DIR}train{idx:02d}*.tfrec\"\n            files = tf.io.gfile.glob(pattern)\n            validation_tfrecords.extend(files)\n        \n        print(f\"Fold {fold_no} ::: Train tfrecords: {len(train_tfrecords)}\")\n        print(f\"Fold {fold_no} ::: Validation tfrecords: {len(validation_tfrecords)}\")\n    \n    \n        print(f\"Number of examples in Train TFRecords: {count_examples(train_tfrecords)}\")\n        \n    \n        # # Compute class stats for the training split only\n        num_pos, num_neg, total_count = compute_class_counts_from_tfrecords(train_tfrecords)\n        label_stats = label_counter(train_tfrecords)\n        print(f\"✅ Label counts from Train TFRecords: {label_stats}\")\n    \n        # _, _, val_total_count = compute_class_counts_from_tfrecords(validation_tfrecords)\n    \n        # steps_per_epoch = math.ceil(total_count / BATCH_SIZES[fold_no])\n        print(f\"Fold {fold_no} ::: total count = {total_count}\")\n        \n        # steps_per_epoch = math.ceil(total_count / (BATCH_SIZES[fold_no] * REPLICAS))\n        steps_per_epoch=int(count_data_items(train_tfrecords)/BATCH_SIZES[fold_no]//REPLICAS)\n\n        \n        print(f\"Fold {fold_no} ::: steps_per_epoch = {steps_per_epoch}\")\n        # validation_steps = math.ceil(val_total_count / BATCH_SIZES[fold_no])\n    \n        zeros_weight = total_count / (2.0 * num_neg)\n        ones_weight  = total_count / (2.0 * num_pos)\n        class_weight_dict = {0: zeros_weight, 1: ones_weight}\n        print(f\"Fold {fold_no} ::: class weights: {class_weight_dict}\")\n\n        train_dataset = get_dataset(\n            tfrecords=train_tfrecords,\n            df_type=DatasetType.TRAINING,\n            batch_size=BATCH_SIZES[fold_no],\n            num_sex_categories=num_sex_categories,\n            num_anatom_sites=num_anatom_sites,\n            mean=MEAN,\n            std=STD,\n            global_mean_age=global_mean_age,\n            global_max_age=global_max_age,\n            sex_lookup_table=sex_lookup_table,\n            anatom_site_lookup_table=anatom_site_lookup_table,\n            has_label=True\n        )\n        print(f\"Fold {fold_no} ::: Train dataset created successfully.\")\n        \n        print(f\"Fold {fold_no} ::: Sample of Training images:\")\n        # display_images(train_dataset, mean=MEAN, std=STD, num_batches=1, images_per_batch=12)\n        display_balanced_images(train_dataset, mean=MEAN, std=STD, num_batches=1, images_per_batch=12)\n        \n        # for (img_batch, meta_batch), labels in train_dataset.take(1):\n        #     print(\"✅ image batch shape:\", img_batch.shape)\n        #     print(\"✅ metadata shape:\", meta_batch.shape)\n        #     print(\"✅ labels shape:\", labels.shape)\n    \n        \n        validation_dataset = get_dataset(\n            tfrecords=validation_tfrecords,\n            df_type=DatasetType.VALIDATION,\n            batch_size=BATCH_SIZES[fold_no],\n            num_sex_categories=num_sex_categories,\n            num_anatom_sites=num_anatom_sites,\n            mean=MEAN,\n            std=STD,\n            global_mean_age=global_mean_age,\n            global_max_age=global_max_age,\n            sex_lookup_table=sex_lookup_table,\n            anatom_site_lookup_table=anatom_site_lookup_table,\n            has_label=True\n        )\n        print(f\"Fold {fold_no} ::: Validation dataset created successfully.\")\n        print(f\"Fold {fold_no} ::: Sample of Validation images:\")\n    \n        # display_images(validation_dataset, mean=MEAN, std=STD, num_batches=1, images_per_batch=12)\n        display_balanced_images(validation_dataset, mean=MEAN, std=STD, num_batches=1, images_per_batch=12)\n    \n        K.clear_session()\n        \n        print(f\"::: Fold {fold_no} ::: TRAINING START\")\n        with strategy.scope():\n            model = build_model(fold_no, BaseModel.DENSENET121, num_additional_features)\n            history = model.fit(\n                train_dataset,\n                epochs=EPOCHS[fold_no],\n                steps_per_epoch=steps_per_epoch,\n                validation_data=validation_dataset,\n                callbacks=get_callbacks(fold_no),\n                verbose=2,\n                class_weight=class_weight_dict\n            )\n        \n        # Evaluate the model\n        print(f\"Fold {fold_no} ::: Starting model evaluation against validation dataset.\")\n        true_labels = []\n        pred_probs = []\n    \n        for batch in validation_dataset:\n            (images, tabular_inputs), labels = batch\n            preds = model.predict((images, tabular_inputs), verbose=0)\n            pred_probs.extend(preds.flatten())\n            true_labels.extend(labels.numpy())\n    \n        true_labels = np.array(true_labels)\n        pred_probs = np.array(pred_probs)\n    \n        # Compute AUC (threshold-independent)\n        auc_value = roc_auc_score(true_labels, pred_probs)\n        \n        # ------------------------------------------\n        # Evaluate multiple predefined thresholds\n        # ------------------------------------------\n        \n        min_val = 0.20\n        max_val = 0.80\n        step_val = 0.05\n        thresholds_to_eval = [round(x, 2) for x in np.arange(min_val, max_val + step_val, step_val)]\n    \n        results_table = []\n        \n        for t in thresholds_to_eval:\n            predicted_labels_t = (pred_probs >= t).astype(int)\n        \n            # Compute confusion matrix for threshold t\n            conf = confusion_matrix(true_labels, predicted_labels_t)\n            tn, fp, fn, tp = conf.ravel()\n        \n            # Calculate precision and recall\n            prec_t = precision_score(true_labels, predicted_labels_t, zero_division=0)\n            rec_t = recall_score(true_labels, predicted_labels_t, zero_division=0)\n    \n            f1 = f1_score(true_labels, predicted_labels_t, zero_division=0)\n            \n            # Collect results in a dictionary\n            results_table.append({\n                'Threshold': t,\n                'AUC': f\"{auc_value:.4f}\",   # same AUC for all thresholds, it's threshold-independent\n                'Precision': f\"{prec_t:.4f}\",\n                'TNs': tn,\n                'FPs': fp,\n                'Recall': f\"{rec_t:.4f}\",\n                'TPs': tp,\n                'FNs': fn,\n                'F1': f\"{f1:.4f}\",\n            })\n        \n        # Convert to a DataFrame for a neat table\n        df_results = pd.DataFrame(results_table)\n        print(\"Evaluation at Fixed Thresholds:\")\n        print(df_results.to_string(index=False))\n        \n        # Compute ROC curve\n        fpr, tpr, thresholds_roc = roc_curve(true_labels, pred_probs)\n    \n        plt.figure(figsize=(6, 6))\n        plt.plot(fpr, tpr, color='r', label=f'ROC curve (AUC = {auc_value:.4f})')\n        plt.plot([0, 1], [0, 1], color='navy', linestyle='--')  # Diagonal line for reference\n        plt.title('ROC Curve')\n        plt.xlabel('False Positive Rate')\n        plt.ylabel('True Positive Rate')\n        plt.tight_layout()  \n        plt.show()\n    \n    \n        f1_scores = []\n        for t in thresholds_to_eval:\n            predicted_labels_t = (pred_probs >= t).astype(int)\n            f1 = f1_score(true_labels, predicted_labels_t, zero_division=0)\n            f1_scores.append(f1)\n        \n        # Find the threshold that gives the highest F1\n        best_idx = np.argmax(f1_scores)\n        best_threshold = thresholds_to_eval[best_idx]\n        best_f1 = f1_scores[best_idx]\n        \n        print(f\"\\n✅ Best Threshold by F1 Score: {best_threshold}\")\n        print(f\"🧮 Best F1 Score: {best_f1:.4f}\")\n    \n        best_preds = (pred_probs >= best_threshold).astype(int)\n        tn, fp, fn, tp = confusion_matrix(true_labels, best_preds).ravel()\n        print(f\"Confusion matrix at best F1 threshold ({best_threshold}):\")\n        print(f\"  TN={tn}, FP={fp}, FN={fn}, TP={tp}\")\n    \n        prec, rec, _ = precision_recall_curve(true_labels, pred_probs)\n        pr_auc = auc(rec, prec)\n        \n        print(f\"📈 PR AUC: {pr_auc:.4f}\")\n        \n        # Plot PR Curve\n        plt.figure(figsize=(6, 6))\n        plt.plot(rec, prec, color='green', label=f'PR Curve (AUC = {pr_auc:.4f})')\n        plt.xlabel('Recall')\n        plt.ylabel('Precision')\n        plt.title('Precision-Recall Curve')\n        plt.legend()\n        plt.grid(True)\n        plt.tight_layout()\n        plt.show()\n    \n    \n        \n        print(f\"::: Fold {fold_no} COMPLETE :::\")\n        print(f\"::: ::::::::::::::::::::::: :::\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T09:13:13.751776Z","iopub.execute_input":"2025-05-27T09:13:13.751993Z","iopub.status.idle":"2025-05-27T09:13:13.774676Z","shell.execute_reply.started":"2025-05-27T09:13:13.751973Z","shell.execute_reply":"2025-05-27T09:13:13.773822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create the test dataset\ntest_batch_size = BATCH_SIZES[0]\n\ntest_dataset = get_dataset(\n    tfrecords=input_tfrec_test,\n    df_type=DatasetType.TEST,\n    batch_size=test_batch_size,\n    num_sex_categories=num_sex_categories,\n    num_anatom_sites=num_anatom_sites,\n    mean=MEAN,\n    std=STD,\n    global_mean_age=global_mean_age,\n    global_max_age=global_max_age,\n    sex_lookup_table=sex_lookup_table,\n    anatom_site_lookup_table=anatom_site_lookup_table,\n    has_label=False\n)\n\nprint(f\"Test dataset created successfully.\")\n\n# for batch in test_dataset.take(1):\n#     print(batch['image_name'])  # should print a list of string tensors\n\n# print(f\"Fold {fold_no} ::: Sample of Validation images:\")\n# display_balanced_images(validation_dataset, mean=MEAN, std=STD, num_batches=1, images_per_batch=12)\n\n# for (image_batch, additional_features_batch) in test_dataset.take(1):\n#     print(f\"Image batch shape: {image_batch.shape}\")  # Expected: (batch_size, 224, 224, 3)\n#     print(f\"Additional features batch shape: {additional_features_batch.shape}\")  # Expected: (batch_size, num_additional_features)\n\n# Paths to the saved models\nmodel_paths = [\n    'MelanomaModel_fold_0_AUC_0.83144.keras',\n    'MelanomaModel_fold_1_AUC_0.89147.keras',\n    'MelanomaModel_fold_2_AUC_0.92443.keras'\n]\n\n# Fold 0\n# Evaluation at Fixed Thresholds:\n#  Threshold    AUC Precision   TNs   FPs Recall  TPs  FNs     F1\n#       0.20 0.8244    0.0184     0 12197 1.0000  229    0 0.0362\n#       0.25 0.8244    0.0185    23 12174 1.0000  229    0 0.0363\n#       0.30 0.8244    0.0186   140 12057 1.0000  229    0 0.0366\n#       0.35 0.8244    0.0193   558 11639 1.0000  229    0 0.0379\n#       0.40 0.8244    0.0202  1358 10839 0.9738  223    6 0.0395\n#       0.45 0.8244    0.0226  2719  9478 0.9563  219   10 0.0441\n#       0.50 0.8244    0.0277  4836  7361 0.9170  210   19 0.0538\n#       0.55 0.8244    0.0382  7159  5038 0.8734  200   29 0.0732\n#       0.60 0.8244    0.0505  8734  3463 0.8035  184   45 0.0949\n#       0.65 0.8244    0.0634  9804  2393 0.7074  162   67 0.1164\n#       0.70 0.8244    0.0804 10596  1601 0.6114  140   89 0.1421\n#       0.75 0.8244    0.1032 11163  1034 0.5197  119  110 0.1722\n#       0.80 0.8244    0.1266 11583   614 0.3886   89  140 0.1910\n#       0.85 0.8244    0.1683 11856   341 0.3013   69  160 0.2160\n\n# Fold 1\n# Evaluation at Fixed Thresholds:\n#  Threshold    AUC Precision  TNs  FPs Recall  TPs  FNs     F1\n#       0.20 0.8432    0.0225 2466 7702 1.0000  177    0 0.0439\n#       0.25 0.8432    0.0248 3213 6955 1.0000  177    0 0.0484\n#       0.30 0.8432    0.0265 3714 6454 0.9944  176    1 0.0517\n#       0.35 0.8432    0.0280 4129 6039 0.9831  174    3 0.0545\n#       0.40 0.8432    0.0295 4484 5684 0.9774  173    4 0.0573\n#       0.45 0.8432    0.0306 4880 5288 0.9435  167   10 0.0593\n#       0.50 0.8432    0.0328 5306 4862 0.9322  165   12 0.0634\n#       0.55 0.8432    0.0353 5767 4401 0.9096  161   16 0.0679\n#       0.60 0.8432    0.0381 6209 3959 0.8870  157   20 0.0731\n#       0.65 0.8432    0.0426 6685 3483 0.8757  155   22 0.0813\n#       0.70 0.8432    0.0478 7198 2970 0.8418  149   28 0.0904\n#       0.75 0.8432    0.0544 7717 2451 0.7966  141   36 0.1018\n#       0.80 0.8432    0.0564 8227 1941 0.6554  116   61 0.1038\n#       0.85 0.8432    0.0677 8804 1364 0.5593   99   78 0.1207\n\n# Fold 2\n# Evaluation at Fixed Thresholds:\n#  Threshold    AUC Precision  TNs  FPs Recall  TPs  FNs     F1\n#       0.20 0.8970    0.0180  472 9705 1.0000  178    0 0.0354\n#       0.25 0.8970    0.0231 2698 7479 0.9944  177    1 0.0452\n#       0.30 0.8970    0.0314 4808 5369 0.9775  174    4 0.0608\n#       0.35 0.8970    0.0380 5820 4357 0.9663  172    6 0.0731\n#       0.40 0.8970    0.0436 6489 3688 0.9438  168   10 0.0833\n#       0.45 0.8970    0.0488 6978 3199 0.9213  164   14 0.0926\n#       0.50 0.8970    0.0548 7415 2762 0.8989  160   18 0.1032\n#       0.55 0.8970    0.0603 7762 2415 0.8708  155   23 0.1128\n#       0.60 0.8970    0.0676 8095 2082 0.8483  151   27 0.1253\n#       0.65 0.8970    0.0775 8438 1739 0.8202  146   32 0.1415\n#       0.70 0.8970    0.0885 8756 1421 0.7753  138   40 0.1589\n#       0.75 0.8970    0.1046 9090 1087 0.7135  127   51 0.1825\n#       0.80 0.8970    0.1211 9357  820 0.6348  113   65 0.2034\n#       0.85 0.8970    0.1431 9644  533 0.5000   89   89 0.2225\n\n\noptimal_thresholds = [0.60 ,0.60 , 0.55]\n\nmodels = [\n    load_model(path)  # No need for custom_objects\n    for path in model_paths\n]\n\n# Step 1: Predict\npredictions = []\nfor idx, model in enumerate(models):\n    print(f\"Making predictions with Model {idx}...\")\n    preds = model.predict(test_dataset, verbose=2)\n    predictions.append(preds.flatten())\n\npredictions = np.array(predictions)\navg_predictions = predictions.mean(axis=0)\naveraged_threshold = np.mean(optimal_thresholds)\nfinal_predictions = (avg_predictions >= averaged_threshold).astype(int)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T09:13:13.775426Z","iopub.execute_input":"2025-05-27T09:13:13.775685Z","iopub.status.idle":"2025-05-27T09:27:47.047451Z","shell.execute_reply.started":"2025-05-27T09:13:13.775664Z","shell.execute_reply":"2025-05-27T09:27:47.046122Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Step 2: Extract image names from test_dataset\nimage_names = []\n\nfor batch in test_dataset:\n    batch_image_names = batch['image_name']\n    decoded_names = [name.numpy().decode('utf-8') for name in batch_image_names]\n    image_names.extend(decoded_names)\n\n# Step 3: Create the submission DataFrame\nsubmission_df = pd.DataFrame({\n    'image_name': image_names,\n    'target': final_predictions\n})\n\n# Step 4: Save to CSV\nsubmission_df.to_csv('submission.csv', index=False)\nprint(\"✅ submission.csv created.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T09:27:47.049022Z","iopub.execute_input":"2025-05-27T09:27:47.049469Z","iopub.status.idle":"2025-05-27T09:27:52.426542Z","shell.execute_reply.started":"2025-05-27T09:27:47.049413Z","shell.execute_reply":"2025-05-27T09:27:52.425471Z"}},"outputs":[],"execution_count":null}]}