{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceType":"competition","sourceId":36363,"databundleVersionId":4050810}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# %% [markdown]\n# # Phase 1: Data Preparation & EDA Pipeline\n# # Professional Research Code - RSNA 2022 Cervical Spine\n\n# %% [code]\n# Install dependencies (Run once)\n!pip install pydicom pylibjpeg pylibjpeg-libjpeg nibabel pandas matplotlib seaborn scikit-learn -q\n\nimport os\nimport logging\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport pydicom\nfrom pathlib import Path\nfrom sklearn.model_selection import train_test_split\nfrom typing import Tuple, List\n\n# Configure Logging\nlogging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')\nlogger = logging.getLogger(__name__)\n\n# Configuration Class\nclass Config:\n    DATA_PATH = \"/kaggle/input/competitions/rsna-2022-cervical-spine-fracture-detection\"\n    TRAIN_IMAGES = os.path.join(DATA_PATH, \"train_images\")\n    TRAIN_CSV = os.path.join(DATA_PATH, \"train.csv\")\n    BBOX_CSV = os.path.join(DATA_PATH, \"train_bounding_boxes.csv\")\n    SEG_PATH = os.path.join(DATA_PATH, \"segmentations\")\n    \n    # Target Definition: C1 or C2 Fracture (Odontoid focus)\n    TARGET_COLS = ['C1', 'C2']\n    OUTPUT_DIR = \"/kaggle/working/phase1_output\"\n    \n    # Image Processing\n    HU_MIN = -1000\n    HU_MAX = 2000\n    IMG_SIZE = 224\n    USE_PRETRAINED = False\n    # Random State\n    SEED = 42\n\nos.makedirs(Config.OUTPUT_DIR, exist_ok=True)\nlogger.info(f\"Configuration loaded. Output dir: {Config.OUTPUT_DIR}\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-03-11T14:24:13.02769Z","iopub.execute_input":"2026-03-11T14:24:13.028015Z","iopub.status.idle":"2026-03-11T14:24:16.329983Z","shell.execute_reply.started":"2026-03-11T14:24:13.027987Z","shell.execute_reply":"2026-03-11T14:24:16.329168Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %% [code]\nclass DICOMEngine:\n    \"\"\"\n    Professional DICOM handling with HU rescaling and compression support.\n    \"\"\"\n    @staticmethod\n    def load_patient_series(patient_uid: str) -> List[np.ndarray]:\n        \"\"\"\n        Loads all slices for a patient, sorts them, and applies HU rescaling.\n        Returns a list of numpy arrays (slices).\n        \"\"\"\n        patient_path = Path(Config.TRAIN_IMAGES) / patient_uid\n        if not patient_path.exists():\n            logger.warning(f\"Patient path not found: {patient_uid}\")\n            return []\n        \n        dicom_files = list(patient_path.glob(\"*.dcm\"))\n        if not dicom_files:\n            return []\n        \n        slices = []\n        for df in dicom_files:\n            try:\n                # Use force=True to handle missing headers, pylibjpeg handles compression\n                ds = pydicom.dcmread(df, force=True) \n                \n                # Handle Pixel Data\n                if hasattr(ds, 'PixelData'):\n                    pixel_array = ds.pixel_array.astype(np.float32)\n                    \n                    # HU Rescaling (Critical for CT)\n                    slope = float(ds.RescaleSlope) if hasattr(ds, 'RescaleSlope') else 1.0\n                    intercept = float(ds.RescaleIntercept) if hasattr(ds, 'RescaleIntercept') else 0.0\n                    hu_array = pixel_array * slope + intercept\n                    \n                    # Store with metadata for sorting\n                    slice_loc = float(ds.SliceLocation) if hasattr(ds, 'SliceLocation') else 0.0\n                    slices.append({'hu': hu_array, 'loc': slice_loc, 'instance': ds.InstanceNumber})\n            except Exception as e:\n                logger.error(f\"Error loading {df.name}: {str(e)}\")\n                continue\n        \n        # Sort by Slice Location (Anatomical order)\n        slices.sort(key=lambda x: x['loc'])\n        \n        return [s['hu'] for s in slices]\n\n    @staticmethod\n    def normalize_hu(image: np.ndarray) -> np.ndarray:\n        \"\"\"Clips HU to bone window and normalizes to [0, 1]\"\"\"\n        image = np.clip(image, Config.HU_MIN, Config.HU_MAX)\n        image = (image - Config.HU_MIN) / (Config.HU_MAX - Config.HU_MIN)\n        return image.astype(np.float32)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-11T14:24:16.331794Z","iopub.execute_input":"2026-03-11T14:24:16.332056Z","iopub.status.idle":"2026-03-11T14:24:16.340776Z","shell.execute_reply.started":"2026-03-11T14:24:16.332028Z","shell.execute_reply":"2026-03-11T14:24:16.340131Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %% [code]\nclass EDAEngine:\n    def __init__(self, df: pd.DataFrame):\n        self.df = df\n        self.setup_targets()\n    \n    def setup_targets(self):\n        \"\"\"Create binary target for C1/C2 (Odontoid) fracture\"\"\"\n        # 1 if C1 OR C2 is fractured, else 0\n        self.df['target_c1_c2'] = (self.df['C1'] == 1) | (self.df['C2'] == 1)\n        self.df['target_c1_c2'] = self.df['target_c1_c2'].astype(int)\n        \n        # Count total fractures per patient\n        vertebrae_cols = ['C1', 'C2', 'C3', 'C4', 'C5', 'C6', 'C7']\n        self.df['fracture_count'] = self.df[vertebrae_cols].sum(axis=1)\n        \n    def generate_report(self) -> pd.DataFrame:\n        \"\"\"Generates statistical summary\"\"\"\n        total = len(self.df)\n        c1_c2_pos = self.df['target_c1_c2'].sum()\n        \n        report = {\n            'Total Patients': total,\n            'C1/C2 Fractures (Positive)': c1_c2_pos,\n            'C1/C2 Fractures (Negative)': total - c1_c2_pos,\n            'Positive Ratio (%)': round(c1_c2_pos / total * 100, 2),\n            'Avg Fractures per Patient': round(self.df['fracture_count'].mean(), 2)\n        }\n        \n        logger.info(\"=== EDA Report ===\")\n        for k, v in report.items():\n            logger.info(f\"{k}: {v}\")\n            \n        return pd.DataFrame([report])\n    \n    def plot_distributions(self, save_path: str):\n        \"\"\"Plots target distribution and fracture counts\"\"\"\n        fig, axes = plt.subplots(1, 2, figsize=(15, 5))\n        \n        # Plot 1: Target Balance\n        sns.countplot(data=self.df, x='target_c1_c2', ax=axes[0], palette='viridis')\n        axes[0].set_title('C1/C2 Fracture Distribution (Binary)')\n        axes[0].set_xlabel('Fracture (1=Yes, 0=No)')\n        \n        # Plot 2: Fracture Count Histogram\n        sns.histplot(data=self.df, x='fracture_count', bins=8, ax=axes[1], color='salmon')\n        axes[1].set_title('Number of Fractured Vertebrae per Patient')\n        \n        plt.tight_layout()\n        plt.savefig(save_path)\n        plt.show()\n        logger.info(f\"Distribution plots saved to {save_path}\")\n\n    def visualize_sample_scans(self, patient_ids: List[str], save_path: str):\n        \"\"\"Visualizes 3 random slices from specific patients\"\"\"\n        fig, axes = plt.subplots(len(patient_ids), 3, figsize=(15, 5 * len(patient_ids)))\n        if len(patient_ids) == 1: axes = np.array([axes]) # Handle single row case\n        \n        for i, pid in enumerate(patient_ids):\n            logger.info(f\"Loading scans for {pid}...\")\n            slices = DICOMEngine.load_patient_series(pid)\n            \n            if not slices:\n                continue\n                \n            # Select 3 representative slices (Top, Middle, Bottom)\n            n = len(slices)\n            indices = [0, n // 2, n - 1] if n >= 3 else list(range(n))\n            \n            for j, idx in enumerate(indices):\n                img = DICOMEngine.normalize_hu(slices[idx])\n                ax = axes[i, j] if len(patient_ids) > 1 else axes[j]\n                ax.imshow(img, cmap='bone')\n                ax.set_title(f'{pid} - Slice {idx}/{n}')\n                ax.axis('off')\n        \n        plt.tight_layout()\n        plt.savefig(save_path)\n        plt.show()\n        logger.info(f\"Sample scans saved to {save_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-11T14:24:16.341533Z","iopub.execute_input":"2026-03-11T14:24:16.341803Z","iopub.status.idle":"2026-03-11T14:24:16.361263Z","shell.execute_reply.started":"2026-03-11T14:24:16.341781Z","shell.execute_reply":"2026-03-11T14:24:16.360582Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %% [code]\nclass DataSplitter:\n    def __init__(self, df: pd.DataFrame):\n        self.df = df\n    \n    def split(self, test_size=0.15, val_size=0.15) -> Tuple[pd.DataFrame, pd.DataFrame, pd.DataFrame]:\n        \"\"\"\n        Stratified split to maintain C1/C2 fracture ratio across sets.\n        \"\"\"\n        # First split: Test set\n        train_temp, test_df = train_test_split(\n            self.df, \n            test_size=test_size, \n            stratify=self.df['target_c1_c2'],\n            random_state=Config.SEED\n        )\n        \n        # Second split: Validation set from remaining training data\n        # Adjust val_size relative to the remaining data\n        actual_val_size = val_size / (1 - test_size)\n        \n        train_df, val_df = train_test_split(\n            train_temp,\n            test_size=actual_val_size,\n            stratify=train_temp['target_c1_c2'],\n            random_state=Config.SEED\n        )\n        \n        logger.info(f\"Split Complete -> Train: {len(train_df)}, Val: {len(val_df)}, Test: {len(test_df)}\")\n        return train_df, val_df, test_df\n\n    def save_metadata(self, train_df, val_df, test_df, output_dir):\n        \"\"\"Saves CSVs with file paths for easy loading in Phase 2\"\"\"\n        for name, df in [('train', train_df), ('val', val_df), ('test', test_df)]:\n            path = os.path.join(output_dir, f\"{name}_metadata.csv\")\n            df.to_csv(path, index=False)\n            logger.info(f\"Saved {name} metadata to {path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-11T14:24:16.362835Z","iopub.execute_input":"2026-03-11T14:24:16.36309Z","iopub.status.idle":"2026-03-11T14:24:16.375376Z","shell.execute_reply.started":"2026-03-11T14:24:16.363061Z","shell.execute_reply":"2026-03-11T14:24:16.374757Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %% [code]\ndef run_phase1():\n    logger.info(\"=== Starting Phase 1 Pipeline ===\")\n    \n    # 1. Load Data\n    df = pd.read_csv(Config.TRAIN_CSV)\n    logger.info(f\"Loaded {len(df)} records from train.csv\")\n    \n    # 2. Initialize EDA\n    eda = EDAEngine(df)\n    eda.generate_report()\n    eda.plot_distributions(os.path.join(Config.OUTPUT_DIR, \"distribution_plots.png\"))\n    \n    # 3. Visualize Samples (Pick 2 positive, 2 negative)\n    pos_samples = df[df['target_c1_c2'] == 1]['StudyInstanceUID'].head(2).tolist()\n    neg_samples = df[df['target_c1_c2'] == 0]['StudyInstanceUID'].head(2).tolist()\n    sample_ids = pos_samples + neg_samples\n    \n    eda.visualize_sample_scans(sample_ids, os.path.join(Config.OUTPUT_DIR, \"sample_scans.png\"))\n    \n    # 4. Split Data\n    splitter = DataSplitter(df)\n    train_df, val_df, test_df = splitter.split()\n    splitter.save_metadata(train_df, val_df, test_df, Config.OUTPUT_DIR)\n    \n    # 5. Test DICOM Loading on one patient\n    test_pid = train_df['StudyInstanceUID'].iloc[0]\n    slices = DICOMEngine.load_patient_series(test_pid)\n    logger.info(f\"Test Load Success: Patient {test_pid} has {len(slices)} slices.\")\n    \n    logger.info(\"=== Phase 1 Complete ===\")\n    return train_df, val_df, test_df\n\n# Execute Pipeline\ntrain_meta, val_meta, test_meta = run_phase1()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-11T14:24:16.376086Z","iopub.execute_input":"2026-03-11T14:24:16.376364Z","iopub.status.idle":"2026-03-11T14:24:39.483835Z","shell.execute_reply.started":"2026-03-11T14:24:16.376336Z","shell.execute_reply":"2026-03-11T14:24:39.483374Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# phase 2","metadata":{}},{"cell_type":"code","source":"# %% [code]\n# =============================================================================\n# PHASE 2: SECTION 0 - SHARED DATASET CLASS & UTILITIES\n# =============================================================================\n# Purpose: Common code used by ALL model training cells\n# Run this ONCE before training any model\n# =============================================================================\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport pydicom\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nimport cv2\nimport timm\nimport os\nimport json\nfrom sklearn.metrics import roc_auc_score\nimport time\nfrom datetime import datetime\nimport gc\n\nprint(\"=\" * 70)\nprint(\" \" * 20 + \"PHASE 2: SHARED UTILITIES\")\nprint(\"=\" * 70)\n\n# =============================================================================\n# Configuration\n# =============================================================================\n\nclass TrainingConfig:\n    METADATA_PATH = \"/kaggle/working/phase1_output\"\n    DATA_PATH = \"/kaggle/input/rsna-2022-cervical-spine-fracture-detection\"\n    IMG_SIZE = 224\n    NUM_WORKERS = 2\n    DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    SEED = 42\n    SAVE_BASE = '/kaggle/working/saved_models'\n\n# Set seeds for reproducibility\ntorch.manual_seed(TrainingConfig.SEED)\nnp.random.seed(TrainingConfig.SEED)\n\nprint(f\"✅ Device: {TrainingConfig.DEVICE}\")\nif torch.cuda.is_available():\n    print(f\"✅ GPU: {torch.cuda.get_device_name(0)}\")\n    print(f\"✅ GPU Memory: {torch.cuda.get_device_properties(0).total_memory / 1e9:.2f} GB\")\n\nos.makedirs(TrainingConfig.SAVE_BASE, exist_ok=True)\nprint(f\"✅ Save directory: {TrainingConfig.SAVE_BASE}\")\nprint(\"=\" * 70)\n\n# =============================================================================\n# Dataset Class (Shared by All Models)\n# =============================================================================\n\nclass CervicalSpineDataset(Dataset):\n    def __init__(self, metadata_path, transform=True, slices_per_patient=8, \n                 max_dicom_files=50):\n        self.df = pd.read_csv(metadata_path)\n        self.transform = transform\n        self.slices_per_patient = slices_per_patient\n        self.max_dicom_files = max_dicom_files\n        self.data_path = Path(\"/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train_images\")\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def load_patient_slices(self, patient_uid):\n        patient_path = self.data_path / patient_uid\n        if not patient_path.exists():\n            return np.zeros((self.slices_per_patient, TrainingConfig.IMG_SIZE, \n                           TrainingConfig.IMG_SIZE))\n        \n        dicom_files = sorted(patient_path.glob(\"*.dcm\"), key=lambda x: int(x.stem))\n        slices = []\n        \n        for dicom_file in dicom_files[:self.max_dicom_files]:\n            try:\n                ds = pydicom.dcmread(dicom_file, force=True)\n                pixel_array = ds.pixel_array.astype(np.float32)\n                slope = float(ds.RescaleSlope) if hasattr(ds, 'RescaleSlope') else 1.0\n                intercept = float(ds.RescaleIntercept) if hasattr(ds, 'RescaleIntercept') else 0.0\n                hu = pixel_array * slope + intercept\n                hu = np.clip(hu, -1000, 2000)\n                hu = (hu - (-1000)) / (2000 - (-1000))\n                hu = cv2.resize(hu, (TrainingConfig.IMG_SIZE, TrainingConfig.IMG_SIZE))\n                slices.append(hu)\n            except:\n                continue\n        \n        if len(slices) == 0:\n            return np.zeros((self.slices_per_patient, TrainingConfig.IMG_SIZE, \n                           TrainingConfig.IMG_SIZE))\n        \n        if len(slices) > self.slices_per_patient:\n            indices = np.linspace(0, len(slices)-1, self.slices_per_patient, dtype=int)\n            slices = [slices[i] for i in indices]\n        else:\n            while len(slices) < self.slices_per_patient:\n                slices.append(slices[-1].copy())\n        \n        return np.array(slices)\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        patient_uid = row['StudyInstanceUID']\n        label = row['target_c1_c2']\n        slices = self.load_patient_slices(patient_uid)\n        \n        if self.transform and np.random.random() > 0.5:\n            slices = np.flip(slices, axis=2).copy()\n        \n        images = torch.from_numpy(slices).float().unsqueeze(1)\n        label = torch.tensor(label, dtype=torch.float32)\n        return images, label\n\n# =============================================================================\n# Loss Function (Shared by All Models)\n# =============================================================================\n\nclass WeightedBCELoss(nn.Module):\n    def __init__(self, pos_weight=3.0):\n        super().__init__()\n        self.pos_weight = pos_weight\n    \n    def forward(self, outputs, labels):\n        outputs = outputs.view(-1)\n        labels = labels.view(-1)\n        bce = nn.BCELoss(reduction='none')(outputs, labels)\n        weights = torch.where(labels == 1, self.pos_weight, 1.0)\n        return (bce * weights).mean()\n\n# =============================================================================\n# Training Loop (Shared by All Models)\n# =============================================================================\n\nclass ModelTrainer:\n    def __init__(self, model, train_loader, val_loader, device, \n                 learning_rate=2e-4, epochs=40, grad_accum=4, model_name='model'):\n        self.model = model.to(device)\n        self.train_loader = train_loader\n        self.val_loader = val_loader\n        self.criterion = WeightedBCELoss(pos_weight=3.0)\n        self.optimizer = optim.AdamW(model.parameters(), lr=learning_rate, \n                                     weight_decay=1e-3)\n        self.scheduler = optim.lr_scheduler.CosineAnnealingLR(self.optimizer, \n                                                              T_max=epochs)\n        self.device = device\n        self.epochs = epochs\n        self.grad_accum = grad_accum\n        self.model_name = model_name\n        self.best_auc = 0\n        self.history = {'train_loss': [], 'val_loss': [], 'val_auc': []}\n    \n    def train_epoch(self):\n        self.model.train()\n        total_loss = 0\n        accumulated_loss = 0\n        \n        for batch_idx, (images, labels) in enumerate(self.train_loader):\n            images, labels = images.to(self.device), labels.to(self.device)\n            \n            outputs = self.model(images)\n            outputs = outputs.view(-1)\n            labels = labels.view(-1)\n            \n            loss = self.criterion(outputs, labels) / self.grad_accum\n            loss.backward()\n            \n            accumulated_loss += loss.item() * self.grad_accum\n            \n            if (batch_idx + 1) % self.grad_accum == 0:\n                torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1.0)\n                self.optimizer.step()\n                self.optimizer.zero_grad()\n                \n                del images, labels, outputs\n                torch.cuda.empty_cache()\n                gc.collect()\n            \n            total_loss += accumulated_loss\n            accumulated_loss = 0\n        \n        if len(self.train_loader) % self.grad_accum != 0:\n            torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1.0)\n            self.optimizer.step()\n            self.optimizer.zero_grad()\n            torch.cuda.empty_cache()\n        \n        return total_loss / len(self.train_loader)\n    \n    def validate(self):\n        self.model.eval()\n        total_loss = 0\n        all_preds = []\n        all_labels = []\n        \n        with torch.no_grad():\n            for images, labels in self.val_loader:\n                images, labels = images.to(self.device), labels.to(self.device)\n                outputs = self.model(images)\n                \n                outputs = outputs.view(-1)\n                labels = labels.view(-1)\n                \n                loss = self.criterion(outputs, labels)\n                total_loss += loss.item()\n                all_preds.extend(outputs.cpu().numpy())\n                all_labels.extend(labels.cpu().numpy())\n                \n                del images, labels, outputs\n        \n        torch.cuda.empty_cache()\n        auc = roc_auc_score(all_labels, all_preds)\n        return total_loss / len(self.val_loader), auc\n    \n    def train(self):\n        print(f\"{'Epoch':<8} {'Train Loss':<12} {'Val Loss':<12} {'Val AUC':<10} {'Time':<10}\")\n        print(\"-\" * 70)\n        for epoch in range(self.epochs):\n            start = time.time()\n            train_loss = self.train_epoch()\n            val_loss, val_auc = self.validate()\n            self.scheduler.step()\n            elapsed = time.time() - start\n            print(f\"{epoch+1:<8} {train_loss:<12.4f} {val_loss:<12.4f} {val_auc:<10.4f} {elapsed:<10.1f}s\")\n            self.history['train_loss'].append(train_loss)\n            self.history['val_loss'].append(val_loss)\n            self.history['val_auc'].append(val_auc)\n            if val_auc > self.best_auc:\n                self.best_auc = val_auc\n                print(f\"  → 💾 Best model saved! AUC: {val_auc:.4f}\")\n            \n            torch.cuda.empty_cache()\n            gc.collect()\n        \n        print(f\"\\n✅ Training Complete! Best AUC: {self.best_auc:.4f}\")\n        return self.history\n    \n    def save_model(self, save_dir=None):\n        if save_dir is None:\n            timestamp = datetime.now().strftime(\"%Y%m%d_%H%M%S\")\n            save_dir = f'{TrainingConfig.SAVE_BASE}/{self.model_name}_{timestamp}'\n        \n        os.makedirs(save_dir, exist_ok=True)\n        \n        # Save model weights\n        torch.save({\n            'model_state_dict': self.model.state_dict(),\n            'model_name': self.model_name,\n            'best_auc': self.best_auc,\n            'epochs': self.epochs\n        }, f'{save_dir}/model_weights.pth')\n        \n        # Save training history\n        with open(f'{save_dir}/training_history.json', 'w') as f:\n            json.dump(self.history, f, indent=2)\n        \n        # Save metrics summary\n        with open(f'{save_dir}/metrics.txt', 'w') as f:\n            f.write(f\"Model: {self.model_name}\\n\")\n            f.write(f\"Timestamp: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\\n\")\n            f.write(f\"Best Validation AUC: {self.best_auc:.4f}\\n\")\n            f.write(f\"Final Validation AUC: {self.history['val_auc'][-1]:.4f}\\n\")\n            f.write(f\"Total Epochs: {self.epochs}\\n\")\n            f.write(f\"Train Loss (Final): {self.history['train_loss'][-1]:.4f}\\n\")\n            f.write(f\"Val Loss (Final): {self.history['val_loss'][-1]:.4f}\\n\")\n        \n        # Save training curves\n        import matplotlib.pyplot as plt\n        fig, axes = plt.subplots(1, 2, figsize=(12, 4))\n        axes[0].plot(self.history['train_loss'], label='Train')\n        axes[0].plot(self.history['val_loss'], label='Val')\n        axes[0].set_xlabel('Epoch')\n        axes[0].set_ylabel('Loss')\n        axes[0].set_title('Loss Curves')\n        axes[0].legend()\n        axes[0].grid(True, alpha=0.3)\n        \n        axes[1].plot(self.history['val_auc'], color='green')\n        axes[1].set_xlabel('Epoch')\n        axes[1].set_ylabel('AUC')\n        axes[1].set_title('Validation AUC')\n        axes[1].grid(True, alpha=0.3)\n        \n        plt.tight_layout()\n        plt.savefig(f'{save_dir}/training_curves.png', dpi=150)\n        plt.close()\n        \n        print(f\"💾 Model saved to: {save_dir}\")\n        return save_dir\n\n# =============================================================================\n# Helper Functions\n# =============================================================================\n\ndef create_dataloaders(batch_size=4, slices_per_patient=8):\n    \"\"\"Create train and validation dataloaders\"\"\"\n    train_dataset = CervicalSpineDataset(\n        f'{TrainingConfig.METADATA_PATH}/train_metadata.csv',\n        transform=True,\n        slices_per_patient=slices_per_patient\n    )\n    val_dataset = CervicalSpineDataset(\n        f'{TrainingConfig.METADATA_PATH}/val_metadata.csv',\n        transform=False,\n        slices_per_patient=slices_per_patient\n    )\n    \n    train_loader = DataLoader(train_dataset, batch_size=batch_size, \n                             shuffle=True, num_workers=TrainingConfig.NUM_WORKERS)\n    val_loader = DataLoader(val_dataset, batch_size=batch_size, \n                           shuffle=False, num_workers=TrainingConfig.NUM_WORKERS)\n    \n    return train_loader, val_loader, len(train_dataset), len(val_dataset)\n\nprint(\"\\n✅ SHARED UTILITIES LOADED SUCCESSFULLY!\")\nprint(\"   - CervicalSpineDataset class\")\nprint(\"   - WeightedBCELoss class\")\nprint(\"   - ModelTrainer class\")\nprint(\"   - create_dataloaders function\")\nprint(\"=\" * 70)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Phase 2 FIX: Pre-cache slices + Faster Training Config\n# Run this BEFORE the main training loop\n# ============================================================\n\n# %% [code] — Fix 1: Update augmentations (fixes warnings)\ndef get_transforms(is_train: bool):\n    if is_train and Config.AUGMENT:\n        return A.Compose([\n            A.RandomRotate90(p=0.3),\n            A.HorizontalFlip(p=0.5),\n            A.Affine(shift_limit=0.05, scale_limit=0.1, rotate_limit=15, p=0.5),  # replaces ShiftScaleRotate\n            A.RandomBrightnessContrast(brightness_limit=0.15, contrast_limit=0.15, p=0.4),\n            A.GaussNoise(p=0.3),                                  # no var_limit\n            A.CoarseDropout(num_holes_range=(1,4), hole_height_range=(16,32),\n                            hole_width_range=(16,32), p=0.3),     # updated API\n            A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n            ToTensorV2(),\n        ])\n    else:\n        return A.Compose([\n            A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n            ToTensorV2(),\n        ])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %% [code] — Fix 2: Pre-cache ALL patient slices as PNG (run ONCE, ~10-15 min)\nimport cv2\nfrom tqdm import tqdm\n\nCACHE_DIR = \"/kaggle/working/slice_cache\"\nos.makedirs(CACHE_DIR, exist_ok=True)\n\ndef extract_and_cache_slice(patient_uid: str) -> str:\n    \"\"\"\n    Extracts the representative slice, saves as PNG, returns cache path.\n    Returns path string (even if it already exists — skips reprocessing).\n    \"\"\"\n    cache_path = f\"{CACHE_DIR}/{patient_uid}.png\"\n    if os.path.exists(cache_path):\n        return cache_path  # Already cached — skip\n\n    patient_path = Path(Config.TRAIN_IMAGES) / patient_uid\n    dcm_files = sorted(patient_path.glob(\"*.dcm\"))\n    \n    if not dcm_files:\n        # Save blank image so we don't retry every epoch\n        blank = np.zeros((Config.IMG_SIZE, Config.IMG_SIZE, 3), dtype=np.uint8)\n        cv2.imwrite(cache_path, blank)\n        return cache_path\n\n    # Sort by InstanceNumber\n    def safe_instance(f):\n        try:\n            return int(pydicom.dcmread(f, stop_before_pixels=True).get(\"InstanceNumber\", 0))\n        except:\n            return 0\n    dcm_files = sorted(dcm_files, key=safe_instance)\n\n    # Use top-third slice (C1/C2 is upper cervical spine)\n    idx = max(0, len(dcm_files) // 6)\n    \n    try:\n        ds  = pydicom.dcmread(dcm_files[idx], force=True)\n        arr = ds.pixel_array.astype(np.float32)\n        slope     = float(getattr(ds, \"RescaleSlope\",     1.0))\n        intercept = float(getattr(ds, \"RescaleIntercept\", 0.0))\n        arr = arr * slope + intercept\n        arr = np.clip(arr, Config.HU_MIN, Config.HU_MAX)\n        arr = (arr - Config.HU_MIN) / (Config.HU_MAX - Config.HU_MIN) * 255.0\n        arr = arr.astype(np.uint8)\n        arr = cv2.resize(arr, (Config.IMG_SIZE, Config.IMG_SIZE))\n        arr_rgb = np.stack([arr, arr, arr], axis=-1)\n        cv2.imwrite(cache_path, arr_rgb)\n    except Exception as e:\n        logger.warning(f\"Cache failed for {patient_uid}: {e}\")\n        blank = np.zeros((Config.IMG_SIZE, Config.IMG_SIZE, 3), dtype=np.uint8)\n        cv2.imwrite(cache_path, blank)\n\n    return cache_path\n\n\ndef cache_all_patients(df_list):\n    \"\"\"Pre-caches slices for all patients across all splits.\"\"\"\n    all_uids = set()\n    for df in df_list:\n        all_uids.update(df[\"StudyInstanceUID\"].tolist())\n    \n    logger.info(f\"Caching {len(all_uids)} patients → {CACHE_DIR}\")\n    failed = 0\n    for uid in tqdm(all_uids, desc=\"Caching slices\"):\n        try:\n            extract_and_cache_slice(uid)\n        except Exception as e:\n            failed += 1\n            logger.warning(f\"Failed: {uid} — {e}\")\n    \n    logger.info(f\"Caching complete. Failed: {failed}/{len(all_uids)}\")\n\n\n# Load splits and run caching\ntrain_df = pd.read_csv(f\"{Config.PHASE1_OUTPUT}/train_metadata.csv\")\nval_df   = pd.read_csv(f\"{Config.PHASE1_OUTPUT}/val_metadata.csv\")\ntest_df  = pd.read_csv(f\"{Config.PHASE1_OUTPUT}/test_metadata.csv\")\n\ncache_all_patients([train_df, val_df, test_df])","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}