eegnet-gnn-features / eegnet_gnn.py
shemalfoy's picture
Upload eegnet_gnn.py with huggingface_hub
7b27d1c verified
Raw History Blame Contribute Delete
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)