{"cells":[{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"#!/usr/bin/env python\n# coding: utf-8\n\n# In[ ]:\n\n\n# --- 1. ENTERPRISE ENVIRONMENT SETUP ---\n# Install decoders for JPEG Lossless DICOM decompression\nget_ipython().system('pip install -q -U python-gdcm pylibjpeg pylibjpeg-libjpeg pylibjpeg-openjpeg')\n\nimport os, glob, pydicom, cv2, gc, numpy as np, pandas as pd\nimport tensorflow as tf\nfrom tensorflow.keras import layers, mixed_precision\nfrom sklearn.model_selection import train_test_split\nimport matplotlib.pyplot as plt\n\n# CRITICAL: Disable layout optimizer to prevent EfficientNetV2 + TimeDistributed crash\ntf.config.optimizer.set_experimental_options({\"layout_optimizer\": False})\nmixed_precision.set_global_policy('mixed_float16')\n\n# GLOBAL CONFIG\nIMG_SIZE = 224\nSEQ_LEN = 16     # Halved depth for VRAM safety (P100 stable)\nBATCH_SIZE = 2   # Conservative batching to prevent 500s OOM failure\nINPUT_DIR = '/kaggle/input/rsna-2022-cervical-spine-fracture-detection'\nDISEASES = ['C1', 'C2', 'C3', 'C4', 'C5', 'C6', 'C7', 'patient_overall']\n\n# --- 2. HIGH-PERFORMANCE DATA GENERATOR ---\nclass RSNADataGenerator(tf.keras.utils.Sequence):\n    def __init__(self, df, base_path, batch_size=2, mode='train', **kwargs):\n        super().__init__(**kwargs)\n        self.df = df\n        self.base_path = base_path\n        self.batch_size = batch_size\n        self.mode = mode\n        self.study_map = {}\n        from concurrent.futures import ThreadPoolExecutor\n        self.executor = ThreadPoolExecutor(max_workers=8) # Persistent pool\n        self._precompute_file_lists()\n\n    def _precompute_file_lists(self):\n        \"\"\"Map study IDs to sorted DICOM file lists using fast scandir.\"\"\"\n        print(f\"Precomputing file lists for {len(self.df)} studies...\")\n        import os\n        base_dir = os.path.join(self.base_path, 'train_images')\n        for study_id in self.df['StudyInstanceUID'].unique():\n            study_path = os.path.join(base_dir, study_id)\n            if os.path.isdir(study_path):\n                # Faster than glob.glob or os.listdir\n                files = [os.path.join(study_path, f.name) for f in os.scandir(study_path) if f.name.endswith('.dcm')]\n                self.study_map[study_id] = sorted(files, key=lambda x: int(os.path.basename(x).split('.')[0]))\n            else:\n                self.study_map[study_id] = []\n\n    def __len__(self):\n        return int(np.ceil(len(self.df) / self.batch_size))\n\n    def __getitem__(self, idx):\n        batch = self.df.iloc[idx * self.batch_size : (idx + 1) * self.batch_size]\n        X, y = [], []\n        \n        for _, row in batch.iterrows():\n            study_id = row['StudyInstanceUID']\n            files = self.study_map.get(study_id, [])\n            \n            if not files:\n                X.append(np.zeros((SEQ_LEN, IMG_SIZE, IMG_SIZE, 3), dtype=np.float32))\n            else:\n                indices = np.linspace(0, len(files) - 1, SEQ_LEN).astype(int)\n                \n                def process_dicom(f_idx):\n                    try:\n                        ds = pydicom.dcmread(files[f_idx])\n                        img = ds.pixel_array.astype(np.float32)\n                        # Bone windowing\n                        lower, upper = 400.0 - 900.0, 400.0 + 900.0\n                        img = np.clip(img, lower, upper)\n                        img = (img - lower) / (upper - lower + 1e-7)\n                        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n                        return np.stack([img]*3, axis=-1)\n                    except:\n                        return np.zeros((IMG_SIZE, IMG_SIZE, 3), dtype=np.float32)\n                \n                # Use persistent class-level executor\n                volume = list(self.executor.map(process_dicom, indices))\n                X.append(volume)\n                \n            if self.mode == 'train':\n                y.append(row[DISEASES].values.astype(np.float32))\n        \n        return np.array(X, dtype=np.float32), np.array(y, dtype=np.float32)\n\n# --- 3. SOTA ARCHITECTURE: EfficientNetV2-M + BI-GRU ---\ndef build_master_model():\n    inputs = layers.Input(shape=(SEQ_LEN, IMG_SIZE, IMG_SIZE, 3))\n    \n    # EfficientNetV2-B0 Backbone (7M params) to strictly prevent any OOM/Kernel memory limits\n    backbone = tf.keras.applications.EfficientNetV2B0(weights='imagenet', include_top=False, pooling='avg')\n    backbone.trainable = False\n    \n    # TimeDistributed extracts spatial features from the 3D sequence\n    x = layers.TimeDistributed(backbone)(inputs)\n    \n    # Bidirectional GRU models fracture continuity across slices\n    x = layers.Bidirectional(layers.GRU(256, return_sequences=False))(x)\n    \n    # AxonFlow Style Vector Head\n    x = layers.Dense(1024, activation='swish')(x)\n    x = layers.Dropout(0.5)(x) # High-dropout for small batch stability\n    \n    outputs = layers.Dense(len(DISEASES), activation='sigmoid', dtype='float32')(x)\n    \n    # Competition Weighted Loss: 7x importance for 'patient_overall'\n    def rsna_weighted_loss(y_true, y_pred):\n        weights = tf.constant([1, 1, 1, 1, 1, 1, 1, 7], dtype=tf.float32)\n        return tf.reduce_mean(tf.keras.losses.binary_crossentropy(y_true, y_pred) * weights)\n\n    model = tf.keras.Model(inputs, outputs)\n    # Use BinaryAccuracy named 'accuracy' explicitly to track exactly how close we are to 100% accuracy\n    model.compile(optimizer=tf.keras.optimizers.Adam(1e-4), loss=rsna_weighted_loss, \n                  metrics=[tf.keras.metrics.BinaryAccuracy(name='accuracy'), tf.keras.metrics.AUC(name='auc')])\n    return model\n\n# --- 4. EXECUTION & CLINICAL AUDIT DASHBOARD ---\ntrain_df = pd.read_csv(f'{INPUT_DIR}/train.csv')\ntrain_df = train_df[train_df['StudyInstanceUID'].isin(os.listdir(f'{INPUT_DIR}/train_images'))]\ntrain_idx, val_idx = train_test_split(train_df, test_size=0.1, random_state=42)\n\nval_gen = RSNADataGenerator(val_idx, INPUT_DIR)\nmodel = build_master_model()\n\n# Fit with Callbacks for SOTA convergence\nclass GarbageCollectorCallback(tf.keras.callbacks.Callback):\n    def on_epoch_end(self, epoch, logs=None):\n        gc.collect()\n        tf.keras.backend.clear_session()\n\nmodel.fit(\n    RSNADataGenerator(train_idx, INPUT_DIR), \n    validation_data=val_gen, \n    epochs=10, \n    callbacks=[\n        tf.keras.callbacks.ModelCheckpoint(\"Best_SOTA_Spine.keras\", save_best_only=True),\n        tf.keras.callbacks.ReduceLROnPlateau(monitor='val_loss', factor=0.5, patience=2),\n        GarbageCollectorCallback()\n    ]\n)\n\n# Visual Dashboard for Management Efficiency Audit\nX_val, y_val = val_gen[0]\npreds = model.predict(X_val, verbose=0)\nfig, axes = plt.subplots(1, 4, figsize=(24, 10))\nfor i in range(4):\n    mid_slice = X_val[i][SEQ_LEN // 2]\n    img_disp = (mid_slice - mid_slice.min()) / (mid_slice.max() - mid_slice.min() + 1e-7)\n    axes[i].imshow(img_disp, cmap='bone')\n    axes[i].set_title(f\"ID: {val_idx.iloc[i]['StudyInstanceUID'][-5:]}\\nPred Overall: {preds[i][-1]:.1%}\", \n                      color='red' if preds[i][-1] > 0.5 else 'lime', weight='bold')\n    axes[i].axis('off')\nplt.show()\n\nmodel.save(\"RSNA_Spine_Final_SOTA.keras\")\n\n"}],"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.10"}},"nbformat":4,"nbformat_minor":5}