Download eegnet_gnn.py from shemalfoy/eegnet-gnn-features: direct link, hf CLI and curl.
- Browser
- Download file 10.3 kB
-
https://huggingface.co/shemalfoy/eegnet-gnn-features/resolve/main/eegnet_gnn.py
- Command line
-
hf download hf://shemalfoy/eegnet-gnn-features/eegnet_gnn.py
-
curl -L -o eegnet_gnn.py https://huggingface.co/shemalfoy/eegnet-gnn-features/resolve/main/eegnet_gnn.py
10.3 kB
| """ | |
| EEGNet-GNN feature extractor (Option C, through Layer 4 -- no VQC head) | |
| Standard EEGNet with one change: Layer 2's depthwise *spatial* conv is swapped for a | |
| **graph convolution** over the electrode montage, so electrodes mix according to how | |
| close they are on the scalp instead of as a flat, order-agnostic channel list. | |
| Temporal conv -> GRAPH conv -> Separable conv -> Avg-pool + flatten -> flat vector | |
| (Layer 1) (Layer 2) (Layer 3) (Layer 4) OUTPUT | |
| IDENTICAL THE SWAP IDENTICAL IDENTICAL | |
| The model stops at Layer 4 and returns the flattened feature vector (B, flat_dim). | |
| Attach your own classifier / VQC downstream. A `spatial=` flag switches between the | |
| graph conv and the original depthwise conv so you can benchmark the swap in isolation. | |
| Dependencies: torch, huggingface_hub, safetensors. | |
| """ | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| try: | |
| # gives the model save_pretrained / push_to_hub / from_pretrained | |
| from huggingface_hub import PyTorchModelHubMixin as _HubMixin | |
| HAS_HFHUB = True | |
| except Exception: | |
| HAS_HFHUB = False | |
| class _HubMixin: # no-op so the file still runs without huggingface_hub installed | |
| def __init_subclass__(cls, **kwargs): # swallow license=/tags= metadata kwargs | |
| super().__init_subclass__() | |
| # --------------------------------------------------------------------------- | |
| # 1. Electrode topology -> adjacency matrix (this is what makes it a GNN) | |
| # --------------------------------------------------------------------------- | |
| # Nominal 2-D scalp positions for the 22 channels of BCI Competition IV-2a. | |
| # x = lateral (left negative), y = anterior->posterior. Numbers follow 10-10 | |
| # naming: 'z' = midline (0), odd = left, even = right, magnitude = laterality. | |
| BCI_IV_2A_COORDS = { | |
| "Fz": (0, 2), | |
| "FC3": (-3, 1), "FC1": (-1, 1), "FCz": (0, 1), "FC2": (2, 1), "FC4": (4, 1), | |
| "C5": (-5, 0), "C3": (-3, 0), "C1": (-1, 0), "Cz": (0, 0), | |
| "C2": (2, 0), "C4": (4, 0), "C6": (6, 0), | |
| "CP3": (-3, -1), "CP1": (-1, -1), "CPz": (0, -1), "CP2": (2, -1), "CP4": (4, -1), | |
| "P1": (-1, -2), "Pz": (0, -2), "P2": (2, -2), | |
| "POz": (0, -3), | |
| } | |
| BCI_IV_2A_ORDER = list(BCI_IV_2A_COORDS.keys()) | |
| def build_adjacency(coords=BCI_IV_2A_COORDS, order=BCI_IV_2A_ORDER, | |
| sigma=1.5, threshold=0.25, knn=None): | |
| """Distance-based adjacency. A_ij = exp(-d_ij^2 / (2 sigma^2)); far pairs -> 0. | |
| sigma : Gaussian width (how quickly connection strength decays with distance) | |
| threshold : edges weaker than this are cut (keeps only true neighbours) | |
| knn : if set, keep each node's k nearest neighbours instead (overrides threshold) | |
| Returns a (N, N) float tensor with zero diagonal (self-loops added at normalise time). | |
| """ | |
| pos = np.array([coords[e] for e in order], dtype=np.float32) # (N, 2) | |
| diff = pos[:, None, :] - pos[None, :, :] | |
| dist = np.sqrt((diff ** 2).sum(-1)) # (N, N) | |
| A = np.exp(-(dist ** 2) / (2 * sigma ** 2)) | |
| np.fill_diagonal(A, 0.0) | |
| if knn is not None: | |
| keep = np.zeros_like(A) | |
| for i in range(A.shape[0]): | |
| nbrs = np.argsort(-A[i])[:knn] | |
| keep[i, nbrs] = A[i, nbrs] | |
| A = np.maximum(keep, keep.T) # symmetrise | |
| else: | |
| A[A < threshold] = 0.0 | |
| return torch.tensor(A, dtype=torch.float32) | |
| def normalize_adjacency(A): | |
| """Symmetric-normalised adjacency with self-loops (Kipf & Welling): | |
| A_hat = D^-1/2 (A + I) D^-1/2.""" | |
| N = A.shape[0] | |
| A = A + torch.eye(N, dtype=A.dtype, device=A.device) | |
| deg = A.sum(1) | |
| d_inv_sqrt = deg.pow(-0.5) | |
| d_inv_sqrt[torch.isinf(d_inv_sqrt)] = 0.0 | |
| D_inv_sqrt = torch.diag(d_inv_sqrt) | |
| return D_inv_sqrt @ A @ D_inv_sqrt | |
| def neighbours_of(electrode, A=None, order=BCI_IV_2A_ORDER): | |
| """List the electrodes an electrode is connected to (for sanity checks).""" | |
| if A is None: | |
| A = build_adjacency() | |
| i = order.index(electrode) | |
| return [order[j] for j in range(len(order)) if A[i, j] > 0 and j != i] | |
| # --------------------------------------------------------------------------- | |
| # 2. THE SWAP --- Layer 2 as a graph convolution | |
| # --------------------------------------------------------------------------- | |
| class GraphConvLayer(nn.Module): | |
| """Topology-aware replacement for EEGNet's depthwise spatial conv. | |
| Input : (B, F_in, N, T) -- F_in temporal-filter maps, N electrodes, T time | |
| Steps : | |
| 1. message passing H = A_hat @ X @ W (each node pulls from its neighbours) | |
| 2. BatchNorm + ELU | |
| 3. node readout collapse N electrodes -> 1 (learned spatial pooling) | |
| Output: (B, F_out, 1, T) -- same shape the depthwise conv would produce. | |
| """ | |
| def __init__(self, in_features, out_features, adjacency, readout="weighted"): | |
| super().__init__() | |
| A_hat = normalize_adjacency(adjacency) | |
| self.register_buffer("A_hat", A_hat) # (N, N), fixed | |
| self.N = A_hat.shape[0] | |
| self.lin = nn.Linear(in_features, out_features, bias=False) # the "W" | |
| self.bn = nn.BatchNorm2d(out_features) | |
| self.act = nn.ELU() | |
| self.readout = readout | |
| if readout == "weighted": # static learned per-node weight | |
| self.node_weight = nn.Parameter(torch.ones(self.N)) | |
| elif readout == "attention": # dynamic, input-dependent | |
| self.att = nn.Linear(out_features, 1) | |
| def forward(self, x): | |
| B, F_in, N, T = x.shape | |
| h = x.permute(0, 3, 2, 1) # (B, T, N, F_in) | |
| h = torch.einsum("ij,btjf->btif", self.A_hat, h) # aggregate from neighbours | |
| h = self.lin(h) # (B, T, N, F_out) | |
| h = h.permute(0, 3, 2, 1) # (B, F_out, N, T) | |
| h = self.act(self.bn(h)) | |
| # collapse the N electrodes -> 1 | |
| if self.readout == "mean": | |
| pooled = h.mean(dim=2, keepdim=True) | |
| elif self.readout == "weighted": | |
| w = torch.softmax(self.node_weight, dim=0) # (N,) | |
| pooled = torch.einsum("bfnt,n->bft", h, w).unsqueeze(2) | |
| elif self.readout == "attention": | |
| s = h.permute(0, 3, 2, 1) # (B, T, N, F_out) | |
| a = torch.softmax(self.att(s), dim=2) # (B, T, N, 1) | |
| pooled = (s * a).sum(2).permute(0, 2, 1).unsqueeze(2) | |
| else: | |
| raise ValueError(f"unknown readout '{self.readout}'") | |
| return pooled # (B, F_out, 1, T) | |
| class DepthwiseSpatial(nn.Module): | |
| """Standard EEGNet Layer 2 (the thing being swapped out) -- kept for A/B comparison. | |
| Learns F2 independent weighted sums across ALL electrodes at once; no topology.""" | |
| def __init__(self, F1, F2, n_channels): | |
| super().__init__() | |
| assert F2 % F1 == 0, "F2 must be a multiple of F1 (depth multiplier D = F2/F1)" | |
| self.conv = nn.Conv2d(F1, F2, (n_channels, 1), groups=F1, bias=False) | |
| self.bn = nn.BatchNorm2d(F2) | |
| self.act = nn.ELU() | |
| def forward(self, x): | |
| return self.act(self.bn(self.conv(x))) # (B, F2, 1, T) | |
| # --------------------------------------------------------------------------- | |
| # 3. Shared piece --- separable conv (Layer 3), IDENTICAL in both branches | |
| # --------------------------------------------------------------------------- | |
| class SeparableConv(nn.Module): | |
| """EEGNet Layer 3: depthwise temporal conv + pointwise mix. Refines temporal features.""" | |
| def __init__(self, in_ch, out_ch, kernel): | |
| super().__init__() | |
| self.depthwise = nn.Conv2d(in_ch, in_ch, (1, kernel), | |
| padding="same", groups=in_ch, bias=False) | |
| self.pointwise = nn.Conv2d(in_ch, out_ch, 1, bias=False) | |
| def forward(self, x): | |
| return self.pointwise(self.depthwise(x)) | |
| # --------------------------------------------------------------------------- | |
| # 4. Feature extractor --- Layers 1-4, returns the flat vector (no head) | |
| # --------------------------------------------------------------------------- | |
| class EEGNetGNN(nn.Module, _HubMixin, | |
| license="mit", | |
| pipeline_tag="feature-extraction", | |
| tags=["eeg", "bci", "motor-imagery", "eegnet", | |
| "graph-neural-network", "feature-extraction", "pytorch"]): | |
| def __init__(self, n_channels=22, n_times=1000, adjacency=None, | |
| F1=8, F2=16, kern_length=64, sep_kern=16, | |
| spatial="graph", readout="weighted", pool=32, dropout=0.25): | |
| super().__init__() | |
| if adjacency is None: | |
| adjacency = build_adjacency() | |
| # Layer 1 -- Temporal conv (IDENTICAL in both branches) | |
| self.temporal = nn.Sequential( | |
| nn.Conv2d(1, F1, (1, kern_length), padding="same", bias=False), | |
| nn.BatchNorm2d(F1), | |
| ) | |
| # Layer 2 -- the swap point | |
| if spatial == "graph": | |
| self.spatial = GraphConvLayer(F1, F2, adjacency, readout=readout) | |
| elif spatial == "depthwise": | |
| self.spatial = DepthwiseSpatial(F1, F2, n_channels) | |
| else: | |
| raise ValueError("spatial must be 'graph' or 'depthwise'") | |
| self.spatial_drop = nn.Dropout(dropout) | |
| # Layer 3 -- Separable conv (IDENTICAL) | |
| self.separable = nn.Sequential( | |
| SeparableConv(F2, F2, sep_kern), | |
| nn.BatchNorm2d(F2), | |
| nn.ELU(), | |
| nn.Dropout(dropout), | |
| ) | |
| # Layer 4 -- Avg pool + flatten (IDENTICAL) -> flat feature vector | |
| self.pool = nn.AvgPool2d((1, pool)) | |
| self.flatten = nn.Flatten() | |
| # expose the output feature size (for whatever head you attach later) | |
| with torch.no_grad(): | |
| self.flat_dim = self.forward(torch.zeros(1, 1, n_channels, n_times)).shape[1] | |
| def forward(self, x): | |
| x = self.temporal(x) # (B, F1, N, T) | |
| x = self.spatial(x) # (B, F2, 1, T) | |
| x = self.spatial_drop(x) | |
| x = self.separable(x) # (B, F2, 1, T) | |
| x = self.pool(x) # (B, F2, 1, T/pool) | |
| return self.flatten(x) # (B, flat_dim) | |