{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install pydicom\n!pip install matplotlib seaborn scikit-learn\n# !pip install tensorflow_addons","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport pydicom\nimport tensorflow as tf\nfrom tensorflow.keras import layers, models, optimizers, callbacks\nfrom tensorflow.keras.applications import EfficientNetB1\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import precision_score, recall_score, f1_score, confusion_matrix\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom datetime import datetime","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Print session info\nprint(f\"Training started at: 2025-06-05 16:52:20 UTC\")\nprint(f\"User: masoudshahrian\")\n\n# Create output directories\nos.makedirs('model_results', exist_ok=True)\nos.makedirs('saved_models', exist_ok=True)\nos.makedirs('logs', exist_ok=True)\n\n# Set random seeds\nnp.random.seed(42)\ntf.random.set_seed(42)\n\n# Enable mixed precision\ntf.keras.mixed_precision.set_global_policy('mixed_bfloat16')\n\n# Define dataset path\ntrain_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/'\n\ndef load_and_process_data(train_path):\n    print(\"Loading and processing data...\")\n    \n    train = pd.read_csv(os.path.join(train_path, 'train.csv'))\n    label = pd.read_csv(os.path.join(train_path, 'train_label_coordinates.csv'))\n    train_desc = pd.read_csv(os.path.join(train_path, 'train_series_descriptions.csv'))\n    \n    def reshape_row(row):\n        data = {'study_id': [], 'condition': [], 'level': [], 'severity': []}\n        for column, value in row.items():\n            if column != 'study_id':\n                parts = column.split('_')\n                condition = ' '.join([word.capitalize() for word in parts[:-2]])\n                level = parts[-2].capitalize() + '/' + parts[-1].capitalize()\n                data['study_id'].append(row['study_id'])\n                data['condition'].append(condition)\n                data['level'].append(level)\n                data['severity'].append(value)\n        return pd.DataFrame(data)\n    \n    new_train_df = pd.concat([reshape_row(row) for _, row in train.iterrows()], ignore_index=True)\n    merged_df = pd.merge(new_train_df, label, on=['study_id', 'condition', 'level'], how='inner')\n    final_merged_df = pd.merge(merged_df, train_desc, on=['series_id', 'study_id'], how='inner')\n    \n    final_merged_df['image_path'] = (\n        train_path + 'train_images/' +\n        final_merged_df['study_id'].astype(str) + '/' +\n        final_merged_df['series_id'].astype(str) + '/' +\n        final_merged_df['instance_number'].astype(str) + '.dcm'\n    )\n    \n    final_merged_df['severity'] = final_merged_df['severity'].map({\n        'Normal/Mild': 'normal_mild',\n        'Moderate': 'moderate',\n        'Severe': 'severe'\n    })\n    \n    valid_data = final_merged_df[final_merged_df['severity'].isin(['normal_mild', 'moderate', 'severe'])]\n    print(f\"Total valid samples: {len(valid_data)}\")\n    \n    return valid_data\n\n@tf.function\ndef load_dicom_tf(path):\n    try:\n        dicom = pydicom.dcmread(path.numpy().decode('utf-8'))\n        data = dicom.pixel_array\n        data = data - np.min(data)\n        if np.max(data) != 0:\n            data = data / np.max(data)\n        data = np.expand_dims(data, axis=-1)\n        return data.astype(np.float32)\n    except Exception as e:\n        print(f\"Error loading DICOM: {e}\")\n        return np.zeros((224, 224, 1), dtype=np.float32)\n\n@tf.function\ndef rotate_image(image, angle):\n    # Convert angle from radians to degrees\n    angle_deg = angle * 180.0 / np.pi\n    \n    # Rotate image using tf.image\n    rotated = tf.image.rot90(\n        image,\n        k=tf.cast(angle_deg / 90, tf.int32)\n    )\n    return rotated\n\n@tf.function\ndef preprocess_image(image_path, label=None, is_training=False):\n    # Load and preprocess image\n    image = tf.py_function(load_dicom_tf, [image_path], tf.float32)\n    image.set_shape([None, None, 1])\n    \n    # Resize image\n    image = tf.image.resize(image, [224, 224])\n    image = tf.image.grayscale_to_rgb(image)\n    \n    if is_training:\n        # Data augmentation for training\n        image = tf.image.random_flip_left_right(image)\n        image = tf.image.random_flip_up_down(image)\n        image = tf.image.random_brightness(image, 0.2)\n        image = tf.image.random_contrast(image, 0.8, 1.2)\n        \n        # Add random rotation\n        random_angle = tf.random.uniform([], -0.5, 0.5)  # Random angle between -0.5 and 0.5 radians\n        image = rotate_image(image, random_angle)\n    \n    # Preprocess for EfficientNet\n    image = tf.keras.applications.efficientnet.preprocess_input(image)\n    \n    return (image, label) if label is not None else image\n\ndef create_dataset(df, batch_size=32, is_training=False):\n    labels = df['severity'].map({'normal_mild': 0, 'moderate': 1, 'severe': 2}).astype(np.int32)\n    dataset = tf.data.Dataset.from_tensor_slices((df['image_path'], labels))\n    \n    dataset = dataset.map(\n        lambda x, y: preprocess_image(x, y, is_training),\n        num_parallel_calls=tf.data.AUTOTUNE\n    )\n    \n    if is_training:\n        dataset = dataset.shuffle(1000)\n        dataset = dataset.repeat()\n    \n    dataset = dataset.batch(batch_size)\n    dataset = dataset.prefetch(tf.data.AUTOTUNE)\n    \n    return dataset\n\ndef build_hybrid_model(input_shape=(224, 224, 3), num_classes=3):\n    efficientnet = EfficientNetB1(\n        include_top=False,\n        weights='imagenet',\n        input_shape=input_shape\n    )\n    \n    for layer in efficientnet.layers[:100]:\n        layer.trainable = False\n    \n    inputs = layers.Input(shape=input_shape)\n    x = efficientnet(inputs)\n    \n    attention = layers.Conv2D(x.shape[-1], 1, activation='sigmoid')(x)\n    x = layers.Multiply()([x, attention])\n    \n    x = layers.GlobalAveragePooling2D()(x)\n    x = layers.Dense(512, activation='relu')(x)\n    x = layers.BatchNormalization()(x)\n    x = layers.Dropout(0.5)(x)\n    x = layers.Dense(256, activation='relu')(x)\n    x = layers.BatchNormalization()(x)\n    x = layers.Dropout(0.3)(x)\n    outputs = layers.Dense(num_classes)(x)\n    \n    return models.Model(inputs, outputs)\n\ndef focal_loss(gamma=2., alpha=4.):\n    def focal_loss_with_logits(y_true, y_pred):\n        y_true = tf.cast(y_true, tf.int32)\n        ce = tf.nn.sparse_softmax_cross_entropy_with_logits(labels=y_true, logits=y_pred)\n        probs = tf.nn.softmax(y_pred)\n        probs = tf.gather(probs, y_true, batch_dims=1)\n        return tf.reduce_mean(alpha * tf.pow(1. - probs, gamma) * ce)\n    return focal_loss_with_logits\n\ndef train_model(model, train_ds, val_ds, series_name, epochs=100, steps_per_epoch=100):\n    optimizer = optimizers.Adam(learning_rate=1e-4)\n    \n    model.compile(\n        optimizer=optimizer,\n        loss=focal_loss(),\n        metrics=['accuracy']\n    )\n    \n    log_dir = f\"logs/{series_name}_{datetime.utcnow().strftime('%Y%m%d-%H%M%S')}\"\n    \n    callbacks_list = [\n        callbacks.EarlyStopping(\n            monitor='val_loss',\n            patience=30,\n            restore_best_weights=True\n        ),\n        callbacks.ReduceLROnPlateau(\n            monitor='val_loss',\n            factor=0.5,\n            patience=5,\n            min_lr=1e-6\n        ),\n        callbacks.TensorBoard(\n            log_dir=log_dir,\n            histogram_freq=1\n        )\n    ]\n    \n    history = model.fit(\n        train_ds,\n        validation_data=val_ds,\n        epochs=epochs,\n        steps_per_epoch=steps_per_epoch,\n        validation_steps=30,\n        callbacks=callbacks_list,\n        verbose=1\n    )\n    \n    return history\n\ndef evaluate_model(model, test_ds, class_names):\n    y_true = []\n    y_pred = []\n    y_pred_probs = []\n    \n    for images, labels in test_ds:\n        predictions = model.predict(images)\n        y_pred_probs.extend(tf.nn.softmax(predictions).numpy())\n        y_pred.extend(np.argmax(predictions, axis=1))\n        y_true.extend(labels.numpy())\n    \n    y_true = np.array(y_true)\n    y_pred = np.array(y_pred)\n    y_pred_probs = np.array(y_pred_probs)\n    \n    cm = confusion_matrix(y_true, y_pred)\n    \n    class_metrics = {}\n    for i, class_name in enumerate(class_names):\n        true_class = (y_true == i)\n        pred_class = (y_pred == i)\n        \n        tp = np.sum((true_class) & (pred_class))\n        fp = np.sum((~true_class) & (pred_class))\n        fn = np.sum((true_class) & (~pred_class))\n        tn = np.sum((~true_class) & (~pred_class))\n        \n        sensitivity = tp / (tp + fn) if (tp + fn) > 0 else 0\n        specificity = tn / (tn + fp) if (tn + fp) > 0 else 0\n        precision = tp / (tp + fp) if (tp + fp) > 0 else 0\n        f1 = 2 * (precision * sensitivity) / (precision + sensitivity) if (precision + sensitivity) > 0 else 0\n        \n        class_metrics[class_name] = {\n            'TP': int(tp),\n            'FP': int(fp),\n            'FN': int(fn),\n            'TN': int(tn),\n            'Sensitivity': sensitivity,\n            'Specificity': specificity,\n            'Precision': precision,\n            'F1-Score': f1\n        }\n    \n    overall_metrics = {\n        'accuracy': np.mean(y_pred == y_true),\n        'weighted_precision': precision_score(y_true, y_pred, average='weighted'),\n        'weighted_recall': recall_score(y_true, y_pred, average='weighted'),\n        'weighted_f1': f1_score(y_true, y_pred, average='weighted')\n    }\n    \n    return {\n        'confusion_matrix': cm,\n        'class_metrics': class_metrics,\n        'overall_metrics': overall_metrics,\n        'predictions': y_pred_probs\n    }\n\ndef plot_results(history, results, class_names, series_name):\n    output_dir = 'model_results'\n    os.makedirs(output_dir, exist_ok=True)\n    \n    fig = plt.figure(figsize=(20, 15))\n    \n    # Training History - Loss\n    plt.subplot(2, 2, 1)\n    plt.plot(history.history['loss'], label='Train Loss')\n    plt.plot(history.history['val_loss'], label='Val Loss')\n    plt.title(f'{series_name} - Training Loss', pad=20)\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.legend()\n    plt.grid(True)\n    \n    # Training History - Accuracy\n    plt.subplot(2, 2, 2)\n    plt.plot(history.history['accuracy'], label='Train Acc')\n    plt.plot(history.history['val_accuracy'], label='Val Acc')\n    plt.title(f'{series_name} - Training Accuracy', pad=20)\n    plt.xlabel('Epoch')\n    plt.ylabel('Accuracy')\n    plt.legend()\n    plt.grid(True)\n    \n    # Confusion Matrix with percentages\n    plt.subplot(2, 2, 3)\n    cm = results['confusion_matrix']\n    cm_sum = np.sum(cm, axis=1, keepdims=True)\n    cm_perc = cm / cm_sum * 100\n    \n    sns.heatmap(cm, annot=np.array([[f'{int(x)}\\n({y:.1f}%)' for x, y in zip(row_true, row_perc)] \n                                   for row_true, row_perc in zip(cm, cm_perc)]),\n                fmt='', cmap='Blues', xticklabels=class_names, yticklabels=class_names)\n    plt.title(f'{series_name} - Confusion Matrix\\nwith Counts and Percentages', pad=20)\n    plt.xlabel('Predicted')\n    plt.ylabel('True')\n    \n    # Metrics Table\n    plt.subplot(2, 2, 4)\n    plt.axis('off')\n    \n    table_data = []\n    table_colors = []\n    metrics_display = ['Sensitivity', 'Specificity', 'Precision', 'F1-Score']\n    \n    header = ['Metrics'] + class_names + ['Overall']\n    table_data.append(header)\n    table_colors.append(['lightgray'] * len(header))\n    \n    for metric in metrics_display:\n        row = [metric]\n        for class_name in class_names:\n            value = results['class_metrics'][class_name][metric]\n            row.append(f'{value:.3f}')\n        if metric == 'Precision':\n            row.append(f'{results[\"overall_metrics\"][\"weighted_precision\"]:.3f}')\n        elif metric == 'Sensitivity':\n            row.append(f'{results[\"overall_metrics\"][\"weighted_recall\"]:.3f}')\n        elif metric == 'F1-Score':\n            row.append(f'{results[\"overall_metrics\"][\"weighted_f1\"]:.3f}')\n        else:\n            row.append('-')\n        table_data.append(row)\n        table_colors.append(['white'] * len(header))\n    \n    acc_row = ['Accuracy'] + ['-'] * len(class_names) + [f'{results[\"overall_metrics\"][\"accuracy\"]:.3f}']\n    table_data.append(acc_row)\n    table_colors.append(['white'] * len(header))\n    \n    table = plt.table(cellText=table_data,\n                     cellColours=table_colors,\n                     cellLoc='center',\n                     loc='center',\n                     bbox=[0.1, 0.1, 0.8, 0.8])\n    \n    table.auto_set_font_size(False)\n    table.set_fontsize(9)\n    table.scale(1.2, 1.5)\n    \n    plt.title(f'{series_name} - Performance Metrics', pad=20)\n    \n    plt.tight_layout(h_pad=1.0, w_pad=1.0)\n    safe_series_name = series_name.lower().replace(\"/\", \"_\").replace(\" \", \"_\")\n    output_file = os.path.join(output_dir, f'{safe_series_name}_detailed_results.png')\n    plt.savefig(output_file, dpi=300, bbox_inches='tight')\n    plt.close()\n    \n    # Print detailed metrics\n    print(f\"\\nDetailed Results for {series_name}:\")\n    print(\"=\" * 50)\n    print(\"Overall Metrics:\")\n    for metric, value in results['overall_metrics'].items():\n        print(f\"{metric}: {value:.4f}\")\n    \n    print(\"\\nPer-Class Metrics:\")\n    for class_name in class_names:\n        print(f\"\\n{class_name}:\")\n        metrics = results['class_metrics'][class_name]\n        print(f\"TP: {metrics['TP']}, FP: {metrics['FP']}, FN: {metrics['FN']}, TN: {metrics['TN']}\")\n        print(f\"Sensitivity: {metrics['Sensitivity']:.4f}\")\n        print(f\"Specificity: {metrics['Specificity']:.4f}\")\n        print(f\"Precision: {metrics['Precision']:.4f}\")\n        print(f\"F1-Score: {metrics['F1-Score']:.4f}\")\n\ndef main():\n    final_merged_df = load_and_process_data(train_path)\n    \n    series_types = ['Sagittal T1', 'Axial T2', 'Sagittal T2/STIR']\n    class_names = ['normal_mild', 'moderate', 'severe']\n    \n    for series_name in series_types:\n        print(f\"\\nProcessing {series_name}\")\n        \n        series_df = final_merged_df[final_merged_df['series_description'] == series_name].copy()\n        if series_df.empty:\n            print(f\"No data for {series_name}\")\n            continue\n        \n        train_df, temp_df = train_test_split(series_df, test_size=0.3, stratify=series_df['severity'])\n        val_df, test_df = train_test_split(temp_df, test_size=0.5, stratify=temp_df['severity'])\n        \n        print(f\"\\nDataset splits for {series_name}:\")\n        print(f\"Train samples: {len(train_df)}\")\n        print(f\"Validation samples: {len(val_df)}\")\n        print(f\"Test samples: {len(test_df)}\")\n        \n        train_ds = create_dataset(train_df, batch_size=32, is_training=True)\n        val_ds = create_dataset(val_df, batch_size=32)\n        test_ds = create_dataset(test_df, batch_size=32)\n        \n        model = build_hybrid_model()\n        history = train_model(model, train_ds, val_ds, series_name)\n        \n        results = evaluate_model(model, test_ds, class_names)\n        plot_results(history, results, class_names, series_name)\n        \n        safe_series_name = series_name.lower().replace(\"/\", \"_\").replace(\" \", \"_\")\n        model_file = os.path.join('saved_models', f'{safe_series_name}_model.h5')\n        model.save(model_file)\n        \n        print(f\"Completed {series_name} at: {datetime.utcnow().strftime('%Y-%m-%d %H:%M:%S')} UTC\")\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# cause of low efficiency commented for improving with sonet3.5\n# !pip install pydicom\n# !pip install matplotlib seaborn scikit-learn\n# import os\n# import numpy as np\n# import pandas as pd\n# import cv2\n# import pydicom\n# import tensorflow as tf\n# from tensorflow.keras import layers, models, optimizers, callbacks\n# from tensorflow.keras.applications import EfficientNetB1\n# from sklearn.model_selection import train_test_split\n# from sklearn.metrics import precision_score, recall_score, f1_score, confusion_matrix\n# from sklearn.utils.class_weight import compute_class_weight\n# import matplotlib.pyplot as plt\n# import seaborn as sns\n# # Set up mixed precision\n# tf.keras.mixed_precision.set_global_policy('mixed_bfloat16')\n\n# # Define dataset path\n# train_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/'\n\n# # Load CSV files\n# train = pd.read_csv(os.path.join(train_path, 'train.csv'))\n# label = pd.read_csv(os.path.join(train_path, 'train_label_coordinates.csv'))\n# train_desc = pd.read_csv(os.path.join(train_path, 'train_series_descriptions.csv'))\n\n# # Reshape train.csv\n# def reshape_row(row):\n#     data = {'study_id': [], 'condition': [], 'level': [], 'severity': []}\n#     for column, value in row.items():\n#         if column != 'study_id':\n#             parts = column.split('_')\n#             condition = ' '.join([word.capitalize() for word in parts[:-2]])\n#             level = parts[-2].capitalize() + '/' + parts[-1].capitalize()\n#             data['study_id'].append(row['study_id'])\n#             data['condition'].append(condition)\n#             data['level'].append(level)\n#             data['severity'].append(value)\n#     return pd.DataFrame(data)\n\n# # Create and merge DataFrames\n# new_train_df = pd.concat([reshape_row(row) for _, row in train.iterrows()], ignore_index=True)\n# merged_df = pd.merge(new_train_df, label, on=['study_id', 'condition', 'level'], how='inner')\n# final_merged_df = pd.merge(merged_df, train_desc, on=['series_id', 'study_id'], how='inner')\n\n# # Create image paths\n# final_merged_df['image_path'] = (\n#     train_path + 'train_images/' +\n#     final_merged_df['study_id'].astype(str) + '/' +\n#     final_merged_df['series_id'].astype(str) + '/' +\n#     final_merged_df['instance_number'].astype(str) + '.dcm'\n# )\n\n# # Map severity labels\n# final_merged_df['severity'] = final_merged_df['severity'].map({\n#     'Normal/Mild': 'normal_mild',\n#     'Moderate': 'moderate',\n#     'Severe': 'severe'\n# })\n\n# # Filter invalid rows\n# final_merged_df = final_merged_df[final_merged_df['severity'].isin(['normal_mild', 'moderate', 'severe'])]\n\n# # Compute class weights\n# class_counts = final_merged_df['severity'].value_counts().sort_index().values\n# class_weights = compute_class_weight('balanced', classes=np.array([0, 1, 2]), \n#                                     y=final_merged_df['severity'].map({'normal_mild': 0, 'moderate': 1, 'severe': 2}))\n# class_weights_dict = dict(enumerate(class_weights))\n\n# # Define Focal Loss\n# def focal_loss(y_true, y_pred, alpha=list(class_weights_dict.values()), gamma=2.0):\n#     y_true = tf.cast(y_true, tf.int32)\n#     ce = tf.nn.sparse_softmax_cross_entropy_with_logits(labels=y_true, logits=y_pred)\n#     probs = tf.nn.softmax(y_pred, axis=-1)\n#     probs = tf.gather(probs, y_true, batch_dims=1)\n#     alpha = tf.gather(alpha, y_true)\n#     modulating_factor = tf.pow(1.0 - probs, gamma)\n#     return tf.reduce_mean(alpha * modulating_factor * ce)\n\n# # Image preprocessing functions\n# def apply_clahe(image, clip_limit=2, tile_grid_size=(16, 16)):\n#     if len(image.shape) == 3 and image.shape[2] == 3:\n#         image = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)\n#     if image.dtype != np.uint8:\n#         image = (image * 255).astype(np.uint8) if np.issubdtype(image.dtype, np.floating) else image.astype(np.uint8)\n#     clahe = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=tile_grid_size)\n#     return clahe.apply(image)\n\n# def apply_bilateral_filter(image, diameter=5, sigma_color=10, sigma_space=10):\n#     if image.dtype == np.float32 or image.dtype == np.float64:\n#         if image.max() <= 1.0:\n#             sigma_color = sigma_color / 255.0\n#         else:\n#             image = (image / 255.0).astype(np.float32)\n#     elif image.dtype != np.uint8:\n#         image = image.astype(np.uint8)\n#     return cv2.bilateralFilter(image, diameter, sigma_color, sigma_space)\n\n# def augment_image(image):\n#     image = tf.image.random_brightness(image, max_delta=5/255.0)\n#     image = tf.image.random_contrast(image, lower=0.9, upper=1.1)\n#     image_uint8 = tf.image.convert_image_dtype(image, tf.uint8)\n#     image_uint8 = tf.image.random_jpeg_quality(image_uint8, 80, 100)\n#     return tf.image.convert_image_dtype(image_uint8, tf.float32)\n\n# def remove_background(image):\n#     if len(image.shape) == 3 and image.shape[2] == 3:\n#         image = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)\n#     if image.dtype != np.uint8:\n#         image = (image * 255).astype(np.uint8) if image.dtype == np.float32 else image.astype(np.uint8)\n#     _, mask = cv2.threshold(image, 0, 255, cv2.THRESH_BINARY + cv2.THRESH_OTSU)\n#     return cv2.bitwise_and(image, image, mask=mask)\n\n# def load_dicom_tf(path):\n#     dicom = pydicom.dcmread(path.numpy().decode('utf-8'))\n#     data = dicom.pixel_array\n#     data = data - np.min(data)\n#     if np.max(data) != 0:\n#         data = data / np.max(data)\n#     return data.astype(np.float32)\n\n# def preprocess_image(image, label=None, is_training=False):\n#     def process_with_opencv(img_tensor):\n#         img_np = img_tensor.numpy().squeeze()\n#         img_np = apply_bilateral_filter(img_np)\n#         img_np = apply_clahe(img_np)\n#         img_np = remove_background(img_np)\n#         if is_training:\n#             img_np = augment_image(tf.expand_dims(img_np, axis=-1)).numpy().squeeze()\n#         img_np = np.expand_dims(img_np, axis=-1)\n#         return img_np.astype(np.float32)\n\n#     image = tf.py_function(load_dicom_tf, [image], tf.float32)\n#     image.set_shape([None, None])\n#     image = tf.py_function(process_with_opencv, [image], tf.float32)\n#     image.set_shape([None, None, 1])\n#     image = tf.image.resize(image, [224, 224])\n#     image = tf.image.grayscale_to_rgb(image)\n#     image = tf.keras.applications.efficientnet.preprocess_input(image)\n#     return (image, label) if label is not None else image\n\n# # Create dataset\n# def create_dataset(df, batch_size=32, is_training=False):\n#     labels = df['severity'].map({'normal_mild': 0, 'moderate': 1, 'severe': 2}).astype(np.int32)\n#     dataset = tf.data.Dataset.from_tensor_slices((df['image_path'], labels))\n#     dataset = dataset.map(lambda x, y: preprocess_image(x, y, is_training), num_parallel_calls=tf.data.AUTOTUNE)\n    \n#     if is_training:\n#         dataset = dataset.apply(tf.data.experimental.rejection_resample(\n#             class_func=lambda image, label: label,\n#             target_dist=list(class_weights_dict.values()),\n#             initial_dist=list(class_weights_dict.values())\n#         )).map(lambda resampled_label, original_sample: original_sample)\n    \n#     dataset = dataset.batch(batch_size).prefetch(tf.data.AUTOTUNE)\n#     return dataset\n\n# # Build U-Net encoder\n# def build_unet_encoder(input_shape=(224, 224, 3)):\n#     inputs = layers.Input(shape=input_shape)\n#     c1 = layers.Conv2D(64, 3, padding='same', activation='relu')(inputs)\n#     c1 = layers.Conv2D(64, 3, padding='same', activation='relu')(c1)\n#     p1 = layers.MaxPooling2D((2, 2))(c1)\n    \n#     c2 = layers.Conv2D(128, 3, padding='same', activation='relu')(p1)\n#     c2 = layers.Conv2D(128, 3, padding='same', activation='relu')(c2)\n#     p2 = layers.MaxPooling2D((2, 2))(c2)\n    \n#     c3 = layers.Conv2D(256, 3, padding='same', activation='relu')(p2)\n#     c3 = layers.Conv2D(256, 3, padding='same', activation='relu')(c3)\n#     p3 = layers.MaxPooling2D((2, 2))(c3)\n    \n#     c4 = layers.Conv2D(512, 3, padding='same', activation='relu')(p3)\n#     c4 = layers.Conv2D(512, 3, padding='same', activation='relu')(c4)\n#     p4 = layers.MaxPooling2D((2, 2))(c4)\n    \n#     c5 = layers.Conv2D(1024, 3, padding='same', activation='relu')(p4)\n#     c5 = layers.Conv2D(1024, 3, padding='same', activation='relu')(c5)\n    \n#     return models.Model(inputs, c5)\n\n# # Build hybrid U-Net + EfficientNetB1 model\n# def build_hybrid_model(input_shape=(224, 224, 3), num_classes=3):\n#     unet_encoder = build_unet_encoder(input_shape)\n#     efficientnet = EfficientNetB1(\n#         include_top=False,\n#         weights='imagenet',\n#         input_shape=input_shape\n#     )\n    \n#     for layer in efficientnet.layers:\n#         if 'block6' in layer.name or 'block7' in layer.name:\n#             layer.trainable = True\n#         else:\n#             layer.trainable = False\n    \n#     inputs = layers.Input(shape=input_shape)\n#     unet_features = unet_encoder(inputs)\n#     eff_features = efficientnet(inputs)\n    \n#     unet_pooled = layers.GlobalAveragePooling2D()(unet_features)\n#     eff_pooled = layers.GlobalAveragePooling2D()(eff_features)\n#     combined = layers.Concatenate()([unet_pooled, eff_pooled])\n    \n#     x = layers.Dense(512, activation='relu')(combined)\n#     x = layers.Dropout(0.5)(x)\n#     outputs = layers.Dense(num_classes, activation='linear')(x)\n    \n#     return models.Model(inputs, outputs)\n\n# # Train model\n# def train_model(model, train_ds, val_ds, series_name):\n#     model.compile(optimizer=optimizers.Adam(learning_rate=0.0001), \n#                   loss=focal_loss, \n#                   metrics=['accuracy'])\n#     callbacks_list = [\n#         callbacks.EarlyStopping(patience=10, restore_best_weights=True),\n#         callbacks.ModelCheckpoint(f'best_{series_name}_hybrid.keras', save_best_only=True)\n#     ]\n#     history = model.fit(train_ds, validation_data=val_ds, epochs=50, callbacks=callbacks_list)\n#     return history\n\n# # Evaluate model\n# def evaluate_model(model, test_ds, class_names):\n#     y_true, y_pred = [], []\n#     for images, labels in test_ds:\n#         y_true.extend(labels.numpy())\n#         preds = model.predict(images)\n#         y_pred.extend(np.argmax(preds, axis=1))\n#     y_true, y_pred = np.array(y_true), np.array(y_pred)\n#     cm = confusion_matrix(y_true, y_pred)\n    \n#     cm_details = {}\n#     for i, class_name in enumerate(class_names):\n#         tp = cm[i, i]\n#         fp = cm[:, i].sum() - tp\n#         fn = cm[i, :].sum() - tp\n#         tn = cm.sum() - (tp + fp + fn)\n#         cm_details[class_name] = {\n#             'True Positives': int(tp),\n#             'False Positives': int(fp),\n#             'False Negatives': int(fn),\n#             'True Negatives': int(tn)\n#         }\n    \n#     return {\n#         'precision': precision_score(y_true, y_pred, average='weighted', zero_division=0),\n#         'recall': recall_score(y_true, y_pred, average='weighted', zero_division=0),\n#         'f1': f1_score(y_true, y_pred, average='weighted', zero_division=0),\n#         'cm': cm,\n#         'cm_details': cm_details\n#     }\n\n# # Plot confusion matrix\n# def plot_confusion_matrix(cm, classes, title):\n#     plt.figure(figsize=(10, 8))\n#     sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', \n#                 xticklabels=classes, yticklabels=classes,\n#                 annot_kws={\"size\": 12}, cbar=True)\n#     plt.title(title, fontsize=14, pad=15)\n#     plt.xlabel('Predicted Label', fontsize=12)\n#     plt.ylabel('True Label', fontsize=12)\n#     plt.xticks(rotation=45, ha='right')\n#     plt.yticks(rotation=0)\n#     plt.tight_layout()\n    \n#     # Save the plot\n#     file_path = f'confusion_matrix_{title.lower().replace(\" \", \"_\")}.png'\n#     directory = os.path.dirname(file_path)\n#     if directory and not os.path.exists(directory):\n#         os.makedirs(directory)\n#     plt.savefig(file_path, dpi=300, bbox_inches='tight')\n#     plt.show()  # Display the plot\n#     plt.close()\n\n# # Main loop for each series\n# series_types = ['Sagittal T1', 'Axial T2', 'Sagittal T2/STIR']\n# class_names = ['normal_mild', 'moderate', 'severe']\n\n# for series_name in series_types:\n#     series_df = final_merged_df[final_merged_df['series_description'] == series_name].copy()\n#     if series_df.empty:\n#         print(f\"No valid data for {series_name}.\")\n#         continue\n    \n#     # Split data\n#     train_df, temp_df = train_test_split(series_df, test_size=0.3, stratify=series_df['severity'], random_state=42)\n#     val_df, test_df = train_test_split(temp_df, test_size=0.6667, stratify=temp_df['severity'], random_state=42)\n    \n#     print(f\"\\nClass distribution for {series_name}:\")\n#     print(\"Train:\", train_df['severity'].value_counts().to_dict())\n#     print(\"Validation:\", val_df['severity'].value_counts().to_dict())\n#     print(\"Test:\", test_df['severity'].value_counts().to_dict())\n    \n#     # Create datasets\n#     train_ds = create_dataset(train_df, batch_size=32, is_training=True)\n#     val_ds = create_dataset(val_df, batch_size=32)\n#     test_ds = create_dataset(test_df, batch_size=32)\n    \n#     # Build and train hybrid model\n#     model = build_hybrid_model()\n#     history = train_model(model, train_ds, val_ds, series_name)\n    \n#     # Evaluate\n#     results = evaluate_model(model, test_ds, class_names)\n#     print(f\"\\nMetrics for {series_name}:\")\n#     print(f\"Precision: {results['precision']:.4f}\")\n#     print(f\"Recall: {results['recall']:.4f}\")\n#     print(f\"F1-Score: {results['f1']:.4f}\")\n    \n#     # Print confusion matrix details\n#     print(f\"\\nConfusion Matrix Details for {series_name}:\")\n#     for class_name, metrics in results['cm_details'].items():\n#         print(f\"\\nClass: {class_name}\")\n#         for metric_name, value in metrics.items():\n#             print(f\"{metric_name}: {value}\")\n            \n#     # Plot confusion matrix\n#     plot_confusion_matrix(results['cm'], class_names, f'{series_name} Confusion Matrix')\n    \n#     # Plot training metrics\n#     plt.figure(figsize=(12, 5))\n#     plt.subplot(1, 2, 1)\n#     plt.plot(history.history['loss'], label='Train Loss')\n#     plt.plot(history.history['val_loss'], label='Val Loss')\n#     plt.title(f'{series_name} Loss')\n#     plt.xlabel('Epoch')\n#     plt.ylabel('Loss')\n#     plt.legend()\n#     plt.subplot(1, 2, 2)\n#     plt.plot(history.history['accuracy'], label='Train Acc')\n#     plt.plot(history.history['val_accuracy'], label='Val Acc')\n#     plt.title(f'{series_name} Accuracy')\n#     plt.xlabel('Epoch')\n#     plt.ylabel('Accuracy')\n#     plt.legend()\n#     file_path = f'{series_name.lower().replace(\" \", \"_\")}_plots.png'\n#     directory = os.path.dirname(file_path)\n#     if directory and not os.path.exists(directory):\n#         os.makedirs(directory)\n#     plt.savefig(file_path, dpi=300, bbox_inches='tight')\n#     plt.show()\n#     plt.close()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}