EEGNet-GNN-Pos feature extractor (F1=6, F2=12)

EEGNet with Layer 2 replaced by a graph convolution over the electrode montage, where the montage geometry is a first-class input rather than something used once to build a fixed adjacency and then discarded.

Temporal conv -> POS-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 a flat feature vector; EEGNetGNNPosClassifier attaches a linear head. Output size is exposed as model.flat_dim (372 with the defaults).

What is position-conditioned

H = FiLM_pos(ร‚ X W) + X W_res, then BatchNorm + ELU, then a signed readout collapsing the electrodes to one. Coordinates enter twice:

  • FiLM โ€” per-node scale and shift predicted from (x, y). ร‚ alone cannot distinguish C3 from C4: their neighbourhoods are isomorphic, so message passing treats both hemispheres identically.
  • Readout โ€” node_weight[:, n] = MLP(gamma(x_n, y_n)), a continuous spatial field sampled at each electrode, rather than a free (F_out, N) table indexed by channel order.

That buys three things a lookup table cannot have: invariance to channel ordering, transfer to a montage with a different electrode count (transfer_to_montage), and a spatial field that can be evaluated between electrodes (pattern_at).

The readout must stay signed. A softmax over electrodes is a non-negative convex combination and cannot build the spatial high-pass filter that isolates focal C3/C4 ERD.

Options

arg default effect
readout "coord" "coord" / "lookup" (free table) / "attention" / "mean"
pos_mode "film" "film" / "bias" (shift only) / "none"
residual True adds un-smoothed node features back after message passing
learn_adjacency False makes ร‚ a trainable parameter
spatial "graph" "depthwise" restores stock EEGNet for a controlled A/B

Usage

from eegnet_gnn_pos import EEGNetGNNPos, EEGNetGNNPosClassifier
import torch

backbone = EEGNetGNNPos(spatial="graph", readout="coord")
model = EEGNetGNNPosClassifier(backbone, n_classes=4)

x = torch.randn(1, 1, 22, 1000)     # (batch, 1, channels, time)
logits = model(x)                   # (1, 4)

Train end-to-end. Do not freeze the backbone unless loading genuinely pretrained weights.

Caveat

readout="coord" has more parameters than the free table it replaces (~2.1k backbone vs ~1.2k). On ~288 trials per subject that is a real cost, and the geometric structure has to earn it. The honest comparison is a three-way A/B against readout="lookup" and spatial="depthwise" with an identical head, which is what Notebook 2 runs.

Downloads last month
11
Safetensors
Model size
2.74k params
Tensor type
F32
ยท
Inference Providers NEW
This model isn't deployed by any Inference Provider. ๐Ÿ™‹ Ask for provider support