{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.11.11"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"}],"dockerImageVersionId":31234,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false},"papermill":{"default_parameters":{},"duration":1844.021902,"end_time":"2025-05-14T06:09:22.260982","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2025-05-14T05:38:38.239080","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport pydicom\nimport numpy as np\nimport pandas as pd\nimport tensorflow as tf\nfrom tensorflow.keras import layers, models, applications\nfrom sklearn.model_selection import train_test_split\nfrom tensorflow.keras.utils import to_categorical\nimport cv2\nfrom tensorflow.keras.callbacks import EarlyStopping, ModelCheckpoint, ReduceLROnPlateau\n\n# ========== CONFIGURATION ==========\nIMG_SIZE = 224\nBATCH_SIZE = 8\nEPOCHS = 50\nWARMUP_EPOCHS = 5\nMIXUP_ALPHA = 0.2\nDATA_DIR = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/'\n\n# ========== LOAD DATA ==========\nprint(\"Loading data...\")\ntrain_df = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv')\nseries_desc_df = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv')\n\n# ========== LABEL ENCODING ==========\nlabel_cols = train_df.columns[1:]\nlabel_map = {'Normal/Mild': 0, 'Moderate': 1, 'Severe': 2}\n\ndef encode_labels(row):\n    return [label_map.get(row[col], 0) for col in label_cols]\n\ntrain_df['encoded_labels'] = train_df.apply(encode_labels, axis=1)\n\ndef get_max_severity(encoded_labels):\n    return max(encoded_labels)\n\ntrain_df['max_severity'] = train_df['encoded_labels'].apply(get_max_severity)\n\n# ========== BALANCED SAMPLING ==========\nMAX_IMAGES = 400000\nIMAGES_PER_STUDY = 3\nMAX_STUDIES = MAX_IMAGES // IMAGES_PER_STUDY\nstudies_per_class = MAX_STUDIES // 3\n\nprint(f\"Sampling up to {studies_per_class} studies per class...\")\ndfs = []\nfor severity in [0, 1, 2]:\n    subset = train_df[train_df['max_severity'] == severity]\n    sampled = subset.sample(n=min(len(subset), studies_per_class), random_state=42)\n    dfs.append(sampled)\n    print(f\"  Class {severity}: {len(sampled)} studies\")\n\nbalanced_train_df = pd.concat(dfs).sample(frac=1, random_state=42).reset_index(drop=True)\nprint(f\"Total balanced dataset: {len(balanced_train_df)} studies\\n\")\n\n# ========== HELPER FUNCTIONS ==========\ndef get_series_ids(study_id):\n    sub_df = series_desc_df[series_desc_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\n# ========== DATA AUGMENTATION ==========\ndef mixup_batch(X1, X2, X3, y, alpha=MIXUP_ALPHA):\n    batch_size = len(X1)\n    indices = np.random.permutation(batch_size)\n    lam = np.random.beta(alpha, alpha, batch_size)\n    \n    X1_mixed = np.array([lam[i] * X1[i] + (1 - lam[i]) * X1[indices[i]] for i in range(batch_size)])\n    X2_mixed = np.array([lam[i] * X2[i] + (1 - lam[i]) * X2[indices[i]] for i in range(batch_size)])\n    X3_mixed = np.array([lam[i] * X3[i] + (1 - lam[i]) * X3[indices[i]] for i in range(batch_size)])\n    y_mixed = np.array([lam[i] * y[i] + (1 - lam[i]) * y[indices[i]] for i in range(batch_size)])\n    \n    return X1_mixed, X2_mixed, X3_mixed, y_mixed\n\ndef augment_image(img):\n    if np.random.rand() < 0.5:\n        angle = np.random.uniform(-15, 15)\n        M = cv2.getRotationMatrix2D((IMG_SIZE//2, IMG_SIZE//2), angle, 1.0)\n        img = cv2.warpAffine(img, M, (IMG_SIZE, IMG_SIZE))\n    \n    if np.random.rand() < 0.5:\n        factor = np.random.uniform(0.8, 1.2)\n        img = np.clip(img * factor, 0, 1)\n    \n    if np.random.rand() < 0.5:\n        img = cv2.flip(img, 1)\n    \n    if np.random.rand() < 0.3:\n        zoom = np.random.uniform(0.9, 1.1)\n        h, w = img.shape[:2]\n        new_h, new_w = int(h * zoom), int(w * zoom)\n        img = cv2.resize(img, (new_w, new_h))\n        \n        if zoom > 1:\n            start_h = (new_h - h) // 2\n            start_w = (new_w - w) // 2\n            img = img[start_h:start_h+h, start_w:start_w+w]\n        else:\n            pad_h = (h - new_h) // 2\n            pad_w = (w - new_w) // 2\n            img = cv2.copyMakeBorder(img, pad_h, h-new_h-pad_h, pad_w, w-new_w-pad_w, cv2.BORDER_CONSTANT)\n    \n    return img\n\n# ========== DICOM LOADING ==========\ndef load_dicom_image(path, augment=False):\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 = img / (np.max(img) + 1e-8)\n        \n        if augment:\n            img = augment_image(img)\n        \n        return img\n    except Exception as e:\n        return np.zeros((IMG_SIZE, IMG_SIZE, 3), dtype=np.float32)\n\ndef load_study_images(study_id, augment=False):\n    views = get_series_ids(study_id)\n    images = []\n    \n    for view in ['Sagittal T1', 'Sagittal T2/STIR', 'Axial T2']:\n        series_id = views[view]\n        if pd.isna(series_id):\n            images.append(np.zeros((IMG_SIZE, IMG_SIZE, 3)))\n        else:\n            series_path = os.path.join(DATA_DIR, str(study_id), str(series_id))\n            if os.path.exists(series_path):\n                instances = sorted(os.listdir(series_path))\n                if instances:\n                    img_path = os.path.join(series_path, instances[len(instances)//2])\n                    images.append(load_dicom_image(img_path, augment=augment))\n                else:\n                    images.append(np.zeros((IMG_SIZE, IMG_SIZE, 3)))\n            else:\n                images.append(np.zeros((IMG_SIZE, IMG_SIZE, 3)))\n    \n    return images\n\n# ========== DATA GENERATOR ==========\nclass DataGenerator(tf.keras.utils.Sequence):\n    def __init__(self, df, batch_size=BATCH_SIZE, shuffle=True, augment=False, mixup=False):\n        self.df = df\n        self.batch_size = batch_size\n        self.shuffle = shuffle\n        self.augment = augment\n        self.mixup = mixup\n        self.indexes = np.arange(len(self.df))\n        self.on_epoch_end()\n\n    def __len__(self):\n        return int(np.floor(len(self.df) / self.batch_size))\n\n    def __getitem__(self, index):\n        batch_ids = self.indexes[index*self.batch_size:(index+1)*self.batch_size]\n        batch_df = self.df.iloc[batch_ids]\n        X1, X2, X3, y = [], [], [], []\n        \n        for _, row in batch_df.iterrows():\n            imgs = load_study_images(row['study_id'], augment=self.augment)\n            X1.append(imgs[0])\n            X2.append(imgs[1])\n            X3.append(imgs[2])\n            y.append(row['encoded_labels'])\n        \n        X1, X2, X3 = np.array(X1), np.array(X2), np.array(X3)\n        y = to_categorical(np.array(y), num_classes=3)\n        \n        if self.mixup and self.augment:\n            X1, X2, X3, y = mixup_batch(X1, X2, X3, y)\n        \n        return (X1, X2, X3), y\n\n    def on_epoch_end(self):\n        if self.shuffle:\n            np.random.shuffle(self.indexes)\n\n# ========== MODEL ARCHITECTURE ==========\ndef create_backbone():\n    base = applications.MobileNetV2(\n        include_top=False, \n        weights='imagenet', \n        input_shape=(IMG_SIZE, IMG_SIZE, 3), \n        pooling='avg'\n    )\n    for layer in base.layers[-65:]:\n        layer.trainable = True\n    return base\n\ndef build_mvcnn():\n    input1 = layers.Input(shape=(IMG_SIZE, IMG_SIZE, 3), name='sagittal_t1')\n    input2 = layers.Input(shape=(IMG_SIZE, IMG_SIZE, 3), name='sagittal_t2')\n    input3 = layers.Input(shape=(IMG_SIZE, IMG_SIZE, 3), name='axial_t2')\n    \n    backbone = create_backbone()\n    feat1 = backbone(input1)\n    feat2 = backbone(input2)\n    feat3 = backbone(input3)\n    \n    merged = layers.Concatenate()([feat1, feat2, feat3])\n    \n    x = layers.Dense(1024, kernel_regularizer=tf.keras.regularizers.l2(1e-4))(merged)\n    x = layers.BatchNormalization()(x)\n    x = layers.Activation('relu')(x)\n    x = layers.Dropout(0.4)(x)\n    \n    x = layers.Dense(768, kernel_regularizer=tf.keras.regularizers.l2(1e-4))(x)\n    x = layers.BatchNormalization()(x)\n    x = layers.Activation('relu')(x)\n    x = layers.Dropout(0.35)(x)\n    \n    x = layers.Dense(512, kernel_regularizer=tf.keras.regularizers.l2(1e-4))(x)\n    x = layers.BatchNormalization()(x)\n    x = layers.Activation('relu')(x)\n    x = layers.Dropout(0.3)(x)\n    \n    x = layers.Dense(256, kernel_regularizer=tf.keras.regularizers.l2(1e-4))(x)\n    x = layers.BatchNormalization()(x)\n    x = layers.Activation('relu')(x)\n    x = layers.Dropout(0.25)(x)\n    \n    output = layers.Dense(len(label_cols) * 3, activation='softmax')(x)\n    output = layers.Reshape((len(label_cols), 3))(output)\n    \n    model = models.Model(inputs=[input1, input2, input3], outputs=output)\n    model.compile(\n        optimizer=tf.keras.optimizers.Adam(learning_rate=2e-4),\n        loss=tf.keras.losses.CategoricalCrossentropy(label_smoothing=0.1),\n        metrics=['accuracy']\n    )\n    return model\n\n# ========== TRAIN/VAL SPLIT ==========\nprint(\"Creating train/validation split...\")\ntrain_ids, val_ids = train_test_split(\n    balanced_train_df, \n    test_size=0.2, \n    random_state=42, \n    stratify=balanced_train_df['max_severity']\n)\n\nprint(f\"Train: {len(train_ids)} studies\")\nprint(f\"Val: {len(val_ids)} studies\\n\")\n\n# ========== DATA GENERATORS ==========\ntrain_gen = DataGenerator(train_ids, augment=True, mixup=True)\nval_gen = DataGenerator(val_ids, augment=False, mixup=False)\n\n# ========== BUILD MODEL ==========\nprint(\"Building model...\")\nmodel = build_mvcnn()\nprint(f\"Total parameters: {model.count_params():,}\\n\")\n\n# ========== LEARNING RATE SCHEDULE ==========\ndef lr_schedule(epoch, lr):\n    if epoch < WARMUP_EPOCHS:\n        return 1e-5 + (2e-4 - 1e-5) * (epoch / WARMUP_EPOCHS)\n    else:\n        progress = (epoch - WARMUP_EPOCHS) / (EPOCHS - WARMUP_EPOCHS)\n        return 2e-4 * 0.5 * (1 + np.cos(np.pi * progress))\n\nlr_callback = tf.keras.callbacks.LearningRateScheduler(lr_schedule, verbose=1)\n\n# ========== CALLBACKS ==========\ncallbacks = [\n    EarlyStopping(\n        monitor='val_accuracy',\n        patience=15,\n        restore_best_weights=True,\n        mode='max',\n        verbose=1\n    ),\n    ModelCheckpoint(\n        'best_model_balanced.keras', \n        monitor='val_accuracy', \n        save_best_only=True, \n        mode='max', \n        verbose=1\n    ),\n    lr_callback\n]\n\n# ========== TRAINING ==========\nprint(\"Starting training with warmup + MixUp...\")\nhistory = model.fit(\n    train_gen,\n    validation_data=val_gen,\n    epochs=EPOCHS,\n    callbacks=callbacks,\n    verbose=1\n)\n\n# ========== SAVE MODEL ==========\nmodel.save('mvc_MobileNetV2_balanced_90pct.keras')\nprint(f\"\\n{'='*60}\")\nprint(\"Training complete!\")\nprint(f\"Best model saved as: 'best_model_balanced.keras'\")\nprint(f\"Final model saved as: 'mvc_MobileNetV2_balanced_90pct.keras'\")\nprint(f\"{'='*60}\")","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.execute_input":"2025-05-14T05:38:42.397668Z","iopub.status.busy":"2025-05-14T05:38:42.397239Z","iopub.status.idle":"2025-05-14T06:09:17.603034Z","shell.execute_reply":"2025-05-14T06:09:17.602418Z"},"papermill":{"duration":1835.210478,"end_time":"2025-05-14T06:09:17.604422","exception":false,"start_time":"2025-05-14T05:38:42.393944","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":0.240311,"end_time":"2025-05-14T06:09:18.086449","exception":false,"start_time":"2025-05-14T06:09:17.846138","status":"completed"},"tags":[]},"outputs":[],"execution_count":null}]}