{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"}],"dockerImageVersionId":31236,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport pywt\nfrom scipy.signal import stft\nfrom scipy.stats import skew, kurtosis\n\n# =========================\n# Basic utils\n# =========================\n\nFS = 200          # sampling rate\nWIN_SEC = 50\nWIN_SAMPLES = FS * WIN_SEC\n\nSAVE_DIR = \"./figures\"\nos.makedirs(SAVE_DIR, exist_ok=True)\n\ndef plot_eeg(df, title=\"\", max_ch=20):\n    fig, axs = plt.subplots(max_ch, 1, figsize=(25, 15), sharex=True)\n\n    for i, ax in enumerate(axs):\n        ax.plot(df.iloc[:, i], color=\"black\", linewidth=0.6)\n        ax.set_ylabel(df.columns[i], rotation=0, labelpad=25)\n        ax.set_yticks([])\n        ax.set_xticks([])\n        ax.spines[:].set_visible(False)\n\n    plt.suptitle(title)\n    plt.tight_layout()\n    plt.savefig(f\"{SAVE_DIR}/{title.replace(' ', '_')}.png\", dpi=200, bbox_inches=\"tight\")\n    plt.close()\n\n# =========================\n# Wavelet denoising\n# =========================\n\ndef maddest(d):\n    return np.mean(np.abs(d - np.mean(d)))\n\ndef visualize_denoise(df, wavelet=\"db8\", level=1, mode=\"per\"):\n    for ch in df.columns:\n        signal = df[ch].values\n        # 小波分解\n        coeffs = pywt.wavedec(signal, wavelet, mode=mode)\n        sigma = (1 / 0.6745) * maddest(coeffs[-level])\n        uthresh = sigma * np.sqrt(2 * np.log(len(df)))\n        \n        # Threshold 前的 coeffs\n        coeffs_before = coeffs.copy()\n        # Threshold 後\n        coeffs_thresholded = [\n            pywt.threshold(c, uthresh, mode='hard') if i>0 else c\n            for i, c in enumerate(coeffs)\n        ]\n        reconstructed = pywt.waverec(coeffs_thresholded, wavelet, mode=mode)[:len(signal)]\n\n        n_levels = len(coeffs)\n        fig, axs = plt.subplots(n_levels + 1, 1, figsize=(20, 2*(n_levels+1)))\n        fig.suptitle(f\"Channel: {ch}\", fontsize=16)\n\n        for i in range(n_levels):\n            axs[i].plot(coeffs_before[i], color='gray', alpha=0.5, label='original coeff')\n            if i > 0:\n                axs[i].plot(coeffs_thresholded[i], color='red', alpha=0.8, label='thresholded coeff')\n            axs[i].set_ylabel(f\"c{i}\")\n            axs[i].legend(loc='upper right')\n\n        axs[-1].plot(signal, color='gray', alpha=0.5, label='original signal')\n        axs[-1].plot(reconstructed, color='blue', alpha=0.8, label='denoised signal')\n        axs[-1].set_ylabel(\"signal\")\n        axs[-1].legend(loc='upper right')\n        plt.tight_layout()\n        plt.savefig(f\"{SAVE_DIR}/denoise_{ch}.png\", dpi=200, bbox_inches=\"tight\")\n        plt.show()\n\ndef denoise(df, wavelet=\"db8\", level=1, mode=\"per\"):\n    out = {}\n\n    for ch in df.columns:\n        print(f'Processing channel {ch}')\n        coeffs = pywt.wavedec(df[ch], wavelet, mode=mode)\n        sigma = (1 / 0.6745) * maddest(coeffs[-level])\n        print(\"len\", len(df))\n        uthresh = sigma * np.sqrt(2 * np.log(len(df)))\n\n        coeffs[1:] = [\n            pywt.threshold(c, uthresh, mode=\"hard\")\n            for c in coeffs[1:]\n        ]\n\n        out[ch] = pywt.waverec(coeffs, wavelet, mode=mode)[:len(df)]\n\n    return pd.DataFrame(out)\n\n\n# =========================\n# Feature functions\n# =========================\n\ndef energy(x):\n    return np.sum(x ** 2)\n\n\ndef shannon_entropy(p):\n    p = p / (np.sum(p) + 1e-12)\n    return -np.sum(p * np.log(p + 1e-12))\n\n\ndef basic_stats(x):\n    return {\n        \"mean\": np.mean(x),\n        \"std\": np.std(x),\n        \"skew\": skew(x),\n        \"kurtosis\": kurtosis(x),\n        \"rms\": np.sqrt(np.mean(x ** 2))\n    }\n\n\n# =========================\n# WPD features\n# =========================\n\ndef wpd_features(x, wavelet=\"db4\", level=4, subband_ratio=(0.5,1.0)):\n    wp = pywt.WaveletPacket(x, wavelet, mode=\"per\", maxlevel=level)\n    nodes = wp.get_level(level, order=\"freq\") #order according to frequency\n    energies = np.array([energy(n.data) for n in nodes])\n    rel_energy = energies / (energies.sum() + 1e-12)\n    features = {\n        \"WPD_total_energy\": energies.sum(),\n        \"WPD_shannon_entropy\": shannon_entropy(rel_energy),\n        \"WPD_log_energy_entropy\": np.sum(np.log(energies + 1e-12)),\n    }\n    # RSWE \n    for i, re in enumerate(rel_energy):\n        features[f\"WPD_p_{i}\"] = re \n    # -----------------------------\n    # RSWE entropy (only for middel/high frequency sub-band)\n    # -----------------------------\n    n_nodes = len(nodes)\n    start_idx = int(n_nodes * subband_ratio[0])\n    end_idx = int(n_nodes * subband_ratio[1])\n    selected_rel_energy = rel_energy[start_idx:end_idx]\n    selected_rel_energy /= selected_rel_energy.sum() + 1e-12  # normalize\n    RSWE_entropy = -np.sum(selected_rel_energy * np.log(selected_rel_energy + 1e-12))\n    features[\"WPD_RSWE_entropy_mid_high\"] = RSWE_entropy\n    \n    return features, energies\n\n\n# =========================\n# DWT + STFT features\n# =========================\n\ndef plot_dwt_spectrum(x, wavelet=\"db4\", level=5, FS=200):\n    coeffs = pywt.wavedec(x, wavelet, level=level, mode=\"per\")\n    \n    plt.figure(figsize=(14, 2.5 * (level + 1)))\n\n    # Approximation\n    A = coeffs[0]\n    N = len(A)\n    fft_A = np.abs(np.fft.rfft(A))**2\n    f = np.linspace(0, FS / (2**(level+1)), len(fft_A))\n\n    plt.subplot(level + 1, 1, 1)\n    plt.plot(f, fft_A)\n    plt.title(f\"A{level}: 0–{FS/(2**(level+1)):.2f} Hz\")\n    plt.ylabel(\"Power\")\n\n    # Details\n    for i, D in enumerate(coeffs[1:], start=1):\n        band_low = FS / (2**(level - i + 2))\n        band_high = FS / (2**(level - i + 1))\n\n        fft_D = np.abs(np.fft.rfft(D))**2\n        f = np.linspace(band_low, band_high, len(fft_D))\n\n        plt.subplot(level + 1, 1, i + 1)\n        plt.plot(f, fft_D)\n        plt.title(f\"D{level - i + 1}: {band_low:.2f}–{band_high:.2f} Hz\")\n        plt.ylabel(\"Power\")\n\n    plt.xlabel(\"Frequency [Hz]\")\n    plt.tight_layout()\n    plt.savefig(f\"{SAVE_DIR}/dwt_fft.png\", dpi=200, bbox_inches=\"tight\")\n    plt.show()\n\n# =========================\n# Main (single case)\n# =========================\n\n#if __name__ == \"__main__\":\n\n# ---- Load label row\ndf_label = pd.read_csv(\n    \"/kaggle/input/hms-harmful-brain-activity-classification/train.csv\"\n)\nrow = df_label.sample(n=1, random_state=60).iloc[0]\nprint(\"Selected row:\\n\", row, \"\\n\")\n\n# ---- Load EEG\neeg = pd.read_parquet(\n    f\"/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/{row['eeg_id']}.parquet\"\n)\nsp = pd.read_parquet(\n    f\"/kaggle/input/hms-harmful-brain-activity-classification/train_spectrograms/{row['spectrogram_id']}.parquet\"   \n)\n\nstart = int(row[\"eeg_label_offset_seconds\"]) * FS\neeg = eeg.iloc[start:start + WIN_SAMPLES]\n\n\nprint(\"EEG shape:\", eeg.shape)\nprint(\"Spectrogram shape:\", sp.shape)\n# ---- Plot raw\nplot_eeg(eeg, title=\"Raw EEG\")\n\n# ---- Denoise\nvisualize_denoise(eeg) # For visualization\neeg_denoised = denoise(eeg, wavelet=\"db8\")\nplot_eeg(eeg_denoised, title=\"Denoised EEG\")\n\neeg_denoised = eeg\n# ---- Single channel analysis\n#eeg_denoised = eeg\nch_name = eeg_denoised.columns[5]\nsignal = eeg_denoised[ch_name].values\n\nprint(f\"\\nUsing channel: {ch_name}\")\n\n# ---- WPD\nwpd_feat, wpd_energy = wpd_features(signal)\nprint(\"\\nWPD features:\")\nfor k, v in wpd_feat.items():\n    print(k, \":\", v)\n\nplt.figure(figsize=(8, 3))\nplt.bar(range(len(wpd_energy)), wpd_energy)\nplt.title(\"WPD Energy Distribution\")\nplt.xlabel(\"Subband\")\nplt.ylabel(\"Energy\")\nplt.show()\n\n# ---- DWT + FFT\nplot_dwt_spectrum(signal)\nprint(\"\\nPipeline finished ✔\")\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-25T13:26:42.943008Z","iopub.execute_input":"2025-12-25T13:26:42.943204Z","iopub.status.idle":"2025-12-25T13:28:04.560723Z","shell.execute_reply.started":"2025-12-25T13:26:42.943182Z","shell.execute_reply":"2025-12-25T13:28:04.559969Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport pywt\nfrom scipy.stats import skew, kurtosis\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\n# =========================\n# Config\n# =========================\nFS = 200\nWIN_SEC = 50\nWIN_SAMPLES = FS * WIN_SEC\nEEG_DIR = \"/kaggle/input/hms-harmful-brain-activity-classification/train_eegs\"\n\n# =========================\n# Basic utils\n# =========================\ndef energy(x):\n    return np.sum(x ** 2)\n\ndef shannon_entropy(p):\n    p = p / (np.sum(p) + 1e-12)\n    return -np.sum(p * np.log(p + 1e-12))\n\ndef basic_stats(x):\n    return {\n        \"mean\": np.mean(x),\n        \"std\": np.std(x),\n        \"skew\": skew(x),\n        \"kurtosis\": kurtosis(x),\n        \"rms\": np.sqrt(np.mean(x ** 2))\n    }\n\n# =========================\n# Wavelet + WPD features\n# =========================\ndef wpd_features(x, wavelet=\"db4\", level=3, subband_ratio=(0.5,1.0)):\n    wp = pywt.WaveletPacket(x, wavelet, mode=\"per\", maxlevel=level)\n    nodes = wp.get_level(level, order=\"freq\")\n    energies = np.array([energy(n.data) for n in nodes])\n    rel_energy = energies / (energies.sum() + 1e-12)\n\n    features = {\n        \"WPD_total_energy\": energies.sum(),\n        \"WPD_shannon_entropy\": shannon_entropy(rel_energy),\n        \"WPD_log_energy_entropy\": np.sum(np.log(energies + 1e-12)),\n    }\n\n    # RSWE per subband\n    for i, re in enumerate(rel_energy):\n        features[f\"WPD_p_{i}\"] = re\n    for i, n in enumerate(nodes):\n        features[f\"WPD_sig_{i}\"] = n.data\n    \n    # RSWE entropy (中高頻 subband)\n    n_nodes = len(nodes)\n    start_idx = int(n_nodes * subband_ratio[0])\n    end_idx = int(n_nodes * subband_ratio[1])\n    selected_rel_energy = rel_energy[start_idx:end_idx]\n    selected_rel_energy /= selected_rel_energy.sum() + 1e-12\n    features[\"WPD_RSWE_entropy_mid_high\"] = -np.sum(selected_rel_energy * np.log(selected_rel_energy + 1e-12))\n    \n    return features\n\ndef dwt_basic_features(x, wavelet=\"db4\", level=5):\n    coeffs = pywt.wavedec(x, wavelet, level=level, mode=\"per\")\n    features = {}\n    for i, coef in enumerate(coeffs):\n        e = energy(coef)\n        stats = basic_stats(coef)\n        prefix = \"A\" if i==0 else f\"D{i}\"\n        features[f\"{prefix}_coef\"] = coef\n        features[f\"{prefix}_energy\"] = e\n        for k,v in stats.items():\n            features[f\"{prefix}_{k}\"] = v\n    return features\n\ndef preprocess_eeg(eeg):\n    eeg = eeg.copy()\n    eeg = eeg.interpolate(limit_direction=\"both\")\n    eeg = eeg.fillna(0)\n    return eeg\n\ndef denoise(df, wavelet=\"db8\", level=1, mode=\"per\"):\n    out = {}\n\n    for ch in df.columns:\n        coeffs = pywt.wavedec(df[ch], wavelet, mode=mode)\n        sigma = (1 / 0.6745) * maddest(coeffs[-level])\n        uthresh = sigma * np.sqrt(2 * np.log(len(df)))\n\n        coeffs[1:] = [\n            pywt.threshold(c, uthresh, mode=\"hard\")\n            for c in coeffs[1:]\n        ]\n\n        out[ch] = pywt.waverec(coeffs, wavelet, mode=mode)[:len(df)]\n\n    return pd.DataFrame(out)\n\ndef get_scalar_feature_columns(X_df):\n    scalar_cols = []\n    for col in X_df.columns:\n        sample_val = X_df[col].iloc[0]\n        if np.isscalar(sample_val):\n            scalar_cols.append(col)\n    return scalar_cols\n\ndef plot_feature_scatter(X, y, f1, f2, figsize=(6,5)):\n    df_plot = pd.DataFrame({\n        f1: X[f1],\n        f2: X[f2],\n        \"label\": y.values\n    })\n\n    plt.figure(figsize=figsize)\n    sns.scatterplot(\n        data=df_plot,\n        x=f1,\n        y=f2,\n        hue=\"label\",\n        palette=\"tab10\",\n        alpha=0.7\n    )\n    plt.title(f\"{f1} vs {f2}\")\n    plt.tight_layout()\n    plt.show()\n\ndef plot_feature_vs_class(X, y, feature, figsize=(6,4)):\n    df_plot = pd.DataFrame({\n        feature: X[feature],\n        \"label\": y.values\n    })\n\n    plt.figure(figsize=figsize)\n    sns.violinplot(\n        data=df_plot,\n        x=\"label\",\n        y=feature,\n        inner=\"quartile\"\n    )\n    plt.title(f\"{feature} vs Class\")\n    plt.xticks(rotation=30)\n    plt.tight_layout()\n    plt.show()\n\n# =========================\n# Load dataset & extract features\n# =========================\ndef prepare_dataset(df_label, eeg_dir=EEG_DIR, channels=None):\n    \"\"\"\n    channels: list of channels to extract features, None = all\n    \"\"\"\n    X = []\n    y = []\n    \n    for _, row in tqdm(df_label.iterrows(), total=len(df_label)):\n        eeg_path = os.path.join(eeg_dir, f\"{row['eeg_id']}.parquet\")\n        if not os.path.exists(eeg_path):\n            continue\n        eeg = pd.read_parquet(eeg_path)\n        start = int(row['eeg_label_offset_seconds'])*FS\n        eeg = eeg.iloc[start:start+WIN_SAMPLES]\n        eeg = preprocess_eeg(eeg)\n        # denoise\n        eeg = denoise(eeg, wavelet=\"db8\")\n        \n        if channels is None:\n            channels_use = eeg.columns.tolist()\n        else:\n            channels_use = [ch for ch in channels if ch in eeg.columns]\n        \n        feat = {}\n        \n        for ch in channels_use:\n            sig = eeg[ch].values\n            #feat.update({f\"{ch}_sig\":sig})\n            # WPD features\n            feat.update({f\"{ch}_{k}\":v for k,v in wpd_features(sig).items()})\n            # DWT basic features\n            feat.update({f\"{ch}_{k}\":v for k,v in dwt_basic_features(sig).items()})\n        X.append(feat)\n        y.append(row['expert_consensus'])\n        \n    X_df = pd.DataFrame(X)\n    y_series = pd.Series(y, name='label')\n    return X_df, y_series\n\n# =========================\n# Train/validation split\n# =========================\ndf_label = pd.read_csv(\"/kaggle/input/hms-harmful-brain-activity-classification/train.csv\")\n\n# Only use 500 samples for fast run\ndf_label = df_label.sample(n=500, random_state=42)\n\nchannels = ['F3', 'F4', 'C3', 'C4', 'Cz', 'Pz']  # or None -> Choose all electrodes\n\nX, y = prepare_dataset(df_label, channels=channels)\n\n# Stratified split\nX_train, X_val, y_train, y_val = train_test_split(\n    X, y, test_size=0.2, random_state=42, stratify=y\n)\n\nprint(\"Train size:\", X_train.shape)\nprint(\"Validation size:\", X_val.shape)\nprint(\"Class distribution (train):\\n\", y_train.value_counts())\nprint(\"Class distribution (val):\\n\", y_val.value_counts())\nprint(X_train.iloc[0])\nprint(y_train.iloc[0])\n\nscalar_cols = get_scalar_feature_columns(X_train)\nprint(\"Number of scalar features:\", len(scalar_cols))\nprint(scalar_cols[:10])\n\nplot_feature_vs_class(X_train, y_train, \"Cz_WPD_shannon_entropy\")\nplot_feature_vs_class(X_train, y_train, \"Cz_WPD_RSWE_entropy_mid_high\")\nplot_feature_vs_class(X_train, y_train, \"Cz_D3_energy\")\nplot_feature_scatter(\n    X_train, y_train,\n    \"Cz_WPD_RSWE_entropy_mid_high\",\n    \"Cz_WPD_shannon_entropy\"\n)\nplot_feature_scatter(\n    X_train, y_train,\n    \"Cz_D3_energy\",\n    \"Cz_D4_energy\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-25T13:28:04.562438Z","iopub.execute_input":"2025-12-25T13:28:04.562741Z","iopub.status.idle":"2025-12-25T13:29:01.008872Z","shell.execute_reply.started":"2025-12-25T13:28:04.562717Z","shell.execute_reply":"2025-12-25T13:29:01.008220Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.preprocessing import LabelEncoder\n\nle = LabelEncoder()\nle.fit(y_train)\n\ny_train_enc = le.transform(y_train)\ny_val_enc   = le.transform(y_val)\n\nprint(le.classes_)\n\ndef build_wpd_tensor(X_df, channels, level=3):\n    \"\"\"\n    Returns:\n        X_tensor: shape (N, C, H, W)\n    \"\"\"\n    n_samples = len(X_df)\n    n_bands = 2 ** level\n    n_channels = len(channels)\n\n    # 用第一筆 sample 決定時間長度\n    first_key = f\"{channels[0]}_WPD_sig_0\"\n    W = len(X_df.iloc[0][first_key])\n\n    X_tensor = np.zeros((n_samples, n_bands, n_channels, W), dtype=np.float32)\n\n    for i in tqdm(range(n_samples)):\n        for h, ch in enumerate(channels):\n            for c in range(n_bands):\n                key = f\"{ch}_WPD_sig_{c}\"\n                X_tensor[i, c, h, :] = X_df.iloc[i][key]\n\n    return X_tensor\n\nX_train = build_wpd_tensor(X_train, channels=channels, level=3)\nX_val   = build_wpd_tensor(X_val,   channels=channels, level=3)\n\nprint(X_train.shape)\nn_classes = y_train.nunique()\nprint(n_classes)\n\ndef one_hot(y, num_classes):\n    y = y.values if isinstance(y, pd.Series) else y\n    return np.eye(num_classes)[y]\n\nprint(y_train_enc)\ny_train_oh = one_hot(y_train_enc, n_classes)\ny_val_oh   = one_hot(y_val_enc,   n_classes)\n\nprint(y_train_oh.shape)\n\nmean = X_train.mean(axis=(0,2,3), keepdims=True)\nstd  = X_train.std(axis=(0,2,3), keepdims=True) + 1e-6\n\nX_train = (X_train - mean) / std\nX_val   = (X_val   - mean) / std\n\nprint(X_train[0])\nprint(y_train[0])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-25T13:29:01.010125Z","iopub.execute_input":"2025-12-25T13:29:01.010420Z","iopub.status.idle":"2025-12-25T13:29:05.624892Z","shell.execute_reply.started":"2025-12-25T13:29:01.010397Z","shell.execute_reply":"2025-12-25T13:29:05.624120Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"efficient=False\nif(efficient):\n    #Efficient RAM friendly dataset preparation\n    import os\n    import numpy as np\n    import pandas as pd\n    import pywt\n    from tqdm import tqdm\n    from sklearn.model_selection import train_test_split\n    from sklearn.preprocessing import LabelEncoder\n    FS = 200\n    WIN_SEC = 50\n    WIN_SAMPLES = FS * WIN_SEC\n    EEG_DIR = \"/kaggle/input/hms-harmful-brain-activity-classification/train_eegs\"\n    \n    channels = ['F3', 'F4', 'C3', 'C4', 'Cz', 'Pz']\n    wavelet = \"db4\"\n    level = 4\n    n_bands = 2 ** level\n    \n    def preprocess_eeg(eeg):\n        eeg = eeg.interpolate(limit_direction=\"both\")\n        eeg = eeg.fillna(0)\n        return eeg\n    \n    def build_wpd_tensor_from_labels(\n        df_label,\n        channels,\n        eeg_dir=EEG_DIR,\n        level=3\n    ):\n        \"\"\"\n        Returns:\n            X: (N, C_freq, H_ch, W_time)\n            y: (N,)\n        \"\"\"\n        n_samples = len(df_label)\n        n_channels = len(channels)\n        n_bands = 2 ** level\n    \n        # 先讀一筆決定 W\n        for _, row in df_label.iterrows():\n            eeg_path = os.path.join(eeg_dir, f\"{row['eeg_id']}.parquet\")\n            if os.path.exists(eeg_path):\n                eeg = pd.read_parquet(eeg_path)\n                start = int(row['eeg_label_offset_seconds']) * FS\n                eeg = eeg.iloc[start:start + WIN_SAMPLES]\n                eeg = preprocess_eeg(eeg)\n                wp = pywt.WaveletPacket(\n                    eeg[channels[0]].values,\n                    wavelet,\n                    mode=\"per\",\n                    maxlevel=level\n                )\n                W = len(wp.get_level(level, order=\"freq\")[0].data)\n                break\n    \n        X = np.zeros(\n            (n_samples, n_bands, n_channels, W),\n            dtype=np.float32\n        )\n        y = np.zeros(n_samples, dtype=np.int64)\n    \n        valid_idx = 0\n    \n        for _, row in tqdm(df_label.iterrows(), total=n_samples):\n            eeg_path = os.path.join(eeg_dir, f\"{row['eeg_id']}.parquet\")\n            if not os.path.exists(eeg_path):\n                continue\n    \n            eeg = pd.read_parquet(eeg_path)\n            start = int(row['eeg_label_offset_seconds']) * FS\n            eeg = eeg.iloc[start:start + WIN_SAMPLES]\n            eeg = preprocess_eeg(eeg)\n    \n            for h, ch in enumerate(channels):\n                sig = eeg[ch].values\n                wp = pywt.WaveletPacket(\n                    sig, wavelet, mode=\"per\", maxlevel=level\n                )\n                nodes = wp.get_level(level, order=\"freq\")\n                for c, node in enumerate(nodes):\n                    X[valid_idx, c, h, :] = node.data\n    \n            y[valid_idx] = row[\"expert_consensus\"]\n            valid_idx += 1\n    \n        return X[:valid_idx], y[:valid_idx]\n    \n    df_label = pd.read_csv(\n        \"/kaggle/input/hms-harmful-brain-activity-classification/train.csv\"\n    )\n    \n    df_label = df_label.sample(n=10000, random_state=42)\n    \n    train_df, val_df = train_test_split(\n        df_label,\n        test_size=0.2,\n        random_state=42,\n        stratify=df_label[\"expert_consensus\"]\n    )\n    X_train, y_train = build_wpd_tensor_from_labels(\n        train_df, channels=channels, level=level\n    )\n    \n    X_val, y_val = build_wpd_tensor_from_labels(\n        val_df, channels=channels, level=level\n    )\n    le = LabelEncoder()\n    y_train_enc = le.fit_transform(y_train)\n    y_val_enc   = le.transform(y_val)\n    \n    n_classes = len(le.classes_)\n    print(le.classes_)\n    mean = X_train.mean(axis=(0,2,3), keepdims=True)\n    std  = X_train.std(axis=(0,2,3), keepdims=True) + 1e-6\n    \n    X_train = (X_train - mean) / std\n    X_val   = (X_val   - mean) / std\n    print(X_train.shape)  # (N, n_bands, n_channels, W)\n    print(X_val.shape)\n    print(y_train_enc.shape)\n    print(y_val_enc.shape)\n    y_train_idx = y_train_enc\n    y_val_idx = y_val_enc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-25T13:29:05.625870Z","iopub.execute_input":"2025-12-25T13:29:05.626109Z","iopub.status.idle":"2025-12-25T13:29:05.639016Z","shell.execute_reply.started":"2025-12-25T13:29:05.626089Z","shell.execute_reply":"2025-12-25T13:29:05.638320Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset\n\nclass EEGTensorDataset(Dataset):\n    def __init__(self, X, y):\n        # numpy -> torch\n        self.X = torch.from_numpy(X).float()\n        self.y = torch.from_numpy(y)\n\n        # 如果 y 是 one-hot，就轉成 float\n        if self.y.ndim == 2:\n            self.y = self.y.float()\n        else:\n            self.y = self.y.long()\n\n    def __len__(self):\n        return self.X.shape[0]\n\n    def __getitem__(self, idx):\n        return self.X[idx], self.y[idx]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-25T13:29:05.639879Z","iopub.execute_input":"2025-12-25T13:29:05.640121Z","iopub.status.idle":"2025-12-25T13:29:09.423785Z","shell.execute_reply.started":"2025-12-25T13:29:05.640101Z","shell.execute_reply":"2025-12-25T13:29:09.423024Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import DataLoader\n\nbatch_size = 16        # EEG 很吃記憶體，16 是穩的\nnum_workers = 2        # Kaggle 通常 2~4 最穩\n\nclasses = sorted(y_train.unique())\nprint(classes)\nlabel2idx = {c: i for i, c in enumerate(classes)}\nidx2label = {i: c for c, i in label2idx.items()}\ny_train_idx = y_train.map(label2idx).values\ny_val_idx   = y_val.map(label2idx).values\nprint(y_train.iloc[0], \"->\", y_train_idx[0])\n\nprint(y_train_idx)\ntrain_ds = EEGTensorDataset(X_train, y_train_idx)\nval_ds   = EEGTensorDataset(X_val,   y_val_idx)\n\ntrain_loader = DataLoader(\n    train_ds,\n    batch_size=batch_size,\n    shuffle=True,\n    num_workers=num_workers,\n    pin_memory=True\n)\n\nval_loader = DataLoader(\n    val_ds,\n    batch_size=batch_size,\n    shuffle=False,\n    num_workers=num_workers,\n    pin_memory=True\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-25T13:29:09.425406Z","iopub.execute_input":"2025-12-25T13:29:09.425843Z","iopub.status.idle":"2025-12-25T13:29:09.443907Z","shell.execute_reply.started":"2025-12-25T13:29:09.425819Z","shell.execute_reply":"2025-12-25T13:29:09.443206Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Model architecture\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom sklearn.metrics import accuracy_score\n\nclass EEGClassifier(nn.Module):\n    def __init__(self, n_freq, n_ch, n_classes, W_time):\n        super().__init__()\n\n        # ===== Temporal convolution =====\n        self.temporal = nn.Sequential(\n            nn.Conv2d(\n                in_channels=n_freq,\n                out_channels=32,\n                kernel_size=(1, 25),\n                padding=(0, 12)\n            ),\n            nn.BatchNorm2d(32),\n            nn.ELU(),\n        )\n\n        # ===== Spatial (channel) convolution =====\n        self.spatial = nn.Sequential(\n            nn.Conv2d(\n                in_channels=32,\n                out_channels=64,\n                kernel_size=(n_ch, 1),\n                groups=1\n            ),\n            nn.BatchNorm2d(64),\n            nn.ELU(),\n        )\n\n        # ===== Pooling + regularization =====\n        self.pool = nn.Sequential(\n            nn.AvgPool2d(kernel_size=(1, 4)),\n            nn.Dropout(0.5)\n        )\n\n        # ===== Classifier =====\n        self.classifier = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(64 * (W_time // 4), 128),\n            nn.ELU(),\n            nn.Dropout(0.5),\n            nn.Linear(128, n_classes)\n        )\n\n    def forward(self, x):\n        # x: (B, C_freq, H_ch, W_time)\n        x = self.temporal(x)\n        x = self.spatial(x)\n        x = self.pool(x)\n        x = self.classifier(x)\n        return x\n\ndef save_checkpoint(model, optimizer, scheduler, epoch, path):\n    state = {\n        \"epoch\": epoch,\n        \"model_state\": model.state_dict(),\n        \"optimizer_state\": optimizer.state_dict(),\n        \"scheduler_state\": scheduler.state_dict()\n    }\n    torch.save(state, path)\n\ndef load_checkpoint(path, model, optimizer=None, scheduler=None):\n    state = torch.load(path, map_location=device)\n    model.load_state_dict(state[\"model_state\"])\n    if optimizer:\n        optimizer.load_state_dict(state[\"optimizer_state\"])\n    if scheduler:\n        scheduler.load_state_dict(state[\"scheduler_state\"])\n    start_epoch = state[\"epoch\"] + 1\n    return model, optimizer, scheduler, start_epoch\n\ndef run_epoch(model, loader, train=True):\n    model.train() if train else model.eval()\n\n    total_loss = 0\n    all_preds, all_targets = [], []\n\n    for X, y in loader:\n        X, y = X.to(device), y.to(device)\n\n        if train:\n            optimizer.zero_grad()\n\n        logits = model(X)\n        loss = criterion(logits, y)\n\n        if train:\n            loss.backward()\n            optimizer.step()\n\n        total_loss += loss.item() * X.size(0)\n\n        preds = logits.argmax(dim=1)\n        all_preds.append(preds.cpu())\n        all_targets.append(y.cpu())\n\n    all_preds = torch.cat(all_preds)\n    all_targets = torch.cat(all_targets)\n\n    acc = accuracy_score(all_targets, all_preds)\n    avg_loss = total_loss / len(loader.dataset)\n\n    return avg_loss, acc\n\n# =========================\n# Checkpoint config\n# =========================\nCHECKPOINT_DIR = \"./checkpoints\"\nos.makedirs(CHECKPOINT_DIR, exist_ok=True)\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = EEGClassifier(\n    n_freq=X_train.shape[1],\n    n_ch=X_train.shape[2],\n    n_classes=n_classes,\n    W_time=X_train.shape[3]\n).to(device)\n\ncriterion = nn.CrossEntropyLoss()\n\noptimizer = optim.AdamW(\n    model.parameters(),\n    lr=1e-3,\n    weight_decay=1e-4\n)\n\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer,\n    mode=\"max\",\n    factor=0.5,\n    patience=3\n)\n\n# =========================\n# Training loop with checkpoint\n# =========================\nbest_val_acc = 0.0\nepochs = 100\n\nfor epoch in range(1, epochs + 1):\n    train_loss, train_acc = run_epoch(model, train_loader, train=True)\n    val_loss,   val_acc   = run_epoch(model, val_loader,   train=False)\n\n    # Scheduler step\n    scheduler.step(val_acc)\n\n    print(\n        f\"[{epoch:02d}] \"\n        f\"Train loss: {train_loss:.4f}, acc: {train_acc:.4f} | \"\n        f\"Val loss: {val_loss:.4f}, acc: {val_acc:.4f}\"\n    )\n\n    # save checkpoint every epoch\n    checkpoint_path = os.path.join(CHECKPOINT_DIR, f\"epoch_{epoch:02d}.pt\")\n    save_checkpoint(model, optimizer, scheduler, epoch, checkpoint_path)\n\n    # save best model separately\n    if val_acc > best_val_acc:\n        best_val_acc = val_acc\n        best_model_path = os.path.join(CHECKPOINT_DIR, \"best_model.pt\")\n        save_checkpoint(model, optimizer, scheduler, epoch, best_model_path)\n        print(f\"  -> New best model saved at epoch {epoch} with val_acc {val_acc:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-25T13:29:09.444819Z","iopub.execute_input":"2025-12-25T13:29:09.445089Z","iopub.status.idle":"2025-12-25T13:29:51.912283Z","shell.execute_reply.started":"2025-12-25T13:29:09.445060Z","shell.execute_reply":"2025-12-25T13:29:51.911450Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# Data extracted from logs\nepochs = list(range(1, 101))\n\ntrain_loss = [\n1.9228,1.7640,1.7403,1.7290,1.6876,1.6700,1.6637,1.6386,1.6086,1.5706,\n1.5355,1.4916,1.4606,1.4262,1.3211,1.2649,1.2177,1.1821,1.1436,1.1101,\n1.0682,1.0166,0.9610,0.9593,0.9368,0.9237,0.9036,0.8793,0.8617,0.8591,\n0.8386,0.8318,0.7829,0.7776,0.7711,0.7549,0.7468,0.7296,0.7237,0.7133,\n0.7211,0.7296,0.7069,0.7141,0.6872,0.7094,0.7073,0.7006,0.6920,0.6974,\n0.7083,0.7047,0.6930,0.6951,0.6899,0.6855,0.7004,0.6944,0.6878,0.7056,\n0.6955,0.6904,0.6989,0.6958,0.6975,0.6926,0.6927,0.6897,0.6852,0.6919,\n0.6914,0.6921,0.6917,0.6863,0.6942,0.7001,0.6899,0.6977,0.6846,0.6816,\n0.6918,0.6849,0.6897,0.6869,0.6825,0.6954,0.6965,0.6918,0.6721,0.6874,\n0.7148,0.6908,0.6846,0.6980,0.6991,0.6880,0.6904,0.6928,0.6905,0.6884\n]\n\nval_loss = [\n1.7915,1.7934,1.8117,1.8031,1.8310,1.9200,1.7590,1.8456,1.7988,1.7966,\n1.7892,1.8563,1.8348,1.8722,1.8153,1.8377,1.7980,2.0231,1.8836,1.9677,\n1.8117,1.9504,1.9201,1.9386,1.9163,1.9971,1.8941,1.9224,1.9432,2.0163,\n1.8921,2.0053,1.9272,1.9614,1.8844,2.0266,1.9555,1.9332,1.9198,1.9674,\n1.9896,1.9988,1.9775,1.9828,1.9972,1.9414,1.8942,2.0668,1.9500,1.9436,\n1.9108,2.0203,2.0462,1.9414,1.9585,2.0709,2.0175,1.9717,1.9623,2.0526,\n2.0447,2.0873,1.9473,2.0305,1.9631,1.9650,2.0515,1.9071,2.0532,1.9957,\n1.9288,2.0210,1.9319,2.0840,1.9919,1.9420,1.9145,1.8974,1.9520,1.9180,\n2.0324,1.9604,1.9219,1.9701,1.9298,2.1080,2.0295,2.1302,2.0586,2.0240,\n1.9118,1.8865,2.0362,1.9341,1.9887,2.1556,2.0294,1.9444,1.9619,1.9693\n]\n\ntrain_acc = [\n0.2080,0.2401,0.2575,0.2715,0.2924,0.2965,0.3098,0.3207,0.3376,0.3543,\n0.3787,0.4002,0.4220,0.4365,0.4784,0.5127,0.5255,0.5463,0.5617,0.5761,\n0.5931,0.6170,0.6385,0.6401,0.6484,0.6498,0.6600,0.6698,0.6793,0.6851,\n0.6834,0.6901,0.7080,0.7092,0.7173,0.7212,0.7229,0.7337,0.7364,0.7380,\n0.7351,0.7312,0.7418,0.7358,0.7494,0.7362,0.7361,0.7432,0.7462,0.7408,\n0.7392,0.7393,0.7459,0.7428,0.7436,0.7505,0.7430,0.7481,0.7469,0.7448,\n0.7435,0.7448,0.7394,0.7402,0.7418,0.7432,0.7414,0.7419,0.7498,0.7416,\n0.7473,0.7443,0.7460,0.7472,0.7457,0.7409,0.7441,0.7411,0.7468,0.7458,\n0.7446,0.7492,0.7457,0.7452,0.7490,0.7462,0.7471,0.7451,0.7570,0.7441,\n0.7350,0.7484,0.7515,0.7399,0.7451,0.7458,0.7481,0.7413,0.7435,0.7497\n]\n\nval_acc = [\n0.1917,0.1867,0.2077,0.1930,0.1803,0.1913,0.2287,0.2087,0.2107,0.2390,\n0.2193,0.2187,0.2027,0.2390,0.2240,0.2200,0.2583,0.2040,0.2553,0.2280,\n0.1973,0.2553,0.2417,0.2400,0.2657,0.2110,0.2477,0.2673,0.2513,0.2433,\n0.2280,0.2470,0.2110,0.2383,0.2410,0.2487,0.2287,0.2243,0.2293,0.2343,\n0.2500,0.2650,0.2270,0.2277,0.2330,0.2410,0.1943,0.2620,0.2287,0.2240,\n0.2197,0.2367,0.2697,0.2237,0.2307,0.2703,0.2463,0.2153,0.2433,0.2777,\n0.2500,0.2683,0.2323,0.2487,0.2387,0.2283,0.2563,0.2150,0.2527,0.2390,\n0.2030,0.2443,0.2370,0.2640,0.2207,0.2170,0.2160,0.2247,0.2340,0.2223,\n0.2497,0.2167,0.2153,0.2243,0.2157,0.2837,0.2507,0.2813,0.2610,0.2417,\n0.2050,0.2130,0.2527,0.2310,0.2267,0.2853,0.2393,0.2220,0.2330,0.2280\n]\n\n# Plot Loss\nplt.figure()\nplt.plot(epochs, train_loss, label=\"Train Loss\")\nplt.plot(epochs, val_loss, label=\"Validation Loss\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.title(\"Training vs Validation Loss\")\nplt.legend()\nplt.show()\n\n# Plot Accuracy\nplt.figure()\nplt.plot(epochs, train_acc, label=\"Train Accuracy\")\nplt.plot(epochs, val_acc, label=\"Validation Accuracy\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Accuracy\")\nplt.title(\"Training vs Validation Accuracy\")\nplt.legend()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-25T13:29:51.913544Z","iopub.execute_input":"2025-12-25T13:29:51.914144Z","iopub.status.idle":"2025-12-25T13:29:52.193643Z","shell.execute_reply.started":"2025-12-25T13:29:51.914115Z","shell.execute_reply":"2025-12-25T13:29:52.193106Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}