{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"}],"dockerImageVersionId":30786,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ------------------------\n# LIBRARY IMPORTS\n# ------------------------\nimport os\nimport cv2\nimport pydicom\nimport numpy as np\nimport pandas as pd\nimport tensorflow as tf\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import f1_score, classification_report\nfrom sklearn.model_selection import train_test_split\nfrom tensorflow.keras import layers, models, applications\nfrom tensorflow.keras.callbacks import EarlyStopping, ModelCheckpoint, ReduceLROnPlateau\nfrom tensorflow.keras.utils import to_categorical\nfrom sklearn.utils.class_weight import compute_class_weight\n\n# ------------------------\n# CONFIGURATION\n# ------------------------\nBATCH_SIZE = 16\nEPOCHS = 40\nLEARNING_RATE = 5e-5\nIMG_SIZE = 256\nDATA_DIR = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images\"\n\n# ------------------------\n# LOAD DATA\n# ------------------------\ntrain_df = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv')\nseries_df = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv')\n\nlabel_cols = train_df.columns[1:]\nlabel_map = {'Normal/Mild': 0, 'Moderate': 1, 'Severe': 2}\ntrain_df['encoded_labels'] = train_df.apply(lambda row: [label_map.get(row[col], 0) for col in label_cols], axis=1)\ntrain_df['max_severity'] = train_df['encoded_labels'].apply(max)\n# Class weights for imbalance\nclass_weights_array = compute_class_weight('balanced', classes=np.unique(train_df['max_severity']), y=train_df['max_severity'])\nclass_weights_dict = {i: class_weights_array[i] for i in range(len(class_weights_array))}\nprint(f\"Class weights: {class_weights_dict}\")\n\n# ------------------------\n# IMAGE LOAD FUNCTIONS\n# ------------------------\ndef get_series_ids(study_id):\n    sub_df = series_df[series_df['study_id'] == study_id]\n    views = {'Sagittal T1': None, 'Sagittal T2/STIR': None, 'Axial T2': None}\n    for view in views:\n        found = sub_df[sub_df['series_description'].str.contains(view, case=False, na=False)]\n        if not found.empty:\n            views[view] = found.iloc[0]['series_id']\n    return views\n\ndef load_dicom_image(path):\n    try:\n        dcm = pydicom.dcmread(path)\n        img = dcm.pixel_array.astype(np.float32)\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n        img = cv2.cvtColor(img, cv2.COLOR_GRAY2RGB)\n        img /= np.max(img) if np.max(img) != 0 else 1.0\n        return img\n    except:\n        return np.zeros((IMG_SIZE, IMG_SIZE, 3), dtype=np.float32)\n\ndef load_study_images(study_id):\n    views = get_series_ids(study_id)\n    images = []\n    for view in ['Sagittal T1', 'Sagittal T2/STIR', 'Axial T2']:\n        img = np.zeros((IMG_SIZE, IMG_SIZE, 3), dtype=np.float32)\n        series_id = views[view]\n        if pd.notna(series_id):\n            series_path = os.path.join(DATA_DIR, str(study_id), str(series_id))\n            if os.path.exists(series_path):\n                files = sorted([f for f in os.listdir(series_path) if f.endswith('.dcm')])\n                if len(files) >= 3:\n                    # Use middle 3 slices averaged for richer context\n                    mid = len(files)//2\n                    slice_imgs = [load_dicom_image(os.path.join(series_path, files[i])) for i in range(mid - 1, mid + 2)]\n                    img = np.mean(slice_imgs, axis=0)\n                elif files:\n                    img = load_dicom_image(os.path.join(series_path, files[len(files)//2]))\n        images.append(img)\n    return images\n\n# ------------------------\n# FOCAL LOSS FOR IMBALANCED CLASSIFICATION\nclass FocalLoss(tf.keras.losses.Loss):\n    def __init__(self, alpha=0.25, gamma=2.0):\n        super(FocalLoss, self).__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n    \n    def call(self, y_true, y_pred):\n        y_pred = tf.clip_by_value(y_pred, 1e-7, 1.0 - 1e-7)\n        ce_loss = -y_true * tf.math.log(y_pred)\n        focal_weight = self.alpha * tf.math.pow(1.0 - y_pred, self.gamma)\n        focal_loss = focal_weight * ce_loss\n        return tf.reduce_mean(tf.reduce_sum(focal_loss, axis=-1))\n\n# ------------------------\n# MIXUP IMPLEMENTATION\n# ------------------------\ndef mixup(x1, x2, x3, y, alpha=0.2):\n    lam = np.random.beta(alpha, alpha)\n    idx = np.random.permutation(len(x1))\n    x1_mix = lam * x1 + (1 - lam) * x1[idx]\n    x2_mix = lam * x2 + (1 - lam) * x2[idx]\n    x3_mix = lam * x3 + (1 - lam) * x3[idx]\n    y_mix = lam * y + (1 - lam) * y[idx]\n    return x1_mix, x2_mix, x3_mix, y_mix\n\n# ------------------------\n# DATA GENERATOR\n# ------------------------\ndata_augmentation = tf.keras.Sequential([\n    layers.RandomFlip(\"horizontal\"),\n    layers.RandomRotation(0.1),\n    layers.RandomZoom(0.1),\n])\n\nclass DataGenerator(tf.keras.utils.Sequence):\n    def __init__(self, df, batch_size=BATCH_SIZE, augment=False, mixup_active=False):\n        self.df = df.reset_index(drop=True)\n        self.batch_size = batch_size\n        self.augment = augment\n        self.mixup_active = mixup_active\n\n    def __len__(self):\n        return len(self.df) // self.batch_size\n\n    def __getitem__(self, idx):\n        batch_df = self.df.iloc[idx*self.batch_size:(idx+1)*self.batch_size]\n        X1, X2, X3, y = [], [], [], []\n        for _, row in batch_df.iterrows():\n            imgs = load_study_images(row['study_id'])\n            if self.augment:\n                imgs = [data_augmentation(img) for img in imgs]\n            X1.append(imgs[0])\n            X2.append(imgs[1])\n            X3.append(imgs[2])\n            y.append(row['encoded_labels'])\n        X1 = np.array(X1, dtype=np.float32)\n        X2 = np.array(X2, dtype=np.float32)\n        X3 = np.array(X3, dtype=np.float32)\n        y_cat = np.stack([to_categorical(sample, num_classes=3) for sample in y]).astype(np.float32)\n        if self.mixup_active:\n            X1, X2, X3, y_cat = mixup(X1, X2, X3, y_cat)\n        return (X1, X2, X3), y_cat\n\n# ------------------------\n# TRAIN/VAL SPLIT & GENERATORS\n# ------------------------\ntrain_df, val_df = train_test_split(train_df, test_size=0.2, stratify=train_df['max_severity'], random_state=42)\ntrain_gen = DataGenerator(train_df, augment=True, mixup_active=True)\nval_gen = DataGenerator(val_df, augment=False, mixup_active=False)\n\n# ------------------------\n# MODEL\n# ------------------------\ninput1 = layers.Input(shape=(IMG_SIZE, IMG_SIZE, 3))\ninput2 = layers.Input(shape=(IMG_SIZE, IMG_SIZE, 3))\ninput3 = layers.Input(shape=(IMG_SIZE, IMG_SIZE, 3))\n\nbase_model = applications.EfficientNetV2M(include_top=False, weights='imagenet', pooling='avg')\n# Fine-tune last 60 layers\nfor layer in base_model.layers[:-60]:\n    layer.trainable = False\nfor layer in base_model.layers[-60:]:\n    layer.trainable = True\n\nf1 = base_model(input1)\nf2 = base_model(input2)\nf3 = base_model(input3)\n\nx = layers.Concatenate()([f1, f2, f3])\n# Deeper classification head with L2 regularization and batch norm\nx = layers.Dense(768, activation='relu', kernel_regularizer=tf.keras.regularizers.l2(1e-4))(x)\nx = layers.BatchNormalization()(x)\nx = layers.Dropout(0.4)(x)\n\nx = layers.Dense(512, activation='relu', kernel_regularizer=tf.keras.regularizers.l2(1e-4))(x)\nx = layers.BatchNormalization()(x)\nx = layers.Dropout(0.35)(x)\n\nx = layers.Dense(256, activation='relu', kernel_regularizer=tf.keras.regularizers.l2(1e-4))(x)\nx = layers.BatchNormalization()(x)\nx = layers.Dropout(0.3)(x)\n\nout = layers.Dense(len(label_cols) * 3, activation='softmax')(x)\nout = layers.Reshape((len(label_cols), 3))(out)\n\nmodel = models.Model(inputs=[input1, input2, input3], outputs=out)\nmodel.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=LEARNING_RATE),\n              loss=FocalLoss(alpha=0.25, gamma=2.0),\n              metrics=['accuracy'])\n\ncallbacks = [\n    EarlyStopping(monitor='val_loss', patience=8, restore_best_weights=True, verbose=1),\n    ModelCheckpoint('best_model_v2.keras', monitor='val_loss', save_best_only=True, verbose=1),\n    ReduceLROnPlateau(monitor='val_loss', factor=0.5, patience=3, min_lr=1e-7, verbose=1)\n]\n\nprint(model.summary())\nhistory = model.fit(train_gen, validation_data=val_gen, epochs=EPOCHS, callbacks=callbacks, verbose=1)\n\n# ------------------------\n# EVALUATION\n# ------------------------\ny_true, y_pred = [], []\nfor i in range(len(val_gen)):\n    try:\n        (X1, X2, X3), y_batch = val_gen[i]\n        preds = model.predict([X1, X2, X3], verbose=0)\n        y_true.extend(np.argmax(y_batch, axis=-1).flatten())\n        y_pred.extend(np.argmax(preds, axis=-1).flatten())\n    except Exception as e:\n        print(f\"Skipped batch {i} due to error: {e}\")\n\nf1 = f1_score(y_true, y_pred, average='macro')\nprint(f\"\\nMacro F1 Score: {f1:.4f}\")\nprint(classification_report(y_true, y_pred))\n\n# ------------------------\n# PLOT TRAINING\n# ------------------------\nplt.figure(figsize=(12, 5))\nplt.subplot(1, 2, 1)\nplt.plot(history.history['loss'], label='Train Loss')\nplt.plot(history.history['val_loss'], label='Val Loss')\nplt.title('Loss')\nplt.legend()\n\nplt.subplot(1, 2, 2)\nplt.plot(history.history['accuracy'], label='Train Acc')\nplt.plot(history.history['val_accuracy'], label='Val Acc')\nplt.title('Accuracy')\nplt.legend()\nplt.tight_layout()\nplt.show()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-18T08:17:43.343026Z","iopub.execute_input":"2025-12-18T08:17:43.343344Z","execution_failed":"2025-12-18T09:30:34.264Z"}},"outputs":[],"execution_count":null}]}