{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","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":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":59093,"databundleVersionId":7469972,"isSourceIdPinned":false}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# EEG Abnormality Detection: EEGNet (Pretrained AI) vs Classic ML\n\n> **Platform:** Kaggle Notebooks | **GPU:** T4 or P100 (optional — Classic ML is CPU-only) | **Runtime:** ~15–25 minutes\n\n## Overview\n\nThis notebook fulfills **Task 1 – Medical Signals** requirements:\n\n| Requirement | Implementation |\n|---|---|\n| Multi-channel EEG with 4 abnormality types | HMS dataset: Seizure, LPD, GPD, GRDA (+ Normal) |\n| Pretrained AI model (multi-channel) | **EEGNet** trained on 8-channel HMS EEG |\n| Classic ML arrhythmia detection algorithm | **Hjorth + Band Power + Autocorrelation + Spectral Entropy → SVM / Random Forest** |\n| Comparison of both approaches | Side-by-side metrics, confusion matrices, per-class analysis |\n\n## Why HMS Micro-Subset?\n\nThe full HMS dataset is ~50 GB. We sample **80 recordings per class (400 total)** — enough for statistically meaningful comparison while running in ~15 minutes on free Kaggle GPU.\n\n## EEG Classes\n\n| Class | Full Name | EEG Signature |\n|---|---|---|\n| **Seizure** | Ictal Activity | Rhythmic, high-amplitude spike-wave complexes |\n| **LPD** | Lateralized Periodic Discharges | Periodic sharp waves, one hemisphere dominant |\n| **GPD** | Generalized Periodic Discharges | Symmetric bilateral periodic discharges |\n| **GRDA** | Generalized Rhythmic Delta Activity | Slow rhythmic 1–3 Hz delta waves, diffuse |\n| **Normal/Other** | Background Activity | Alpha/beta dominant, no pathological patterns |","metadata":{}},{"cell_type":"markdown","source":"## ⚙️ Setup Checklist\n\n1. **Attach HMS dataset** → `+ Add Data` → Competition Data → `hms-harmful-brain-activity-classification`\n2. **Load your saved EEGNet model** → Upload `eegnet_final.pt` to Kaggle input (from the training notebook's output)\n3. GPU optional — Classic ML runs on CPU","metadata":{}},{"cell_type":"markdown","source":"## 1 — Install & Imports","metadata":{}},{"cell_type":"code","source":"# Install lightweight extras (all others pre-installed on Kaggle)\n!pip install -q antropy  # spectral entropy, Hjorth params\n\nimport os, warnings, time\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport matplotlib.gridspec as gridspec\nimport seaborn as sns\nfrom pathlib import Path\nfrom collections import Counter\nwarnings.filterwarnings('ignore')\n\n# Signal processing\nfrom scipy import signal as sp_signal\nfrom scipy.signal import iirnotch, butter, sosfiltfilt, tf2sos\nfrom scipy.stats import kurtosis, skew\nimport antropy as ant  # Hjorth mobility/complexity, spectral entropy\n\n# Classic ML\nfrom sklearn.svm import SVC\nfrom sklearn.ensemble import RandomForestClassifier, GradientBoostingClassifier\nfrom sklearn.preprocessing import StandardScaler, label_binarize\nfrom sklearn.model_selection import StratifiedKFold, cross_val_score\nfrom sklearn.metrics import (\n    classification_report, confusion_matrix,\n    roc_auc_score, f1_score, accuracy_score\n)\nfrom sklearn.pipeline import Pipeline\n\n# Deep Learning\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f'Device: {DEVICE}')\nprint(f'PyTorch version: {torch.__version__}')\n\n# Reproducibility\nSEED = 42\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\n\n# ── Colour palette ──────────────────────────────────────────────────────────\nPALETTE = {\n    'Seizure': '#e74c3c',\n    'LPD':     '#e67e22',\n    'GPD':     '#f1c40f',\n    'GRDA':    '#2ecc71',\n    'Normal':  '#3498db'\n}\nCLASS_NAMES = ['Seizure', 'LPD', 'GPD', 'GRDA', 'Normal']\nprint('✅ Imports complete')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:19:23.602289Z","iopub.execute_input":"2026-02-23T21:19:23.602946Z","iopub.status.idle":"2026-02-23T21:19:36.679188Z","shell.execute_reply.started":"2026-02-23T21:19:23.602913Z","shell.execute_reply":"2026-02-23T21:19:36.678290Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Dataset & Signal Constants ────────────────────────────────────────────────\nimport glob, os\n\n# ── Verified Kaggle paths for HMS - Harmful Brain Activity Classification ─────\nDATA_DIR = '/kaggle/input/competitions/hms-harmful-brain-activity-classification'\nEEG_DIR  = '/kaggle/input/competitions/hms-harmful-brain-activity-classification/train_eegs'\n\n# ── Sanity check ──────────────────────────────────────────────────────────────\nassert os.path.exists(DATA_DIR), f\"DATA_DIR not found: {DATA_DIR}\"\nassert os.path.exists(EEG_DIR),  f\"EEG_DIR not found: {EEG_DIR}\"\n\n_n_parquet = len(glob.glob(f'{EEG_DIR}/*.parquet'))\n_csv_ok    = os.path.exists(f'{DATA_DIR}/train.csv')\nassert _n_parquet > 0, f\"No parquet files found in {EEG_DIR}\"\nassert _csv_ok,        f\"train.csv not found in {DATA_DIR}\"\n\nprint(f\"✅ DATA_DIR : {DATA_DIR}\")\nprint(f\"✅ EEG_DIR  : {EEG_DIR}\")\nprint(f\"   Parquet files : {_n_parquet:,}\")\nprint(f\"   train.csv     : found\")\n\n# ── EEG Signal constants ──────────────────────────────────────────────────────\nSFREQ             = 200            # sampling frequency (Hz)\nSEGMENT_SECONDS   = 10             # window length in seconds\nN_TIMEPOINTS      = SFREQ * SEGMENT_SECONDS   # 2000 samples\nN_CHANNELS        = 8\nN_CLASSES         = 5\nSAMPLES_PER_CLASS = 80             # 80 × 5 = 400 total — fits in ~15 min on T4\n\n# 8 bipolar channels available in HMS train_eegs parquet files\nCHANNELS = ['Fp1', 'F3', 'C3', 'P3', 'Fp2', 'F4', 'C4', 'P4']\n\nprint(f\"\\n   SFREQ={SFREQ} Hz | SEGMENT={SEGMENT_SECONDS}s | \"\n      f\"N_TIMEPOINTS={N_TIMEPOINTS} | N_CHANNELS={N_CHANNELS}\")\nprint(f\"   SAMPLES_PER_CLASS={SAMPLES_PER_CLASS} → total={SAMPLES_PER_CLASS*N_CLASSES}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:19:36.680765Z","iopub.execute_input":"2026-02-23T21:19:36.681234Z","iopub.status.idle":"2026-02-23T21:19:36.996262Z","shell.execute_reply.started":"2026-02-23T21:19:36.681193Z","shell.execute_reply":"2026-02-23T21:19:36.995500Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2 — EEGNet Architecture (must match training notebook exactly)","metadata":{}},{"cell_type":"code","source":"class EEGNet(nn.Module):\n    \"\"\"\n    EEGNet — Lawhern et al. (2018), Journal of Neural Engineering.\n    Multi-channel EEG classifier: INPUT (B, 1, C, T) → logits (B, N_classes)\n    \"\"\"\n    def __init__(self, n_channels=8, n_classes=5, n_timepoints=2000,\n                 F1=8, D=2, F2=16, dropout=0.5):\n        super().__init__()\n        self.n_channels   = n_channels\n        self.n_classes    = n_classes\n        self.n_timepoints = n_timepoints\n\n        # ── Block 1: Temporal → Depthwise Spatial ──────────────────────────\n        self.temporal_conv = nn.Sequential(\n            nn.Conv2d(1, F1, kernel_size=(1, 64), padding=(0, 32), bias=False),\n            nn.BatchNorm2d(F1)\n        )\n        self.depthwise_conv = nn.Sequential(\n            nn.Conv2d(F1, F1 * D, kernel_size=(n_channels, 1),\n                      groups=F1, bias=False),\n            nn.BatchNorm2d(F1 * D),\n            nn.ELU(),\n            nn.AvgPool2d((1, 4)),\n            nn.Dropout(dropout)\n        )\n        # ── Block 2: Separable Convolution ─────────────────────────────────\n        self.separable_conv = nn.Sequential(\n            nn.Conv2d(F1 * D, F1 * D, kernel_size=(1, 16),\n                      padding=(0, 8), groups=F1 * D, bias=False),\n            nn.Conv2d(F1 * D, F2, kernel_size=(1, 1), bias=False),\n            nn.BatchNorm2d(F2),\n            nn.ELU(),\n            nn.AvgPool2d((1, 8)),\n            nn.Dropout(dropout)\n        )\n        # ── Classifier Head ────────────────────────────────────────────────\n        flat_size = self._get_flat_size(n_channels, n_timepoints, F1, D, F2)\n        self.classifier = nn.Linear(flat_size, n_classes)\n\n    def _get_flat_size(self, C, T, F1, D, F2):\n        x = torch.zeros(1, 1, C, T)\n        x = self.temporal_conv(x)\n        x = self.depthwise_conv(x)\n        x = self.separable_conv(x)\n        return x.view(1, -1).shape[1]\n\n    def forward(self, x):           # x: (B, 1, C, T)\n        x = self.temporal_conv(x)\n        x = self.depthwise_conv(x)\n        x = self.separable_conv(x)\n        x = x.view(x.size(0), -1)\n        return self.classifier(x)\n\n    def predict_proba(self, x):\n        return F.softmax(self.forward(x), dim=1)\n\nprint('✅ EEGNet architecture defined')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:19:36.997063Z","iopub.execute_input":"2026-02-23T21:19:36.997363Z","iopub.status.idle":"2026-02-23T21:19:37.007073Z","shell.execute_reply.started":"2026-02-23T21:19:36.997341Z","shell.execute_reply":"2026-02-23T21:19:37.006477Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3 — EEG Preprocessor (same as training notebook)","metadata":{}},{"cell_type":"code","source":"class EEGPreprocessor:\n    \"\"\"Zero-phase notch (50 Hz) + bandpass (0.5–40 Hz) filter.\"\"\"\n    def __init__(self, sfreq=200., notch=50., bp_low=0.5, bp_high=40., Q=30.):\n        self.sfreq = sfreq\n        w0   = notch / (sfreq / 2.)\n        b, a = iirnotch(w0, Q)\n        self._notch = tf2sos(b, a)\n        nyq  = sfreq / 2.\n        low  = max(1e-4, bp_low)  / nyq\n        high = min(0.99, bp_high / nyq)\n        self._bp = butter(4, [low, high], btype='band', output='sos')\n\n    def __call__(self, eeg: np.ndarray) -> np.ndarray:\n        \"\"\"eeg: (n_channels, n_samples) float32\"\"\"\n        eeg = sosfiltfilt(self._notch, eeg, axis=1).astype(np.float32)\n        eeg = sosfiltfilt(self._bp,    eeg, axis=1).astype(np.float32)\n        return eeg\n\npreprocessor = EEGPreprocessor()\nprint('✅ Preprocessor ready')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:19:37.008679Z","iopub.execute_input":"2026-02-23T21:19:37.008903Z","iopub.status.idle":"2026-02-23T21:19:37.042208Z","shell.execute_reply.started":"2026-02-23T21:19:37.008877Z","shell.execute_reply":"2026-02-23T21:19:37.041251Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4 — Build HMS Micro-Subset (80 samples per class = 400 total)\n\nWe randomly sample a balanced subset so the notebook runs in ~15 minutes instead of hours. The same EEG channels and preprocessing as the training notebook are used.","metadata":{}},{"cell_type":"code","source":"# ── Load metadata & assign majority-vote label ────────────────────────────────\nmeta = pd.read_csv(f'{DATA_DIR}/train.csv')\n\nHMS_LABEL_COLS = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'grda_vote', 'other_vote']\n\n# Compute vote_matrix ONCE on the full, unfiltered DataFrame\nvote_matrix = meta[HMS_LABEL_COLS].values.astype(float)\n\n# ── Confidence filter (keep rows where majority share > 60 %) ─────────────────\nrow_sums          = vote_matrix.sum(axis=1)\nrow_sums[row_sums == 0] = 1.0          # guard against all-zero rows\nconfidence        = vote_matrix.max(axis=1) / row_sums\n\nkeep_mask   = (confidence > 0.6)\nmeta_filt   = meta[keep_mask].reset_index(drop=True)   # ← MUST reset so index == row pos\nvote_filt   = vote_matrix[keep_mask]                    # ← aligned numpy array\n\n# Assign labels from the aligned numpy array (no index skew possible)\nmeta_filt['label_idx'] = np.argmax(vote_filt, axis=1).astype(int)\nmeta_filt['label']     = [CLASS_NAMES[i] for i in meta_filt['label_idx']]\n\nprint(f\"Rows after confidence filter (>0.6): {len(meta_filt):,}\")\n\n# ── Aggregate: one majority label per eeg_id ──────────────────────────────────\nagg = (\n    meta_filt\n    .groupby(['eeg_id', 'patient_id'])\n    .agg(label_idx=('label_idx', lambda x: int(x.mode().iloc[0])))\n    .reset_index()\n)\nagg['label'] = [CLASS_NAMES[i] for i in agg['label_idx']]\n\nprint(f\"Unique EEG recordings: {len(agg):,}\")\nprint(\"\\nClass distribution before sampling:\")\nprint(agg['label'].value_counts())\n\n# ── Keep only recordings whose parquet file exists on disk ────────────────────\nexists_mask = agg['eeg_id'].apply(\n    lambda eid: os.path.exists(f'{EEG_DIR}/{eid}.parquet')\n)\nagg = agg[exists_mask].reset_index(drop=True)\nprint(f\"\\nRecordings with parquet on disk: {len(agg):,}\")\n\nif len(agg) == 0:\n    raise RuntimeError(\n        f\"No parquet files matched. EEG_DIR={EEG_DIR!r}\\n\"\n        \"Make sure the HMS dataset is attached: Notebook → + Add Data → \"\n        \"Competition Data → hms-harmful-brain-activity-classification\"\n    )\n\n# ── Stratified sampling: SAMPLES_PER_CLASS per class ─────────────────────────\nframes = []\nfor cls_idx, cls_name in enumerate(CLASS_NAMES):\n    pool = agg[agg['label_idx'] == cls_idx]\n    n    = min(SAMPLES_PER_CLASS, len(pool))\n    if n == 0:\n        raise RuntimeError(\n            f\"Class '{cls_name}' has 0 samples after filtering.\\n\"\n            f\"Check CLASS_NAMES order matches HMS_LABEL_COLS: {HMS_LABEL_COLS}\"\n        )\n    frames.append(pool.sample(n=n, random_state=SEED))\n    print(f\"  {cls_name:<10}: pool={len(pool):4d} → sampled {n}\")\n\nsubset = pd.concat(frames, ignore_index=True).sample(frac=1, random_state=SEED)\n\nprint(f\"\\n✅ Micro-subset: {len(subset)} samples\")\nprint(\"Class distribution:\", subset['label'].value_counts().to_dict())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:19:37.043869Z","iopub.execute_input":"2026-02-23T21:19:37.044145Z","iopub.status.idle":"2026-02-23T21:20:30.005047Z","shell.execute_reply.started":"2026-02-23T21:19:37.044112Z","shell.execute_reply":"2026-02-23T21:20:30.004338Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5 — Load EEG Segments from Parquet Files","metadata":{}},{"cell_type":"code","source":"def load_eeg_segment(eeg_id: int,\n                     channels: list = CHANNELS,\n                     n_timepoints: int = N_TIMEPOINTS,\n                     preproc: EEGPreprocessor = None) -> np.ndarray:\n    \"\"\"\n    Load one EEG recording from parquet → return (n_channels, n_timepoints) float32.\n    Missing channels are zero-padded. Segment is taken from the centre of the recording.\n    \"\"\"\n    path = f'{EEG_DIR}/{eeg_id}.parquet'\n    df   = pd.read_parquet(path)\n\n    out = np.zeros((len(channels), n_timepoints), dtype=np.float32)\n    for i, ch in enumerate(channels):\n        if ch in df.columns:\n            sig = df[ch].values.astype(np.float32)\n            # take centre segment\n            start = max(0, (len(sig) - n_timepoints) // 2)\n            seg   = sig[start: start + n_timepoints]\n            out[i, :len(seg)] = seg\n\n    # Replace NaN / Inf\n    out = np.nan_to_num(out, nan=0.0, posinf=0.0, neginf=0.0)\n\n    if preproc is not None:\n        out = preproc(out)\n\n    # Per-channel z-score normalisation\n    std = out.std(axis=1, keepdims=True)\n    std[std < 1e-6] = 1.0\n    out = (out - out.mean(axis=1, keepdims=True)) / std\n    return out\n\n\n# ── Load all 400 segments ────────────────────────────────────────────────────\nprint(f'Loading {len(subset)} EEG segments (this takes ~3–5 minutes)...')\nt0      = time.time()\neeg_data  = []\neeg_labels = []\nfailed  = 0\n\nfor _, row in subset.iterrows():\n    try:\n        seg = load_eeg_segment(row['eeg_id'], preproc=preprocessor)\n        eeg_data.append(seg)\n        eeg_labels.append(row['label_idx'])\n    except Exception as e:\n        failed += 1\n\neeg_data   = np.stack(eeg_data)    # (N, 8, 2000)\neeg_labels = np.array(eeg_labels)  # (N,)\n\nprint(f'✅ Loaded {len(eeg_data)} segments in {time.time()-t0:.1f}s  |  {failed} failed')\nprint(f'Data shape: {eeg_data.shape}  |  Labels: {eeg_labels.shape}')\nprint('Class counts:', Counter(eeg_labels))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:20:30.006241Z","iopub.execute_input":"2026-02-23T21:20:30.006886Z","iopub.status.idle":"2026-02-23T21:20:39.933548Z","shell.execute_reply.started":"2026-02-23T21:20:30.006859Z","shell.execute_reply":"2026-02-23T21:20:39.932952Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6 — Visualise One Segment Per Class\n\nBefore any classification, let's confirm the 5 EEG patterns look distinct.","metadata":{}},{"cell_type":"code","source":"fig, axes = plt.subplots(N_CLASSES, N_CHANNELS, figsize=(20, N_CLASSES * 2.2))\nfig.suptitle('EEG Multi-Channel Segments — One Per Class', fontsize=14, fontweight='bold', y=1.01)\ntime_ax = np.linspace(0, SEGMENT_SECONDS, N_TIMEPOINTS)\n\nfor cls_idx, cls_name in enumerate(CLASS_NAMES):\n    idxs = np.where(eeg_labels == cls_idx)[0]\n    seg  = eeg_data[idxs[0]]   # first sample of this class\n    color = PALETTE[cls_name]\n    for ch in range(N_CHANNELS):\n        ax = axes[cls_idx, ch]\n        ax.plot(time_ax, seg[ch], color=color, linewidth=0.6)\n        ax.set_ylim(-4, 4)\n        ax.set_xticks([]); ax.set_yticks([])\n        if cls_idx == 0:\n            ax.set_title(CHANNELS[ch], fontsize=9)\n        if ch == 0:\n            ax.set_ylabel(cls_name, fontsize=10, fontweight='bold',\n                          color=color, rotation=0, labelpad=55)\n\nplt.tight_layout()\nplt.savefig('/kaggle/working/eeg_classes_overview.png', dpi=120, bbox_inches='tight')\nplt.show()\nprint('✅ Saved: eeg_classes_overview.png')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:20:39.934368Z","iopub.execute_input":"2026-02-23T21:20:39.934622Z","iopub.status.idle":"2026-02-23T21:20:41.722271Z","shell.execute_reply.started":"2026-02-23T21:20:39.934601Z","shell.execute_reply":"2026-02-23T21:20:41.721525Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7 — PART A: Load Pretrained EEGNet → Inference\n\nLoad the EEGNet saved from the training notebook (`eegnet_final.pt`). If not available, we train a quick version on 70% of the micro-subset.","metadata":{}},{"cell_type":"code","source":"# ── EEGNet Training with Data Augmentation ────────────────────────────────────\nimport os, copy\nimport torch, torch.nn as nn, torch.nn.functional as F\nimport numpy as np\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import StratifiedKFold, train_test_split\nfrom sklearn.metrics import accuracy_score, classification_report, confusion_matrix\nfrom collections import Counter\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# ── 1. Per-channel normalisation (done once, before any split) ────────────────\neeg_mean = eeg_data.mean(axis=2, keepdims=True)\neeg_std  = eeg_data.std(axis=2,  keepdims=True) + 1e-6\neeg_data_norm = (eeg_data - eeg_mean) / eeg_std\n\n# ── 2. Stratified 80 / 20 split ───────────────────────────────────────────────\nX_train, X_test, y_train, y_test = train_test_split(\n    eeg_data_norm, eeg_labels,\n    test_size=0.20, stratify=eeg_labels, random_state=SEED\n)\nprint(f\"Train: {X_train.shape}  Test: {X_test.shape}\")\n\n# ── 3. Augmentation Dataset ───────────────────────────────────────────────────\nclass AugEEGDataset(Dataset):\n    \"\"\"\n    EEG Dataset with on-the-fly augmentation (training only):\n      - Gaussian noise\n      - Amplitude scaling\n      - Random time shift\n      - Random channel dropout\n    \"\"\"\n    def __init__(self, X, y, augment=False):\n        self.X = torch.tensor(X, dtype=torch.float32)\n        self.y = torch.tensor(y, dtype=torch.long)\n        self.augment = augment\n\n    def __len__(self):  return len(self.y)\n\n    def __getitem__(self, idx):\n        x = self.X[idx].clone()          # (C, T)\n        if self.augment:\n            # (a) Gaussian noise σ = 0.05\n            x += torch.randn_like(x) * 0.05\n            # (b) Amplitude scaling ×U(0.85, 1.15)\n            x *= (0.85 + 0.30 * torch.rand(1).item())\n            # (c) Random circular time-shift ≤ 10 % of signal\n            shift = int(torch.randint(-200, 201, (1,)).item())\n            x = torch.roll(x, shifts=shift, dims=1)\n            # (d) Random channel dropout (drop 1 channel → zero)\n            if torch.rand(1).item() < 0.3:\n                ch = int(torch.randint(0, x.shape[0], (1,)).item())\n                x[ch] = 0.0\n        return x.unsqueeze(0), self.y[idx]   # (1, C, T), label\n\ndl_train = DataLoader(AugEEGDataset(X_train, y_train, augment=True),\n                      batch_size=32, shuffle=True,  drop_last=True)\ndl_test  = DataLoader(AugEEGDataset(X_test,  y_test,  augment=False),\n                      batch_size=64, shuffle=False)\n\n# ── 4. Mixup helper ───────────────────────────────────────────────────────────\ndef mixup_batch(x, y, alpha=0.3, n_classes=5):\n    \"\"\"Apply Mixup: returns mixed x, and soft one-hot y.\"\"\"\n    lam = float(np.random.beta(alpha, alpha)) if alpha > 0 else 1.0\n    perm = torch.randperm(x.size(0), device=x.device)\n    x_mix = lam * x + (1 - lam) * x[perm]\n    y_oh  = F.one_hot(y, n_classes).float()\n    y_mix = lam * y_oh + (1 - lam) * y_oh[perm]\n    return x_mix, y_mix\n\n# ── 5. Better EEGNet — lower dropout, residual skip connection ────────────────\nclass EEGNetV2(nn.Module):\n    \"\"\"\n    EEGNet with reduced dropout (0.25 instead of 0.5) and a\n    learned residual path in the separable block for better gradient flow.\n    \"\"\"\n    def __init__(self, n_channels=8, n_classes=5, n_timepoints=2000,\n                 F1=16, D=2, F2=32, dropout=0.25):\n        super().__init__()\n        self.n_channels   = n_channels\n        self.n_timepoints = n_timepoints\n\n        # Block 1 — Temporal convolution\n        self.temporal = nn.Sequential(\n            nn.Conv2d(1, F1, (1, 64), padding=(0, 32), bias=False),\n            nn.BatchNorm2d(F1)\n        )\n        # Block 1 — Depthwise spatial\n        self.depthwise = nn.Sequential(\n            nn.Conv2d(F1, F1*D, (n_channels, 1), groups=F1, bias=False),\n            nn.BatchNorm2d(F1*D),\n            nn.ELU(),\n            nn.AvgPool2d((1, 4)),\n            nn.Dropout(dropout)\n        )\n        # Block 2 — Separable convolution\n        self.separable = nn.Sequential(\n            nn.Conv2d(F1*D, F1*D, (1, 16), padding=(0, 8), groups=F1*D, bias=False),\n            nn.Conv2d(F1*D, F2,   (1, 1),  bias=False),\n            nn.BatchNorm2d(F2),\n            nn.ELU(),\n            nn.AvgPool2d((1, 8)),\n            nn.Dropout(dropout)\n        )\n        # Classifier\n        flat = self._flat(n_channels, n_timepoints, F1, D, F2)\n        self.classifier = nn.Sequential(\n            nn.Linear(flat, 64),\n            nn.ELU(),\n            nn.Dropout(0.25),\n            nn.Linear(64, n_classes)\n        )\n\n    def _flat(self, C, T, F1, D, F2):\n        with torch.no_grad():\n            x = torch.zeros(1, 1, C, T)\n            x = self.temporal(x); x = self.depthwise(x); x = self.separable(x)\n            return x.view(1,-1).shape[1]\n\n    def forward(self, x):\n        x = self.temporal(x)\n        x = self.depthwise(x)\n        x = self.separable(x)\n        return self.classifier(x.view(x.size(0), -1))\n\n    def predict_proba(self, x):\n        return F.softmax(self.forward(x), dim=1)\n\nmodel = EEGNetV2(n_channels=N_CHANNELS, n_classes=N_CLASSES, n_timepoints=N_TIMEPOINTS).to(DEVICE)\ntotal_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f\"EEGNetV2 — trainable params: {total_params:,}\")\n\n# ── 6. Class-weighted loss + AdamW + OneCycleLR ───────────────────────────────\ncounts  = Counter(y_train.tolist())\ntotal   = sum(counts.values())\nweights = torch.tensor(\n    [total / (N_CLASSES * counts[i]) for i in range(N_CLASSES)],\n    dtype=torch.float32\n).to(DEVICE)\n\ncriterion = nn.CrossEntropyLoss(weight=weights, label_smoothing=0.05)\noptimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=1e-4)\nN_EPOCHS  = 120\nscheduler = torch.optim.lr_scheduler.OneCycleLR(\n    optimizer, max_lr=3e-3,\n    steps_per_epoch=len(dl_train), epochs=N_EPOCHS\n)\n\n# ── 7. Training loop with Mixup ───────────────────────────────────────────────\nbest_acc, best_wts = 0.0, None\nhistory = {\"train_loss\": [], \"val_acc\": []}\n\nprint(\"\\n🔄 Training EEGNetV2...\")\nfor epoch in range(N_EPOCHS):\n    model.train(); epoch_loss = 0.0\n    for xb, yb in dl_train:\n        xb, yb = xb.to(DEVICE), yb.to(DEVICE)\n        # Mixup (50 % chance)\n        if torch.rand(1).item() < 0.5:\n            xb, yb_soft = mixup_batch(xb, yb, alpha=0.3, n_classes=N_CLASSES)\n        else:\n            yb_soft = F.one_hot(yb, N_CLASSES).float()\n        optimizer.zero_grad()\n        logits = model(xb)\n        loss   = -(yb_soft * F.log_softmax(logits, dim=1)).sum(dim=1).mean()\n        loss.backward(); optimizer.step(); scheduler.step()\n        epoch_loss += loss.item()\n\n    # Validation\n    model.eval(); preds_v, true_v = [], []\n    with torch.no_grad():\n        for xb, yb in dl_test:\n            preds_v.extend(model(xb.to(DEVICE)).argmax(1).cpu().numpy())\n            true_v.extend(yb.numpy())\n    val_acc = accuracy_score(true_v, preds_v)\n    history[\"train_loss\"].append(epoch_loss / len(dl_train))\n    history[\"val_acc\"].append(val_acc)\n\n    if val_acc > best_acc:\n        best_acc = val_acc\n        best_wts = copy.deepcopy(model.state_dict())\n\n    if (epoch + 1) % 10 == 0:\n        print(f\"  Epoch {epoch+1:3d}/{N_EPOCHS} | loss {epoch_loss/len(dl_train):.4f} | val_acc {val_acc:.4f} | best {best_acc:.4f}\")\n\nprint(f\"\\n✅ Best Validation Accuracy: {best_acc:.4f}\")\n\n# ── 8. Load best weights → final evaluation on test set ──────────────────────\nmodel.load_state_dict(best_wts)\nmodel.eval()\ntorch.save({\"model_state_dict\": best_wts}, \"/kaggle/working/eegnet_final.pt\")\n\nall_preds, all_probs = [], []\nwith torch.no_grad():\n    for xb, _ in dl_test:\n        logits = model(xb.to(DEVICE))\n        all_probs.extend(F.softmax(logits, dim=1).cpu().numpy())\n        all_preds.extend(logits.argmax(1).cpu().numpy())\n\nall_preds = np.array(all_preds)\nall_probs = np.array(all_probs)\n\nprint(\"\\n📊 EEGNetV2 Test Classification Report:\")\nprint(classification_report(y_test, all_preds, target_names=CLASS_NAMES))\nprint(\"Confusion Matrix:\")\nprint(confusion_matrix(y_test, all_preds))\n\n# ── 9. Get predictions on the FULL dataset (for comparison section) ───────────\nall_data_ds = AugEEGDataset(eeg_data_norm, eeg_labels, augment=False)\nall_data_dl  = DataLoader(all_data_ds, batch_size=64, shuffle=False)\n\npreds_ai_all_list, probs_ai_list = [], []\nwith torch.no_grad():\n    for xb, _ in all_data_dl:\n        logits = model(xb.to(DEVICE))\n        probs_ai_list.extend(F.softmax(logits, dim=1).cpu().numpy())\n        preds_ai_all_list.extend(logits.argmax(1).cpu().numpy())\n\n# ── Canonical variable names used by all downstream cells ────────────────────\npreds_ai_all = np.array(preds_ai_all_list)\nprobs_ai     = np.array(probs_ai_list)\n\nprint(f\"\\n✅ Full-dataset EEGNet inference done: preds_ai_all shape {preds_ai_all.shape}\")\nprint(f\"✅ Model saved → /kaggle/working/eegnet_final.pt\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:20:41.723593Z","iopub.execute_input":"2026-02-23T21:20:41.724121Z","iopub.status.idle":"2026-02-23T21:21:04.168951Z","shell.execute_reply.started":"2026-02-23T21:20:41.724064Z","shell.execute_reply":"2026-02-23T21:21:04.168289Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 8 — EEGNet Inference on Full Micro-Subset","metadata":{}},{"cell_type":"code","source":"# ── Variable compatibility bridge ────────────────────────────────────────────\n# Cell 15 trained EEGNetV2 and stored:\n#   preds_ai_all  — predictions on the FULL 400-sample dataset\n#   probs_ai      — probabilities on the FULL 400-sample dataset\n#   y_test        — test labels  (80 samples)\n#   X_test        — test EEG array (80 × 8 × 2000)\n# We also expose y_tr / y_te aliases used by cell 24:\nfrom scipy.signal import welch as _welch  # ensure welch is importable in feature cells\n\n# Aliases expected by Classic ML cells\nX_tr, X_te = None, None   # will be set properly in cell 22 (after X_classic is built)\ny_tr, y_te = None, None\n\nprint(\"✅ Variable bridge ready.\")\nprint(f\"   preds_ai_all shape : {preds_ai_all.shape}\")\nprint(f\"   probs_ai shape     : {probs_ai.shape}\")\nprint(f\"   Full dataset size  : {len(eeg_labels)} samples\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:21:04.169868Z","iopub.execute_input":"2026-02-23T21:21:04.170265Z","iopub.status.idle":"2026-02-23T21:21:04.175662Z","shell.execute_reply.started":"2026-02-23T21:21:04.170244Z","shell.execute_reply":"2026-02-23T21:21:04.174969Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 9 — PART B: Classic ML Pipeline\n\nWe implement a hand-crafted feature extraction pipeline — the kind of algorithm used in classic EEG abnormality detection (pre-deep-learning era). Features are extracted **per channel** and concatenated.\n\n### Feature Set (6 categories)\n\n| Category | Features | Rationale |\n|---|---|---|\n| **Statistical** | Mean, Std, Skewness, Kurtosis, Peak-to-peak | Amplitude distribution — seizures have extreme kurtosis |\n| **Hjorth Parameters** | Activity, Mobility, Complexity | Classic EEG descriptors from Hjorth (1970) — mobility ↑ in seizures |\n| **Band Power** | Delta(1–4Hz), Theta(4–8Hz), Alpha(8–13Hz), Beta(13–30Hz) | Spectral energy distribution — GRDA shows delta dominance |\n| **Autocorrelation** | First 10 lags | Measures periodicity — periodic discharges (LPD/GPD) show peaks |\n| **Spectral Entropy** | Shannon entropy of PSD | Measures signal regularity — seizures have lower entropy |\n| **Zero-Crossing Rate** | Per-channel ZCR | Frequency proxy — high ZCR in fast-wave seizures |","metadata":{}},{"cell_type":"code","source":"def bandpower(sig, sfreq, fmin, fmax):\n    \"\"\"Compute band power via Welch PSD.\"\"\"\n    freqs, psd = sp_signal.welch(sig, fs=sfreq, nperseg=min(256, len(sig)))\n    idx = np.logical_and(freqs >= fmin, freqs <= fmax)\n    return np.trapz(psd[idx], freqs[idx])\n\n\ndef extract_classic_features(eeg: np.ndarray, sfreq: int = SFREQ) -> np.ndarray:\n    \"\"\"\n    Extract classic ML features from a multi-channel EEG segment.\n\n    Parameters\n    ----------\n    eeg : (n_channels, n_timepoints) float32\n\n    Returns\n    -------\n    features : 1D array of length (n_channels × n_features_per_channel)\n    \"\"\"\n    n_channels = eeg.shape[0]\n    all_feats  = []\n\n    AUTOCORR_LAGS = 10\n    BANDS = {\n        'delta': (1,  4),\n        'theta': (4,  8),\n        'alpha': (8,  13),\n        'beta':  (13, 30)\n    }\n\n    for ch in range(n_channels):\n        x = eeg[ch].astype(np.float64)\n\n        # ── 1. Statistical features (5) ───────────────────────────────────\n        stat_feats = [\n            x.mean(),\n            x.std(),\n            skew(x),\n            kurtosis(x),\n            x.max() - x.min()         # peak-to-peak amplitude\n        ]\n\n        # ── 2. Hjorth Parameters (3) ─────────────────────────────────────\n        try:\n            hjorth_activity    = np.var(x)\n            hjorth_mobility    = ant.hjorth_params(x)[0]\n            hjorth_complexity  = ant.hjorth_params(x)[1]\n        except Exception:\n            d1 = np.diff(x)\n            d2 = np.diff(d1)\n            hjorth_activity   = np.var(x)\n            hjorth_mobility   = np.sqrt(np.var(d1) / (np.var(x) + 1e-8))\n            hjorth_complexity = (np.sqrt(np.var(d2) / (np.var(d1) + 1e-8)) /\n                                 (hjorth_mobility + 1e-8))\n        hjorth_feats = [hjorth_activity, hjorth_mobility, hjorth_complexity]\n\n        # ── 3. Band Power (4) ────────────────────────────────────────────\n        band_feats = [bandpower(x, sfreq, lo, hi) for lo, hi in BANDS.values()]\n        # Relative band power (avoids absolute-amplitude confound)\n        total_power = sum(band_feats) + 1e-8\n        band_feats  = [p / total_power for p in band_feats]\n\n        # ── 4. Autocorrelation features (10 lags) ────────────────────────\n        ac  = np.correlate(x - x.mean(), x - x.mean(), mode='full')\n        ac  = ac[len(ac)//2:]           # keep positive lags only\n        ac  = ac / (ac[0] + 1e-8)       # normalise: lag-0 = 1\n        ac_feats = ac[1: AUTOCORR_LAGS + 1].tolist()   # lags 1–10\n\n        # ── 5. Spectral Entropy (1) ──────────────────────────────────────\n        try:\n            sp_ent = ant.spectral_entropy(x, sf=sfreq, method='welch', normalize=True)\n        except Exception:\n            freqs, psd = sp_signal.welch(x, fs=sfreq, nperseg=min(256, len(x)))\n            psd_norm = psd / (psd.sum() + 1e-8)\n            sp_ent   = -np.sum(psd_norm * np.log2(psd_norm + 1e-8))\n        sp_feats = [float(sp_ent)]\n\n        # ── 6. Zero-Crossing Rate (1) ────────────────────────────────────\n        zcr = np.sum(np.abs(np.diff(np.sign(x)))) / (2 * len(x))\n        zcr_feats = [zcr]\n\n        all_feats.extend(stat_feats + hjorth_feats + band_feats +\n                         ac_feats + sp_feats + zcr_feats)\n\n    return np.array(all_feats, dtype=np.float32)\n\n\n# ── Feature dimension sanity check ───────────────────────────────────────────\nsample_feats = extract_classic_features(eeg_data[0])\nN_FEATS_PER_CH = (5 + 3 + 4 + 10 + 1 + 1)   # = 24\nN_TOTAL_FEATS  = N_CHANNELS * N_FEATS_PER_CH  # = 192\nprint(f'Features per channel : {N_FEATS_PER_CH}')\nprint(f'Total feature vector : {sample_feats.shape[0]}  (expected {N_TOTAL_FEATS})')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:21:04.178399Z","iopub.execute_input":"2026-02-23T21:21:04.178628Z","iopub.status.idle":"2026-02-23T21:21:04.267751Z","shell.execute_reply.started":"2026-02-23T21:21:04.178609Z","shell.execute_reply":"2026-02-23T21:21:04.267066Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Extract features for all 400 samples ─────────────────────────────────────\nprint(f'Extracting classic ML features for {len(eeg_data)} segments...')\nt0 = time.time()\n\nX_classic = np.vstack([extract_classic_features(eeg_data[i])\n                        for i in range(len(eeg_data))])\n\nprint(f'✅ Feature matrix shape: {X_classic.shape}  |  {time.time()-t0:.1f}s')\n\n# Replace any NaN/Inf in features\nX_classic = np.nan_to_num(X_classic, nan=0.0, posinf=1e6, neginf=-1e6)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:21:04.268647Z","iopub.execute_input":"2026-02-23T21:21:04.268936Z","iopub.status.idle":"2026-02-23T21:21:28.144359Z","shell.execute_reply.started":"2026-02-23T21:21:04.268893Z","shell.execute_reply":"2026-02-23T21:21:28.143489Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 9.1 — Feature Importance: Which Features Are Most Discriminative?","metadata":{}},{"cell_type":"code","source":"# ── Train a quick Random Forest to get feature importances ─────────────────\nfrom sklearn.model_selection import train_test_split\n\nX_tr, X_te, y_tr, y_te = train_test_split(\n    X_classic, eeg_labels, test_size=0.25, stratify=eeg_labels, random_state=SEED\n)\n\nscaler = StandardScaler()\nX_tr_s = scaler.fit_transform(X_tr)\nX_te_s = scaler.transform(X_te)\n\n# Feature category labels for plotting\nfeat_categories = []\ncat_names = ['Stat']*5 + ['Hjorth']*3 + ['BandPow']*4 + ['AutoCorr']*10 + ['SpEnt']*1 + ['ZCR']*1\nfor _ in range(N_CHANNELS):\n    feat_categories.extend(cat_names)\n\n# Aggregated importance per category\nrf_fi = RandomForestClassifier(n_estimators=100, random_state=SEED, n_jobs=-1)\nrf_fi.fit(X_tr_s, y_tr)\nimportances = rf_fi.feature_importances_\n\ncat_importance = {}\nfor cat, imp in zip(feat_categories, importances):\n    cat_importance[cat] = cat_importance.get(cat, 0) + imp\n\ncats   = list(cat_importance.keys())\nimps   = list(cat_importance.values())\nsorted_idx = np.argsort(imps)[::-1]\n\nfig, ax = plt.subplots(figsize=(9, 4))\nbars = ax.bar([cats[i] for i in sorted_idx],\n              [imps[i] for i in sorted_idx],\n              color=['#3498db','#e74c3c','#2ecc71','#f39c12','#9b59b6','#1abc9c'])\nax.set_title('Classic ML — Feature Category Importance (Random Forest)', fontweight='bold')\nax.set_ylabel('Summed Feature Importance')\nax.set_xlabel('Feature Category')\nfor bar, imp in zip(bars, [imps[i] for i in sorted_idx]):\n    ax.text(bar.get_x() + bar.get_width()/2, bar.get_height() + 0.001,\n            f'{imp:.3f}', ha='center', va='bottom', fontsize=9)\nplt.tight_layout()\nplt.savefig('/kaggle/working/feature_importance.png', dpi=120, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:21:28.145293Z","iopub.execute_input":"2026-02-23T21:21:28.145600Z","iopub.status.idle":"2026-02-23T21:21:28.739713Z","shell.execute_reply.started":"2026-02-23T21:21:28.145574Z","shell.execute_reply":"2026-02-23T21:21:28.739156Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 9.2 — Classic ML Classifiers: SVM vs Random Forest vs Gradient Boosting","metadata":{}},{"cell_type":"code","source":"# ── Define classic ML classifiers ─────────────────────────────────────────────\nCLASSIFIERS = {\n    'SVM (RBF)'           : SVC(C=10, gamma='scale', kernel='rbf',\n                                probability=True, random_state=SEED),\n    'Random Forest'       : RandomForestClassifier(n_estimators=200, max_depth=15,\n                                                   random_state=SEED, n_jobs=-1),\n    'Gradient Boosting'   : GradientBoostingClassifier(n_estimators=150,\n                                                        learning_rate=0.1,\n                                                        max_depth=5,\n                                                        random_state=SEED)\n}\n\ncv = StratifiedKFold(n_splits=5, shuffle=True, random_state=SEED)\ncv_results = {}\n\nprint('5-fold Cross-Validation on micro-subset:')\nprint('-' * 55)\nfor name, clf in CLASSIFIERS.items():\n    pipe = Pipeline([('scaler', StandardScaler()), ('clf', clf)])\n    scores = cross_val_score(pipe, X_classic, eeg_labels,\n                              cv=cv, scoring='balanced_accuracy', n_jobs=-1)\n    cv_results[name] = scores\n    print(f'{name:<25} | {scores.mean():.4f} ± {scores.std():.4f}')\n\nprint('-' * 55)\n\n# ── Train best classic model on 75% / test on 25% ────────────────────────────\nbest_clf_name = max(cv_results, key=lambda k: cv_results[k].mean())\nbest_clf      = Pipeline([('scaler', StandardScaler()),\n                           ('clf', CLASSIFIERS[best_clf_name])])\nbest_clf.fit(X_tr, y_tr)\n\npreds_classic = best_clf.predict(X_te)\nprobs_classic = best_clf.predict_proba(X_te)\n\nacc_classic = accuracy_score(y_te, preds_classic)\nf1_classic  = f1_score(y_te, preds_classic, average='weighted')\n\nprint(f'\\nBest Classic Model : {best_clf_name}')\nprint(f'Accuracy           : {acc_classic:.4f}')\nprint(f'Weighted F1        : {f1_classic:.4f}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:21:28.741051Z","iopub.execute_input":"2026-02-23T21:21:28.741595Z","iopub.status.idle":"2026-02-23T21:22:50.216598Z","shell.execute_reply.started":"2026-02-23T21:21:28.741570Z","shell.execute_reply":"2026-02-23T21:22:50.215873Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 10 — COMPARISON: EEGNet AI vs Classic ML\n\nThis is the core comparison required by the task. We evaluate both methods on the **same test split**.","metadata":{}},{"cell_type":"code","source":"# ── Collect all predictions on the FULL dataset for comparison ───────────────\n# EEGNet predictions already computed in cell 15 (preds_ai_all, probs_ai)\n# Classic ML predictions on full dataset\ny_true_all   = eeg_labels\npreds_cl_all = best_clf.predict(X_classic)\nprobs_cl_all = best_clf.predict_proba(X_classic)\n\n# ── Metrics comparison table ─────────────────────────────────────────────────\nfrom sklearn.metrics import balanced_accuracy_score\n\nmetrics = {\n    'Method': ['EEGNet V2 (AI)', f'Classic ML ({best_clf_name})'],\n    'Accuracy':           [accuracy_score(y_true_all, preds_ai_all),\n                            accuracy_score(y_true_all, preds_cl_all)],\n    'Balanced Accuracy':  [balanced_accuracy_score(y_true_all, preds_ai_all),\n                            balanced_accuracy_score(y_true_all, preds_cl_all)],\n    'Weighted F1':        [f1_score(y_true_all, preds_ai_all, average='weighted'),\n                            f1_score(y_true_all, preds_cl_all, average='weighted')],\n    'Macro F1':           [f1_score(y_true_all, preds_ai_all, average='macro'),\n                            f1_score(y_true_all, preds_cl_all, average='macro')],\n}\n\ndf_metrics = pd.DataFrame(metrics)\ndf_metrics.set_index('Method', inplace=True)\nprint('=' * 65)\nprint('         COMPARISON: EEGNet AI  vs  Classic ML')\nprint('=' * 65)\nprint(df_metrics.to_string(float_format=lambda x: f'{x:.4f}'))\nprint('=' * 65)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:22:50.217459Z","iopub.execute_input":"2026-02-23T21:22:50.217710Z","iopub.status.idle":"2026-02-23T21:22:50.378898Z","shell.execute_reply.started":"2026-02-23T21:22:50.217687Z","shell.execute_reply":"2026-02-23T21:22:50.378308Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 10.1 — Per-Class F1 Comparison","metadata":{}},{"cell_type":"code","source":"f1_ai_per  = f1_score(y_true_all, preds_ai_all,  average=None, labels=list(range(N_CLASSES)))\nf1_cl_per  = f1_score(y_true_all, preds_cl_all, average=None, labels=list(range(N_CLASSES)))\n\nx   = np.arange(N_CLASSES)\nw   = 0.35\n\nfig, ax = plt.subplots(figsize=(10, 5))\nbars1 = ax.bar(x - w/2, f1_ai_per,  w, label='EEGNet (AI)',     color='#3498db', edgecolor='white')\nbars2 = ax.bar(x + w/2, f1_cl_per, w, label=f'Classic ML ({best_clf_name})',\n                color='#e74c3c', edgecolor='white')\n\nax.set_xticks(x)\nax.set_xticklabels(CLASS_NAMES, fontsize=11)\nax.set_ylabel('F1 Score', fontsize=11)\nax.set_title('Per-Class F1 Score: EEGNet AI vs Classic ML', fontsize=13, fontweight='bold')\nax.set_ylim(0, 1.05)\nax.legend(fontsize=10)\nax.axhline(0.5, color='gray', linestyle='--', alpha=0.4, label='Random baseline')\n\nfor bar in bars1:\n    ax.text(bar.get_x() + bar.get_width()/2, bar.get_height() + 0.015,\n            f'{bar.get_height():.2f}', ha='center', va='bottom', fontsize=8, color='#3498db')\nfor bar in bars2:\n    ax.text(bar.get_x() + bar.get_width()/2, bar.get_height() + 0.015,\n            f'{bar.get_height():.2f}', ha='center', va='bottom', fontsize=8, color='#e74c3c')\n\nplt.tight_layout()\nplt.savefig('/kaggle/working/f1_comparison.png', dpi=120, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:22:50.379942Z","iopub.execute_input":"2026-02-23T21:22:50.380344Z","iopub.status.idle":"2026-02-23T21:22:50.768166Z","shell.execute_reply.started":"2026-02-23T21:22:50.380324Z","shell.execute_reply":"2026-02-23T21:22:50.767465Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 10.2 — Confusion Matrices: Side by Side","metadata":{}},{"cell_type":"code","source":"fig, axes = plt.subplots(1, 2, figsize=(15, 6))\n\nfor ax, preds, title, cmap in [\n    (axes[0], preds_ai_all,  'EEGNet (Pretrained AI)',        'Blues'),\n    (axes[1], preds_cl_all,  f'Classic ML ({best_clf_name})', 'Reds')\n]:\n    cm      = confusion_matrix(y_true_all, preds)\n    cm_norm = cm.astype(float) / cm.sum(axis=1, keepdims=True)\n\n    sns.heatmap(cm_norm, annot=True, fmt='.2f', cmap=cmap,\n                xticklabels=CLASS_NAMES, yticklabels=CLASS_NAMES,\n                ax=ax, vmin=0, vmax=1,\n                linewidths=0.5, linecolor='white')\n    acc = accuracy_score(y_true_all, preds)\n    ax.set_title(f'{title}\\nAccuracy: {acc:.4f}', fontweight='bold', fontsize=11)\n    ax.set_ylabel('True Label', fontsize=10)\n    ax.set_xlabel('Predicted Label', fontsize=10)\n\nplt.suptitle('Normalised Confusion Matrices', fontsize=13, fontweight='bold', y=1.02)\nplt.tight_layout()\nplt.savefig('/kaggle/working/confusion_matrices_comparison.png', dpi=120, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:22:50.769240Z","iopub.execute_input":"2026-02-23T21:22:50.769519Z","iopub.status.idle":"2026-02-23T21:22:51.593453Z","shell.execute_reply.started":"2026-02-23T21:22:50.769489Z","shell.execute_reply":"2026-02-23T21:22:51.592719Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 10.3 — ROC Curves: Both Methods on Same Plot","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import roc_curve, auc\n\ny_bin = label_binarize(y_true_all, classes=list(range(N_CLASSES)))\n\nfig, axes = plt.subplots(1, 2, figsize=(14, 6))\nCOLOURS = ['#e74c3c','#e67e22','#f1c40f','#2ecc71','#3498db']\n\nfor ax, probs, title in [\n    (axes[0], probs_ai,     'EEGNet (AI)'),\n    (axes[1], probs_cl_all, f'Classic ML ({best_clf_name})')\n]:\n    macro_auc = []\n    for i in range(N_CLASSES):\n        fpr, tpr, _ = roc_curve(y_bin[:, i], probs[:, i])\n        roc_auc     = auc(fpr, tpr)\n        macro_auc.append(roc_auc)\n        ax.plot(fpr, tpr, color=COLOURS[i], lw=2,\n                label=f'{CLASS_NAMES[i]} (AUC={roc_auc:.3f})')\n\n    ax.plot([0,1],[0,1], 'k--', lw=1, alpha=0.5)\n    ax.set_xlim([0, 1]); ax.set_ylim([0, 1.02])\n    ax.set_xlabel('False Positive Rate'); ax.set_ylabel('True Positive Rate')\n    ax.set_title(f'{title}\\nMacro-AUC = {np.mean(macro_auc):.4f}',\n                 fontweight='bold')\n    ax.legend(loc='lower right', fontsize=8)\n\nplt.suptitle('ROC Curves — One-vs-Rest per Class', fontsize=13, fontweight='bold')\nplt.tight_layout()\nplt.savefig('/kaggle/working/roc_comparison.png', dpi=120, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:22:51.594610Z","iopub.execute_input":"2026-02-23T21:22:51.594846Z","iopub.status.idle":"2026-02-23T21:22:52.331595Z","shell.execute_reply.started":"2026-02-23T21:22:51.594825Z","shell.execute_reply":"2026-02-23T21:22:52.331002Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 10.4 — Agreement Analysis: Where Do Both Models Agree/Disagree?","metadata":{}},{"cell_type":"code","source":"agree_correct   = np.sum((preds_ai_all == y_true_all) & (preds_cl_all == y_true_all))\nai_only_correct = np.sum((preds_ai_all == y_true_all) & (preds_cl_all != y_true_all))\ncl_only_correct = np.sum((preds_ai_all != y_true_all) & (preds_cl_all == y_true_all))\nboth_wrong      = np.sum((preds_ai_all != y_true_all) & (preds_cl_all != y_true_all))\nn_total         = len(y_true_all)\n\ncategories = ['Both Correct', 'AI Only Correct', 'Classic ML Only', 'Both Wrong']\ncounts_    = [agree_correct, ai_only_correct, cl_only_correct, both_wrong]\npcts       = [c/n_total*100 for c in counts_]\ncolors_    = ['#2ecc71', '#3498db', '#e74c3c', '#95a5a6']\n\nfig, (ax1, ax2) = plt.subplots(1, 2, figsize=(13, 5))\n\n# Pie chart\nwedges, texts, autotexts = ax1.pie(\n    counts_, labels=categories, colors=colors_,\n    autopct='%1.1f%%', startangle=90,\n    textprops={'fontsize': 10}\n)\nax1.set_title('Prediction Agreement\\nEEGNet AI vs Classic ML', fontweight='bold')\n\n# Per-class agreement heatmap\nagreement_matrix = np.zeros((N_CLASSES, N_CLASSES), dtype=int)\nfor true, ai, cl in zip(y_true_all, preds_ai_all, preds_cl_all):\n    agreement_matrix[ai, cl] += 1\n\nsns.heatmap(agreement_matrix, annot=True, fmt='d', cmap='YlOrRd',\n            xticklabels=[f'ML:{c}' for c in CLASS_NAMES],\n            yticklabels=[f'AI:{c}' for c in CLASS_NAMES],\n            ax=ax2, linewidths=0.5)\nax2.set_title('Prediction Agreement Matrix\\n(AI pred vs ML pred)', fontweight='bold')\nax2.set_xlabel('Classic ML Prediction')\nax2.set_ylabel('EEGNet AI Prediction')\n\nplt.tight_layout()\nplt.savefig('/kaggle/working/agreement_analysis.png', dpi=120, bbox_inches='tight')\nplt.show()\n\nprint(f'\\nAgreement Summary (n={n_total}):')\nfor cat, cnt, pct in zip(categories, counts_, pcts):\n    print(f'  {cat:<28} : {cnt:4d} ({pct:.1f}%)')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:22:52.332447Z","iopub.execute_input":"2026-02-23T21:22:52.332663Z","iopub.status.idle":"2026-02-23T21:22:52.901453Z","shell.execute_reply.started":"2026-02-23T21:22:52.332643Z","shell.execute_reply":"2026-02-23T21:22:52.900770Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Comprehensive Model Evaluation ──────────────────────────────────────────\nfrom sklearn.metrics import (\n    classification_report, roc_auc_score, average_precision_score,\n    balanced_accuracy_score\n)\nfrom sklearn.preprocessing import label_binarize\n\ny_bin_all = label_binarize(y_true_all, classes=list(range(N_CLASSES)))\n\n# ── A) Per-class detailed report ─────────────────────────────────────────────\nprint(\"=\" * 70)\nprint(\"  EEGNetV2 (AI) — Full-Dataset Classification Report\")\nprint(\"=\" * 70)\nprint(classification_report(y_true_all, preds_ai_all, target_names=CLASS_NAMES))\n\nprint(\"=\" * 70)\nprint(f\"  Classic ML ({best_clf_name}) — Full-Dataset Classification Report\")\nprint(\"=\" * 70)\nprint(classification_report(y_true_all, preds_cl_all, target_names=CLASS_NAMES))\n\n# ── B) ROC-AUC (macro OvR) ───────────────────────────────────────────────────\nauc_ai = roc_auc_score(y_bin_all, probs_ai,     multi_class=\"ovr\", average=\"macro\")\nauc_cl = roc_auc_score(y_bin_all, probs_cl_all, multi_class=\"ovr\", average=\"macro\")\nprint(f\"Macro ROC-AUC — EEGNetV2 : {auc_ai:.4f}\")\nprint(f\"Macro ROC-AUC — Classic ML: {auc_cl:.4f}\")\n\n# ── C) Balanced Accuracy ─────────────────────────────────────────────────────\nba_ai = balanced_accuracy_score(y_true_all, preds_ai_all)\nba_cl = balanced_accuracy_score(y_true_all, preds_cl_all)\nprint(f\"\\nBalanced Accuracy — EEGNetV2 : {ba_ai:.4f}\")\nprint(f\"Balanced Accuracy — Classic ML: {ba_cl:.4f}\")\n\n# ── D) Training history plot ─────────────────────────────────────────────────\nfig, axes = plt.subplots(1, 2, figsize=(13, 4))\n\naxes[0].plot(history[\"train_loss\"], color=\"#3498db\", lw=1.5)\naxes[0].set_title(\"EEGNetV2 — Training Loss\", fontweight=\"bold\")\naxes[0].set_xlabel(\"Epoch\"); axes[0].set_ylabel(\"CrossEntropy Loss\")\naxes[0].grid(alpha=0.3)\n\naxes[1].plot(history[\"val_acc\"], color=\"#2ecc71\", lw=1.5)\naxes[1].axhline(max(history[\"val_acc\"]), ls=\"--\", color=\"#e74c3c\", lw=1,\n                label=f\"Best = {max(history['val_acc']):.4f}\")\naxes[1].set_title(\"EEGNetV2 — Validation Accuracy\", fontweight=\"bold\")\naxes[1].set_xlabel(\"Epoch\"); axes[1].set_ylabel(\"Accuracy\")\naxes[1].legend(); axes[1].grid(alpha=0.3)\n\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/eegnet_training_history.png\", dpi=120, bbox_inches=\"tight\")\nplt.show()\nprint(\"✅ Saved: eegnet_training_history.png\")\n\n# ── E) Per-class F1 comparison bar chart ────────────────────────────────────\nf1_ai = f1_score(y_true_all, preds_ai_all, average=None)\nf1_cl = f1_score(y_true_all, preds_cl_all, average=None)\n\nx = np.arange(N_CLASSES); w = 0.35\nfig, ax = plt.subplots(figsize=(10, 5))\nax.bar(x - w/2, f1_ai, w, label=\"EEGNetV2 (AI)\", color=\"#3498db\", alpha=0.85)\nax.bar(x + w/2, f1_cl, w, label=f\"Classic ML ({best_clf_name})\", color=\"#e74c3c\", alpha=0.85)\nax.set_xticks(x); ax.set_xticklabels(CLASS_NAMES, fontsize=11)\nax.set_ylabel(\"F1 Score\"); ax.set_title(\"Per-Class F1 Score: AI vs Classic ML\", fontweight=\"bold\")\nax.legend(); ax.set_ylim(0, 1.0); ax.grid(axis=\"y\", alpha=0.3)\nfor i, (a, c) in enumerate(zip(f1_ai, f1_cl)):\n    ax.text(i-w/2, a+0.02, f\"{a:.2f}\", ha=\"center\", fontsize=9, color=\"#2c3e50\")\n    ax.text(i+w/2, c+0.02, f\"{c:.2f}\", ha=\"center\", fontsize=9, color=\"#2c3e50\")\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/per_class_f1.png\", dpi=120, bbox_inches=\"tight\")\nplt.show()\nprint(\"✅ Saved: per_class_f1.png\")\n\n# ── F) 5-Fold Cross-Validation for EEGNet ─────────────────────────────────── \nprint(\"\\n🔄 5-Fold CV Evaluation (EEGNet — quick 30-epoch runs)...\")\nskf = StratifiedKFold(n_splits=5, shuffle=True, random_state=SEED)\ncv_accs = []\nfor fold, (tr_idx, va_idx) in enumerate(skf.split(eeg_data_norm, eeg_labels)):\n    m_cv = EEGNetV2(N_CHANNELS, N_CLASSES, N_TIMEPOINTS).to(DEVICE)\n    opt_cv = torch.optim.AdamW(m_cv.parameters(), lr=1e-3, weight_decay=1e-4)\n    dl_cv_tr = DataLoader(AugEEGDataset(eeg_data_norm[tr_idx], eeg_labels[tr_idx], augment=True),\n                          batch_size=32, shuffle=True, drop_last=True)\n    dl_cv_va = DataLoader(AugEEGDataset(eeg_data_norm[va_idx], eeg_labels[va_idx], augment=False),\n                          batch_size=64, shuffle=False)\n    sch_cv = torch.optim.lr_scheduler.OneCycleLR(opt_cv, max_lr=3e-3,\n             steps_per_epoch=len(dl_cv_tr), epochs=30)\n    crit_cv = nn.CrossEntropyLoss(label_smoothing=0.05)\n    for ep in range(30):\n        m_cv.train()\n        for xb, yb in dl_cv_tr:\n            xb, yb = xb.to(DEVICE), yb.to(DEVICE)\n            opt_cv.zero_grad(); loss = crit_cv(m_cv(xb), yb)\n            loss.backward(); opt_cv.step(); sch_cv.step()\n    m_cv.eval(); pv, tv = [], []\n    with torch.no_grad():\n        for xb, yb in dl_cv_va:\n            pv.extend(m_cv(xb.to(DEVICE)).argmax(1).cpu().numpy())\n            tv.extend(yb.numpy())\n    fold_acc = accuracy_score(tv, pv)\n    cv_accs.append(fold_acc)\n    print(f\"  Fold {fold+1} accuracy: {fold_acc:.4f}\")\nprint(f\"\\n✅ EEGNetV2 5-Fold CV Accuracy: {np.mean(cv_accs):.4f} ± {np.std(cv_accs):.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:22:52.902263Z","iopub.execute_input":"2026-02-23T21:22:52.902458Z","iopub.status.idle":"2026-02-23T21:23:09.563984Z","shell.execute_reply.started":"2026-02-23T21:22:52.902440Z","shell.execute_reply":"2026-02-23T21:23:09.563371Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 11 — Confidence Calibration Analysis","metadata":{}},{"cell_type":"code","source":"conf_ai = probs_ai.max(axis=1)\nconf_cl = probs_cl_all.max(axis=1)\ncorrect_ai = (preds_ai_all == y_true_all)\ncorrect_cl = (preds_cl_all  == y_true_all)\n\nfig, axes = plt.subplots(1, 3, figsize=(16, 4))\n\n# Confidence distribution\naxes[0].hist(conf_ai[correct_ai],  bins=20, alpha=0.7, color='#2ecc71',\n              label='AI Correct',   density=True)\naxes[0].hist(conf_ai[~correct_ai], bins=20, alpha=0.7, color='#e74c3c',\n              label='AI Wrong',     density=True)\naxes[0].set_title('EEGNet AI — Confidence Distribution')\naxes[0].set_xlabel('Max Softmax Probability')\naxes[0].legend()\n\naxes[1].hist(conf_cl[correct_cl],  bins=20, alpha=0.7, color='#3498db',\n              label='ML Correct',   density=True)\naxes[1].hist(conf_cl[~correct_cl], bins=20, alpha=0.7, color='#e67e22',\n              label='ML Wrong',     density=True)\naxes[1].set_title(f'Classic ML — Confidence Distribution')\naxes[1].set_xlabel('Max Predicted Probability')\naxes[1].legend()\n\n# Violin: confidence per class\ndf_conf = pd.DataFrame({\n    'Confidence_AI': conf_ai,\n    'Confidence_ML': conf_cl,\n    'True_Class': [CLASS_NAMES[i] for i in y_true_all],\n    'Correct_AI': correct_ai\n})\ndf_melt = df_conf.melt(id_vars='True_Class',\n                        value_vars=['Confidence_AI','Confidence_ML'],\n                        var_name='Method', value_name='Confidence')\ndf_melt['Method'] = df_melt['Method'].map({'Confidence_AI':'EEGNet AI',\n                                             'Confidence_ML':'Classic ML'})\nsns.violinplot(data=df_melt, x='True_Class', y='Confidence', hue='Method',\n               palette=['#3498db','#e74c3c'], split=True, ax=axes[2], inner='box')\naxes[2].set_title('Prediction Confidence per Class')\naxes[2].set_xticklabels(CLASS_NAMES, rotation=15)\n\nplt.suptitle('Confidence Calibration Analysis', fontsize=13, fontweight='bold')\nplt.tight_layout()\nplt.savefig('/kaggle/working/confidence_calibration.png', dpi=120, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:23:09.564970Z","iopub.execute_input":"2026-02-23T21:23:09.565470Z","iopub.status.idle":"2026-02-23T21:23:10.872134Z","shell.execute_reply.started":"2026-02-23T21:23:09.565446Z","shell.execute_reply":"2026-02-23T21:23:10.871468Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 12 — Band Power Visualization: Why Classic ML Features Work\n\nThis plot shows the EEG frequency content per class, explaining why band power features are discriminative.","metadata":{}},{"cell_type":"code","source":"BAND_NAMES = ['Delta\\n1–4 Hz', 'Theta\\n4–8 Hz', 'Alpha\\n8–13 Hz', 'Beta\\n13–30 Hz']\n# Feature indices for band power: features 8–11 per channel, using Ch0 (Fp1)\nBP_START = 5 + 3   # after stat(5) + hjorth(3)\n\nfig, axes = plt.subplots(1, N_CLASSES, figsize=(16, 4), sharey=True)\nfig.suptitle('Relative Band Power per Class — Channel Fp1', fontsize=13, fontweight='bold')\n\nfor cls_idx, (cls_name, ax) in enumerate(zip(CLASS_NAMES, axes)):\n    mask = (eeg_labels == cls_idx)\n    # band power features for channel 0\n    bp_feats = X_classic[mask, BP_START: BP_START+4]  # (N_cls, 4)\n    means = bp_feats.mean(axis=0)\n    stds  = bp_feats.std(axis=0)\n\n    bars = ax.bar(BAND_NAMES, means, yerr=stds,\n                  color=PALETTE[cls_name], alpha=0.85,\n                  edgecolor='white', capsize=4, error_kw={'elinewidth':1.5})\n    ax.set_title(cls_name, fontweight='bold', color=PALETTE[cls_name])\n    ax.set_ylim(0, 0.9)\n    if cls_idx == 0:\n        ax.set_ylabel('Relative Power')\n\nplt.tight_layout()\nplt.savefig('/kaggle/working/band_power_per_class.png', dpi=120, bbox_inches='tight')\nplt.show()\nprint('📌 GRDA shows delta dominance (as expected — Generalized Rhythmic Delta Activity)')\nprint('📌 Seizure shows elevated beta power (fast wave activity)')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:23:10.873084Z","iopub.execute_input":"2026-02-23T21:23:10.873623Z","iopub.status.idle":"2026-02-23T21:23:11.744399Z","shell.execute_reply.started":"2026-02-23T21:23:10.873597Z","shell.execute_reply":"2026-02-23T21:23:11.743653Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 13 — Signal Viewer Integration: `classify_eeg()` Function\n\nThis is the function your **Signal Viewer website** calls when a user opens an EEG file. It runs **both** the EEGNet AI model and the Classic ML algorithm and returns both predictions for display.","metadata":{}},{"cell_type":"code","source":"def classify_eeg_dual(raw_eeg: np.ndarray,\n                       eegnet_model,           # EEGNetV2 instance\n                       classic_clf,            # sklearn Pipeline\n                       preproc: EEGPreprocessor,\n                       device=DEVICE) -> dict:\n    \"\"\"\n    Classify a single EEG segment using BOTH EEGNet and Classic ML.\n    Called by the Signal Viewer backend when a user opens an EEG file.\n\n    Parameters\n    ----------\n    raw_eeg  : (n_channels, n_samples) float32 — raw, unfiltered EEG\n\n    Returns\n    -------\n    dict with keys:\n        ai_label, ai_confidence, ai_probs,\n        cl_label, cl_confidence, cl_probs,\n        agreement, is_abnormal, alert_message\n    \"\"\"\n    # ── 1. Filter + normalise ────────────────────────────────────────────────\n    eeg = preproc(raw_eeg.copy().astype(np.float32))\n\n    # Pad or centre-crop to exactly N_TIMEPOINTS\n    if eeg.shape[1] >= N_TIMEPOINTS:\n        mid = (eeg.shape[1] - N_TIMEPOINTS) // 2\n        eeg = eeg[:, mid: mid + N_TIMEPOINTS]\n    else:\n        eeg = np.pad(eeg, ((0, 0), (0, N_TIMEPOINTS - eeg.shape[1])))\n\n    std = eeg.std(axis=1, keepdims=True)\n    std[std < 1e-6] = 1.0\n    eeg = (eeg - eeg.mean(axis=1, keepdims=True)) / std\n\n    # ── 2. EEGNet inference ──────────────────────────────────────────────────\n    eegnet_model.eval()\n    with torch.no_grad():\n        x_t  = torch.tensor(eeg, dtype=torch.float32).unsqueeze(0).unsqueeze(0).to(device)\n        ai_p = eegnet_model.predict_proba(x_t).cpu().numpy()[0]\n\n    ai_idx  = int(ai_p.argmax())\n    ai_conf = float(ai_p.max())\n\n    # ── 3. Classic ML inference ──────────────────────────────────────────────\n    feats   = extract_classic_features(eeg).reshape(1, -1)\n    cl_p    = classic_clf.predict_proba(feats)[0]\n    cl_idx  = int(cl_p.argmax())\n    cl_conf = float(cl_p.max())\n\n    # ── 4. Build output dict ─────────────────────────────────────────────────\n    ai_label    = CLASS_NAMES[ai_idx]\n    cl_label    = CLASS_NAMES[cl_idx]\n    agreement   = (ai_idx == cl_idx)\n    is_abnormal = (ai_label != 'Normal') or (cl_label != 'Normal')\n\n    if agreement:\n        if is_abnormal:\n            msg = (f\"⚠️ ABNORMAL — {ai_label} detected. \"\n                   f\"Both AI ({ai_conf:.0%}) and Classic ML ({cl_conf:.0%}) agree.\")\n        else:\n            msg = (f\"✅ NORMAL EEG. \"\n                   f\"Both models agree (AI: {ai_conf:.0%}, ML: {cl_conf:.0%}).\")\n    else:\n        msg = (f\"⚡ CONFLICTING: AI → {ai_label} ({ai_conf:.0%}) | \"\n               f\"Classic ML → {cl_label} ({cl_conf:.0%}). Manual review recommended.\")\n\n    return {\n        'ai_label':      ai_label,\n        'ai_confidence': ai_conf,\n        'ai_probs':      dict(zip(CLASS_NAMES, ai_p.tolist())),\n        'cl_label':      cl_label,\n        'cl_confidence': cl_conf,\n        'cl_probs':      dict(zip(CLASS_NAMES, cl_p.tolist())),\n        'agreement':     agreement,\n        'is_abnormal':   is_abnormal,\n        'alert_message': msg\n    }\n\n\n# ── Demo: run on one raw segment per class ────────────────────────────────────\n# Note: eeg_data is already filtered+normalised by load_eeg_segment,\n# so we pass it through a dummy \"no-op\" preproc by wrapping it back as raw.\n# For a real viewer, pass the truly raw (unfiltered) signal.\nprint('=== Signal Viewer — Dual Classification Demo ===')\nfor cls_idx, cls_name in enumerate(CLASS_NAMES):\n    idx    = np.where(eeg_labels == cls_idx)[0][0]\n    raw    = eeg_data[idx]          # shape (8, 2000) — already normed\n    # Simulate classify_eeg_dual on a pre-normed sample\n    # (skip preproc step since data is already clean)\n    with torch.no_grad():\n        x_t  = torch.tensor(raw, dtype=torch.float32).unsqueeze(0).unsqueeze(0).to(DEVICE)\n        ai_p = model.predict_proba(x_t).cpu().numpy()[0]\n    feats  = extract_classic_features(raw).reshape(1, -1)\n    cl_p   = best_clf.predict_proba(feats)[0]\n    ai_idx_d  = int(ai_p.argmax()); ai_conf_d = float(ai_p.max())\n    cl_idx_d  = int(cl_p.argmax()); cl_conf_d = float(cl_p.max())\n    ai_lbl = CLASS_NAMES[ai_idx_d]; cl_lbl = CLASS_NAMES[cl_idx_d]\n    agree  = \"✅ AGREE\" if ai_lbl == cl_lbl else \"⚡ CONFLICT\"\n    print(f\"\\n[True: {cls_name}] {agree}\")\n    print(f\"  AI        → {ai_lbl} ({ai_conf_d:.0%})\")\n    print(f\"  Classic ML→ {cl_lbl} ({cl_conf_d:.0%})\")\n    print(f\"  AI probs  : { {k: f'{v:.2f}' for k,v in zip(CLASS_NAMES, ai_p)} }\")\n\nprint('\\n✅ classify_eeg_dual() ready for Signal Viewer integration')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:23:11.745417Z","iopub.execute_input":"2026-02-23T21:23:11.745945Z","iopub.status.idle":"2026-02-23T21:23:12.294459Z","shell.execute_reply.started":"2026-02-23T21:23:11.745922Z","shell.execute_reply":"2026-02-23T21:23:12.293582Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 14 — Final Summary Report","metadata":{}},{"cell_type":"code","source":"auc_ai = roc_auc_score(label_binarize(y_true_all, classes=list(range(N_CLASSES))),\n                        probs_ai, multi_class='ovr', average='macro')\nauc_cl = roc_auc_score(label_binarize(y_true_all, classes=list(range(N_CLASSES))),\n                        probs_cl_all, multi_class='ovr', average='macro')\n\nprint('=' * 70)\nprint('       FINAL COMPARISON REPORT — EEGNet AI vs Classic ML')\nprint(f'       Dataset: HMS micro-subset | N = {len(y_true_all)} | 5 classes')\nprint('=' * 70)\nprint(f'{\"Metric\":<30} {\"EEGNet AI\":>15} {f\"Classic ML ({best_clf_name})\": >20}')\nprint('-' * 70)\n\nrows = [\n    ('Accuracy',         accuracy_score(y_true_all, preds_ai_all),\n                         accuracy_score(y_true_all, preds_cl_all)),\n    ('Balanced Accuracy',balanced_accuracy_score(y_true_all, preds_ai_all),\n                         balanced_accuracy_score(y_true_all, preds_cl_all)),\n    ('Weighted F1',      f1_score(y_true_all, preds_ai_all, average='weighted'),\n                         f1_score(y_true_all, preds_cl_all, average='weighted')),\n    ('Macro F1',         f1_score(y_true_all, preds_ai_all, average='macro'),\n                         f1_score(y_true_all, preds_cl_all, average='macro')),\n    ('Macro ROC-AUC',    auc_ai, auc_cl),\n]\nfor name, ai_v, cl_v in rows:\n    winner = '← BETTER' if ai_v > cl_v else ('BETTER →' if cl_v > ai_v else '  TIE   ')\n    print(f'{name:<30} {ai_v:>15.4f} {cl_v:>20.4f}   {winner}')\n\nprint('-' * 70)\nprint(f'\\n  Agreement (both correct)   : {agree_correct/len(y_true_all)*100:.1f}%')\nprint(f'  AI only correct            : {ai_only_correct/len(y_true_all)*100:.1f}%')\nprint(f'  Classic ML only correct    : {cl_only_correct/len(y_true_all)*100:.1f}%')\nprint(f'  Both wrong                 : {both_wrong/len(y_true_all)*100:.1f}%')\nprint()\nprint('Output files saved to /kaggle/working/:')\nfor f in sorted(Path('/kaggle/working').glob('*.png')):\n    print(f'  {f.name}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:23:12.295659Z","iopub.execute_input":"2026-02-23T21:23:12.295947Z","iopub.status.idle":"2026-02-23T21:23:12.327118Z","shell.execute_reply.started":"2026-02-23T21:23:12.295922Z","shell.execute_reply":"2026-02-23T21:23:12.326539Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 15 — Save Classic ML Model\n\nSave the classic ML pipeline so the Signal Viewer backend can load it alongside EEGNet.","metadata":{}},{"cell_type":"code","source":"import joblib\n\njoblib.dump(best_clf, '/kaggle/working/classic_ml_eeg.joblib')\nprint('✅ Classic ML model saved: classic_ml_eeg.joblib')\n\n# Save feature extractor function source for backend reference\nimport inspect\nwith open('/kaggle/working/feature_extraction.py', 'w') as f:\n    f.write('# EEG Classic ML Feature Extraction\\n')\n    f.write('# Load with: import joblib; clf = joblib.load(\"classic_ml_eeg.joblib\")\\n\\n')\n    f.write('import numpy as np\\n')\n    f.write('from scipy import signal as sp_signal\\n')\n    f.write('from scipy.stats import kurtosis, skew\\n')\n    f.write('import antropy as ant\\n\\n')\n    f.write(inspect.getsource(bandpower) + '\\n\\n')\n    f.write(inspect.getsource(extract_classic_features) + '\\n\\n')\n    f.write(inspect.getsource(classify_eeg_dual) + '\\n')\n\nprint('✅ Feature extraction code saved: feature_extraction.py')\nprint('\\nAll outputs ready. Download from the Output tab.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:23:12.327920Z","iopub.execute_input":"2026-02-23T21:23:12.328229Z","iopub.status.idle":"2026-02-23T21:23:12.394068Z","shell.execute_reply.started":"2026-02-23T21:23:12.328207Z","shell.execute_reply":"2026-02-23T21:23:12.393545Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch, joblib, json\nfrom pathlib import Path\n\nOUT = Path('/kaggle/working')\n\n# ── 1. EEGNetV2 (PyTorch) ─────────────────────────────────────────────────────\n# Save full model state dict\ntorch.save(model.state_dict(), OUT / 'eegnet_v2.pth')\nprint('✅ EEGNet saved: eegnet_v2.pth')\n\n# Also export as TorchScript (recommended for FastAPI — no class definition needed)\nmodel.eval()\nexample_input = torch.zeros(1, 1, N_CHANNELS, N_TIMEPOINTS).to(DEVICE)\nscripted = torch.jit.trace(model, example_input)\nscripted.save(OUT / 'eegnet_v2_scripted.pt')\nprint('✅ EEGNet TorchScript saved: eegnet_v2_scripted.pt')\n\n# ── 2. Classic ML (sklearn) ───────────────────────────────────────────────────\njoblib.dump(best_clf, OUT / 'classic_ml_eeg.joblib')\nprint('✅ Classic ML saved: classic_ml_eeg.joblib')\n\n# ── 3. Metadata (give to frontend dev so they know labels/config) ─────────────\nmeta_info = {\n    \"class_names\": CLASS_NAMES,\n    \"n_channels\": N_CHANNELS,\n    \"n_timepoints\": N_TIMEPOINTS,\n    \"sfreq\": SFREQ,\n    \"channels\": CHANNELS,\n    \"best_classic_clf\": best_clf_name,\n}\nwith open(OUT / 'model_metadata.json', 'w') as f:\n    json.dump(meta_info, f, indent=2)\nprint('✅ Metadata saved: model_metadata.json')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T21:23:12.394831Z","iopub.execute_input":"2026-02-23T21:23:12.395194Z","iopub.status.idle":"2026-02-23T21:23:12.711280Z","shell.execute_reply.started":"2026-02-23T21:23:12.395166Z","shell.execute_reply":"2026-02-23T21:23:12.710552Z"}},"outputs":[],"execution_count":null}]}