{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":59093,"databundleVersionId":7469972},{"sourceType":"kernelVersion","sourceId":233130142}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# HMS - Harmful Brain Activity Classification\n## Improved Version: Patient-ID Split + Residual CNN + KL-Divergence Loss\n\n**Author:** Prithwi  \n**Based on:** deepak412/hms-dataexplore-cnn  \n\n### Key Changes from Baseline\n| Aspect | Baseline | This Notebook |\n|---|---|---|\n| Split strategy | By `eeg_id` (data leakage) | By `patient_id` (zero leakage) |\n| EEG channels | 3 raw channels | 8 differential pairs (double banana montage) |\n| Labels | Hard one-hot | Soft probability from vote distribution |\n| Loss | Categorical cross-entropy | KL-Divergence (matches Kaggle metric) |\n| Architecture | Plain sequential CNN | Residual CNN with BatchNorm |\n| Data used | 20,000 / 106,800 rows | All 106,800 rows |\n| LR scheduling | None (LR=1e-5 fixed) | ReduceLROnPlateau + EarlyStopping |\n| Model saving | Manual only | Auto ModelCheckpoint on best val_loss |","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\n\n# Prevent TensorFlow from grabbing all GPU memory at startup\nfor gpu in tf.config.list_physical_devices('GPU'):\n    tf.config.experimental.set_memory_growth(gpu, True)\nprint(\"Memory growth enabled for all GPUs.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-04T05:39:59.943741Z","iopub.execute_input":"2026-04-04T05:39:59.944191Z","iopub.status.idle":"2026-04-04T05:40:14.455359Z","shell.execute_reply.started":"2026-04-04T05:39:59.944166Z","shell.execute_reply":"2026-04-04T05:40:14.454726Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 1. Imports","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport warnings\nimport os\nimport gc\nwarnings.filterwarnings('ignore')\n\nimport tensorflow as tf\nfrom tensorflow import keras\nfrom tensorflow.keras import layers, Model\nfrom tensorflow.keras.callbacks import ReduceLROnPlateau, EarlyStopping, ModelCheckpoint\nfrom tensorflow.keras.utils import to_categorical\n\nfrom sklearn.metrics import classification_report, confusion_matrix, accuracy_score\n\nprint(\"TensorFlow version:\", tf.__version__)\ngpus = tf.config.list_physical_devices('GPU')\nprint(\"GPUs available:\", gpus)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-04T05:40:14.456321Z","iopub.execute_input":"2026-04-04T05:40:14.456798Z","iopub.status.idle":"2026-04-04T05:40:15.508718Z","shell.execute_reply.started":"2026-04-04T05:40:14.456778Z","shell.execute_reply":"2026-04-04T05:40:15.507827Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Configuration","metadata":{}},{"cell_type":"code","source":"# ─── Paths ────────────────────────────────────────────────────────────────────\nBASE_DIR   = \"/kaggle/input/hms-harmful-brain-activity-classification/\"\nWORK_DIR   = \"/kaggle/working/\"\nMODEL_PATH = os.path.join(WORK_DIR, \"hms_best_model.keras\")\n\n# ─── EEG Settings ─────────────────────────────────────────────────────────────\nFS         = 200          # Sampling rate: 200 Hz\nWINDOW_SEC = 50           # Each labeled window = 50 seconds\nN_SAMPLES  = FS * WINDOW_SEC   # = 10,000 time steps per window\n\n# Double Banana Montage: 8 differential pairs\n# Each pair = electrode_A minus electrode_B\n# This removes common-mode noise and captures spatial gradients between electrodes\n# Left temporal chain + Right temporal chain\nDIFF_PAIRS = [\n    ('Fp1', 'F7'),   # Frontal-left to temporal-left\n    ('F7',  'T3'),   # Temporal chain left\n    ('T3',  'T5'),   # Mid-temporal to parietal-temporal left\n    ('T5',  'O1'),   # Parietal-temporal to occipital left\n    ('Fp2', 'F8'),   # Frontal-right to temporal-right\n    ('F8',  'T4'),   # Temporal chain right\n    ('T4',  'T6'),   # Mid-temporal to parietal-temporal right\n    ('T6',  'O2'),   # Parietal-temporal to occipital right\n]\nN_CHANNELS = len(DIFF_PAIRS)   # 8 channels\n\n# ─── Target Classes ───────────────────────────────────────────────────────────\nCLASSES   = ['seizure', 'lpd', 'gpd', 'lrda', 'grda', 'other']\nN_CLASSES = len(CLASSES)\nVOTE_COLS = [f\"{c}_vote\" for c in CLASSES]\n\n# ─── Split Ratios (Patient Level) ─────────────────────────────────────────────\nVAL_RATIO  = 0.15\nTEST_RATIO = 0.15\n# Train gets the remaining 70%\n\n# ─── Training Hyperparameters ─────────────────────────────────────────────────\nBATCH_SIZE    = 64\nEPOCHS        = 50\nLEARNING_RATE = 1e-4    # Start higher, ReduceLROnPlateau will decay it\n\nSEED = 42\nnp.random.seed(SEED)\ntf.random.set_seed(SEED)\n\nprint(f\"Channels: {N_CHANNELS}, Samples per window: {N_SAMPLES}\")\nprint(f\"Classes: {CLASSES}\")\nprint(f\"Vote columns: {VOTE_COLS}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-04T05:40:15.509656Z","iopub.execute_input":"2026-04-04T05:40:15.510446Z","iopub.status.idle":"2026-04-04T05:40:15.517896Z","shell.execute_reply.started":"2026-04-04T05:40:15.510419Z","shell.execute_reply":"2026-04-04T05:40:15.517049Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Load Metadata","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv(f\"{BASE_DIR}train.csv\")\nprint(\"Shape:\", df.shape)\nprint(\"\\nColumns:\", df.columns.tolist())\nprint(\"\\nNull values:\\n\", df.isnull().sum())\ndf.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-04T05:40:15.519754Z","iopub.execute_input":"2026-04-04T05:40:15.519969Z","iopub.status.idle":"2026-04-04T05:40:15.889023Z","shell.execute_reply.started":"2026-04-04T05:40:15.519952Z","shell.execute_reply":"2026-04-04T05:40:15.888335Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Unique patients:\",   df['patient_id'].nunique())\nprint(\"Unique eeg_ids:\",    df['eeg_id'].nunique())\nprint(\"Total labeled rows:\", len(df))\nprint(\"\\nClass distribution (expert_consensus):\")\nprint(df['expert_consensus'].value_counts())\nprint(\"\\nAvg rows per patient:\", len(df) / df['patient_id'].nunique())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-04T05:40:15.889707Z","iopub.execute_input":"2026-04-04T05:40:15.889898Z","iopub.status.idle":"2026-04-04T05:40:15.909551Z","shell.execute_reply.started":"2026-04-04T05:40:15.889883Z","shell.execute_reply":"2026-04-04T05:40:15.908826Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Smart sampling ────────────────────────────────────────────────────────\n# Keep only the highest confidence window per eeg_id\n# Reduces 106,800 rows → ~17,089 rows (one per EEG recording)\n# Quality improves because low-confidence noisy windows are dropped\n\ndf['max_vote'] = df[['seizure_vote','lpd_vote','gpd_vote',\n                      'lrda_vote','grda_vote','other_vote']].max(axis=1)\n\ndf_sampled = (df.sort_values('max_vote', ascending=False)\n                .drop_duplicates(subset='eeg_id', keep='first')\n                .reset_index(drop=True))\n\nprint(f\"Before sampling : {len(df):,} rows\")\nprint(f\"After sampling  : {len(df_sampled):,} rows\")\nprint(f\"Unique patients : {df_sampled['patient_id'].nunique():,}\")\nprint(f\"\\nClass distribution after sampling:\")\nprint(df_sampled['expert_consensus'].value_counts())\n\n# Replace df with sampled version — everything after this uses df_sampled\ndf = df_sampled","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-04T05:40:15.910310Z","iopub.execute_input":"2026-04-04T05:40:15.910581Z","iopub.status.idle":"2026-04-04T05:40:15.957742Z","shell.execute_reply.started":"2026-04-04T05:40:15.910559Z","shell.execute_reply":"2026-04-04T05:40:15.957088Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Patient-ID Based Split\n\n### Why this matters\nThe original notebook splits rows randomly by `eeg_id`. One patient can have ~55 labeled windows on average. With a random split, the same patient's windows appear in **train, val, AND test simultaneously**. The model learns patient-specific brain patterns and then is evaluated on different windows from the same patient — this is data leakage, and it inflates test accuracy by giving a falsely optimistic picture of generalization.\n\n**Our approach:** Group all rows belonging to each patient, then split at the patient level. A patient's data goes to exactly one of train / val / test. The model never sees any signal from a test patient during training.","metadata":{}},{"cell_type":"code","source":"# Get all unique patients and shuffle\nall_patients = df['patient_id'].unique()\nnp.random.shuffle(all_patients)\n\nn_total = len(all_patients)\nn_test  = int(n_total * TEST_RATIO)\nn_val   = int(n_total * VAL_RATIO)\nn_train = n_total - n_test - n_val\n\ntrain_patients = set(all_patients[:n_train])\nval_patients   = set(all_patients[n_train : n_train + n_val])\ntest_patients  = set(all_patients[n_train + n_val:])\n\ndf_train = df[df['patient_id'].isin(train_patients)].reset_index(drop=True)\ndf_val   = df[df['patient_id'].isin(val_patients)].reset_index(drop=True)\ndf_test  = df[df['patient_id'].isin(test_patients)].reset_index(drop=True)\n\nprint(f\"Patients  → train: {len(train_patients):,}, val: {len(val_patients):,}, test: {len(test_patients):,}\")\nprint(f\"Rows      → train: {len(df_train):,}, val: {len(df_val):,}, test: {len(df_test):,}\")\n\n# Verify zero overlap\nassert len(train_patients & val_patients)  == 0, \"Train/Val overlap!\"\nassert len(train_patients & test_patients) == 0, \"Train/Test overlap!\"\nassert len(val_patients   & test_patients) == 0, \"Val/Test overlap!\"\nprint(\"\\n✓ Zero patient overlap between all three splits\")\n\n# Show class distribution in each split\nprint(\"\\nClass distribution per split:\")\nfor name, split_df in [('Train', df_train), ('Val', df_val), ('Test', df_test)]:\n    counts = split_df['expert_consensus'].value_counts(normalize=True).round(3)\n    print(f\"  {name}: {counts.to_dict()}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-04T05:40:15.958400Z","iopub.execute_input":"2026-04-04T05:40:15.958655Z","iopub.status.idle":"2026-04-04T05:40:15.976975Z","shell.execute_reply.started":"2026-04-04T05:40:15.958636Z","shell.execute_reply":"2026-04-04T05:40:15.976349Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. EEG Feature Extraction\n\n### Differential Pairs (Double Banana Montage)\nInstead of using raw electrode voltages (which contain large shared baseline drift), we compute the **difference between adjacent electrodes**. This:\n- Cancels common-mode noise shared between electrodes\n- Highlights local voltage gradients (where brain activity actually manifests)\n- Is the standard clinical EEG interpretation method\n\n### Soft Labels from Vote Distribution\nInstead of converting votes to hard one-hot labels (losing annotator uncertainty), we normalize the 6 vote columns into a probability distribution. A sample where 4 annotators said Seizure and 1 said LPD gets target `[0.8, 0.2, 0, 0, 0, 0]` instead of `[1, 0, 0, 0, 0, 0]`.","metadata":{}},{"cell_type":"code","source":"import math\n\nclass EEGDataGenerator(keras.utils.Sequence):\n    \"\"\"\n    Loads EEG data batch-by-batch from disk.\n    Never loads the full dataset into RAM.\n    Memory usage = one batch at a time (~10-50 MB vs ~34 GB).\n    \"\"\"\n\n    def __init__(self, df_split, batch_size=32, shuffle=True):\n        self.df          = df_split.reset_index(drop=True)\n        self.batch_size  = batch_size\n        self.shuffle     = shuffle\n        self.indices     = np.arange(len(self.df))\n        self._eeg_cache  = {}\n        self._cache_keys = []\n        self.MAX_CACHE   = 200  # keep at most 50 EEG files in RAM at once\n        if self.shuffle:\n            np.random.shuffle(self.indices)\n\n    def __len__(self):\n        \"\"\"Number of batches per epoch.\"\"\"\n        return math.ceil(len(self.df) / self.batch_size)\n\n    def __getitem__(self, batch_idx):\n        \"\"\"Load and return one batch.\"\"\"\n        start         = batch_idx * self.batch_size\n        end           = min(start + self.batch_size, len(self.df))\n        batch_indices = self.indices[start:end]\n\n        X_batch = np.zeros((len(batch_indices), N_CHANNELS, N_SAMPLES, 1),\n                           dtype=np.float32)\n        y_batch = np.zeros((len(batch_indices), N_CLASSES),\n                           dtype=np.float32)\n\n        for i, idx in enumerate(batch_indices):\n            row    = self.df.iloc[idx]\n            eeg_id = row['eeg_id']\n            offset = int(row['eeg_label_offset_seconds'])\n\n            arr    = self._get_eeg(eeg_id)\n\n            start_t = FS * offset\n            window  = arr[:, start_t : start_t + N_SAMPLES]\n\n            if window.shape[1] < N_SAMPLES:\n                pad    = np.zeros((N_CHANNELS, N_SAMPLES - window.shape[1]),\n                                  dtype=np.float32)\n                window = np.concatenate([window, pad], axis=1)\n\n            X_batch[i, :, :, 0] = window\n            y_batch[i]          = self._soft_label(row)\n\n        return X_batch, y_batch\n\n    def _get_eeg(self, eeg_id):\n        \"\"\"Return cached normalized array, or load from disk.\"\"\"\n        if eeg_id not in self._eeg_cache:\n            if len(self._cache_keys) >= self.MAX_CACHE:\n                oldest = self._cache_keys.pop(0)\n                del self._eeg_cache[oldest]\n\n            eeg_df = pd.read_parquet(\n                f\"{BASE_DIR}train_eegs/{eeg_id}.parquet\"\n            )\n            arr                     = self._extract_diff(eeg_df)\n            arr                     = self._normalize(arr)\n            self._eeg_cache[eeg_id] = arr\n            self._cache_keys.append(eeg_id)\n\n        return self._eeg_cache[eeg_id]\n\n    def _extract_diff(self, eeg_df):\n        \"\"\"Compute 8 differential pairs (double banana montage).\"\"\"\n        channels = []\n        for ch_a, ch_b in DIFF_PAIRS:\n            diff = (eeg_df[ch_a].to_numpy(dtype=np.float32)\n                  - eeg_df[ch_b].to_numpy(dtype=np.float32))\n            diff = np.nan_to_num(diff, nan=0.0)\n            diff = np.clip(diff, -1024, 1024)\n            channels.append(diff)\n        return np.stack(channels, axis=0)   # (8, T_total)\n\n    def _normalize(self, arr):\n        \"\"\"Z-score normalize each channel independently.\"\"\"\n        mean = arr.mean(axis=1, keepdims=True)\n        std  = arr.std(axis=1,  keepdims=True)\n        return (arr - mean) / (std + 1e-6)\n\n    def _soft_label(self, row):\n        \"\"\"Convert vote counts to probability distribution.\"\"\"\n        votes = row[VOTE_COLS].values.astype(np.float32)\n        total = votes.sum()\n        if total == 0:\n            return np.ones(N_CLASSES, dtype=np.float32) / N_CLASSES\n        return votes / total\n\n    def on_epoch_end(self):\n        \"\"\"Shuffle indices after each epoch.\"\"\"\n        if self.shuffle:\n            np.random.shuffle(self.indices)\n\nprint(\"EEGDataGenerator class defined successfully.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-04T05:40:15.977675Z","iopub.execute_input":"2026-04-04T05:40:15.977947Z","iopub.status.idle":"2026-04-04T05:40:15.992582Z","shell.execute_reply.started":"2026-04-04T05:40:15.977921Z","shell.execute_reply":"2026-04-04T05:40:15.992020Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_gen = EEGDataGenerator(df_train, batch_size=BATCH_SIZE, shuffle=True)\nval_gen   = EEGDataGenerator(df_val,   batch_size=BATCH_SIZE, shuffle=False)\ntest_gen  = EEGDataGenerator(df_test,  batch_size=BATCH_SIZE, shuffle=False)\n\nprint(f\"Train batches per epoch : {len(train_gen):,}\")\nprint(f\"Val   batches per epoch : {len(val_gen):,}\")\nprint(f\"Test  batches per epoch : {len(test_gen):,}\")\nprint(f\"Batch size              : {BATCH_SIZE}\")\nprint(f\"Memory per batch        : {BATCH_SIZE * N_CHANNELS * N_SAMPLES * 4 / 1e6:.1f} MB\")\nprint(f\"\\nSample check — first batch:\")\nxb, yb = train_gen[0]\nprint(f\"  X batch shape : {xb.shape}\")\nprint(f\"  y batch shape : {yb.shape}\")\nprint(f\"  X min/max     : {xb.min():.2f} / {xb.max():.2f}\")\nprint(f\"  y sample      : {yb[0].round(3)}  (should sum to 1.0: {yb[0].sum():.1f})\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-04T05:40:15.993202Z","iopub.execute_input":"2026-04-04T05:40:15.993446Z","iopub.status.idle":"2026-04-04T05:40:17.968067Z","shell.execute_reply.started":"2026-04-04T05:40:15.993426Z","shell.execute_reply":"2026-04-04T05:40:17.967444Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Convert generators to tf.data.Dataset for faster prefetching\ndef generator_to_dataset(gen):\n    dataset = tf.data.Dataset.from_generator(\n        lambda: gen,\n        output_signature=(\n            tf.TensorSpec(shape=(None, N_CHANNELS, N_SAMPLES, 1), dtype=tf.float32),\n            tf.TensorSpec(shape=(None, N_CLASSES), dtype=tf.float32)\n        )\n    )\n    return dataset.prefetch(tf.data.AUTOTUNE)\n\ntrain_ds = generator_to_dataset(train_gen)\nval_ds   = generator_to_dataset(val_gen)\ntest_ds  = generator_to_dataset(test_gen)\n\nprint(\"tf.data datasets created with prefetching enabled.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-04T05:40:17.970279Z","iopub.execute_input":"2026-04-04T05:40:17.970510Z","iopub.status.idle":"2026-04-04T05:40:18.506743Z","shell.execute_reply.started":"2026-04-04T05:40:17.970492Z","shell.execute_reply":"2026-04-04T05:40:18.506018Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Model Architecture — Residual CNN\n\n### Why Residual Connections?\nA plain sequential CNN (as in the baseline) suffers from the **vanishing gradient problem** as it gets deeper — gradients shrink as they backpropagate through many layers, making early layers learn very slowly. Residual connections (skip connections) add the input of a block directly to its output:\n```\noutput = F(input) + input\n```\nThis means gradients can flow directly through the skip path without degrading, allowing much deeper and more powerful networks.\n\n### Architecture Overview\n```\nInput (N, 8, 10000, 1)\n    ↓\nEntry Conv: 32 filters, kernel (1,16), stride (1,4)  → time: 10000→2500\n    ↓\nResBlock-1 (64 filters, stride 2)   → time: 2500→1250\nResBlock-2 (64 filters, stride 1)\n    ↓\nResBlock-3 (128 filters, stride 2)  → time: 1250→625\nResBlock-4 (128 filters, stride 1)\n    ↓\nResBlock-5 (256 filters, stride 2)  → time: 625→313\nResBlock-6 (256 filters, stride 1)\n    ↓\nResBlock-7 (512 filters, stride 2)  → time: 313→157\n    ↓\nCross-channel block (mixes information across EEG channels)\n    ↓\nGlobalAveragePooling2D\n    ↓\nDense(256) → Dropout(0.4) → Dense(128) → Dropout(0.3) → Dense(6, softmax)\n```","metadata":{}},{"cell_type":"code","source":"def residual_block(x, filters, kernel_size=(1, 8), strides=(1, 1), dropout_rate=0.1):\n    shortcut = x\n\n    x = layers.Conv2D(filters, kernel_size, strides=strides, padding='same',\n                      kernel_initializer='he_normal',\n                      kernel_regularizer=keras.regularizers.l2(1e-4))(x)\n    x = layers.BatchNormalization()(x)\n    x = layers.Activation('relu')(x)\n    x = layers.Dropout(dropout_rate)(x)\n\n    x = layers.Conv2D(filters, kernel_size, padding='same',\n                      kernel_initializer='he_normal',\n                      kernel_regularizer=keras.regularizers.l2(1e-4))(x)\n    x = layers.BatchNormalization()(x)\n\n    if strides != (1, 1) or shortcut.shape[-1] != filters:\n        shortcut = layers.Conv2D(filters, (1, 1), strides=strides, padding='same',\n                                 kernel_initializer='he_normal')(shortcut)\n        shortcut = layers.BatchNormalization()(shortcut)\n\n    x = layers.Add()([x, shortcut])\n    x = layers.Activation('relu')(x)\n    return x\n\n\ndef build_residual_cnn(input_shape):\n    inp = layers.Input(shape=input_shape, name='eeg_input')\n\n    # Entry block\n    x = layers.Conv2D(32, (1, 16), strides=(1, 4), padding='same',\n                      kernel_initializer='he_normal')(inp)\n    x = layers.BatchNormalization()(x)\n    x = layers.Activation('relu')(x)\n\n    # Stage 1\n    x = residual_block(x, 64,  kernel_size=(1, 8), strides=(1, 2))\n    x = residual_block(x, 64,  kernel_size=(1, 8), strides=(1, 1))\n\n    # Stage 2\n    x = residual_block(x, 128, kernel_size=(1, 8), strides=(1, 2))\n    x = residual_block(x, 128, kernel_size=(1, 8), strides=(1, 1))\n\n    # Stage 3\n    x = residual_block(x, 256, kernel_size=(1, 4), strides=(1, 2))\n\n    # Global pooling\n    x = layers.GlobalAveragePooling2D()(x)\n\n    # Classifier head — simpler than before\n    x = layers.Dense(128, activation='relu',\n                     kernel_initializer='he_normal')(x)\n    x = layers.Dropout(0.3)(x)\n    out = layers.Dense(N_CLASSES, activation='softmax',\n                       kernel_initializer='glorot_uniform')(x)\n\n    return Model(inputs=inp, outputs=out)\n\n\n# Build fresh\nINPUT_SHAPE = (N_CHANNELS, N_SAMPLES, 1)\nmodel = build_residual_cnn(INPUT_SHAPE)\n\n# Immediately verify weights are valid — no NaN at init\ndummy      = np.zeros((2, N_CHANNELS, N_SAMPLES, 1), dtype=np.float32)\ndummy_pred = model(dummy, training=False).numpy()\nprint(\"Dummy prediction on zeros:\", dummy_pred)\nprint(\"Any NaN in init weights:\", np.isnan(dummy_pred).any())\n\nmodel.summary()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-04T05:40:18.507407Z","iopub.execute_input":"2026-04-04T05:40:18.507635Z","iopub.status.idle":"2026-04-04T05:40:21.161545Z","shell.execute_reply.started":"2026-04-04T05:40:18.507619Z","shell.execute_reply":"2026-04-04T05:40:21.160827Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7. Loss Function: KL-Divergence\n\nThe Kaggle HMS competition scores submissions using **KL-Divergence** between the predicted probability distribution and the true annotator vote distribution. This is defined as:\n\n```\nKL(P || Q) = Σ P(x) * log(P(x) / Q(x))\n```\n\nWhere P = annotator vote distribution (y_true) and Q = model predictions (y_pred).\n\n**Lower KL-divergence = better.** Using this as our training loss directly aligns training with the evaluation metric — a fundamental principle of good ML practice.","metadata":{}},{"cell_type":"code","source":"def kl_divergence_loss(y_true, y_pred):\n    y_true   = tf.cast(y_true, tf.float32)\n    y_pred   = tf.clip_by_value(y_pred,   1e-6, 1.0)\n    y_true_c = tf.clip_by_value(y_true,   1e-6, 1.0)\n    y_pred   = y_pred   / tf.reduce_sum(y_pred,   axis=-1, keepdims=True)\n    y_true_c = y_true_c / tf.reduce_sum(y_true_c, axis=-1, keepdims=True)\n    kl = tf.reduce_sum(y_true_c * tf.math.log(y_true_c / y_pred), axis=-1)\n    kl = tf.clip_by_value(kl, 0.0, 100.0)\n    return tf.reduce_mean(kl)\n\n\ndef soft_accuracy(y_true, y_pred):\n    true_class = tf.argmax(y_true, axis=1)\n    pred_class = tf.argmax(y_pred, axis=1)\n    return tf.reduce_mean(tf.cast(tf.equal(true_class, pred_class), tf.float32))\n\n\noptimizer = keras.optimizers.Adam(\n    learning_rate=1e-4,\n    clipnorm=1.0\n)\n\nmodel.compile(\n    optimizer=optimizer,\n    loss=kl_divergence_loss,\n    metrics=[soft_accuracy]\n)\n\n# ── SAFE SANITY CHECK (Section 7) ───────────────────────────────────────────\n\nprint(\"Running SAFE sanity check with small batch...\")\n\nSMALL_BATCH_SIZE = 8   # Very safe for Tesla T4\n\n# Get one batch and take only first few samples\nxb_full, yb_full = train_gen[0]\nxb = xb_full[:SMALL_BATCH_SIZE]\nyb = yb_full[:SMALL_BATCH_SIZE]\n\nprint(f\"Testing with batch size = {SMALL_BATCH_SIZE}\")\nprint(f\"X shape: {xb.shape}\")\n\nxb_tensor = tf.constant(xb, dtype=tf.float32)\n\ntry:\n    pred_test = model(xb_tensor, training=False)\n    pred_np = pred_test.numpy()\n    \n    loss_test = kl_divergence_loss(tf.constant(yb, dtype=tf.float32), pred_test).numpy()\n    \n    print(f\"\\nSanity Check Results:\")\n    print(f\"  Prediction min/max : {pred_np.min():.6f} / {pred_np.max():.6f}\")\n    print(f\"  KL Loss            : {loss_test:.4f}\")\n    \n    if np.isnan(loss_test) or np.isinf(loss_test):\n        print(\"✗ LOSS IS NaN/Inf — do not start full training yet\")\n    else:\n        print(\"✓ Sanity check PASSED — model and loss are working\")\n        \nexcept Exception as e:\n    print(f\"Error in sanity check: {e}\")\n    print(\"→ Model is too heavy for current batch size\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-04T05:40:21.162813Z","iopub.execute_input":"2026-04-04T05:40:21.163077Z","iopub.status.idle":"2026-04-04T05:40:22.347193Z","shell.execute_reply.started":"2026-04-04T05:40:21.163060Z","shell.execute_reply":"2026-04-04T05:40:22.346434Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 8. Callbacks","metadata":{}},{"cell_type":"code","source":"callbacks = [\n    # Stop training when val_loss stops improving for 20 epochs\n    EarlyStopping(\n        monitor='val_loss',\n        patience=20,\n        verbose=1,\n        restore_best_weights=True\n    ),\n\n    # Halve the learning rate when val_loss plateaus for 8 epochs\n    # This allows fine-grained convergence after broad learning\n    ReduceLROnPlateau(\n        monitor='val_loss',\n        factor=0.5,\n        patience=8,\n        verbose=1,\n        min_lr=1e-7\n    ),\n\n    # Save the best model to disk automatically\n    ModelCheckpoint(\n        MODEL_PATH,\n        monitor='val_loss',\n        save_best_only=True,\n        verbose=1\n    )\n]\n\nprint(\"Callbacks configured:\")\nprint(\"  EarlyStopping    — patience=20, restores best weights\")\nprint(\"  ReduceLROnPlateau — factor=0.5, patience=8, min_lr=1e-7\")\nprint(f\"  ModelCheckpoint  — saves best model to {MODEL_PATH}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-04T05:40:22.348057Z","iopub.execute_input":"2026-04-04T05:40:22.348546Z","iopub.status.idle":"2026-04-04T05:40:22.353580Z","shell.execute_reply.started":"2026-04-04T05:40:22.348517Z","shell.execute_reply":"2026-04-04T05:40:22.352860Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── DEBUG: Check data quality before training ──────────────────────────────\nprint(\"Checking first 5 batches for NaN/Inf...\")\n\nnan_found = False\nfor batch_idx in range(5):\n    xb, yb = train_gen[batch_idx]\n    \n    x_nan = np.isnan(xb).sum()\n    x_inf = np.isinf(xb).sum()\n    y_nan = np.isnan(yb).sum()\n    y_sum = yb.sum(axis=1)  # each row should sum to 1.0\n    \n    print(f\"\\nBatch {batch_idx}:\")\n    print(f\"  X shape     : {xb.shape}\")\n    print(f\"  X NaN count : {x_nan}\")\n    print(f\"  X Inf count : {x_inf}\")\n    print(f\"  X min/max   : {xb.min():.4f} / {xb.max():.4f}\")\n    print(f\"  y NaN count : {y_nan}\")\n    print(f\"  y row sums  : min={y_sum.min():.4f}, max={y_sum.max():.4f}\")\n    \n    if x_nan > 0 or x_inf > 0 or y_nan > 0:\n        nan_found = True\n        print(\"  *** NaN/Inf FOUND IN THIS BATCH ***\")\n\nif not nan_found:\n    print(\"\\n✓ No NaN or Inf found in first 5 batches — data is clean\")\nelse:\n    print(\"\\n✗ NaN/Inf found — data pipeline needs fixing\")\n\n# ── DEBUG: Check model output on one batch ─────────────────────────────────\nprint(\"\\nChecking model output on one batch...\")\nxb, yb   = train_gen[0]\nxb_tensor = tf.constant(xb)\npred      = model(xb_tensor, training=False)\npred_np   = pred.numpy()\n\nprint(f\"  Prediction shape  : {pred_np.shape}\")\nprint(f\"  Prediction NaN    : {np.isnan(pred_np).sum()}\")\nprint(f\"  Prediction Inf    : {np.isinf(pred_np).sum()}\")\nprint(f\"  Prediction min/max: {pred_np.min():.6f} / {pred_np.max():.6f}\")\nprint(f\"  Row sums (should be ~1.0): {pred_np.sum(axis=1)[:5]}\")\n\n# ── DEBUG: Check loss on one batch ─────────────────────────────────────────\nprint(\"\\nChecking loss on one batch...\")\nloss_val = kl_divergence_loss(\n    tf.constant(yb),\n    tf.constant(pred_np)\n).numpy()\nprint(f\"  KL loss on batch 0: {loss_val}\")\n\nif np.isnan(loss_val):\n    print(\"  *** LOSS IS NaN — problem is in loss function or predictions ***\")\nelse:\n    print(\"  ✓ Loss is a valid number\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-04T05:40:22.354437Z","iopub.execute_input":"2026-04-04T05:40:22.355097Z","iopub.status.idle":"2026-04-04T05:40:32.141193Z","shell.execute_reply.started":"2026-04-04T05:40:22.355079Z","shell.execute_reply":"2026-04-04T05:40:32.140496Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 9. Training","metadata":{}},{"cell_type":"code","source":"# ── TRAINING CELL (Section 9) ───────────────────────────────────────────────\n\nprint(f\"Starting training...\")\nprint(f\"  Batch size      : {BATCH_SIZE}\")\nprint(f\"  Max epochs      : {EPOCHS}\")\nprint(f\"  Initial LR      : {LEARNING_RATE}\")\nprint(f\"  Steps per epoch : {len(train_gen):,}\")\nprint(f\"  Validation steps: {len(val_gen):,}\")\nprint()\n\nhistory = model.fit(\n    train_gen,                    # Use the generator (more stable than train_ds for now)\n    validation_data=val_gen,\n    \n    # THESE TWO LINES ARE CRITICAL\n    steps_per_epoch=len(train_gen),      # Tells Keras exactly how many batches = 1 epoch\n    validation_steps=len(val_gen),       # Same for validation\n    \n    epochs=EPOCHS,\n    callbacks=callbacks,\n    verbose=1\n)\n\nprint(\"\\nTraining complete.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-04T05:40:32.141902Z","iopub.execute_input":"2026-04-04T05:40:32.142116Z","execution_failed":"2026-04-06T19:43:30.025Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 10. Training Curves","metadata":{}},{"cell_type":"code","source":"fig, axes = plt.subplots(1, 2, figsize=(14, 5))\n\n# ── Accuracy ──────────────────────────────────────────────────────────────────\naxes[0].plot(history.history['soft_accuracy'],     label='Train', linewidth=1.5)\naxes[0].plot(history.history['val_soft_accuracy'], label='Val',   linewidth=1.5)\naxes[0].set_title('Accuracy (argmax of soft labels)', fontsize=13)\naxes[0].set_xlabel('Epoch')\naxes[0].set_ylabel('Accuracy')\naxes[0].legend()\naxes[0].grid(True, alpha=0.3)\n\n# ── KL-Divergence Loss ────────────────────────────────────────────────────────\naxes[1].plot(history.history['loss'],     label='Train KL-Loss', linewidth=1.5)\naxes[1].plot(history.history['val_loss'], label='Val KL-Loss',   linewidth=1.5)\naxes[1].set_title('KL-Divergence Loss (lower = better)', fontsize=13)\naxes[1].set_xlabel('Epoch')\naxes[1].set_ylabel('KL Divergence')\naxes[1].legend()\naxes[1].grid(True, alpha=0.3)\n\nplt.suptitle('Training History — Patient-ID Split', fontsize=14, fontweight='bold')\nplt.tight_layout()\nplt.savefig(os.path.join(WORK_DIR, 'training_curves.png'), dpi=150, bbox_inches='tight')\nplt.show()\n\nbest_epoch = np.argmin(history.history['val_loss'])\nprint(f\"Best epoch: {best_epoch+1}\")\nprint(f\"Best val KL-loss: {history.history['val_loss'][best_epoch]:.4f}\")\nprint(f\"Best val accuracy: {history.history['val_soft_accuracy'][best_epoch]*100:.2f}%\")","metadata":{"trusted":true,"execution":{"execution_failed":"2026-04-06T19:43:30.027Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 11. Evaluation on Test Set","metadata":{}},{"cell_type":"code","source":"test_kl, test_acc = model.evaluate(test_ds, verbose=0)\nprint(\"=\" * 50)\nprint(f\"TEST SET RESULTS (Patient-ID Split)\")\nprint(\"=\" * 50)\nprint(f\"KL-Divergence Loss : {test_kl:.4f}  (lower is better)\")\nprint(f\"Accuracy           : {test_acc * 100:.2f}%\")\nprint(\"=\" * 50)","metadata":{"trusted":true,"execution":{"execution_failed":"2026-04-06T19:43:30.028Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"y_pred_prob = model.predict(test_ds, verbose=1)\n\n# Reconstruct true labels from generator\ny_true_all = np.vstack([test_gen[i][1] for i in range(len(test_gen))])\ny_true_cls = np.argmax(y_true_all, axis=1)\ny_pred_cls = np.argmax(y_pred_prob, axis=1)\n\n# Trim to same length (last batch can be smaller)\nmin_len    = min(len(y_true_cls), len(y_pred_cls))\ny_true_cls = y_true_cls[:min_len]\ny_pred_cls = y_pred_cls[:min_len]\n\nprint(\"Classification Report:\")\nprint(classification_report(\n    y_true_cls, y_pred_cls,\n    target_names=[c.upper() for c in CLASSES],\n    digits=4\n))","metadata":{"trusted":true,"execution":{"execution_failed":"2026-04-06T19:43:30.028Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cm = confusion_matrix(y_true_cls, y_pred_cls)\ncm_norm = cm.astype(float) / cm.sum(axis=1, keepdims=True)\nfig, axes = plt.subplots(1, 2, figsize=(14, 5))\nclass_labels = [c.upper() for c in CLASSES]\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues',\n            xticklabels=class_labels, yticklabels=class_labels, ax=axes[0])\naxes[0].set_title('Confusion Matrix (Counts)')\naxes[0].set_ylabel('True Label')\naxes[0].set_xlabel('Predicted Label')\nsns.heatmap(cm_norm, annot=True, fmt='.2f', cmap='Blues',\n            xticklabels=class_labels, yticklabels=class_labels, ax=axes[1])\naxes[1].set_title('Confusion Matrix (Normalized)')\naxes[1].set_ylabel('True Label')\naxes[1].set_xlabel('Predicted Label')\nplt.suptitle('Test Set Results — Patient-ID Split (No Data Leakage)', fontsize=13)\nplt.tight_layout()\nplt.savefig(os.path.join(WORK_DIR, 'confusion_matrix.png'), dpi=150, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"execution_failed":"2026-04-06T19:43:30.028Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Per-class KL-Divergence (test set):\")\nfor i, cls in enumerate(CLASSES):\n    mask = y_true_cls == i\n    if mask.sum() == 0:\n        continue\n    p = np.clip(y_true_all[:min_len][mask], 1e-7, 1.0)\n    q = np.clip(y_pred_prob[:min_len][mask], 1e-7, 1.0)\n    kl = np.mean(np.sum(p * np.log(p / q), axis=1))\n    print(f\"  {cls.upper():8s}: KL={kl:.4f}  (n={mask.sum()})\")","metadata":{"trusted":true,"execution":{"execution_failed":"2026-04-06T19:43:30.028Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 12. Save Final Model","metadata":{}},{"cell_type":"code","source":"final_save_path = os.path.join(WORK_DIR, 'hms_final_patient_split.keras')\nmodel.save(final_save_path)\nprint(f\"Model saved to: {final_save_path}\")\nprint(f\"Best checkpoint also at: {MODEL_PATH}\")","metadata":{"trusted":true,"execution":{"execution_failed":"2026-04-06T19:43:30.028Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 13. Summary\n\n### Architecture\n- Entry Conv: 32 filters, kernel (1,16), stride (1,4)\n- Stage 1: 2× ResBlock(64 filters, kernel (1,8))\n- Stage 2: 2× ResBlock(128 filters, kernel (1,8))\n- Stage 3: 2× ResBlock(256 filters, kernel (1,4))\n- Stage 4: 1× ResBlock(512 filters, kernel (1,4))\n- Cross-channel block: ResBlock(256 filters, kernel (8,1))\n- GlobalAveragePooling2D\n- Dense(256) → Dropout(0.4) → Dense(128) → Dropout(0.3) → Dense(6, softmax)\n\n### Key Design Decisions\n1. **Patient-ID split** — eliminates data leakage; gives honest generalization metrics\n2. **Double banana montage** — 8 differential pairs; standard clinical EEG representation\n3. **KL-divergence loss on soft labels** — directly optimizes the Kaggle metric; respects annotator uncertainty\n4. **Residual connections + BatchNorm** — enables deeper networks; stabilizes training\n5. **ReduceLROnPlateau** — adaptive learning rate decay; squeezes out final performance\n6. **All 106,800 rows** — 5× more data than baseline's arbitrary 20,000 cap","metadata":{}}]}