Architecture-data matching for EEG-EMG decoding: compact deep models match classical spectral decoders on the WAY-EEG-GAL grasp-and-lift dataset.
The 29 matches
- [1] § Materials and methods › Input representations ↔ code/02_classical_baseline.py, lines 1–61 · score 0.99 · magnitude squared coherence, 15–30 Hz, 30–45 Hz, 8–12 Hz, 55–95 Hz, low gamma
- [2] § Materials and methods › Decoders ↔ code/run_deep_models_local.py, lines 161–191 · score 0.94 · attention weighted edges, fully connected, Graph attention network, channel tokens, temporal CNN, Conv1d
- [3] § Materials and methods › Decoders ↔ code/run_deep_models_local.py, lines 128–158 · score 0.93 · learnable CLS token, Conv1d patch embedding, positional embedding, encoder layers, Transformer encoder, feed
- [4] § Materials and methods › Decoders ↔ code/02b_classical_eval.py, lines 1–23 · score 0.92 · random forest, linear SVM, gradient boosting, max_depth, max_iter, logistic regression
- [5] § Materials and methods › Decoders ↔ code/revision_ablations.py, lines 123–154 · score 0.89 · convolution attention hybrids, DBConformer, MFTNet, EEG Conformer, encoder layers, Transformer encoder
- [6] § Materials and methods › Decoders ↔ code/07_figures.py, lines 70–103 · score 0.89 · gradient boosting, hidden layers, max_depth, max_iter, logistic regression, class weight
- [7] § Materials and methods › Decoders ↔ code/04_neural_baselines.py, lines 1–25 · score 0.87 · flattened raw downsampled, hand crafted spectral, channel band, band power, log, coherence
- [8] § Materials and methods › Input representations ↔ code/revision_ablations.py, lines 247–288 · score 0.84 · 15–30 Hz, 30–45 Hz, 8–12 Hz, 55–95 Hz, IQR, log10
- [9] § Materials and methods › Decoders ↔ code/revision_ablations.py, lines 190–227 · score 0.80 · cross entropy loss, AdamW, weight decay, LOSO folds, PyTorch, seed
- [10] § Materials and methods › Preprocessing and trial epoching ↔ code/01_build_dataset.py, lines 1–52 · score 0.79 · 1–100 Hz, Interference EMG, 20–500 Hz, rectification, notch, decimated
- [11] § Materials and methods › Decoders ↔ code/run_deep_models_local.py, lines 87–125 · score 0.79 · separable temporal convolution, depthwise spatial convolution, EEGNet, kernel, ELU, blocks
- [12] § Materials and methods › Interpretability ↔ extract_deep_interpretability.py, lines 1–36 · score 0.79 · model native interpretability, deep decoders, channel attention, LOSO folds, EEGNet, filters
- [13] § Results › Robustness analyses › Deep-model interpretability sanity check ↔ extract_deep_interpretability.py, lines 124–242 · score 0.78 · depthwise spatial weights, attention rollout, temporal filters, EEG channels, EEGNet, heads
- [14] § Materials and methods › Decoders ↔ code/regen_fig3.py, lines 108–140 · score 0.76 · cross entropy loss, AdamW, weight decay, PyTorch, seed, batch
- [15] § Materials and methods › Preprocessing and trial epoching ↔ code/01_build_dataset.py, lines 1–52 · score 0.75 · Skipping ICA, sustained hold, notch, artifact, decimation, pipeline
- [16] § Results › Fused-model feature importance suggests sensorimotor EEG involvement; modality ablation (Section 3.5.3) shows dominant EMG dependence ↔ code/02_classical_baseline.py, lines 1–61 · score 0.74 · low gamma, high gamma, band power, EEG channel, AD, CED
- [17] § Materials and methods › Decoders ↔ code/07_figures.py, lines 47–68 · score 0.71 · MLP_spectral, MLP_temporal, MLP_graph, log, classical
- [18] § Materials and methods › Preprocessing and trial epoching ↔ code/07_figures.py, lines 142–168 · score 0.70 · Pipeline overview, raw temporal, channel pooled, WAY EEG GAL, Friedman, Wilcoxon
- [19] § Materials and methods › Interpretability ↔ code/06_interpretability.py, lines 1–24 · score 0.69 · coefficient magnitude, rest classes, model native, HGBM, classical, surface
- [20] § Materials and methods › Preprocessing and trial epoching ↔ code/01_build_dataset.py, lines 142–215 · score 0.67 · LEDOff, fallback window, LEDOn, cropped, lift, Preprocessing
- [21] § Materials and methods › Interpretability ↔ code/revision_ablations.py, lines 1–50 · score 0.67 · channel attention matrices, model native interpretability, style, convolutional, heads, folds
- [22] § Results › Confusion structure ↔ code/regen_figS3.py, lines 112–156 · score 0.65 · confusion matrices, logistic regression, best classical, suede, predictions, sandpaper
- [23] § Results › Confusion structure ↔ code/07_figures.py, lines 70–103 · score 0.65 · confusion matrices, logistic regression, best classical, suede, predictions, sandpaper
- [24] § Materials and methods › Decoders ↔ code/revision_ablations.py, lines 424–456 · score 0.63 · deep model EMG, bandwidth harmonization, EMG bandwidth, baseline, 250 Hz, 95 Hz
- [25] § Results › Confusion structure ↔ code/regen_fig3.py, lines 157–209 · score 0.59 · confusion matrices, logistic regression, EEGNet, suede, predictions, sandpaper
- [26] § Results › Confusion structure ↔ code/regen_figS3.py, lines 112–156 · score 0.59 · confusion matrices, logistic regression, EEGNet, suede, predictions, sandpaper
- [27] § Materials and methods › Dataset ↔ code/run_deep_models_local.py, lines 1–65 · score 0.53 · WAY EEG GAL, grasp, suede, sandpaper, lift, silk
- [28] § Materials and methods › Evaluation protocol and statistical analysis ↔ code/05_compare.py, lines 1–64 · score 0.52 · paired Wilcoxon, Cohen, dz, Balanced accuracy, Friedman, score
- [29] § Materials and methods › Evaluation protocol and statistical analysis ↔ code/revision_ablations.py, lines 247–288 · score 0.52 · F1 score, modality ablation, scaler, Balanced accuracy, fitted, macro
Paper
Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC
The paper is loaded when this pane is shown.
The authors' code
Python · 528 lines · 28 KB · MIT · 6 matches
- """
- revision_ablations.py
- =====================
- Follow-up ablations requested by the Stanford AI reviewer. Five experiments
- in one resumable script:
- (i) 5-seed robustness for cnn / transformer / gnn (means +/- SD across seeds)
- (ii) modality ablation: EEG-only, EMG-only, fused for HGBM, CNN, GNN
- (iii) compact convolution-attention hybrid (Conformer-style) as a tighter
- attention comparator distinct from the 4-layer Transformer
- (iv) partial-crossing conditioning analyses (decode weight at fixed surface,
- decode surface at fixed weight)
- (v) model-native interpretability dump: GNN per-channel attention weights,
- CNN gradient saliency over time-channel for representative trials
- (vi) EMG-bandwidth harmonisation: 95-Hz lowpass on EMG inputs to deep models
- so they see the same bandwidth as classical features (20-95 Hz)
- (vii) early-hold (first 500 ms) vs late-hold (last 500 ms) decoding to test
- whether tactile-evoked cortical signatures concentrate near contact
- Setup (same venv used for run_complete_local.py):
- .venv\Scripts\activate
- python code\revision_ablations.py
- Default run: all five experiments, ~ 45-90 min on RTX 5080. Use flags to
- restrict (e.g. --skip seeds modality):
- python code\revision_ablations.py --only seeds
- python code\revision_ablations.py --only modality conformer
- python code\revision_ablations.py --skip interp
- Outputs (all under <root>/results/):
- seed_robustness.csv per (model, task, seed, fold) row
- modality_ablation.csv per (model, task, modality, fold) row
- conformer_hybrid.csv per (task, fold) row (model = conformer)
- crossing_conditioning.csv per (task_within_factor, fold) row
- gnn_attention_weights.npy per-fold attention matrices [12, n_layers, n_heads, 37, 37]
- cnn_saliency_examples.npy [n_examples, 37, 500] gradient saliency
- revision_summary.txt human-readable summary block
- """
- from __future__ import annotations
- import argparse, sys, time
- from pathlib import Path
- import numpy as np, pandas as pd
- import torch, torch.nn as nn, torch.nn.functional as F
- from sklearn.ensemble import HistGradientBoostingClassifier
- from sklearn.metrics import balanced_accuracy_score, f1_score
- from sklearn.preprocessing import StandardScaler
- from sklearn.utils.class_weight import compute_class_weight
- from torch.utils.data import DataLoader, TensorDataset
- # --- Paths (auto-detect; same convention as run_complete_local.py) -----------
- SCRIPT_DIR = Path(__file__).resolve().parent
- ROOT = SCRIPT_DIR.parent if SCRIPT_DIR.name == "code" else SCRIPT_DIR
- CACHE = ROOT / "cache_v2"
- RESULTS = ROOT / "results"; RESULTS.mkdir(exist_ok=True)
- N_EEG, N_EMG, N_CHANNELS, N_SAMPLES, N_CLASSES = 32, 5, 37, 500, 3
- # ---- Models (same as run_complete_local.py) --------------------------------
- class EEGNet(nn.Module):
- def __init__(self, n_channels=N_CHANNELS, F1=16, D=2, F2=32, dropout=0.4):
- super().__init__()
- self.conv1 = nn.Conv2d(1, F1, (1, 64), padding=(0, 32), bias=False)
- self.bn1 = nn.BatchNorm2d(F1)
- self.conv2 = nn.Conv2d(F1, F1*D, (n_channels, 1), groups=F1, bias=False)
- self.bn2 = nn.BatchNorm2d(F1*D)
- self.pool1 = nn.AvgPool2d((1, 4)); self.drop1 = nn.Dropout(dropout)
- self.conv3 = nn.Conv2d(F1*D, F2, (1, 16), padding=(0, 8), bias=False)
- self.bn3 = nn.BatchNorm2d(F2)
- self.pool2 = nn.AvgPool2d((1, 8)); self.drop2 = nn.Dropout(dropout)
- with torch.no_grad():
- d = torch.zeros(1, 1, n_channels, N_SAMPLES)
- h = self.bn2(self.conv2(self.bn1(self.conv1(d))))
- h = self.pool2(self.bn3(self.conv3(self.pool1(h))))
- self.flat = h.numel()
- self.fc = nn.Linear(self.flat, N_CLASSES)
- def forward(self, x):
- x = x.unsqueeze(1)
- x = self.bn1(self.conv1(x))
- x = F.elu(self.bn2(self.conv2(x))); x = self.drop1(self.pool1(x))
- x = F.elu(self.bn3(self.conv3(x))); x = self.drop2(self.pool2(x))
- return self.fc(x.flatten(1))
- class EEGTransformer(nn.Module):
- def __init__(self, n_channels=N_CHANNELS, n_tokens=25, d_model=128, n_heads=8, n_layers=4, dropout=0.2):
- super().__init__()
- stride = N_SAMPLES // n_tokens
- self.patch = nn.Conv1d(n_channels, d_model, kernel_size=stride, stride=stride)
- self.cls = nn.Parameter(torch.randn(1, 1, d_model)*0.02)
- self.pos = nn.Parameter(torch.randn(1, n_tokens+1, d_model)*0.02)
- layer = nn.TransformerEncoderLayer(d_model=d_model, nhead=n_heads,
- dim_feedforward=d_model*2, dropout=dropout, batch_first=True, activation="gelu")
- self.encoder = nn.TransformerEncoder(layer, num_layers=n_layers)
- self.norm = nn.LayerNorm(d_model)
- self.fc = nn.Linear(d_model, N_CLASSES)
- def forward(self, x):
- x = self.patch(x).transpose(1, 2)
- cls = self.cls.expand(x.size(0), -1, -1)
- x = torch.cat([cls, x], dim=1) + self.pos
- x = self.encoder(x)
- return self.fc(self.norm(x[:, 0]))
- class EEGGNN(nn.Module):
- def __init__(self, n_channels=N_CHANNELS, embed_dim=64, n_heads=4, n_layers=2, dropout=0.2):
- super().__init__()
- self.stem = nn.Sequential(
- nn.Conv1d(1, 16, 16, stride=4), nn.ELU(),
- nn.Conv1d(16, 32, 8, stride=4), nn.ELU(),
- nn.AdaptiveAvgPool1d(1))
- self.proj = nn.Linear(32, embed_dim)
- layer = nn.TransformerEncoderLayer(d_model=embed_dim, nhead=n_heads,
- dim_feedforward=embed_dim*2, dropout=dropout, batch_first=True, activation="gelu")
- self.gat = nn.TransformerEncoder(layer, num_layers=n_layers)
- self.norm = nn.LayerNorm(embed_dim)
- self.fc = nn.Linear(embed_dim, N_CLASSES)
- def forward(self, x):
- B, C, T = x.shape
- h = self.stem(x.reshape(B*C, 1, T)).squeeze(-1)
- h = self.proj(h.view(B, C, 32))
- h = self.gat(h)
- return self.fc(self.norm(h.mean(1)))
- class ConformerHybrid(nn.Module):
- """Compact convolution-attention hybrid: an EEGNet stem (~4k params) followed by
- a single self-attention block (~25k params). ~30k params total -- aligned with
- the EEG-Conformer / DBConformer / EEG-MFTNet family at clinical sample size."""
- def __init__(self, n_channels=N_CHANNELS, F1=16, D=2, embed_dim=64, n_heads=4, dropout=0.2):
- super().__init__()
- self.conv1 = nn.Conv2d(1, F1, (1, 64), padding=(0, 32), bias=False)
- self.bn1 = nn.BatchNorm2d(F1)
- self.conv2 = nn.Conv2d(F1, F1*D, (n_channels, 1), groups=F1, bias=False)
- self.bn2 = nn.BatchNorm2d(F1*D)
- self.pool = nn.AvgPool2d((1, 8))
- self.drop = nn.Dropout(dropout)
- # tokenize the (F1*D)-channel time stream into 12 tokens
- with torch.no_grad():
- d = torch.zeros(1, 1, n_channels, N_SAMPLES)
- h = self.pool(self.bn2(self.conv2(self.bn1(self.conv1(d)))))
- self._L = h.shape[-1]
- self.proj = nn.Linear(F1*D, embed_dim)
- layer = nn.TransformerEncoderLayer(d_model=embed_dim, nhead=n_heads,
- dim_feedforward=embed_dim*2, dropout=dropout, batch_first=True, activation="gelu")
- self.attn = nn.TransformerEncoder(layer, num_layers=1)
- self.norm = nn.LayerNorm(embed_dim)
- self.fc = nn.Linear(embed_dim, N_CLASSES)
- def forward(self, x):
- x = x.unsqueeze(1)
- x = F.elu(self.bn2(self.conv2(self.bn1(self.conv1(x)))))
- x = self.drop(self.pool(x))
- # x: (B, F1*D, 1, L) -> tokens (B, L, F1*D)
- x = x.squeeze(2).transpose(1, 2)
- x = self.proj(x)
- x = self.attn(x)
- return self.fc(self.norm(x.mean(1)))
- MODELS = {"cnn": EEGNet, "transformer": EEGTransformer, "gnn": EEGGNN, "conformer": ConformerHybrid}
- # ---- Data loading ----------------------------------------------------------
- def load_all_data(task, modality="fused"):
- """modality in {fused, eeg, emg}. eeg drops EMG channels; emg keeps only EMG."""
- Xs, ys, subj = [], [], []
- for p in range(1, 13):
- d = np.load(CACHE / f"P{p}.npz")
- v = d["valid_mask"]
- eeg = d["eeg"][v]; emg = d["emg"][v]
- if modality == "eeg":
- x = eeg.astype(np.float32) # (n, 32, 500)
- elif modality == "emg":
- x = emg.astype(np.float32) # (n, 5, 500)
- else:
- x = np.concatenate([eeg, emg], axis=1).astype(np.float32) # (n, 37, 500)
- y = (d["y_weight"] if task == "weight" else d["y_surface"])[v]
- Xs.append(x); ys.append(y); subj.append(np.full(len(x), p, dtype=np.int8))
- return np.concatenate(Xs), np.concatenate(ys), np.concatenate(subj)
- def load_metadata():
- """Loads CurW and CurS per trial across all 12 subjects (used for conditioning)."""
- rows = []
- for p in range(1, 13):
- d = np.load(CACHE / f"P{p}.npz")
- v = d["valid_mask"]
- n = int(v.sum())
- rows.append(pd.DataFrame({
- "subj": np.full(n, p),
- "weight": d["y_weight"][v],
- "surface": d["y_surface"][v]
- }))
- return pd.concat(rows, ignore_index=True)
- # ---- Deep training one fold ------------------------------------------------
- def deep_loso_fold(model_class, X, y, subj, held, *, device, n_epochs, batch_size, lr, seed, n_channels=N_CHANNELS):
- torch.manual_seed(seed); np.random.seed(seed)
- test = subj == held
- Xtr, ytr = X[~test], y[~test].astype(np.int64)
- Xte, yte = X[test], y[test].astype(np.int64)
- mu = Xtr.mean(0, keepdims=True); sd = Xtr.std(0, keepdims=True) + 1e-6
- Xtr = ((Xtr-mu)/sd).astype(np.float32); Xte = ((Xte-mu)/sd).astype(np.float32)
- Xtr_t, ytr_t = torch.from_numpy(Xtr), torch.from_numpy(ytr)
- Xte_t = torch.from_numpy(Xte)
- cw = compute_class_weight("balanced", classes=np.array([0,1,2]), y=ytr)
- loss_fn = nn.CrossEntropyLoss(weight=torch.tensor(cw, dtype=torch.float32).to(device))
- if model_class in (EEGNet, ConformerHybrid):
- model = model_class(n_channels=n_channels).to(device)
- elif model_class is EEGGNN:
- model = model_class(n_channels=n_channels).to(device)
- elif model_class is EEGTransformer:
- model = model_class(n_channels=n_channels).to(device)
- else:
- model = model_class().to(device)
- opt = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4)
- loader = DataLoader(TensorDataset(Xtr_t, ytr_t), batch_size=batch_size, shuffle=True)
- model.train()
- for _ in range(n_epochs):
- for xb, yb in loader:
- xb, yb = xb.to(device), yb.to(device)
- opt.zero_grad()
- loss_fn(model(xb), yb).backward()
- opt.step()
- model.eval()
- preds = []
- with torch.no_grad():
- for i in range(0, len(Xte_t), batch_size):
- preds.append(model(Xte_t[i:i+batch_size].to(device)).argmax(1).cpu().numpy())
- yp = np.concatenate(preds)
- return dict(held_out_subject=int(held), n_test=int(test.sum()),
- balanced_accuracy=float(balanced_accuracy_score(yte, yp)),
- f1_macro=float(f1_score(yte, yp, average="macro"))), model
- # ---- Experiment 1: seed robustness -----------------------------------------
- def exp_seeds(device, models=("cnn","transformer","gnn"), seeds=(0,1,2,3,4),
- n_epochs=30, batch_size=64, lr=1e-3):
- rows = []
- for task in ("weight","surface"):
- X, y, subj = load_all_data(task, "fused")
- for m in models:
- for s in seeds:
- for h in range(1, 13):
- t0 = time.time()
- res, _ = deep_loso_fold(MODELS[m], X, y, subj, h,
- device=device, n_epochs=n_epochs,
- batch_size=batch_size, lr=lr, seed=s)
- rows.append({"experiment":"seeds","model":m,"task":task,"seed":s,**res,"sec":round(time.time()-t0,1)})
- print(f"[seeds] {m}/{task}/seed{s}/P{h:02d}: bal={res['balanced_accuracy']:.3f}", flush=True)
- pd.DataFrame(rows).to_csv(RESULTS / "seed_robustness.csv", index=False)
- print(f"saved {RESULTS/'seed_robustness.csv'}")
- # ---- Experiment 2: modality ablation ---------------------------------------
- def exp_modality(device, n_epochs=30, batch_size=64, lr=1e-3, seed=0):
- rows = []
- for task in ("weight","surface"):
- for modality in ("eeg","emg","fused"):
- n_ch = {"eeg": N_EEG, "emg": N_EMG, "fused": N_CHANNELS}[modality]
- X, y, subj = load_all_data(task, modality)
- # HGBM on flattened mean+std+IQR per channel + log-power per channel
- from scipy.signal import welch
- Xf = []
- for x in X:
- ch_mean = x.mean(axis=1); ch_std = x.std(axis=1)
- ch_iqr = np.percentile(x, 75, axis=1) - np.percentile(x, 25, axis=1)
- feats = [ch_mean, ch_std, ch_iqr]
- for lo, hi in [(8,12),(15,30),(30,45),(55,95)]:
- f, pxx = welch(x, fs=500, nperseg=128, axis=-1)
- mask = (f >= lo) & (f <= hi)
- feats.append(np.log10(pxx[..., mask].mean(axis=-1) + 1e-12))
- Xf.append(np.concatenate(feats))
- Xf = np.stack(Xf)
- sc = StandardScaler().fit(Xf); Xs = sc.transform(Xf)
- for h in range(1, 13):
- test = subj == h
- clf = HistGradientBoostingClassifier(max_iter=200, max_depth=4,
- learning_rate=0.05, class_weight="balanced",
- random_state=0)
- clf.fit(Xs[~test], y[~test])
- yp = clf.predict(Xs[test])
- rows.append({"experiment":"modality","model":"HGBM","modality":modality,
- "task":task,"held_out_subject":h,
- "balanced_accuracy": balanced_accuracy_score(y[test], yp),
- "f1_macro": f1_score(y[test], yp, average="macro")})
- for m in ("cnn", "gnn"):
- for h in range(1, 13):
- res, _ = deep_loso_fold(MODELS[m], X, y, subj, h, device=device,
- n_epochs=n_epochs, batch_size=batch_size,
- lr=lr, seed=seed, n_channels=n_ch)
- rows.append({"experiment":"modality","model":m,"modality":modality,
- "task":task,**res})
- print(f"[modality] {m}/{task}/{modality}/P{h:02d}: bal={res['balanced_accuracy']:.3f}", flush=True)
- pd.DataFrame(rows).to_csv(RESULTS / "modality_ablation.csv", index=False)
- print(f"saved {RESULTS/'modality_ablation.csv'}")
- # ---- Experiment 3: Conformer-style hybrid ----------------------------------
- def exp_conformer(device, n_epochs=30, batch_size=64, lr=1e-3, seed=0):
- rows = []
- for task in ("weight","surface"):
- X, y, subj = load_all_data(task, "fused")
- for h in range(1, 13):
- res, _ = deep_loso_fold(ConformerHybrid, X, y, subj, h,
- device=device, n_epochs=n_epochs,
- batch_size=batch_size, lr=lr, seed=seed)
- rows.append({"experiment":"conformer","model":"conformer","task":task,**res})
- print(f"[conformer] {task}/P{h:02d}: bal={res['balanced_accuracy']:.3f}", flush=True)
- pd.DataFrame(rows).to_csv(RESULTS / "conformer_hybrid.csv", index=False)
- print(f"saved {RESULTS/'conformer_hybrid.csv'}")
- # ---- Experiment 4: partial-crossing conditioning ---------------------------
- def exp_crossing(device, n_epochs=30, batch_size=64, lr=1e-3, seed=0):
- rows = []
- md = load_metadata()
- X_full, y_full, subj_full = load_all_data("weight", "fused")
- # Decode weight conditioned on a fixed surface (the surface with most coverage)
- for fix_surf in [0, 2]: # sandpaper, silk
- keep = md.surface.values == fix_surf
- X_sub = X_full[keep]; y_sub = y_full[keep]; subj_sub = subj_full[keep]
- for m in ("cnn", "gnn"):
- for h in range(1, 13):
- if (subj_sub == h).sum() < 5: continue
- res, _ = deep_loso_fold(MODELS[m], X_sub, y_sub, subj_sub, h,
- device=device, n_epochs=n_epochs,
- batch_size=batch_size, lr=lr, seed=seed)
- rows.append({"experiment":"conditioning","model":m,
- "decode":"weight","fixed_factor":"surface","fixed_value":int(fix_surf),
- "task":"weight",**res})
- # Decode surface conditioned on a fixed weight
- X_full2, y_full2, subj_full2 = load_all_data("surface", "fused")
- for fix_w in [0, 1, 2]:
- keep = md.weight.values == fix_w
- X_sub = X_full2[keep]; y_sub = y_full2[keep]; subj_sub = subj_full2[keep]
- for m in ("cnn", "gnn"):
- for h in range(1, 13):
- if (subj_sub == h).sum() < 5: continue
- res, _ = deep_loso_fold(MODELS[m], X_sub, y_sub, subj_sub, h,
- device=device, n_epochs=n_epochs,
- batch_size=batch_size, lr=lr, seed=seed)
- rows.append({"experiment":"conditioning","model":m,
- "decode":"surface","fixed_factor":"weight","fixed_value":int(fix_w),
- "task":"surface",**res})
- pd.DataFrame(rows).to_csv(RESULTS / "crossing_conditioning.csv", index=False)
- print(f"saved {RESULTS/'crossing_conditioning.csv'}")
- # ---- Experiment 5: model-native interpretability ---------------------------
- def exp_interp(device, n_epochs=30, batch_size=64, lr=1e-3, seed=0):
- """Hook attention weights from the GNN (per-fold) and gradient saliency for the CNN."""
- X, y, subj = load_all_data("weight", "fused")
- attn_per_fold = []
- for h in range(1, 13):
- torch.manual_seed(seed); np.random.seed(seed)
- test = subj == h
- Xtr, ytr = X[~test], y[~test].astype(np.int64)
- mu = Xtr.mean(0, keepdims=True); sd = Xtr.std(0, keepdims=True) + 1e-6
- Xtr = ((Xtr-mu)/sd).astype(np.float32)
- Xte = ((X[test]-mu)/sd).astype(np.float32)
- cw = compute_class_weight("balanced", classes=np.array([0,1,2]), y=ytr)
- loss_fn = nn.CrossEntropyLoss(weight=torch.tensor(cw, dtype=torch.float32).to(device))
- model = EEGGNN().to(device)
- opt = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4)
- Xtr_t = torch.from_numpy(Xtr); ytr_t = torch.from_numpy(ytr)
- loader = DataLoader(TensorDataset(Xtr_t, ytr_t), batch_size=batch_size, shuffle=True)
- model.train()
- for _ in range(n_epochs):
- for xb, yb in loader:
- xb, yb = xb.to(device), yb.to(device)
- opt.zero_grad(); loss_fn(model(xb), yb).backward(); opt.step()
- model.eval()
- # Hook attention from the first GAT layer
- attn_buffer = []
- def hook(module, inputs, output):
- # MultiheadAttention returns (attn_output, attn_weights) when need_weights=True
- pass
- # Compute attention manually by re-running the model with need_weights=True path
- Xte_t = torch.from_numpy(Xte).to(device)
- with torch.no_grad():
- B, C, T = Xte_t.shape
- h_ = model.stem(Xte_t.reshape(B*C, 1, T)).squeeze(-1)
- h_ = model.proj(h_.view(B, C, 32))
- # iterate the encoder layers and capture attention
- attns = []
- x_ = h_
- for layer in model.gat.layers:
- # nn.TransformerEncoderLayer self-attn
- x_norm = layer.norm1(x_)
- attn_out, attn_w = layer.self_attn(x_norm, x_norm, x_norm,
- need_weights=True, average_attn_weights=False)
- x_ = x_ + layer.dropout1(attn_out)
- x_ = x_ + layer.dropout2(layer.linear2(layer.dropout(layer.activation(layer.linear1(layer.norm2(x_))))))
- attns.append(attn_w.mean(0).cpu().numpy()) # [n_heads, C, C] averaged over batch
- attn_per_fold.append(np.stack(attns)) # [n_layers, n_heads, C, C]
- np.save(RESULTS / "gnn_attention_weights.npy", np.stack(attn_per_fold))
- print(f"saved {RESULTS/'gnn_attention_weights.npy'}", " shape:", np.stack(attn_per_fold).shape)
- # CNN saliency: train a single fold (P1 held out), compute per-class saliency on a few test trials
- torch.manual_seed(seed); np.random.seed(seed)
- test = subj == 1
- Xtr, ytr = X[~test], y[~test].astype(np.int64)
- Xte, yte = X[test], y[test].astype(np.int64)
- mu = Xtr.mean(0, keepdims=True); sd = Xtr.std(0, keepdims=True) + 1e-6
- Xtr = ((Xtr-mu)/sd).astype(np.float32); Xte = ((Xte-mu)/sd).astype(np.float32)
- cw = compute_class_weight("balanced", classes=np.array([0,1,2]), y=ytr)
- loss_fn = nn.CrossEntropyLoss(weight=torch.tensor(cw, dtype=torch.float32).to(device))
- cnn = EEGNet().to(device)
- opt = torch.optim.AdamW(cnn.parameters(), lr=lr, weight_decay=1e-4)
- loader = DataLoader(TensorDataset(torch.from_numpy(Xtr), torch.from_numpy(ytr)), batch_size=batch_size, shuffle=True)
- cnn.train()
- for _ in range(n_epochs):
- for xb, yb in loader:
- xb, yb = xb.to(device), yb.to(device)
- opt.zero_grad(); loss_fn(cnn(xb), yb).backward(); opt.step()
- cnn.eval()
- # Pick 2 trials per class
- examples_idx = []
- for c in range(3):
- cls_idx = np.where(yte == c)[0][:2]
- examples_idx.extend(cls_idx)
- saliencies = []
- for idx in examples_idx:
- x = torch.from_numpy(Xte[idx:idx+1]).to(device).requires_grad_(True)
- logits = cnn(x)
- c = int(yte[idx])
- cnn.zero_grad()
- logits[0, c].backward()
- sal = x.grad.abs().squeeze(0).cpu().numpy() # (37, 500)
- saliencies.append(sal)
- np.save(RESULTS / "cnn_saliency_examples.npy", np.stack(saliencies))
- print(f"saved {RESULTS/'cnn_saliency_examples.npy'}", " shape:", np.stack(saliencies).shape)
- # ---- Experiment 6: EMG-bandwidth harmonisation -----------------------------
- def exp_bandwidth(device, n_epochs=30, batch_size=64, lr=1e-3, seed=0):
- """Apply a 95-Hz lowpass to the deep-model EMG input and re-train CNN/GNN/HGBM
- so all decoders see the same EMG bandwidth (20-95 Hz). Compare the resulting
- LOSO accuracies to the unharmonised baseline."""
- from scipy.signal import butter, filtfilt
- rows = []
- # 95-Hz lowpass at 500 Hz: butter order 4
- b95, a95 = butter(4, 95.0 / 250.0, btype="low")
- for task in ("weight","surface"):
- # Load fused (32 EEG + 5 EMG)
- Xs, ys, subj = [], [], []
- for p in range(1, 13):
- d = np.load(CACHE / f"P{p}.npz")
- v = d["valid_mask"]
- eeg = d["eeg"][v]
- emg = d["emg"][v]
- # Apply 95-Hz LP to EMG channels along time axis
- emg_lp = filtfilt(b95, a95, emg, axis=-1).astype(np.float32)
- x = np.concatenate([eeg, emg_lp], axis=1).astype(np.float32)
- y = (d["y_weight"] if task == "weight" else d["y_surface"])[v]
- Xs.append(x); ys.append(y); subj.append(np.full(len(x), p, dtype=np.int8))
- X = np.concatenate(Xs); y = np.concatenate(ys); subj = np.concatenate(subj)
- for m in ("cnn", "gnn"):
- for h in range(1, 13):
- res, _ = deep_loso_fold(MODELS[m], X, y, subj, h, device=device,
- n_epochs=n_epochs, batch_size=batch_size,
- lr=lr, seed=seed)
- rows.append({"experiment":"bandwidth","model":m,"task":task,
- "harmonisation":"emg_lowpass_95hz",**res})
- print(f"[bandwidth] {m}/{task}/P{h:02d}: bal={res['balanced_accuracy']:.3f}", flush=True)
- pd.DataFrame(rows).to_csv(RESULTS / "bandwidth_harmonisation.csv", index=False)
- print(f"saved {RESULTS/'bandwidth_harmonisation.csv'}")
- # ---- Experiment 7: early-hold vs late-hold ---------------------------------
- def exp_earlyhold(device, n_epochs=30, batch_size=64, lr=1e-3, seed=0):
- """Decode from the first 500 ms (early-hold) vs the last 500 ms (late-hold)
- of the sustained-hold window. The cached window is already 1 s; we slice
- in half along the time axis."""
- rows = []
- for task in ("weight","surface"):
- X, y, subj = load_all_data(task, "fused") # (n, 37, 500)
- halves = {"early": X[..., :250], "late": X[..., 250:]}
- for label, X_half in halves.items():
- # Pad/repeat to 500 samples so the existing CNN/GNN architectures fit unchanged
- X_pad = np.concatenate([X_half, X_half], axis=-1) # tile to 500
- for m in ("cnn", "gnn"):
- for h in range(1, 13):
- res, _ = deep_loso_fold(MODELS[m], X_pad, y, subj, h, device=device,
- n_epochs=n_epochs, batch_size=batch_size,
- lr=lr, seed=seed)
- rows.append({"experiment":"earlyhold","model":m,"task":task,
- "phase":label, **res})
- print(f"[earlyhold] {m}/{task}/{label}/P{h:02d}: bal={res['balanced_accuracy']:.3f}", flush=True)
- pd.DataFrame(rows).to_csv(RESULTS / "earlyhold_vs_latehold.csv", index=False)
- print(f"saved {RESULTS/'earlyhold_vs_latehold.csv'}")
- # ---- Summary ---------------------------------------------------------------
- def summarize():
- summary = []
- for name, fn in [
- ("seed_robustness.csv", lambda df: df.groupby(["model","task"])["balanced_accuracy"].agg(["mean","std","count"]).round(3)),
- ("modality_ablation.csv", lambda df: df.groupby(["model","modality","task"])["balanced_accuracy"].agg(["mean","std"]).round(3)),
- ("conformer_hybrid.csv", lambda df: df.groupby(["task"])["balanced_accuracy"].agg(["mean","std"]).round(3)),
- ("crossing_conditioning.csv", lambda df: df.groupby(["model","decode","fixed_factor","fixed_value"])["balanced_accuracy"].agg(["mean","std","count"]).round(3)),
- ("bandwidth_harmonisation.csv", lambda df: df.groupby(["model","task"])["balanced_accuracy"].agg(["mean","std"]).round(3)),
- ("earlyhold_vs_latehold.csv", lambda df: df.groupby(["model","task","phase"])["balanced_accuracy"].agg(["mean","std"]).round(3)),
- ]:
- path = RESULTS / name
- if path.exists():
- df = pd.read_csv(path)
- summary.append(f"\n=== {name} ===\n{fn(df).to_string()}")
- out = "\n".join(summary)
- (RESULTS / "revision_summary.txt").write_text(out)
- print("\n" + out)
- print(f"\nsaved {RESULTS/'revision_summary.txt'}")
- # ---- Main ------------------------------------------------------------------
- def main():
- ap = argparse.ArgumentParser()
- ap.add_argument("--only", nargs="+", choices=["seeds","modality","conformer","crossing","interp","bandwidth","earlyhold"])
- ap.add_argument("--skip", nargs="+", default=[], choices=["seeds","modality","conformer","crossing","interp","bandwidth","earlyhold"])
- ap.add_argument("--epochs", type=int, default=30)
- ap.add_argument("--batch-size", type=int, default=64)
- ap.add_argument("--lr", type=float, default=1e-3)
- args = ap.parse_args()
- device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
- print("Using device:", device)
- todo = ["seeds","modality","conformer","crossing","interp","bandwidth","earlyhold"]
- if args.only:
- todo = [t for t in todo if t in args.only]
- todo = [t for t in todo if t not in args.skip]
- print("Plan:", todo)
- if "seeds" in todo: exp_seeds(device, n_epochs=args.epochs, batch_size=args.batch_size, lr=args.lr)
- if "modality" in todo: exp_modality(device, n_epochs=args.epochs, batch_size=args.batch_size, lr=args.lr)
- if "conformer" in todo: exp_conformer(device, n_epochs=args.epochs, batch_size=args.batch_size, lr=args.lr)
- if "crossing" in todo: exp_crossing(device, n_epochs=args.epochs, batch_size=args.batch_size, lr=args.lr)
- if "interp" in todo: exp_interp(device, n_epochs=args.epochs, batch_size=args.batch_size, lr=args.lr)
- if "bandwidth" in todo: exp_bandwidth(device, n_epochs=args.epochs, batch_size=args.batch_size, lr=args.lr)
- if "earlyhold" in todo: exp_earlyhold(device, n_epochs=args.epochs, batch_size=args.batch_size, lr=args.lr)
- summarize()
- if __name__ == "__main__":
- main()
revision_ablations.py at commit 2caa855, under MIT · at the source
Overview
- Biomedical Engineering Postgraduate Program, Anhembi Morumbi University, São José dos Campos, Brazil
- Department of Kinesiology, California State University San Marcos (CSUSM), San Marcos, CA, United States
- Neurometra, Carlsbad, CA, United States
Abstract
Modern deep learning has broadened the tools available for non-invasive neural decoding, but its advantage over well-engineered classical pipelines remains unclear at clinical neural-engineering sample sizes. We compared four classical decoders, three multilayer perceptron (MLP) variants, and three deep-learning architectures (an EEGNet-style compact convolutional network, a four-layer Transformer encoder trained from scratch, and a graph attention network) on the public WAY-EEG-GAL grasp-and-lift dataset (12 participants, 3,528 trials). Models were evaluated using leave-one-subject-out (LOSO) cross-validation to decode object weight (165, 330, and 660 g) and grasp-surface friction (sandpaper, suede, and silk). After Benjamini–Hochberg false discovery rate (BH-FDR) correction within the primary/
Reproduced under the paper's license (CC BY), from the paper cited above.
Repositories
Its files are read in the Code ↔ Paper reader above, with 29 matches between paragraphs and lines of code.
osmar235/eeg-emg-architecture-data-matching
2caa855503f5697eb11bdee065cc40b2694df736, 22 June 2026Availability: 1 check, the latest on 27 September 2026: the link answers
- 27 September 2026: the link answers
17 files
- code/
01_build_dataset.py , Python, 249 lines, 3 matches - code/
02_classical_baseline.py , Python, 196 lines, 2 matches - code/
02b_classical_eval.py , Python, 73 lines, 1 match - code/
03_deep_models.py , Python, 186 lines - code/
04_neural_baselines.py , Python, 133 lines, 1 match - code/
05_compare.py , Python, 80 lines, 1 match - code/
06_interpretability.py , Python, 107 lines, 1 match - code/
07_figures.py , Python, 175 lines, 4 matches - code/
code_01_ML_neuro.ipynb , Jupyter, 399 lines - code/
regen_fig3.py , Python, 212 lines, 2 matches - code/
regen_figS3.py , Python, 159 lines, 2 matches - code/
revision_ablations.py , Python, 528 lines, 6 matches - code/
run_complete_local.py , Python, 394 lines - code/
run_deep_models_local.py , Python, 346 lines, 4 matches - extract_deep_interpretab
ility.py , Python, 245 lines, 2 matches - LICENSE, License, 21 lines
- README.md, Text, 191 lines
Zenodo 19966634
Availability: 1 check, the latest on 27 September 2026: the link answers (HTTP 200)
- 27 September 2026: the link answers (HTTP 200)
15 files
- code/
01_build_dataset.py , Python, 249 lines - code/
02_classical_baseline.py , Python, 196 lines - code/
02b_classical_eval.py , Python, 73 lines - code/
03_deep_models.py , Python, 186 lines - code/
04_neural_baselines.py , Python, 133 lines - code/
05_compare.py , Python, 80 lines - code/
06_interpretability.py , Python, 107 lines - code/
07_figures.py , Python, 175 lines - code/
code_01_ML_neuro.ipynb , Jupyter, 399 lines - code/
regen_fig3.py , Python, 212 lines - code/
revision_ablations.py , Python, 528 lines - code/
run_complete_local.py , Python, 394 lines - code/
run_deep_models_local.py , Python, 346 lines - LICENSE, License, 21 lines
- README.md, Text, 191 lines
The paper's code and data availability statement is in the Data section.
Tracing map
Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.
What the map holds:
- 2 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
- 28 scripts, each with its path and the digest of its content;
- 29 matches between paragraphs of the paper and lines of the code (method lexical-v1);
- neither the text of the paper nor the code itself.
Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.
Data
No dataset and no data link were found in the paper.
Data availability statement
The raw WAY-EEG-GAL dataset analyzed in this study is publicly available from Luciw et al. (2014) and via PhysioNet. All analysis code, derived per-subject feature tensors, model checkpoints, and figure-generation scripts are available on GitHub (https://
Reproduced under the paper's license (CC BY), from the paper cited above.
Versions
The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.
Version 1, 27 September 2026: the first record
Recorded: type, language, journal, volume, pages, dates, 1 author, 8 keywords, 1 funder, 28 references.
Cite
This paper
Pinto Neto, O. (2026). Architecture-data matching for EEG-EMG decoding: compact deep models match classical spectral decoders on the WAY-EEG-GAL grasp-and-lift dataset. Frontiers in neuroscience, 20, 1874302. https://
BibTeX
@article{pintoneto2026ar
author = {Pinto Neto, Osmar},
title = {{Architecture-data matching for EEG-EMG decoding: compact deep models match classical spectral decoders on the WAY-EEG-GAL grasp-and-lift dataset}},
journal = {Frontiers in neuroscience},
year = {2026},
month = jul,
volume = {20},
pages = {1874302},
publisher = {Frontiers Media SA},
issn = {1662-4548},
doi = {10.3389/
url = {https://
pmid = {42548755},
pmcid = {PMC13429724}
}
RIS
TY - JOUR
AU - Pinto Neto, Osmar
TI - Architecture-data matching for EEG-EMG decoding: compact deep models match classical spectral decoders on the WAY-EEG-GAL grasp-and-lift dataset
T2 - Frontiers in neuroscience
J2 - Front Neurosci
PY - 2026
DA - 2026/
VL - 20
SP - 1874302
SN - 1662-4548
PB - Frontiers Media SA
DO - 10.3389/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.3389/
"type": "article-journal",
"title": "Architecture-data matching for EEG-EMG decoding: compact deep models match classical spectral decoders on the WAY-EEG-GAL grasp-and-lift dataset",
"container-title": "Frontiers in neuroscience",
"author": [
{
"family": "Pinto Neto",
"given": "Osmar"
}
],
"container-title-short":
"volume": "20",
"page": "1874302",
"DOI": "10.3389/
"PMID": "42548755",
"PMCID": "PMC13429724",
"ISSN": "1662-4548",
"publisher": "Frontiers Media SA",
"URL": "https://
"language": "en",
"issued": {
"date-parts": [
[
2026,
7,
20
]
]
}
}
The tracing map gets a citation of its own once an author has validated it and it has a DOI.
Similar papers
The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.
- [1] doi:10.1371/journal.pone.0347671 [code]
- RMETNet: A cross-subject motor imagery EEG signal classification model based on TSLANet and riemannian geometry features.Journal: PloS oneIn common: PyTorch, scikit-learn, pandas, 3 other tools, methods / tools, EEG, 4 references
- [2] doi:10.3390/s26103065 [code]
- Subject-Wise Depression Screening from Eight-Channel Resting-State EEG Using Asymmetry-Aware Spectral Features and Connectivity Ablation.Journal: Sensors (Basel, Switzerland)In common: PyTorch, scikit-learn, pandas, 3 other tools, other, EEG, 3 references
- [3] doi:10.1038/s41598-026-68186-2 [code]
- NeuroStream: spectral-spatio-temporal
deep learning for visual stimulus classification from EEG. Journal: Scientific reportsIn common: PyTorch, scikit-learn, pandas, 3 other tools, methods / tools, EEG, 3 references - [4] doi:10.1038/s41746-026-02778-0 [code]
- Trust-gated synthetic EEG augmentation reduces performance drops when generalizing to new patients.Journal: NPJ digital medicineIn common: PyTorch, scikit-learn, pandas, 3 other tools, EEG, 3 references
- [5] doi:10.3390/bios16080437 [code]
- Deep Learning for Ear-EEG-Based Brain-Computer Interface: A Systematic Comparison and Design Insights.Journal: BiosensorsIn common: PyTorch, scikit-learn, pandas, 2 other tools, EEG, 3 references
- [6] doi:10.3389/fpsyg.2026.1774068 [code]
- Analysis of cognitive mechanisms in phoneme perception and pronunciation errors among Korean language learners.Journal: Frontiers in psychologyIn common: PyTorch, scikit-learn, pandas, 3 other tools, EEG, 2 references
- [7] doi:10.1038/s41586-026-10658-6 [code]
- An AI system to help scientists write expert-level empirical software.Journal: NatureIn common: JAX, PyTorch, scikit-learn, 4 other tools, methods / tools
- [8] doi:10.1371/journal.pone.0346575 [code]
- Statistically valid explainable black-box machine learning: applications in sex classification across species using brain imaging.Journal: PloS oneIn common: JAX, PyTorch, scikit-learn, 4 other tools, methods / tools
- [9] doi:10.3390/e28030310
- Entropy-Based Dual-Teacher Distillation for Efficient Motor Imagery EEG Classification.Journal: Entropy (Basel, Switzerland)In common: methods / tools, EEG, 5 references
- [10] doi:10.3390/s26051730 [code]
- SFE-GAT: Structure-Feature Evolution Graph Attention Network for Motor Imagery Decoding.Journal: Sensors (Basel, Switzerland)In common: PyTorch, SciPy, Matplotlib, 1 other tool, EEG, 3 references
Contribute
The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.
Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 2 repositories of the authors' code, each at its verified commit and with its license, 28 scripts, and 29 matches between paragraphs and code (see the Code and Map sections). It then receives a DOI on Zenodo, with you (your ORCID iD) and OSCR as its creators; the code itself is not deposited.
The map's fingerprint: sha256:045cf659d4524d18…
Add the badge to its README
The badge links the code to this page. Copy one of these into the README of the paper's code: only you decide where it goes, and nothing is changed for you.
Markdown
[, paste the snippet at the top, then “Commit changes…” and, to review it first, “Create a new branch and start a pull request”. You open the pull request; OSCR asks for no permission.
Request its removal
To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).
Discussion, reproductions, activity
Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.
Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.
Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.
