{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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":"gpu","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install ipywidgets\n!pip install pymupdf\n!pip install python-docx","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:17:00.714077Z","iopub.execute_input":"2025-05-01T00:17:00.714996Z","iopub.status.idle":"2025-05-01T00:17:09.565161Z","shell.execute_reply.started":"2025-05-01T00:17:00.714963Z","shell.execute_reply":"2025-05-01T00:17:09.564441Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nos.environ[\"backend\"] = \"jax\"\n\nimport tensorflow as tf\n\nimport jax.numpy as jnp\nfrom jax import device_put\n\nimport keras_cv\nimport keras\nfrom keras import ops\nfrom keras import backend as K\nfrom keras import layers, models, Model\nfrom keras.callbacks import EarlyStopping, ModelCheckpoint\nfrom keras import mixed_precision\n\nfrom sklearn.ensemble import RandomForestClassifier\nfrom sklearn.model_selection import train_test_split, StratifiedKFold\nfrom sklearn.metrics import classification_report, confusion_matrix, accuracy_score\n\nfrom scipy.signal import butter, filtfilt, welch, spectrogram\nfrom scipy.stats import skew, kurtosis\nfrom scipy.signal import lfilter\n\nfrom mne.preprocessing import ICA\nfrom mne import create_info, EpochsArray\nfrom mne.io import RawArray\nfrom mne.channels import make_standard_montage\n\nfrom docx import Document\nfrom docx.shared import Inches\nfrom io import BytesIO\n\nimport cv2\nimport pandas as pd\nimport numpy as np\nimport shutil\nfrom glob import glob\nfrom tqdm.notebook import tqdm\nfrom joblib import Parallel, delayed\nimport pywt\nimport random\nimport math\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:17:09.566885Z","iopub.execute_input":"2025-05-01T00:17:09.567157Z","iopub.status.idle":"2025-05-01T00:17:32.287887Z","shell.execute_reply.started":"2025-05-01T00:17:09.567135Z","shell.execute_reply":"2025-05-01T00:17:32.287345Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"TensorFlow:\", tf.__version__)\nprint(\"Keras:\", keras.__version__)\nprint(\"KerasCV:\", keras_cv.__version__)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:17:32.288539Z","iopub.execute_input":"2025-05-01T00:17:32.288998Z","iopub.status.idle":"2025-05-01T00:17:32.293321Z","shell.execute_reply.started":"2025-05-01T00:17:32.288979Z","shell.execute_reply":"2025-05-01T00:17:32.292443Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"癫痫发作检测的数据预处理与配置设置\n\n在本代码单元中，我们配置重要设置项，包括随机种子（random seed）和批量大小（batch size），并对数据集进行预处理以用于模型训练。我们加载数据集，为每个脑电图（EEG）ID选择一个样本，并确保数据已准备好进行后续处理。\n","metadata":{}},{"cell_type":"code","source":"class CFG:\n    verbose = 1  # Verbosity\n    pandas_warning = None #Pandas Chained Assignment Warning Mode, default = 'warn'\n    seed = 42  # Random seed\n    batch_size = 32\n    epochs = 20\n\npd.options.mode.chained_assignment = CFG.pandas_warning\nkeras.utils.set_random_seed(CFG.seed)\n\nBASE_PATH = \"/kaggle/input/hms-harmful-brain-activity-classification\"\nEEG_DIR = \"/tmp/dataset/hms-hbac/numpy_eegs\"\nos.makedirs(EEG_DIR, exist_ok=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:17:32.294043Z","iopub.execute_input":"2025-05-01T00:17:32.294303Z","iopub.status.idle":"2025-05-01T00:17:32.311883Z","shell.execute_reply.started":"2025-05-01T00:17:32.294274Z","shell.execute_reply":"2025-05-01T00:17:32.311344Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(f'{BASE_PATH}/train.csv')\ndisplay(df.info())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:17:32.313712Z","iopub.execute_input":"2025-05-01T00:17:32.313889Z","iopub.status.idle":"2025-05-01T00:17:32.620572Z","shell.execute_reply.started":"2025-05-01T00:17:32.313871Z","shell.execute_reply":"2025-05-01T00:17:32.619849Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# As each file points to only one EEG ID, we will select only 1 sample from metadata.\ndf = df.groupby(\"eeg_id\").head(1).reset_index(drop=True)\ndf.head(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:17:32.621369Z","iopub.execute_input":"2025-05-01T00:17:32.621625Z","iopub.status.idle":"2025-05-01T00:17:32.654883Z","shell.execute_reply.started":"2025-05-01T00:17:32.621608Z","shell.execute_reply":"2025-05-01T00:17:32.654236Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"癫痫发作检测的类别标注与数据平衡\n\n本代码单元主要集中于为癫痫发作检测分配二分类标签，将数据集划分为发作类（seizure）和非发作类（non-seizure），并通过欠采样方法平衡数据。我们确保训练模型时数据具有1:1的类别分布。","metadata":{}},{"cell_type":"code","source":"# As we are working on only seizure / non-seizure classification, we will not use default class labels of the dataset.\ndf['is_seizure'] = df['expert_consensus'] == 'Seizure'\n\n# Train + Valid + Test (We will use train_test_split to split the data)\ndf['eeg_path'] = f'{BASE_PATH}/train_eegs/'+df['eeg_id'].astype(str)+'.parquet'\n# df['eeg2_path'] = f'{EEG_DIR}/'+df['eeg_id'].astype(str)+'.npy'\ndf['spec_path'] = f'{BASE_PATH}/train_spectrograms/'+df['spectrogram_id'].astype(str)+'.parquet'\n\ndisplay(df.head(2))\nprint(f\"Total Data: {len(df)}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:17:32.655724Z","iopub.execute_input":"2025-05-01T00:17:32.655954Z","iopub.status.idle":"2025-05-01T00:17:32.690337Z","shell.execute_reply.started":"2025-05-01T00:17:32.655938Z","shell.execute_reply":"2025-05-01T00:17:32.689577Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Filter for 'Seizure' class and a subset of 'non-Seizure' classes\nseizure_df = df[df['expert_consensus'] == 'Seizure']\nnon_seizure_df = df[df['expert_consensus'] != 'Seizure']\n\n# Undersample non-Seizure data (e.g., 1:1 ratio)\nnon_seizure_sampled = non_seizure_df.sample(n=len(seizure_df), random_state=CFG.seed)\n\n# Combine and shuffle the new dataset\ndf = pd.concat([seizure_df, non_seizure_sampled]).sample(frac=1, random_state=CFG.seed)\ndf.sort_values(by=[\"eeg_id\", \"eeg_label_offset_seconds\"], ignore_index=True, inplace=True)\n\nclass_counts = df['expert_consensus'].value_counts()\ndisplay(class_counts)\nclass_counts.plot(kind='bar', title='Class Distribution')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:17:32.691148Z","iopub.execute_input":"2025-05-01T00:17:32.691446Z","iopub.status.idle":"2025-05-01T00:17:32.980387Z","shell.execute_reply.started":"2025-05-01T00:17:32.691422Z","shell.execute_reply":"2025-05-01T00:17:32.979640Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"检查数据样本：EEG 和谱图文件\n\n在这个单元中，我们从数据集中随机采样 EEG 和谱图文件，以检查它们的形状、特征范围和整体特征。这有助于我们在模型训练之前理解数据的结构和分布。","metadata":{}},{"cell_type":"code","source":"random_samples = np.random.randint(0, len(df) - 1, size=(10))\nfor i in random_samples:\n    reading_spec = pd.read_parquet(f\"{BASE_PATH}/train_spectrograms/{df['spectrogram_id'].iloc[i]}.parquet\")\n    reading_eeg = pd.read_parquet(f\"{BASE_PATH}/train_eegs/{df['eeg_id'].iloc[i]}.parquet\")\n    print(f\"Rows in EEG File: {reading_eeg.shape[0]}\\t Rows in Spectrogram File: {reading_spec.shape[0]}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:17:32.981230Z","iopub.execute_input":"2025-05-01T00:17:32.981514Z","iopub.status.idle":"2025-05-01T00:17:34.244113Z","shell.execute_reply.started":"2025-05-01T00:17:32.981491Z","shell.execute_reply":"2025-05-01T00:17:34.243310Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Information regarding a single spectrogram\ntrain_df_samples = len(df) # There are total 106800 training samples\nprint(f\"Total Samples: {train_df_samples}\")\n\nrandom_index = np.random.randint(0, train_df_samples - 1)\n\nrandom_spec_id = df['spectrogram_id'].iloc[random_index]\nreading_spec = pd.read_parquet(f\"{BASE_PATH}/train_spectrograms/{random_spec_id}.parquet\")\nrandom_eeg_id = df['eeg_id'].iloc[random_index]\nreading_eeg = pd.read_parquet(f\"{BASE_PATH}/train_eegs/{random_eeg_id}.parquet\")\n\nprint(\"\\nSPECTROGRAM INFORMATION:\")\nprint(f\"Shape of spectrogram: {reading_spec.shape}\")\nprint(f\"Spectrogram data min: {reading_spec.iloc[:, 1:].values.min()}\")\nprint(f\"Spectrogram data max: {reading_spec.iloc[:, 1:].values.max()}\")\nprint(f\"Spectrogram data mean: {reading_spec.iloc[:, 1:].values.mean()}\")\nprint(\"Features of Spectrogram:\\n\")\ndisplay(reading_spec.columns)\n\nprint(\"\\nEEG INFORMATION:\")\nprint(f\"Shape of EEG: {reading_eeg.shape}\")\nprint(f\"EEG data min: {reading_eeg.iloc[:, :-1].values.min()}\")\nprint(f\"EEG data max: {reading_eeg.iloc[:, :-1].values.max()}\")\nprint(f\"EEG data mean: {reading_eeg.iloc[:, :-1].values.mean()}\")\nprint(\"Features of Spectrogram:\\n\")\ndisplay(reading_eeg.columns)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:17:34.245161Z","iopub.execute_input":"2025-05-01T00:17:34.245481Z","iopub.status.idle":"2025-05-01T00:17:34.326478Z","shell.execute_reply.started":"2025-05-01T00:17:34.245449Z","shell.execute_reply":"2025-05-01T00:17:34.325870Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"癫痫检测数据集中示例图形的预览\n\n在这个单元格中，我们展示了示例图形，这些图形代表了我们的夺取检测模型的数据集可能看起来的样子。在深入研究实际数据集之前，我们先预览一下您将在实际数据集中遇到的图形类型和数据可视化。","metadata":{}},{"cell_type":"code","source":"import fitz  # PyMuPDF\nfrom PIL import Image\nfrom io import BytesIO\n\n# Folder containing the PDF files\nfolder_path = f\"{BASE_PATH}/example_figures\"\n\n# List all files in the folder\npdf_files = [f for f in os.listdir(folder_path) if f.endswith('.pdf')]\n\n# Loop through all PDF files and display them\nfor pdf_file in pdf_files:\n    # Construct the full path of the PDF file\n    pdf_path = os.path.join(folder_path, pdf_file)\n    \n    # Open the PDF file\n    doc = fitz.open(pdf_path)\n    \n    # Loop through the pages and display each one\n    for page_num in range(doc.page_count):\n        page = doc.load_page(page_num)  # Load the page\n        pix = page.get_pixmap()  # Convert page to a pixmap (image)\n\n        # Convert to a PIL Image\n        img = Image.open(BytesIO(pix.tobytes(\"png\")))\n\n        # Display the image in the notebook\n        display(img)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:17:34.327201Z","iopub.execute_input":"2025-05-01T00:17:34.327481Z","iopub.status.idle":"2025-05-01T00:17:38.975448Z","shell.execute_reply.started":"2025-05-01T00:17:34.327463Z","shell.execute_reply":"2025-05-01T00:17:38.974620Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"频谱图的数据完整性与信号质量检查\n\n在本代码单元中，我们验证频谱图时间数据的一致性，检查信号随时间变化的质量，并对频谱图进行可视化，以确保数据干净且适合用于模型训练。","metadata":{}},{"cell_type":"code","source":"reading_spec = pd.read_parquet(f\"{BASE_PATH}/train_spectrograms/{df['spectrogram_id'].iloc[np.random.randint(0, train_df_samples - 1)]}.parquet\")\n\n# Step 1: Verify time consistency\ndef check_time_consistency(spectrogram_data, metadata, spectrogram_id):\n    \"\"\"\n    Verify if the 'time' column in the spectrogram data matches the metadata offsets.\n    \"\"\"\n    sub_metadata = metadata[metadata[\"spectrogram_id\"] == spectrogram_id]\n    metadata_times = sub_metadata[\"spectrogram_label_offset_seconds\"].values\n    data_times = spectrogram_data[\"time\"].values\n\n    # Check for missing or mismatched times\n    missing_times = set(metadata_times) - set(data_times)\n    extra_times = set(data_times) - set(metadata_times)\n\n    print(\"Time Consistency Check:\")\n    if missing_times:\n        print(f\"  Missing time points in spectrogram data: {missing_times}\")\n    if extra_times:\n        print(f\"  Extra time points in spectrogram data: {extra_times}\")\n    if not missing_times and not extra_times:\n        print(\"  All time points match!\")\n    return missing_times, extra_times\n\n# missing_times, extra_times = check_time_consistency(reading_spec, train_df, random_id)\n\n# Step 2: Check for signal quality across time\ndef check_signal_quality(df, min_variance=1e-5):\n    \"\"\"\n    Analyze the signal quality over time for anomalies or low variance.\n    \"\"\"\n    print(\"\\nSignal Quality Check:\")\n    time_series = df[\"time\"].diff().dropna()\n\n    # Check for irregular time gaps\n    if time_series.std() > 0.1:  # Adjust threshold as needed\n        print(f\"  Irregular time gaps detected! Std Dev: {time_series.std()}\")\n    else:\n        print(\"  Time gaps are consistent.\")\n\n    # Check variance in frequency bands over time\n    variances = df.iloc[:, 1:].var(axis=0)  # Exclude 'time' column\n    low_variance_cols = variances[variances < min_variance]\n    if not low_variance_cols.empty:\n        print(f\"  Low variance detected in columns: {low_variance_cols.index.tolist()}\")\n    else:\n        print(\"  All frequency bands have sufficient variance.\")\n\ncheck_signal_quality(reading_spec)\n    \nimport matplotlib.colors as colors\n\n# Step 3: Visualize time-based trends\ndef plot_spectrogram_over_time(df):\n    \"\"\"\n    Plot the spectrogram over time to visually inspect for noise or artifacts.\n    \"\"\"\n\n    # Extract time and spectrogram data\n    time = df[\"time\"].values\n    spectrogram_data = df.iloc[:, 1:].values  # Excluding the time column\n\n    # Check the min, max, and mean values to understand the scale\n    spec_min = spectrogram_data.min()\n    spec_max = spectrogram_data.max()\n    print(f\"Spectrogram data min: {spec_min}\")\n    print(f\"Spectrogram data max: {spec_max}\")\n    print(f\"Spectrogram data mean: {spectrogram_data.mean()}\")\n\n    colormaps = ['hsv', \"viridis\", 'gist_ncar', 'gist_ncar_r', 'gist_rainbow', 'gist_rainbow_r']\n    plt.figure(figsize=(12, len(colormaps) * 6))\n    for i, cm in enumerate(colormaps):\n        plt.subplot(len(colormaps), 1, i + 1)\n        plt.imshow(spectrogram_data.T, aspect=\"auto\", cmap=cm,\n                   norm=colors.LogNorm(),\n                   extent=[time.min(), time.max(), 0, spectrogram_data.shape[0] - 1])\n        plt.colorbar(label=\"Amplitude\")\n        plt.xlabel(\"Time (seconds)\")\n        plt.ylabel(\"Frequency Bands\")\n        plt.title(f\"Spectrogram Over Time (Color Map: {cm})\")\n    plt.show()\n\nplot_spectrogram_over_time(reading_spec)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:17:38.976493Z","iopub.execute_input":"2025-05-01T00:17:38.976794Z","iopub.status.idle":"2025-05-01T00:17:42.963049Z","shell.execute_reply.started":"2025-05-01T00:17:38.976773Z","shell.execute_reply":"2025-05-01T00:17:42.961834Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"脑电图（EEG）信号质量分析与滤波\n\n本代码单元对脑电图信号质量进行详细分析，包括方差检查、频段功率计算、电极间相关性分析以及噪声信号的滤波处理。我们的目标是提高数据质量，以实现有效的癫痫发作检测。","metadata":{}},{"cell_type":"code","source":"# Bandpass filter function (to remove noise outside 0.5-50 Hz range)\ndef bandpass_filter(data, lowcut=0.5, highcut=50, fs=200, order=5):\n    nyquist = 0.5 * fs\n    low = lowcut / nyquist\n    high = highcut / nyquist\n    b, a = butter(order, [low, high], btype=\"band\")\n    return filtfilt(b, a, data, axis=0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:17:42.963870Z","iopub.execute_input":"2025-05-01T00:17:42.964083Z","iopub.status.idle":"2025-05-01T00:17:42.968685Z","shell.execute_reply.started":"2025-05-01T00:17:42.964067Z","shell.execute_reply":"2025-05-01T00:17:42.967922Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load a sample EEG file (update with actual file paths)\nreading_eeg = pd.read_parquet(f\"{BASE_PATH}/train_eegs/{df['eeg_id'].iloc[np.random.randint(0, train_df_samples - 1)]}.parquet\")\neeg_data = reading_eeg.drop(columns=[\"EKG\"])\n\n# Sampling frequency (Hz)\nFS = 200\n\n# Analysis Results Dictionary\nresults = {}\n\n### 1. Check Signal Variance\nvariances = eeg_data.var()\nmean_variance = variances.mean()\nstd_variance = variances.std()\nhigh_variance_threshold = mean_variance + 2 * std_variance\nlow_variance_threshold = mean_variance - 2 * std_variance\n\n# Identify electrodes with high/low variance\nhigh_variance_electrodes = variances[variances > high_variance_threshold]\nlow_variance_electrodes = variances[variances < low_variance_threshold]\n\nresults[\"high_variance\"] = high_variance_electrodes\nresults[\"low_variance\"] = low_variance_electrodes\n\n### 2. Frequency Analysis (FFT)\nfor column in eeg_data.columns:\n    signal = eeg_data[column]\n    freq = np.fft.rfftfreq(len(signal), d=1 / FS)  # Frequency axis\n    fft_magnitude = np.abs(np.fft.rfft(signal))  # Magnitudes\n\n    # Calculate power in specific frequency bands\n    power = { \n        \"delta\": np.sum(fft_magnitude[(freq >= 0.5) & (freq < 4)]),\n        \"theta\": np.sum(fft_magnitude[(freq >= 4) & (freq < 8)]),\n        \"alpha\": np.sum(fft_magnitude[(freq >= 8) & (freq < 13)]),\n        \"beta\": np.sum(fft_magnitude[(freq >= 13) & (freq < 30)]),\n        \"gamma\": np.sum(fft_magnitude[(freq >= 30) & (freq < 50)]),\n        \"high_freq_noise\": np.sum(fft_magnitude[freq >= 50]),\n    }\n\n    # Print power in each band for the first few channels\n    if column in eeg_data.columns[:3]:\n        print(f\"Power for {column}: {power}\")\n    \n    # Record high-frequency noise if it dominates\n    if power[\"high_freq_noise\"] > 0.2 * sum(power.values()):\n        results[f\"{column}_high_freq_noise\"] = True\n    else:\n        results[f\"{column}_high_freq_noise\"] = False\n\n### 3. Signal Consistency Check\n# Compute correlation between electrodes\ncorr_matrix = eeg_data.corr()\ncorrelated_channels = (corr_matrix > 0.9).sum(axis=1) - 1  # Exclude self-correlation\n\n# Mark channels with high correlation to others (possible redundancy)\nresults[\"highly_correlated_electrodes\"] = correlated_channels[correlated_channels > 5]\n\n# Visualize correlation matrix\nplt.figure(figsize=(30, 30))\nsns.set(font_scale=1.5)\nsns.heatmap(corr_matrix, annot=True, cmap=\"coolwarm\")\nplt.title(\"Correlation Matrix of EEG Signals\")\nplt.show()\n\n### 4. Filtered Signal Inspection\nfor column in eeg_data.columns[:3]:  # Check a few electrodes\n    filtered_signal = bandpass_filter(eeg_data[column], 0.5, 50, FS)\n    plt.figure(figsize=(10, 3))\n    plt.plot(filtered_signal, label=f\"Filtered {column}\")\n    plt.title(f\"Filtered EEG Signal: {column}\")\n    plt.xlabel(\"Time (samples)\")\n    plt.ylabel(\"Amplitude\")\n    plt.legend()\n    plt.show()\n\n### Summary Report\nprint(\"\\n=== EEG Signal Noise Analysis Summary ===\")\nprint(f\"Mean Variance: {mean_variance}\")\nprint(f\"High Variance Threshold: {high_variance_threshold}\")\nprint(f\"Low Variance Threshold: {low_variance_threshold}\")\nprint(f\"Electrodes with High Variance:\\n{high_variance_electrodes}\")\nprint(f\"Electrodes with Low Variance:\\n{low_variance_electrodes}\")\nprint(f\"Electrodes with High Variance:\\n{results['high_variance']}\")\nprint(f\"Electrodes with Low Variance:\\n{results['low_variance']}\")\nprint(f\"Electrodes with High-Frequency Noise:\")\nfor key, val in results.items():\n    if \"high_freq_noise\" in key and val:\n        print(f\"  {key.replace('_high_freq_noise', '')}\")\n\nprint(\"\\nHighly Correlated Electrodes (likely redundant):\")\nprint(results[\"highly_correlated_electrodes\"])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:17:42.971407Z","iopub.execute_input":"2025-05-01T00:17:42.971610Z","iopub.status.idle":"2025-05-01T00:17:45.079441Z","shell.execute_reply.started":"2025-05-01T00:17:42.971594Z","shell.execute_reply":"2025-05-01T00:17:45.078709Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"脑电图（EEG）信号的时域与频域分析\n\n本代码单元对脑电图信号的时域和频域特性进行可视化，比较发作样本与非发作样本。在绘图之前，信号经过裁剪和填充缺失值（NaN）的预处理，以提高信号质量。","metadata":{}},{"cell_type":"code","source":"def plot_temporal_signals(eeg_df, num_samples=3, electrodes=[\"Fp1\", \"F3\", \"C3\"]):\n    \"\"\"\n    Plots temporal EEG signals for a random selection of seizure and non-seizure samples.\n    \n    Parameters:\n        eeg_df (DataFrame): Dataframe containing EEG data and metadata.\n        num_samples (int): Number of samples to visualize for seizure and non-seizure each.\n        electrodes (list): List of electrode names to plot.\n    \"\"\"\n    seizure_samples = eeg_df[eeg_df['is_seizure']].sample(n=num_samples, random_state=42)\n    non_seizure_samples = eeg_df[~eeg_df['is_seizure']].sample(n=num_samples, random_state=42)\n\n    # Total number of samples to plot\n    total_samples = len(seizure_samples) + len(non_seizure_samples)\n    \n    # Adjust the number of rows in the subplot grid\n    fig, axes = plt.subplots(total_samples, len(electrodes), figsize=(15, 5 * total_samples))\n\n    for idx, sample in enumerate([*seizure_samples.iterrows(), *non_seizure_samples.iterrows()]):\n        sample_data = sample[1]\n        eeg_path = sample_data['eeg_path']\n        eeg_data = pd.read_parquet(eeg_path)\n\n        for col_idx, electrode in enumerate(electrodes):\n            signal = eeg_data[electrode]\n\n            # Handle single subplot case when there is only 1 row\n            ax = axes[idx, col_idx] if total_samples > 1 else axes[col_idx]\n            ax.plot(signal)\n            ax.set_title(f\"{'Seizure' if sample_data['is_seizure'] else 'Non-Seizure'} - {electrode}\")\n            ax.set_xlabel(\"Time (ms)\")\n            ax.set_ylabel(\"Amplitude\")\n\n    plt.tight_layout()\n    plt.savefig(\"/kaggle/working/Temporal Consistency Analysis.png\")\n    plt.show()\n\n# Execute the plotting\nplot_temporal_signals(df, num_samples=3, electrodes=[\"Fp1\", \"F3\", \"C3\"])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:17:45.080104Z","iopub.execute_input":"2025-05-01T00:17:45.080327Z","iopub.status.idle":"2025-05-01T00:17:51.492807Z","shell.execute_reply.started":"2025-05-01T00:17:45.080303Z","shell.execute_reply":"2025-05-01T00:17:51.492127Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def preprocess_signal(signal, clip_threshold=1e6):\n    \"\"\"\n    Preprocess the EEG signal by clipping and replacing NaN values.\n\n    Parameters:\n        signal (array-like): The raw EEG signal.\n        clip_threshold (float): The threshold to clip spikes in the signal.\n\n    Returns:\n        array-like: The preprocessed EEG signal.\n    \"\"\"\n    # Replace NaN values with the mean of the signal\n    signal = np.nan_to_num(signal, nan=np.nanmean(signal))\n\n    # Clip large spikes in the signal\n    signal = np.clip(signal, -clip_threshold, clip_threshold)\n\n    return signal\n\n# Update the frequency analysis function\ndef plot_frequency_analysis_with_preprocessing(eeg_df, num_samples=3, electrodes=[\"Fp1\", \"F3\", \"C3\"], fs=200):\n    seizure_samples = eeg_df[eeg_df['is_seizure']].sample(n=num_samples, random_state=42)\n    non_seizure_samples = eeg_df[~eeg_df['is_seizure']].sample(n=num_samples, random_state=42)\n\n    total_samples = len(seizure_samples) + len(non_seizure_samples)\n    fig, axes = plt.subplots(total_samples, len(electrodes), figsize=(15, 5 * total_samples))\n\n    for idx, sample in enumerate([*seizure_samples.iterrows(), *non_seizure_samples.iterrows()]):\n        sample_data = sample[1]\n        eeg_path = sample_data['eeg_path']\n        eeg_data = pd.read_parquet(eeg_path)\n\n        for col_idx, electrode in enumerate(electrodes):\n            signal = preprocess_signal(eeg_data[electrode])\n\n            freqs, psd = welch(signal, fs=fs, nperseg=fs * 2)\n\n            ax = axes[idx, col_idx] if total_samples > 1 else axes[col_idx]\n            ax.plot(freqs, 10 * np.log10(psd + 1e-12))  # Add small constant to avoid log10(0)\n            ax.set_title(f\"{'Seizure' if sample_data['is_seizure'] else 'Non-Seizure'} - {electrode}\")\n            ax.set_xlabel(\"Frequency (Hz)\")\n            ax.set_ylabel(\"Power (dB/Hz)\")\n\n    plt.tight_layout()\n    plt.savefig(\"/kaggle/working/Frequency Analysis.png\")\n    plt.show()\n\n# Run the updated analysis\nplot_frequency_analysis_with_preprocessing(df, num_samples=3, electrodes=[\"Fp1\", \"F3\", \"C3\"], fs=200)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:17:51.493680Z","iopub.execute_input":"2025-05-01T00:17:51.493929Z","iopub.status.idle":"2025-05-01T00:17:57.129698Z","shell.execute_reply.started":"2025-05-01T00:17:51.493910Z","shell.execute_reply":"2025-05-01T00:17:57.128998Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"脑电图（EEG）数据分析与报告生成 📊\n在本节中，我们对脑电图数据进行全面分析，以生成具有洞察力的可视化图表、统计摘要以及结构清晰的报告 📝。具体内容包括：\n\n对原始脑电图信号进行可视化检查。\n\n进行功率谱密度（PSD）分析以探索频率成分。\n\n对信号进行预处理和独立成分分析（ICA）分解。\n\n进行频带特定功率分析以及信噪比（SNR）计算。\n\n分析统计特性以及分段（epoch）级别的噪声。\n\n生成包含所有分析结果的综合性Word报告。\n\n让我们开始深入分析脑电图数据吧！","metadata":{}},{"cell_type":"code","source":"# Create a folder for saving output if it doesn't exist\noutput_dir = '/kaggle/working/EEG Analysis'\nif not os.path.exists(output_dir):\n    os.makedirs(output_dir)\n\n# Create a Word document\ndoc = Document()\ndoc.add_heading('EEG Data Analysis', 0)\n\n# Function to add figures to the Word document\ndef add_figure_to_doc(fig, caption):\n    img_stream = BytesIO()\n    fig.savefig(img_stream, format='png')\n    img_stream.seek(0)  # Rewind to the beginning of the image\n    doc.add_picture(img_stream, width=Inches(5))  # Add image to the Word doc\n    doc.add_paragraph(caption)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:17:57.130520Z","iopub.execute_input":"2025-05-01T00:17:57.130822Z","iopub.status.idle":"2025-05-01T00:17:57.156558Z","shell.execute_reply.started":"2025-05-01T00:17:57.130799Z","shell.execute_reply":"2025-05-01T00:17:57.155835Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sampling_rate = 200  # Replace with your actual EEG sampling rate\n\n# Channels\nchannels = eeg_data.columns\n\n# 1. Visual Inspection: Plot raw EEG signals\ndef plot_raw_signals(eeg_data, channels, duration=5, fs=200):\n    time = np.arange(eeg_data.shape[0]) / fs\n    fig, axes = plt.subplots(5, 1, figsize=(15, 10))\n    for i, channel in enumerate(channels[:5]):  # Plot first 5 channels\n        axes[i].plot(time[:fs * duration], eeg_data[channel].iloc[:fs * duration])\n        axes[i].set_title(f\"Raw EEG Signal - {channel}\")\n        axes[i].set_xlabel(\"Time (s)\")\n        axes[i].set_ylabel(\"Amplitude\")\n    plt.tight_layout()\n    add_figure_to_doc(fig, 'Raw EEG Signal for First 5 Channels')\n    plt.show()\n\nplot_raw_signals(eeg_data, channels)\n\n# 2. Power Spectral Density (PSD) Analysis\ndef plot_psd(eeg_data, channels, fs=200):\n    fig, axes = plt.subplots(5, 1, figsize=(15, 10))\n    for i, channel in enumerate(channels[:5]):  # Plot first 5 channels\n        freqs, psd = welch(eeg_data[channel], fs=fs, nperseg=fs * 2)\n        axes[i].semilogy(freqs, psd)\n        axes[i].set_title(f\"Power Spectral Density - {channel}\")\n        axes[i].set_xlabel(\"Frequency (Hz)\")\n        axes[i].set_ylabel(\"Power (dB/Hz)\")\n    plt.tight_layout()\n    add_figure_to_doc(fig, 'Power Spectral Density for First 5 Channels')\n    plt.show()\n\nplot_psd(eeg_data, channels)\n\n# 3. Check for Missing Data or Flat Channels\ndef check_missing_flat_channels(eeg_data):\n    channel_stats = eeg_data.describe().T\n    flat_channels = channel_stats.loc[channel_stats['std'] == 0].index.tolist()\n    doc.add_paragraph(f\"Flat Channels: {flat_channels}\")\n    print(f\"Flat Channels: {flat_channels}\")\n    return flat_channels\n\nflat_channels = check_missing_flat_channels(eeg_data)\n\n# 4. Band-Specific Analysis\ndef compute_band_power(eeg_data, fs=200):\n    bands = {\"Delta\": (0.5, 4), \"Theta\": (4, 8), \"Alpha\": (8, 13), \"Beta\": (13, 30), \"Gamma\": (30, 40)}\n    band_power = {}\n    for band, (low, high) in bands.items():\n        band_power[band] = eeg_data.apply(lambda x: welch(x, fs=fs, nperseg=fs * 2)[1][\n            (low <= welch(x, fs=fs, nperseg=fs * 2)[0]) & \n            (welch(x, fs=fs, nperseg=fs * 2)[0] <= high)].sum(), axis=0)\n    band_power = pd.DataFrame(band_power)\n    band_power.plot(kind='bar', figsize=(12, 6), title=\"Relative Band Power\")\n    plt.ylabel(\"Power\")\n    plt.xlabel(\"Channels\")\n    plt.tight_layout()\n    add_figure_to_doc(plt, 'Band-Specific Power Analysis')\n    plt.show()\n\ncompute_band_power(eeg_data)\n\n# 5. Independent Component Analysis (ICA)\ndef run_ica(eeg_data, sampling_rate=200, highpass_freq=1.0):\n    n_channels = len(channels)\n    info = create_info(ch_names=list(channels), sfreq=sampling_rate, ch_types=[\"eeg\"] * n_channels)\n    \n    # Create random 3D coordinates for each channel\n    # This is just for visualization purposes, ideally you'd use a proper montage\n    info.set_montage('standard_1020')  # Standard system of EEG recordings\n    \n    # Create a RawArray from the EEG data\n    raw_array = RawArray(eeg_data.values.T, info)\n    \n    # Apply high-pass filter to the data\n    raw_array.filter(l_freq=highpass_freq, h_freq=None)  # High-pass filter with 1.0 Hz lower bound\n    \n    ica = ICA(n_components=n_channels, random_state=42)\n    ica.fit(raw_array)  # Now fitting to RawArray\n    \n    # Plot ICA components and save the figure manually\n    ica_fig = ica.plot_components(picks=None, show=False)  # This returns a figure\n    ica_fig.savefig(os.path.join(output_dir, 'ica_components.png'))  # Save ICA components plot to file\n    \n    # Add the ICA components plot to the document\n    add_figure_to_doc(ica_fig, 'Independent Component Analysis (ICA) Components')\n    return ica\n\nica = run_ica(eeg_data)\n\n# 6. Correlation Between Channels\ndef plot_channel_correlation(eeg_data):\n    corr_matrix = eeg_data.corr()\n    fig, ax = plt.subplots(figsize=(10, 8))\n    cax = ax.imshow(corr_matrix, cmap='viridis', interpolation='none')\n    fig.colorbar(cax, ax=ax, label='Correlation')\n    ax.set_title(\"Channel Correlation Matrix\")\n    ax.set_xlabel(\"Channels\")\n    ax.set_ylabel(\"Channels\")\n    plt.tight_layout()\n    add_figure_to_doc(fig, 'Channel Correlation Matrix')\n    plt.show()\n\nplot_channel_correlation(eeg_data)\n\n# 7. Signal-to-Noise Ratio (SNR)\ndef compute_snr(eeg_data):\n    signal_power = (eeg_data ** 2).mean(axis=0)\n    noise_power = eeg_data.diff().dropna() ** 2\n    noise_power = noise_power.mean(axis=0)\n    snr = 10 * np.log10(signal_power / noise_power)\n    doc.add_paragraph(\"Signal-to-Noise Ratio (SNR) for each channel:\")\n    doc.add_paragraph(snr.to_string())\n    print(\"SNR for each channel:\")\n    print(snr)\n    return snr\n\nsnr = compute_snr(eeg_data)\n\n# 8. Epoch-Level Noise Analysis\ndef check_epoch_noise(eeg_data, epoch_duration=2, fs=200):\n    epoch_samples = fs * epoch_duration\n    num_epochs = eeg_data.shape[0] // epoch_samples\n    variances = []\n    for i in range(num_epochs):\n        epoch = eeg_data.iloc[i * epoch_samples:(i + 1) * epoch_samples]\n        variances.append(epoch.var(axis=0))\n    variances = pd.DataFrame(variances)\n    fig, ax = plt.subplots(figsize=(10, 6))\n    ax.boxplot(variances.values)\n    ax.set_title(\"Epoch-Level Variance Across Channels\")\n    ax.set_xlabel(\"Channels\")\n    ax.set_ylabel(\"Variance\")\n    plt.tight_layout()\n    add_figure_to_doc(fig, 'Epoch-Level Variance Across Channels')\n    plt.show()\n\ncheck_epoch_noise(eeg_data)\n\n# 9. Statistical Properties\ndef plot_statistical_properties(eeg_data):\n    stats = {\n        \"Mean\": eeg_data.mean(),\n        \"Std\": eeg_data.std(),\n        \"Skewness\": eeg_data.apply(skew),\n        \"Kurtosis\": eeg_data.apply(kurtosis),\n    }\n    stats_df = pd.DataFrame(stats)\n    stats_df.plot(kind='bar', figsize=(15, 6), subplots=True, layout=(2, 2), title=\"Statistical Properties\")\n    plt.tight_layout()\n    add_figure_to_doc(plt, 'Statistical Properties')\n    plt.show()\n\nplot_statistical_properties(eeg_data)\n\n# Save the Word document\ndoc.save(os.path.join(output_dir, 'EEG_Analysis_Report.docx'))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:17:57.157363Z","iopub.execute_input":"2025-05-01T00:17:57.157617Z","iopub.status.idle":"2025-05-01T00:18:14.937785Z","shell.execute_reply.started":"2025-05-01T00:17:57.157592Z","shell.execute_reply":"2025-05-01T00:18:14.937209Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"信噪比（SNR）计算与特征提取 📊\n在本节中，我们计算每个脑电图（EEG）文件的信噪比（SNR），过滤掉信噪比较低的通道，并使用分类模型来识别在癫痫发作检测中最重要的通道。具体步骤包括：\n\n计算每个EEG文件的信噪比。\n\n筛选出信噪比值较高的通道，以获取更可靠的特征。\n\n对发作和非发作数据进行平衡处理，以便用于模型训练。\n\n基于优化后的信噪比特征训练随机森林分类器。\n\n展示特征重要性，以识别分类中的关键通道。\n\n让我们从EEG数据中提取有意义的见解，并确定在癫痫发作检测中最重要的通道！","metadata":{}},{"cell_type":"code","source":"# Function to compute SNR for one file\ndef compute_snr_for_file(eeg_file, clip_threshold=1e6):\n    eeg_data = pd.read_parquet(eeg_file)\n    eeg_data.drop(columns=[\"EKG\"], inplace=True)\n    \n    # Handle NaNs or Infs\n    if eeg_data.isnull().values.any() or np.isinf(eeg_data.values).any():\n        eeg_data = eeg_data.fillna(0)\n        eeg_data = eeg_data.replace([np.inf, -np.inf], 0)\n    \n    eeg_data = eeg_data.clip(-clip_threshold, clip_threshold)  # Clip extreme values\n\n    # Compute signal and noise power\n    signal_power = (eeg_data ** 2).mean(axis=0)\n    if np.isnan(signal_power).any():\n        print(f\"Warning: NaN values detected in signal power for {eeg_file}.\")\n    \n    noise_power = eeg_data.diff().dropna() ** 2\n    noise_power = noise_power.mean(axis=0)\n    if np.isnan(noise_power).any():\n        print(f\"Warning: NaN values detected in noise power for {eeg_file}.\")\n    \n    # Add small epsilon to avoid log(0)\n    epsilon = 1e-10\n    snr = 10 * np.log10((signal_power + epsilon) / (noise_power + epsilon))\n    \n    return snr\n\n# Randomly sample EEG files for analysis\nsample_files = df.sample(100, random_state=CFG.seed)['eeg_path']\n\n# Compute SNR for each file\nsnr_results = []\nfor file in sample_files:\n    snr = compute_snr_for_file(file)\n    snr_results.append(snr)\n\n# Aggregate SNR results across files\nsnr_df = pd.DataFrame(snr_results).mean(axis=0)\nprint(\"Average SNR for Channels:\")\nprint(snr_df)\n\n# Filter channels with high SNR\nsnr_threshold = 0  # Adjust based on results\nhigh_snr_channels = snr_df[snr_df > snr_threshold].index.tolist()\nprint(\"High-SNR Channels:\", high_snr_channels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:18:14.938534Z","iopub.execute_input":"2025-05-01T00:18:14.938718Z","iopub.status.idle":"2025-05-01T00:18:19.662254Z","shell.execute_reply.started":"2025-05-01T00:18:14.938703Z","shell.execute_reply":"2025-05-01T00:18:19.661495Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Select a balanced subset of seizure and non-seizure data\nseizure_files = df[df['is_seizure']]['eeg_path'].sample(50, random_state=CFG.seed)\nnon_seizure_files = df[~df['is_seizure']]['eeg_path'].sample(50, random_state=CFG.seed)\nfiles_to_process = pd.concat([seizure_files, non_seizure_files])\n\n# Prepare data for classification\nX, y = [], []\nfor file in files_to_process:\n    eeg_data = pd.read_parquet(file)[high_snr_channels]\n    X.append(eeg_data.mean(axis=0))  # Aggregate features (e.g., mean)\n    y.append(df.loc[df['eeg_path'] == file, 'is_seizure'].values[0])\n\n# Convert to DataFrame\nX = pd.DataFrame(X, columns=high_snr_channels)\ny = np.array(y)\n\n# Train-test split\nX_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=CFG.seed)\n\n# Train a Random Forest Classifier\nrf_model = RandomForestClassifier(random_state=CFG.seed)\nrf_model.fit(X_train, y_train)\n\n# Feature importance\nimportances = pd.Series(rf_model.feature_importances_, index=high_snr_channels)\nimportances.sort_values(ascending=False).plot(kind='bar', figsize=(10, 6), title=\"Channel Importance\")\nplt.savefig(\"Feature Importance.png\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:18:19.663113Z","iopub.execute_input":"2025-05-01T00:18:19.663449Z","iopub.status.idle":"2025-05-01T00:18:23.320163Z","shell.execute_reply.started":"2025-05-01T00:18:19.663409Z","shell.execute_reply":"2025-05-01T00:18:23.319349Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"脑电图（EEG）数据预处理与特征工程 🧹\n在本节中，我们对EEG数据进行详细的预处理，以确保为后续分析提供高质量的特征：\n\n通道选择：仅选择排名靠前的通道进行分析。\n\n数据清洗：处理缺失值（NaN）、无穷大值，并对EEG数据进行归一化。\n\n小波去噪：应用小波变换对信号进行去噪处理。\n\n方差与信噪比（SNR）检查：过滤掉方差较低或SNR值为负的数据。\n\n数据降采样：将数据缩减为30秒的样本，以保持数据一致性。\n\n保存处理后的数据：将清洗后的数据转换为NumPy格式并保存，以便后续使用。\n\n处理完成后，我们检查专家共识数据和患者数据的分布情况。","metadata":{}},{"cell_type":"code","source":"# Selected top channels based on analysis\nSELECTED_CHANNELS = ['Fp1', 'O2', 'T6', 'Fz', 'F4', 'T3', 'Cz', 'T5', 'C4', 'P3']\n\n# Define thresholds\nLOW_VARIANCE_THRESHOLD = 0.001\nHIGH_VARIANCE_THRESHOLD = 10\n\ndef wavelet_denoise(data, wavelet=\"db4\", level=2):\n    coeffs = pywt.wavedec(data, wavelet, axis=0)\n    threshold = np.median(np.abs(coeffs[-1])) / 0.6745  # Universal threshold\n    coeffs = [pywt.threshold(c, threshold, mode=\"soft\") for c in coeffs]\n    return pywt.waverec(coeffs, wavelet, axis=0)\n\ndef calculate_snr(signal):\n    mean_signal = np.mean(signal)\n    noise = signal - mean_signal\n    snr = 10 * np.log10(np.var(signal) / np.var(noise))\n    return snr\n\n# Define a function to process a single eeg_id\ndef evaluate_eeg(eeg_id):\n    eeg_path = f\"{BASE_PATH}/train_eegs/{eeg_id}.parquet\"\n    try:\n        # Load EEG data\n        eeg = pd.read_parquet(eeg_path)\n        \n        # Select relevant channels\n        chan_eeg = eeg[SELECTED_CHANNELS]\n        \n        # Replace NaN and infinite values\n        chan_eeg = chan_eeg.replace([np.inf, -np.inf], 0).fillna(0)\n\n        # Apply bandpass filtering\n        chan_eeg = bandpass_filter(chan_eeg)\n\n        # Wavelet denoising\n        chan_eeg = wavelet_denoise(chan_eeg)\n\n        # Normalize data with safeguard for zero standard deviation\n        stds = chan_eeg.std(axis=0)\n        stds[stds == 0] = 1e-8  # Replace zero stds with a small constant\n        normalized_eeg = (chan_eeg - chan_eeg.mean(axis=0)) / stds\n        normalized_eeg = np.nan_to_num(normalized_eeg)  # Ensure no NaN or inf values\n\n        # Variance checks\n        variances = normalized_eeg.var(axis=0)\n        if (variances < LOW_VARIANCE_THRESHOLD).any() or (variances > HIGH_VARIANCE_THRESHOLD).any():\n            # print(f\"Discarded {eeg_id} due to variance thresholds after preprocessing.\")\n            return None\n\n        # SNR checks\n        snr_values = np.apply_along_axis(calculate_snr, axis=0, arr=normalized_eeg)\n        if (snr_values < 0).any():\n            # print(f\"Discarded {eeg_id} due to negative SNR after preprocessing.\")\n            return None\n        \n        # Undersample to 6,000 rows\n        # The EEGs were recored at 200 samples / second\n        # 200 * 30 = 6000 Samples\n        # Therefore we will pick first 6000 samples which will consolidate 30 seconds of recording\n        sub_eeg = normalized_eeg[:6000, :]\n        \n        # Convert to NumPy and save\n        # eeg_np = sub_eeg.to_numpy(dtype='float', na_value=0)\n        return sub_eeg  # Return eeg_id if successfully processed\n    \n    except Exception as e:\n        # print(f\"Error processing {eeg_id}: {e}\")\n        return None\n\ndef process_and_save_eeg(eeg_id):\n    processed = evaluate_eeg(eeg_id)\n    if processed is not None:\n        np.save(f\"{EEG_DIR}/{eeg_id}.npy\", processed)\n        return eeg_id  # Return ID if successfully processed\n    return None\n\nif __name__ == \"__main__\":\n    # Collect all processed EEG IDs\n    results = Parallel(n_jobs=-1, backend=\"loky\")(\n        delayed(process_and_save_eeg)(eeg_id)\n        for eeg_id in tqdm(df[\"eeg_id\"], desc=\"Processing EEG files\")\n    )\n\n    # Filter out None values and save successfully processed IDs\n    accepted_eegs = [eeg_id for eeg_id in results if eeg_id is not None]\n    print(f\"Successfully processed {len(accepted_eegs)} EEG files.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:18:23.320966Z","iopub.execute_input":"2025-05-01T00:18:23.321193Z","iopub.status.idle":"2025-05-01T00:20:14.045481Z","shell.execute_reply.started":"2025-05-01T00:18:23.321161Z","shell.execute_reply":"2025-05-01T00:20:14.044705Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"accepted_df = df[df[\"eeg_id\"].isin(accepted_eegs)]\nclass_counts = accepted_df['expert_consensus'].value_counts()\npatients = accepted_df[\"patient_id\"].value_counts()\ndisplay(class_counts)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:20:14.046542Z","iopub.execute_input":"2025-05-01T00:20:14.046788Z","iopub.status.idle":"2025-05-01T00:20:14.062054Z","shell.execute_reply.started":"2025-05-01T00:20:14.046765Z","shell.execute_reply":"2025-05-01T00:20:14.061511Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"数据划分与增强以用于模型训练 \n在本节中，我们为模型训练执行以下关键步骤：\n\n分层K折划分：将数据划分为5个折，确保每个折中发作（seizure）和非发作（non-seizure）标签的分布均衡。\n\n数据增强：通过添加高斯噪声和应用随机时间偏移来增强EEG数据，从而提高模型的鲁棒性。\n\n数据预处理：执行必要的预处理步骤，例如重塑数据和添加通道维度。\n\nTensorFlow数据集：将数据转换为高效的TensorFlow数据集，从而支持并行处理以加快训练速度。\n\n这些准备好的数据集已准备好用于在EEG数据上训练深度学习模型。","metadata":{}},{"cell_type":"code","source":"# Encode the \"is_seizure\" label (True -> 1, False -> 0)\naccepted_df[\"is_seizure\"] = accepted_df[\"is_seizure\"].astype(int)\n\n# Convert to NumPy arrays\neeg_ids = accepted_df[\"eeg_id\"].to_numpy()\nlabels = accepted_df[\"is_seizure\"].to_numpy()\n\n# Number of folds\nn_splits = 5\n\nskf = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=CFG.seed)\nsplits = list(skf.split(eeg_ids, labels))\n\n# Placeholder to store datasets for each fold\nfold_datasets = []\n\n# Function to load EEG data from .npy files\ndef load_eeg(eeg_id, label):\n    eeg_id_str = tf.strings.as_string(eeg_id)\n    eeg_path = tf.strings.join([EEG_DIR, tf.strings.join([eeg_id_str, \".npy\"], separator=\"\")], separator=\"/\")\n    eeg_data = tf.numpy_function(np.load, [eeg_path], tf.float64)\n    eeg_data = tf.cast(eeg_data, tf.float32)\n    eeg_data = tf.ensure_shape(eeg_data, [6000, 10])\n    return eeg_data, label\n\n# Function to apply data augmentation\ndef augment_data(eeg_data, label):\n    # Add Gaussian noise\n    noise = tf.random.normal(tf.shape(eeg_data), mean=0.0, stddev=0.02)\n    eeg_data = eeg_data + noise\n\n    # Random time shifting\n    time_shift = tf.random.uniform([], minval=-50, maxval=50, dtype=tf.int32)\n    eeg_data = tf.roll(eeg_data, shift=time_shift, axis=0)\n\n    return eeg_data, label\n\n# Function to preprocess the data\ndef preprocess_data(eeg_data, label):\n    eeg_data = tf.expand_dims(eeg_data, axis=-1)  # Add a channel dimension: (6000, 10) -> (6000, 10, 1)\n    return eeg_data, label\n\n# Function to create TensorFlow datasets for a fold\ndef create_fold_dataset(train_idx, test_idx):\n    # Get train and test EEG IDs and labels\n    train_ids, train_labels = eeg_ids[train_idx], labels[train_idx]\n    test_ids, test_labels = eeg_ids[test_idx], labels[test_idx]\n\n    # Create TensorFlow datasets\n    train_dataset = tf.data.Dataset.from_tensor_slices((train_ids, train_labels))\n    test_dataset = tf.data.Dataset.from_tensor_slices((test_ids, test_labels))\n\n    # Map the loading and preprocessing functions\n    train_dataset = train_dataset.map(\n        lambda eeg_id, label: load_eeg(eeg_id, label), num_parallel_calls=tf.data.AUTOTUNE\n    )\n    train_dataset = train_dataset.map(preprocess_data, num_parallel_calls=tf.data.AUTOTUNE)\n    test_dataset = test_dataset.map(\n        lambda eeg_id, label: load_eeg(eeg_id, label), num_parallel_calls=tf.data.AUTOTUNE\n    )\n    test_dataset = test_dataset.map(preprocess_data, num_parallel_calls=tf.data.AUTOTUNE)\n\n    # Shuffle, batch, and prefetch\n    batch_size = CFG.batch_size\n    train_dataset = train_dataset.shuffle(buffer_size=1000).batch(batch_size).prefetch(tf.data.AUTOTUNE)\n    test_dataset = test_dataset.batch(batch_size).prefetch(tf.data.AUTOTUNE)\n\n    return train_dataset, test_dataset\n\n# Loop through each fold and create datasets\nfor i, (train_idx, test_idx) in enumerate(splits):\n    print(f\"Creating datasets for Fold {i + 1}\")\n    train_dataset, test_dataset = create_fold_dataset(train_idx, test_idx)\n    fold_datasets.append((train_dataset, test_dataset))\n\nprint(f\"Successfully created datasets for {n_splits} folds.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:20:14.062819Z","iopub.execute_input":"2025-05-01T00:20:14.063055Z","iopub.status.idle":"2025-05-01T00:20:16.315601Z","shell.execute_reply.started":"2025-05-01T00:20:14.063034Z","shell.execute_reply":"2025-05-01T00:20:16.314691Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"用于EEG分类的CNN-LSTM模型 🧠\n\n该模型专为EEG数据分类而设计，特别是用于区分发作事件和非发作事件，采用了卷积神经网络（CNN）和长短期记忆网络（LSTM）的结合。以下是各层及其功能的详细分解。\n\n模型概述\n\nCNN-LSTM架构有效地结合了空间特征提取（通过CNN层）和时间依赖性捕捉（通过LSTM层）。这种混合方法使模型能够有效地处理EEG数据的时空特性。模型的工作原理如下：\n\n输入层：接收形状为 (6000, 10, 1) 的EEG数据输入，表示6000个时间步长和10个通道。\nCNN块（空间特征提取）：使用卷积层提取空间特征，并使用最大池化层对空间维度进行下采样。\n重塑：将CNN块的输出展平，以便为LSTM层的序列处理做准备。\nLSTM块（时间依赖性捕捉）：使用双向LSTM层捕捉EEG信号中的时间依赖性。\n全连接层（分类）：使用密集层将提取的特征分类为发作或非发作类别。\n\n各层分解与功能\n\n输入层\n该层接收形状为 (6000, 10, 1) 的EEG数据。第一个维度表示6000个时间步长，第二个维度表示10个EEG通道，第三个维度表示单个特征（因为数据在这种情况下是一维的）。\nCNN块（空间特征提取）\n该块包括两个卷积层，后面跟着批归一化层和最大池化层：\nConv2D（卷积-1）：应用32个滤波器，核大小为 (1, 5)，使模型能够从EEG数据中捕捉空间特征。使用'relu'激活函数引入非线性。\nBatchNormalization（归一化-1）：对卷积层的输出进行归一化，通过减少内部协变量偏移来稳定和加速训练。\nMaxPooling2D（池化-1）：在宽度维度上将特征图下采样2倍，降低空间复杂性并聚焦于最显著的特征。\nConv2D（卷积-2）：应用64个滤波器，核大小为 (1, 3)，帮助模型捕捉更精细的空间细节。\nBatchNormalization（归一化-2）：与第一个批归一化步骤类似，该层确保更稳定的学习。\nMaxPooling2D（池化-2）：将空间维度进一步缩小2倍，得到最终输出形状为 (6000, 2, 64)。这对于聚焦网络于关键特征至关重要。\n\n时间下采样（为LSTM重塑）\n将CNN块的输出从 (6000, 2, 64) 重塑为 (6000, 128)，为LSTM层做准备。这种转换确保数据与LSTM的序列学习兼容，这对于捕捉时间依赖性至关重要。\nLSTM块（时间依赖性捕捉）\nLSTM层旨在捕捉EEG信号中的长程时间依赖性：\n双向LSTM（Bi-LSTM-1）：该层有128个单元，并在正向和反向两个方向上处理EEG数据，提高了模型从过去和未来上下文学习的能力。\nDropout（Dropout-1）：使用0.2的dropout率，通过在训练期间随机丢弃神经元来防止过拟合。\n双向LSTM（Bi-LSTM-2）：第二个具有64个单元的Bi-LSTM层进一步增强了模型捕捉数据中序列模式的能力。\n\n全连接层（分类）\n全连接层用于分类：\nDense（Dense-1）：一个具有64个单元和'relu'激活的全连接层，用于学习数据的高级表示。\nDropout（Dropout-2）：另一个dropout层，速率为0.2，以防止过拟合。\nDense（Dense-2）：第二个具有32个单元的密集层，进一步细化学习的特征。\n输出层：最终的密集层具有与类别数相等的单元（在本例中为2，用于发作与非发作），并使用'softmax'激活来生成类别概率。\n\n⚙️ 超参数分解\n\n优化器：Adam\n功能：Adam优化器是一种自适应学习率优化算法，结合了RMSprop和动量的优点。它根据梯度的一阶和二阶矩估计为每个参数计算自适应学习率。\n使用原因：Adam因其效率、可扩展性以及在稀疏梯度下表现良好的能力而被选中，这在EEG分类中特别有用，因为数据可能嘈杂且稀疏。\n损失函数：稀疏分类交叉熵\n功能：该损失函数用于标签以整数形式提供的多类分类问题。它计算真实标签与预测概率之间的交叉熵损失。\n使用原因：稀疏分类交叉熵被使用，因为目标标签（发作与非发作）是整数，使其适合于该分类任务。\n指标：准确率\n准确率是分类任务的标准指标。它计算正确预测的比例与总预测数的比例。","metadata":{}},{"cell_type":"code","source":"def create_cnn_lstm(input_shape=(6000, 10, 1), num_classes=2):\n    model = models.Sequential(name=\"CNN-BiLSTM-Model\")\n\n    # Input Layer\n    model.add(layers.Input(shape=input_shape, name=\"Input-Layer\"))\n    \n    # 1. CNN Block (Spatial Feature Extraction)\n    model.add(layers.Conv2D(name=\"Convolution-1\", filters=32, kernel_size=(1, 5), activation='relu', padding='same'))  # Added padding\n    model.add(layers.BatchNormalization(name=\"Normalization-1\"))\n    model.add(layers.MaxPooling2D(name=\"Pooling-1\", pool_size=(1, 2)))  # Reduces the width by a factor of 2\n\n    model.add(layers.Conv2D(name=\"Convolution-2\", filters=64, kernel_size=(1, 3), activation='relu', padding='same'))  # Added padding\n    model.add(layers.BatchNormalization(name=\"Normalization-2\"))\n    model.add(layers.MaxPooling2D(name=\"Pooling-2\", pool_size=(1, 2)))  # Reduces the width by a factor of 2\n    \n    # After Conv2D and MaxPooling layers, the shape is reduced to:\n    # Height = 6000 (unchanged), Width = 10 -> (10 / 2 / 2) = 2, Depth = 64\n    # So, the shape after this block is (6000, 2, 64)\n    \n    # 2. Temporal Downsampling (Reshape for LSTM)\n    # The output shape after CNN and Pooling is (6000, 2, 64), so we reshape it to (6000, 128)\n    model.add(layers.Reshape(name=\"Reshape-Conv-Output\", target_shape=(6000, 128)))  # 6000 time steps, 128 features\n    \n    # 3. LSTM Block (Temporal Dependency Capture)\n    model.add(layers.Bidirectional(layers.LSTM(128, return_sequences=True), name=\"Bi-LSTM-1\"))\n    model.add(layers.Dropout(0.2, name=\"Dropout-1\"))\n    model.add(layers.Bidirectional(layers.LSTM(64), name=\"Bi-LSTM-2\"))\n\n    # 4. Fully Connected Layers (Classification)\n    model.add(layers.Dense(64, name=\"Dense-1\", activation='relu'))\n    model.add(layers.Dropout(0.2, name=\"Dropout-2\"))\n    model.add(layers.Dense(32, name=\"Dense-2\", activation='relu'))\n    model.add(layers.Dense(num_classes, name=\"Output-Layer\", activation='softmax'))\n\n    # Compile the model\n    model.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy'])\n    return model\n\n# Instantiate the model\nmodel = create_cnn_lstm()\nmodel.summary()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:20:16.316599Z","iopub.execute_input":"2025-05-01T00:20:16.317086Z","iopub.status.idle":"2025-05-01T00:20:17.334716Z","shell.execute_reply.started":"2025-05-01T00:20:16.317066Z","shell.execute_reply":"2025-05-01T00:20:17.333995Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"使用回调函数优化模型训练 🧠\n\n在本代码单元中，我们实现了以下用于优化模型训练的关键回调函数：\n\n ModelCheckpoint（模型检查点）：根据验证损失保存模型，确保在整个训练过程中保留性能最佳的模型。\n\nEarlyStopping（提前停止）：监控验证损失，如果在定义的轮次（epoch）内没有改善，则停止训练，从而防止过拟合并节省资源。\n这些回调函数将有助于优化训练过程并提升模型性能。","metadata":{}},{"cell_type":"code","source":"ckpt_cb = tf.keras.callbacks.ModelCheckpoint(\n    'cnn_lstm.keras',                # File path to save the model\n    monitor=\"val_loss\",\n    save_best_only=True,             # Save regardless of validation loss\n    save_weights_only=False,         # Save the full model (not just weights)\n    save_freq='epoch',               # Save at the end of each epoch (which is the default)\n    verbose=0                        # Display a message when saving\n)\n\n# Early stopping callback\nerst_cb = EarlyStopping(\n    monitor=\"val_loss\",  # Monitor validation loss\n    patience=3,          # Number of epochs to wait before stopping\n    restore_best_weights=True,  # Restore the model with the best weights\n    verbose=CFG.verbose\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:20:17.335527Z","iopub.execute_input":"2025-05-01T00:20:17.335797Z","iopub.status.idle":"2025-05-01T00:20:17.340049Z","shell.execute_reply.started":"2025-05-01T00:20:17.335773Z","shell.execute_reply":"2025-05-01T00:20:17.339361Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"跨折（Cross-Folds）的模型训练与评估 🧠\n\n在本节中，我们执行以下关键步骤来对模型进行多折训练和评估：\n\n数据划分：我们将训练数据集拆分为一个较小的训练集和一个验证集（占 20%），以便在训练过程中监控模型的性能。\n\n模型训练：使用当前折的数据集对模型进行指定轮次（epochs）的训练，并利用回调函数来监控和保存最佳模型。\n\n损失与准确率跟踪：在每一折训练结束后，我们绘制训练和验证的损失与准确率曲线，以可视化模型在各个轮次中的学习进展。\n\n测试：在每一折中，我们在测试数据集上评估模型，以报告其最终准确率，从而帮助我们了解模型在未见数据上的表现。\n这一过程有助于评估模型在多次数据划分下的泛化能力和稳健性。","metadata":{}},{"cell_type":"code","source":"# Hyperparameters\nepochs = 10\n\n# Loop through each fold\nfor fold_idx, (train_dataset, test_dataset) in enumerate(fold_datasets):\n    print(f\"\\nTraining Fold {fold_idx + 1}/{len(fold_datasets)}\")\n    \n    # Split train_dataset into train and validation datasets\n    validation_split = 0.2  # 20% for validation\n    total_train_size = len(train_dataset)\n    val_size = int(total_train_size * validation_split)\n    \n    train_dataset = train_dataset.skip(val_size)\n    val_dataset = train_dataset.take(val_size)\n    \n    # Train the model\n    training = model.fit(\n        train_dataset,\n        validation_data=val_dataset,\n        epochs=epochs,\n        callbacks=[ckpt_cb, erst_cb],\n        verbose=CFG.verbose,\n    )\n\n    # Extract loss and accuracy values from the history object\n    train_loss = training.history['loss']\n    val_loss = training.history['val_loss']\n    train_acc = training.history['accuracy']\n    val_acc = training.history['val_accuracy']\n    \n    # Plotting the training and validation loss\n    plt.figure(figsize=(12, 6))\n    plt.subplot(1, 2, 1)\n    plt.plot(range(1, len(train_loss) + 1), train_loss, label='Training Loss')\n    plt.plot(range(1, len(val_loss) + 1), val_loss, label='Validation Loss')\n    plt.title('Loss over Epochs')\n    plt.xlabel('Epochs')\n    plt.ylabel('Loss')\n    plt.legend()\n    \n    # Plotting the training and validation accuracy\n    plt.subplot(1, 2, 2)\n    plt.plot(range(1, len(train_acc) + 1), train_acc, label='Training Accuracy')\n    plt.plot(range(1, len(val_acc) + 1), val_acc, label='Validation Accuracy')\n    plt.title('Accuracy over Epochs')\n    plt.xlabel('Epochs')\n    plt.ylabel('Accuracy')\n    plt.legend()\n    \n    # Display the plots\n    plt.tight_layout()\n    plt.show()\n    \n    # Evaluate on the testing dataset\n    test_loss, test_accuracy = model.evaluate(test_dataset, verbose=0)\n    print(f\"Fold {fold_idx + 1} - Test Accuracy: {test_accuracy:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T00:20:17.340880Z","iopub.execute_input":"2025-05-01T00:20:17.341160Z","iopub.status.idle":"2025-05-01T01:53:39.810862Z","shell.execute_reply.started":"2025-05-01T00:20:17.341144Z","shell.execute_reply":"2025-05-01T01:53:39.810162Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"模型评估、指标计算与最终模型保存 \n本节专注于评估模型在所有折的测试集上的性能。我们计算重要的分类指标，并使用混淆矩阵对结果进行可视化。\n\n模型评估：在所有折的测试数据集上评估模型的性能，收集真实标签和预测标签以供分析。\n\n混淆矩阵：可视化混淆矩阵，以评估模型在真正例、假正例、真负例和假负例方面的性能。\n\n性能指标：计算关键指标，如精确率（Precision）、灵敏度（Sensitivity，又称召回率）、特异度（Specificity）和平均准确率（Mean Accuracy），以更好地了解模型在癫痫发作检测方面的性能。\n\n模型保存：将训练好的模型保存为“cnn_lstm.keras”，以便将来使用。\n这些评估和指标提供了模型准确分类癫痫发作事件能力的概述。随后，模型被保存以供进一步部署和分析。","metadata":{}},{"cell_type":"code","source":"# Collect true and predicted labels for the test set across all folds\ny_true_all = []\ny_pred_all = []\n\nprint(\"Evaluation Running...\")\n\n# Loop through each fold's test dataset\nfor fold_idx, (train_dataset, test_dataset) in enumerate(fold_datasets):\n    print(f\"\\nEvaluating Fold {fold_idx + 1}/{len(fold_datasets)}\")\n    \n    y_true = []\n    y_pred = []\n    \n    for eeg_data, labels in test_dataset:\n        predictions = model.predict(eeg_data, verbose=0)\n        y_true.extend(labels.numpy())\n        y_pred.extend(tf.argmax(predictions, axis=1).numpy())\n    \n    # Append results for all folds\n    y_true_all.extend(y_true)\n    y_pred_all.extend(y_pred)\n\nprint(\"Evaluation Ends!\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T01:53:39.822205Z","iopub.execute_input":"2025-05-01T01:53:39.822509Z","iopub.status.idle":"2025-05-01T01:55:25.138628Z","shell.execute_reply.started":"2025-05-01T01:53:39.822491Z","shell.execute_reply":"2025-05-01T01:55:25.137813Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cm = confusion_matrix(y_true_all, y_pred_all)\n# Plot confusion matrix heatmap\nplt.figure(figsize=(8, 5))\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues')\nplt.xlabel(\"Predicted\")\nplt.ylabel(\"True\")\nplt.show()\n\n# Print the average accuracy & specificity model acheived\nTN, FP, FN, TP = cm.ravel()\nprecision = TP / (TP + FP)\nsensitivity = TP / (TP + FN)\nspecificity = TN / (TN + FP)\nacc = accuracy_score(y_true_all, y_pred_all)\nprint(f\"\\033[1mModel Precision: {round(precision, 2) * 100}%\\033[0m\")\nprint(f\"\\033[1mModel Sensitivity: {round(sensitivity, 2) * 100}%\\033[0m\")\nprint(f\"\\033[1mModel Specificity: {round(specificity, 2) * 100}%\\033[0m\")\nprint(f\"\\033[1mModel Mean Accuracy: {round(acc, 2) * 100}%\\n\\033[0m\")\n\n# Generate a classification report\nprint(\"\\t\\t\\tCLASSIFICATION REPORT:\")\nprint(classification_report(y_true_all, y_pred_all, target_names=[\"Non-Seizure\", \"Seizure\"]))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T01:55:25.139511Z","iopub.execute_input":"2025-05-01T01:55:25.139787Z","iopub.status.idle":"2025-05-01T01:55:25.335122Z","shell.execute_reply.started":"2025-05-01T01:55:25.139759Z","shell.execute_reply":"2025-05-01T01:55:25.334306Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.save(\"cnn_lstm.keras\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T01:55:25.336036Z","iopub.execute_input":"2025-05-01T01:55:25.336294Z","iopub.status.idle":"2025-05-01T01:55:25.431758Z","shell.execute_reply.started":"2025-05-01T01:55:25.336274Z","shell.execute_reply":"2025-05-01T01:55:25.430904Z"}},"outputs":[],"execution_count":null}]}