{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":10988536,"sourceType":"datasetVersion","datasetId":6839310}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# !pip install -q kaggle\n# !mkdir -p ~/.kaggle\n# !cp kaggle.json ~/.kaggle/\n# !chmod 600 ~/.kaggle/kaggle.json\n# !kaggle competitions download -c rsna-2024-lumbar-spine-degenerative-classification\n# !unzip -qq /content/rsna-2024-lumbar-spine-degenerative-classification.zip\n\nimport cv2\nimport os\nimport time\nimport json\nimport glob\nimport random\nimport collections\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom tqdm import tqdm\nfrom copy import deepcopy\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import precision_score, recall_score, f1_score, confusion_matrix\nimport joblib\nimport tensorflow as tf\nfrom tensorflow.keras import layers, models, applications, optimizers, callbacks\nfrom tensorflow.keras import backend as K\nfrom sklearn.utils.class_weight import compute_class_weight","metadata":{"trusted":true,"execution":{"execution_failed":"2025-03-12T22:00:03.106Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define the path to the training data\ntrain_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/'\ntrain_path2 = '/kaggle/input/lumbar-2024-imagesdicom2jpg/'  # مسیر فرضی برای دیتاست JPG\n\n# Load CSV files containing metadata and labels\ntrain = pd.read_csv(os.path.join(train_path, 'train.csv'))\nlabel = pd.read_csv(os.path.join(train_path, 'train_label_coordinates.csv'))\ntrain_desc = pd.read_csv(os.path.join(train_path, 'train_series_descriptions.csv'))\ntest_desc = pd.read_csv(os.path.join(train_path, 'test_series_descriptions.csv'))\nsub = pd.read_csv(os.path.join(train_path, 'sample_submission.csv'))\n\n# Display the first few rows of each dataframe\nprint(\"Test Descriptions:\")\nprint(test_desc.head(5))\n\nprint(\"\\nTrain Data:\")\nprint(train.head(5))\n\nprint(\"\\nTrain Series Descriptions:\")\nprint(train_desc.head(5))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Generate Image Paths (تغییر به JPG)\n# ---------------------\ndef generate_image_paths(df, data_dir):\n    image_paths = []\n    for study_id, series_id in zip(df['study_id'], df['series_id']):\n        study_dir = os.path.join(data_dir, str(study_id))\n        series_dir = os.path.join(study_dir, str(series_id))\n        images = sorted([f for f in os.listdir(series_dir) if f.endswith('.jpg')])\n        image_paths.extend([os.path.join(series_dir, img) for img in images])\n    return image_paths\n\n# Generate image paths for training and testing datasets\ntrain_image_paths = generate_image_paths(train_desc, os.path.join(train_path2, 'train_images'))\ntest_image_paths = generate_image_paths(test_desc, os.path.join(train_path2, 'test_images'))\n\n# Example usage\nprint(\"\\nSample Train Image Path:\", train_image_paths[2])\nprint(\"Number of Train Descriptions:\", len(train_desc))\nprint(\"Number of Train Image Paths:\", len(train_image_paths))\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Display JPG Images\n# ---------------------\ndef display_jpg_images(image_paths, num_images=3):\n    plt.figure(figsize=(15, 5))\n    for i, path in enumerate(image_paths[:num_images]):\n        img = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n        plt.subplot(1, num_images, i+1)\n        plt.imshow(img, cmap=plt.cm.bone)\n        plt.title(f\"Image {i+1}\")\n        plt.axis('off')\n    plt.show()\n\n# Display the first three JPG images from training data\ndisplay_jpg_images(train_image_paths)\n\n# Display JPG Images with Coordinates\n# -------------------------------------\ndef display_jpg_with_coordinates(image_paths, label_df):\n    fig, axs = plt.subplots(1, len(image_paths), figsize=(18, 6))\n    for idx, path in enumerate(image_paths):\n        study_id = int(path.split('/')[-3])\n        series_id = int(path.split('/')[-2])\n        filtered_labels = label_df[\n            (label_df['study_id'] == study_id) &\n            (label_df['series_id'] == series_id)\n        ]\n        img = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n        axs[idx].imshow(img, cmap='gray')\n        axs[idx].set_title(f\"Study ID: {study_id}, Series ID: {series_id}\")\n        axs[idx].axis('off')\n        for _, row in filtered_labels.iterrows():\n            axs[idx].plot(row['x'], row['y'], 'ro', markersize=5)\n    plt.tight_layout()\n    plt.show()\n\ndef load_jpg_files(path_to_folder):\n    files = [os.path.join(path_to_folder, f) for f in os.listdir(path_to_folder) if f.endswith('.jpg')]\n    files.sort(key=lambda x: int(os.path.splitext(os.path.basename(x))[0].split('-')[-1]))\n    return files\n\n# Example: Display JPG images with coordinates for a specific study\nstudy_id = \"100206310\"\nstudy_folder = os.path.join(train_path2, 'train_images', study_id)\nimage_paths = []\nfor series_folder in os.listdir(study_folder):\n    series_folder_path = os.path.join(study_folder, series_folder)\n    jpg_files = load_jpg_files(series_folder_path)\n    if jpg_files:\n        image_paths.append(jpg_files[0])\ndisplay_jpg_with_coordinates(image_paths, label)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Data Reshaping and Merging\n# --------------------------\ndef reshape_row(row):\n    data = {'study_id': [], 'condition': [], 'level': [], 'severity': []}\n    for column, value in row.items():\n        if column not in ['study_id', 'series_id', 'instance_number', 'x', 'y', 'series_description']:\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\nnew_train_df = pd.concat([reshape_row(row) for _, row in train.iterrows()], ignore_index=True)\nprint(\"\\nReshaped Train Data:\")\nprint(new_train_df.head(5))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Merge DataFrames\n# -----------------\nmerged_df = pd.merge(new_train_df, label, on=['study_id', 'condition', 'level'], how='inner')\nfinal_merged_df = pd.merge(merged_df, train_desc, on=['series_id', 'study_id'], how='inner')\nfinal_merged_df['row_id'] = (\n    final_merged_df['study_id'].astype(str) + '_' +\n    final_merged_df['condition'].str.lower().str.replace(' ', '_') + '_' +\n    final_merged_df['level'].str.lower().str.replace('/', '_')\n)\nfinal_merged_df['image_path'] = (\n    os.path.join(train_path2, '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) + '.jpg'\n)\nprint(\"\\nUpdated Final Merged DataFrame:\")\nprint(final_merged_df.head(5))\n\nnormal_mild_count = final_merged_df[final_merged_df[\"severity\"] == \"Normal/Mild\"].shape[0]\nmoderate_count = final_merged_df[final_merged_df[\"severity\"] == \"Moderate\"].shape[0]\nsevere_count = final_merged_df[final_merged_df[\"severity\"] == \"Severe\"].shape[0]\nprint(f\"\\nNormal/Mild Count: {normal_mild_count}\")\nprint(f\"Moderate Count: {moderate_count}\")\nprint(f\"Severe Count: {severe_count}\")\n\nbase_path = '/content/test_images/'\n\ndef get_image_paths(row):\n    series_path = os.path.join(base_path, str(row['study_id']), str(row['series_id']))\n    if os.path.exists(series_path):\n        return [os.path.join(series_path, f) for f in os.listdir(series_path) if f.endswith('.jpg')]\n    return []\n\ncondition_mapping = {\n    'Sagittal T1': {'left': 'left_neural_foraminal_narrowing', 'right': 'right_neural_foraminal_narrowing'},\n    'Axial T2': {'left': 'left_subarticular_stenosis', 'right': 'right_subarticular_stenosis'},\n    'Sagittal T2/STIR': 'spinal_canal_stenosis'\n}\n\nexpanded_rows = []\nfor index, row in test_desc.iterrows():\n    image_paths = get_image_paths(row)\n    conditions = condition_mapping.get(row['series_description'], {})\n    if isinstance(conditions, str):\n        conditions = {'left': conditions, 'right': conditions}\n    for side, condition in conditions.items():\n        for image_path in image_paths:\n            expanded_rows.append({\n                'study_id': row['study_id'],\n                'series_id': row['series_id'],\n                'series_description': row['series_description'],\n                'image_path': image_path,\n                'condition': condition,\n                'row_id': f\"{row['study_id']}_{condition}\"\n            })\n\nexpanded_test_desc = pd.DataFrame(expanded_rows)\nprint(\"\\nExpanded Test Descriptions:\")\nprint(expanded_test_desc.head(5))\n\nfinal_merged_df['severity'] = final_merged_df['severity'].map({\n    'Normal/Mild': 'normal_mild',\n    'Moderate': 'moderate',\n    'Severe': 'severe'\n})\n\n\n\n\n\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data = final_merged_df\ntest_data = expanded_test_desc\n\nprint(\"\\nSample Train Data:\")\nprint(train_data.head(10))\nprint(\"\\nSample Test Data:\")\nprint(test_data.head(10))\nprint(\"\\nTrain Data Shape:\", train_data.shape)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Verify File Paths\n# -----------------\ndef check_exists(path):\n    return os.path.exists(path)\n\ndef check_study_id(row):\n    study_id = row['study_id']\n    path = os.path.join(train_path2, 'train_images', str(study_id))\n    return check_exists(path)\n\ndef check_series_id(row):\n    study_id = row['study_id']\n    series_id = row['series_id']\n    path = os.path.join(train_path2, 'train_images', str(study_id), str(series_id))\n    return check_exists(path)\n\ndef check_image_exists(row):\n    image_path = row['image_path']\n    return check_exists(image_path)\n\ntrain_data['study_id_exists'] = train_data.apply(check_study_id, axis=1)\ntrain_data['series_id_exists'] = train_data.apply(check_series_id, axis=1)\ntrain_data['image_exists'] = train_data.apply(check_image_exists, axis=1)\n\ntrain_data = train_data[\n    train_data['study_id_exists'] &\n    train_data['series_id_exists'] &\n    train_data['image_exists']\n]\nprint(\"\\nTrain Data Shape after Filtering:\", train_data.shape)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load and Display Sample Images\n# ------------------------------\ndef load_jpg(path):\n    img = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    img = img - np.min(img)\n    if np.max(img) != 0:\n        img = img / np.max(img)\n    return (img * 255).astype(np.uint8)\n\nimages = []\nrow_ids = []\nselected_indices = random.sample(range(len(train_data)), 2)\nfor i in selected_indices:\n    image = load_jpg(train_data.iloc[i]['image_path'])\n    images.append(image)\n    row_ids.append(train_data.iloc[i]['row_id'])\n\nfig, ax = plt.subplots(1, 2, figsize=(8, 4))\nfor i in range(2):\n    ax[i].imshow(images[i], cmap='gray')\n    ax[i].set_title(f'Row ID: {row_ids[i]}', fontsize=8)\n    ax[i].axis('off')\nplt.tight_layout()\nplt.show()\n\ntrain_data = train_data.dropna()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define Visualization Functions\n# -------------------------------\ndef plot_confusion_matrix(cm, classes, title):\n    plt.figure(figsize=(8, 6))\n    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',\n                xticklabels=classes, yticklabels=classes, cbar=False)\n    plt.title(title)\n    plt.xlabel('Predicted Label')\n    plt.ylabel('True Label')\n    plt.xticks(rotation=45)\n    plt.yticks(rotation=45)\n    plt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Data Loading and Preprocessing Functions\n# ----------------------------------------\ndef load_jpg_tf(path):\n    img = tf.io.read_file(path)\n    img = tf.image.decode_jpeg(img, channels=1)\n    img = tf.cast(img, tf.float32)\n    img = img - tf.reduce_min(img)\n    img = img / tf.reduce_max(img) if tf.reduce_max(img) != 0 else img\n    return img\n\nclass_counts = [37626, 7950, 3081]\nclass_weights = compute_class_weight('balanced', classes=np.array([0, 1, 2]), y=np.repeat([0, 1, 2], class_counts))\nclass_weights = dict(enumerate(class_weights))\n\ndef focal_loss(y_true, y_pred, alpha=list(class_weights.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\ndef apply_clahe(image, clipLimit=2, tileGridSize=(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=clipLimit, tileGridSize=tileGridSize)\n    return clahe.apply(image)\n\ndef 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\ndef augment_image_normalized(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\ndef 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\ndef preprocess_image(image, label=None, is_training=False):\n    def opencv_process(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_normalized(tf.expand_dims(img_np, axis=-1)).numpy().squeeze()\n        return np.expand_dims(img_np, axis=-1).astype(np.float32)\n\n    image = tf.expand_dims(image, axis=-1)\n    image = tf.py_function(opencv_process, [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","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\ndef create_dataset(df, batch_size=64, is_test=False, is_training=False):\n    if is_test:\n        def load_wrapper(path):\n            image = load_jpg_tf(path)\n            return preprocess_image(image)\n        dataset = tf.data.Dataset.from_tensor_slices(df['image_path'])\n        dataset = dataset.map(load_wrapper, num_parallel_calls=tf.data.AUTOTUNE)\n    else:\n        def load_wrapper(path, label):\n            image = load_jpg_tf(path)\n            return preprocess_image(image, label, is_training)\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(load_wrapper, num_parallel_calls=tf.data.AUTOTUNE)\n        if is_training:\n            dataset = dataset.shuffle(buffer_size=1000)\n    dataset = dataset.batch(batch_size).prefetch(tf.data.AUTOTUNE)\n    return dataset\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Train EfficientNetB2 model\n# --------------------------\ntf.keras.mixed_precision.set_global_policy('mixed_bfloat16')\n\ndef weighted_log_loss(y_true, y_pred):\n    weights = tf.gather([1.0, 2.0, 4.0], tf.cast(y_true, tf.int32))\n    loss = tf.keras.losses.sparse_categorical_crossentropy(y_true, y_pred)\n    return tf.reduce_mean(loss * weights)\n\ndef build_efficientnet_b2(num_classes=3):\n    base_model = applications.EfficientNetB2(\n        include_top=False,\n        weights='imagenet',\n        input_shape=(224, 224, 3)\n    )\n    for layer in base_model.layers:\n        if 'block7' in layer.name:\n            layer.trainable = True\n        else:\n            layer.trainable = False\n    inputs = tf.keras.Input(shape=(224, 224, 3))\n    x = base_model(inputs)\n    x = layers.GlobalAveragePooling2D()(x)\n    outputs = layers.Dense(num_classes, activation='linear')(x)\n    return tf.keras.Model(inputs, outputs)\n\ndef train_model(model, train_dataset, val_dataset, series_name):\n    model.compile(\n        optimizer=optimizers.Adam(0.0001),\n        loss=weighted_log_loss,\n        metrics=['accuracy']\n    )\n    callbacks_list = [\n        callbacks.EarlyStopping(patience=10, restore_best_weights=True),\n        callbacks.ModelCheckpoint(f'best_{series_name}.keras', save_best_only=True),\n    ]\n    history = model.fit(\n        train_dataset,\n        validation_data=val_dataset,\n        epochs=5,\n        callbacks=callbacks_list\n    )\n    return history\n\ndef evaluate_model(model, dataset):\n    y_true = []\n    y_pred = []\n    for batch in dataset:\n        images, labels = batch[0], batch[1]\n        y_true.extend(labels.numpy())\n        preds = model.predict(images)\n        y_pred.extend(np.argmax(preds, axis=1))\n    y_true = np.array(y_true)\n    y_pred = np.array(y_pred)\n    valid_indices = ~np.isnan(y_true)\n    y_true = y_true[valid_indices]\n    y_pred = y_pred[valid_indices]\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': confusion_matrix(y_true, y_pred)\n    }\n\nseries_models = {\n    'Sagittal T1': build_efficientnet_b2(),\n    'Axial T2': build_efficientnet_b2(),\n    'Sagittal T2/STIR': build_efficientnet_b2()\n}\n\nclass_names = ['normal_mild', 'moderate', 'severe']\nresults = {}\n\nfor series_name, model in series_models.items():\n    series_df = final_merged_df[\n        (final_merged_df['series_description'] == series_name) &\n        (final_merged_df['severity'].isin(class_names))\n    ].copy()\n    if series_df.empty:\n        print(f\"Skipping {series_name} - no valid data.\")\n        continue\n    train_df, temp_df = train_test_split(\n        series_df,\n        test_size=0.3,\n        stratify=series_df['severity'],\n        random_state=42\n    )\n    val_df, test_df = train_test_split(\n        temp_df,\n        test_size=0.6667,\n        stratify=temp_df['severity'],\n        random_state=42\n    )\n\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"    print(f\"\\nClass distribution for {series_name}:\")\n    print(\"Train:\", train_df['severity'].value_counts())\n    print(\"Validation:\", val_df['severity'].value_counts())\n    print(\"Test:\", test_df['severity'].value_counts())\n    train_ds = create_dataset(train_df, batch_size=64, is_test=False, is_training=True)\n    val_ds = create_dataset(val_df, batch_size=64, is_test=False)\n    test_ds = create_dataset(test_df, batch_size=64, is_test=False)\n    history = train_model(model, train_ds, val_ds, series_name)\n    result = evaluate_model(model, test_ds)\n    results[series_name] = result\n    plot_confusion_matrix(result['cm'], class_names, f'{series_name} Confusion Matrix')\n    print(f\"\\nMetrics for {series_name}:\")\n    print(f\"Precision: {result['precision']:.4f}\")\n    print(f\"Recall: {result['recall']:.4f}\")\n    print(f\"F1-Score: {result['f1']:.4f}\")\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.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.legend()\n    plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ResNet50 Section (بازنویسی برای JPG)\n# -------------------------------------\ndef build_resnet50(num_classes=3):\n    base_model = applications.ResNet50(\n        include_top=False,\n        weights='imagenet',\n        input_shape=(224, 224, 3)\n    )\n    for layer in base_model.layers:\n        if 'conv4_block' in layer.name or 'conv5_block' in layer.name:\n            layer.trainable = True\n        else:\n            layer.trainable = False\n    inputs = tf.keras.Input(shape=(224, 224, 3))\n    x = base_model(inputs)\n    x = layers.GlobalAveragePooling2D()(x)\n    outputs = layers.Dense(num_classes, activation='linear')(x)\n    return tf.keras.Model(inputs, outputs)\n\nseries_models_resnet = {\n    'Sagittal T1': build_resnet50(),\n    'Axial T2': build_resnet50(),\n    'Sagittal T2/STIR': build_resnet50()\n}\n\nresults_resnet = {}\n\nfor series_name, model in series_models_resnet.items():\n    series_df = final_merged_df[\n        (final_merged_df['series_description'] == series_name) &\n        (final_merged_df['severity'].isin(class_names))\n    ].copy()\n    if series_df.empty:\n        print(f\"Skipping {series_name} - no valid data.\")\n        continue\n    train_df, temp_df = train_test_split(\n        series_df,\n        test_size=0.3,\n        stratify=series_df['severity'],\n        random_state=42\n    )\n    val_df, test_df = train_test_split(\n        temp_df,\n        test_size=0.6667,\n        stratify=temp_df['severity'],\n        random_state=42\n    )\n    print(f\"\\nClass distribution for {series_name} (ResNet50):\")\n    print(\"Train:\", train_df['severity'].value_counts())\n    print(\"Validation:\", val_df['severity'].value_counts())\n    print(\"Test:\", test_df['severity'].value_counts())\n    train_ds = create_dataset(train_df, batch_size=64, is_test=False, is_training=True)\n    val_ds = create_dataset(val_df, batch_size=64, is_test=False)\n    test_ds = create_dataset(test_df, batch_size=64, is_test=False)\n    history = train_model(model, train_ds, val_ds, series_name)  # Using weighted_log_loss here too\n    result = evaluate_model(model, test_ds)\n    results_resnet[series_name] = result\n    plot_confusion_matrix(result['cm'], class_names, f'{series_name} Confusion Matrix (ResNet50)')\n    print(f\"\\nMetrics for {series_name} (ResNet50):\")\n    print(f\"Precision: {result['precision']:.4f}\")\n    print(f\"Recall: {result['recall']:.4f}\")\n    print(f\"F1-Score: {result['f1']:.4f}\")\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 (ResNet50)')\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 (ResNet50)')\n    plt.legend()\n    plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}