OSCR

MEG-GPT: A transformer-based foundation model for magnetoencephalography data.

Code ↔ Paper

13 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

The 13 matches
  1. [1] § Methods › The foundation model: MEG-GPT › Transformer and prediction head ↔ osl_foundation/inference/layers.py, lines 237–342 · score 0.70 · receptive field, attention matrix, latent sequence length, attends, unpatched, masked
  2. [2] § Methods › The foundation model: MEG-GPT › Transformer and prediction head ↔ osl_foundation/models/meg_gpt.py, lines 406–444 · score 0.65 · Feed Forward layer, residual connections, MEG GPT
  3. [3] § Results › MEG-GPT captures bursting dynamics in MEG data ↔ examples/paper/analysis/4_bursting_behaviour.py, lines 295–367 · score 0.62 · bursting behaviour, AR model, lifetime, interval, MEG GPT
  4. [4] § Results › MEG-GPT extracts features that enhance decoding performance ↔ examples/paper/decoding/7_decoding_performances.py, lines 87–138 · score 0.62 · decoding accuracy, zero shot, Fine tuning, baseline, model
  5. [5] § Methods › The foundation model: MEG-GPT › Masked multi-head attention ↔ osl_foundation/models/meg_gpt.py, lines 280–322 · score 0.61 · unpatched sequences, attention layers, sequence length, MEG GPT, patching, Decoder
  6. [6] § Methods › The foundation model: MEG-GPT › Masked multi-head attention ↔ osl_foundation/inference/layers.py, lines 237–342 · score 0.58 · receptive field, sequence length, attend, unpatched, patching, layers
  7. [7] § Results › MEG-GPT extracts features that enhance decoding performance ↔ examples/paper/decoding/7_decoding_performances.py, lines 87–138 · score 0.57 · decoding accuracy, zero shot, fine tuned, baseline
  8. [8] § Methods › The foundation model: MEG-GPT › Generating new data ↔ osl_foundation/utils/sampling.py, lines 91–133 · score 0.56 · cumulative probability exceeds, smallest, tokens
  9. [9] § Results › The tokeniser reconstructs MEG data with high accuracy and generalises to unseen data ↔ examples/paper/analysis/1_tokeniser_reconstructs.py, lines 134–183 · score 0.56 · tokeniser reconstructions, variance explained, Wakeman Henson, PVE, Cam, training
  10. [10] § Results › MEG-GPT captures spatial and spectral characteristics of real data ↔ examples/paper/analysis/3_subject_fingerprints.py, lines 166–285 · score 0.54 · power spectral density, spectral features, Welch, GPT, MEG
  11. [11] § Results › The tokeniser reconstructs MEG data with high accuracy and generalises to unseen data ↔ examples/paper/analysis/1_tokeniser_reconstructs.py, lines 134–183 · score 0.53 · variance explained, Wakeman Henson, PVE, reconstruction, tokeniser, Cam
  12. [12] § Results › MEG-GPT captures bursting dynamics in MEG data ↔ examples/paper/analysis/4_bursting_behaviour.py, lines 507–556 · score 0.53 · Wavelet transform, AR model, Location, bursting, HMM, MEG GPT
  13. [13] § Methods › Tokeniser › Training › Annealing ↔ osl_foundation/inference/layers.py, lines 345–386 · score 0.52 · Gumbel Softmax, annealing, argmax, weighted, token, training

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 · 1,134 lines · 40 KB · MIT · 3 matches

  1. from typing import Union
  2. import numpy as np
  3. import tensorflow as tf
  4. import tensorflow_probability as tfp
  5. def rnn_layer(
  6. rnn_type: str, rnn_n_units: int, return_sequences: bool = True
  7. ) -> tf.keras.layers.Layer:
  8. """
  9. Create an RNN layer.
  10. Parameters
  11. ----------
  12. rnn_type : str
  13. Type of RNN layer. Options are 'gru', 'lstm'.
  14. rnn_n_units : int
  15. Number of units in the RNN layer.
  16. return_sequences : bool, optional
  17. Whether to return sequences.
  18. Returns
  19. -------
  20. rnn_layer : tf.keras.layers.Layer
  21. RNN layer.
  22. """
  23. if rnn_type == "gru":
  24. return tf.keras.layers.GRU(
  25. rnn_n_units, return_sequences=return_sequences, stateful=False
  26. )
  27. elif rnn_type == "lstm":
  28. return tf.keras.layers.LSTM(
  29. rnn_n_units, return_sequences=return_sequences, stateful=False
  30. )
  31. else:
  32. raise ValueError(f"Unknown RNN type: {rnn_type}")
  33. class IdentityLayer(tf.keras.layers.Layer):
  34. """
  35. Identity layer.
  36. This layer directly returns the input tensor.
  37. """
  38. def call(self, inputs, **kwargs):
  39. return inputs
  40. class MSELossLayer(tf.keras.layers.Layer):
  41. """
  42. Layer for computing the mean squared error loss.
  43. This is a wrapper around tf.keras.losses.MeanSquaredError.
  44. """
  45. def __init__(self, **kwargs):
  46. super().__init__(**kwargs)
  47. self.loss_fn = lambda y_true, y_pred: tf.reduce_mean(tf.square(y_true - y_pred))
  48. def call(self, y_true, y_pred, **kwargs):
  49. loss = self.loss_fn(y_true, y_pred)
  50. self.add_loss(loss)
  51. return tf.expand_dims(loss, -1)
  52. class SinusoidalPositionalEncodingLayer(tf.keras.layers.Layer):
  53. """
  54. Layer for generating sinusoidal positional encoding for position embeddings.
  55. Implemented as in "Attention is All You Need" (Vaswani et al., 2017), and partly
  56. adpated from the `huggingface/transformers` library and TensorFlow
  57. (https://www.tensorflow.org/text/tutorials/transformer).
  58. Parameters
  59. ----------
  60. sequence_length : int
  61. Length of the sequence.
  62. """
  63. def __init__(self, sequence_length: int, **kwargs):
  64. super().__init__(**kwargs)
  65. self.sequence_length = sequence_length
  66. def call(self, inputs, **kwargs):
  67. # Validation
  68. embedding_dim = tf.shape(inputs)[-1]
  69. tf.debugging.assert_equal(
  70. tf.math.floormod(embedding_dim, 2),
  71. tf.constant(0, dtype=tf.int32),
  72. message="embedding dimension must be even for positional encoding.",
  73. )
  74. # Precompute scaling factor
  75. embedding_dim_float = tf.cast(embedding_dim, tf.float32)
  76. denominator = tf.math.pow(
  77. 10000.0,
  78. 2 * (tf.range(embedding_dim_float) // 2) / embedding_dim_float,
  79. )
  80. # Get position indices
  81. position_indices = tf.range(self.sequence_length, dtype=tf.float32)[
  82. :, tf.newaxis
  83. ]
  84. # position_indices.shape = (sequence_length, 1)
  85. # Compute positional encoding
  86. angle_rads = position_indices / denominator
  87. # angle_rads.shape = (sequence_length, embedding_dim)
  88. pos_encoding = tf.TensorArray(dtype=tf.float32, size=embedding_dim)
  89. for i in range(embedding_dim):
  90. if i % 2 == 0:
  91. pos_encoding = pos_encoding.write(i, tf.sin(angle_rads[:, i]))
  92. else:
  93. pos_encoding = pos_encoding.write(i, tf.cos(angle_rads[:, i]))
  94. pos_encoding = tf.transpose(pos_encoding.stack())
  95. # pos_encoding.shape = (sequence_length, embedding_dim)
  96. # Reshape the positions to match shapes of other embeddings
  97. positions = tf.expand_dims(
  98. tf.expand_dims(pos_encoding, axis=0), axis=0
  99. ) # positions.shape = (1, 1, sequence_length, embedding_dim)
  100. return positions
  101. class RotaryPositionEmbeddingLayer(tf.keras.layers.Layer):
  102. """
  103. Layer for generating Rotary Position Embedding (RoPE). Implemented as in
  104. "RoFormer: Enhanced Transformer with Rotary Position Embedding" (Su et al., 2022), and
  105. adpated from the PyTorch and Keras 3 libraries (see below for the reference).
  106. References:
  107. - https://pytorch.org/torchtune/0.2/_modules/torchtune/modules/position_embeddings.html#RotaryPositionalEmbeddings
  108. - https://github.com/keras-team/keras-hub/blob/v0.17.0/keras_hub/src/layers/modeling/rotary_embedding.py
  109. Parameters
  110. ----------
  111. embedding_dim : int
  112. Dimension of the rotary position embedding. When using multi-head attention blocks,
  113. the dimension is usually `embedding_dim // n_heads`.
  114. max_sequence_length : int
  115. Maximum expected length of the sequence. If the model dimensions of the key and
  116. query vectors differ, the sequence length should be the maximum of the two.
  117. base : int, optional
  118. The base value for computing the rotation angles.
  119. """
  120. def __init__(
  121. self, embedding_dim: int, max_sequence_length: int, base: int = 10000, **kwargs
  122. ):
  123. super().__init__(**kwargs)
  124. # Validation
  125. if embedding_dim % 2 != 0:
  126. raise ValueError(
  127. "embedding dimension must be even for positional encoding."
  128. )
  129. self.embedding_dim = embedding_dim
  130. self.max_sequence_length = max_sequence_length
  131. self.base = base
  132. self.rotation_values = self._compute_rotation_matrix()
  133. def _compute_rotation_matrix(self) -> tf.Tensor:
  134. # Compute pre-defined theta angles
  135. theta = 1.0 / (
  136. self.base
  137. ** (
  138. tf.range(0, self.embedding_dim, 2, dtype=tf.float32)[
  139. : (self.embedding_dim // 2)
  140. ]
  141. / self.embedding_dim
  142. )
  143. )
  144. # Create position indices
  145. seq_idx = tf.range(self.max_sequence_length, dtype=tf.float32)
  146. # NOTE: Having 0 as a starting index is okay, since we don't have to rotate the first position.
  147. # Compute outer product of position index and theta
  148. idx_theta = tf.einsum("i,j->ij", seq_idx, theta)
  149. # idx_theta.shape = (sequence_length, embedding_dim // 2)
  150. # Compute elements of the rotation matrix
  151. rotation_values = tf.stack([tf.cos(idx_theta), tf.sin(idx_theta)], axis=-1)
  152. # rotation_values.shape = (sequence_length, embedding_dim // 2, 2)
  153. return rotation_values
  154. def call(self, inputs, **kwargs):
  155. # Get input sequence length
  156. sequence_length = tf.shape(inputs)[3]
  157. # inputs.shape = (batch_size, n_heads, n_channels, sequence_length, key_dim)
  158. # Validation
  159. tf.debugging.assert_greater_equal(
  160. self.max_sequence_length,
  161. sequence_length,
  162. message="input sequence length should not exceed max_sequence_length.",
  163. )
  164. tf.debugging.assert_equal(
  165. tf.shape(inputs)[-1],
  166. self.embedding_dim,
  167. message="the model dimension of the input tensor must match the embedding dimension.",
  168. )
  169. # Reshape the rotation values to be compabitle with the input shape
  170. rot_vals = tf.reshape(
  171. self.rotation_values[:sequence_length],
  172. (1, 1, 1, sequence_length, self.embedding_dim // 2, 2),
  173. )
  174. # NOTE: We clip the rotation values to the sequence length, in case sequence_length is
  175. # less than max_sequence_length as the dimension for key and query differs.
  176. # Match the input shape to the rope shape
  177. rot_shape = tf.concat(
  178. [tf.shape(inputs)[:-1], [self.embedding_dim // 2, 2]], axis=0
  179. )
  180. x_reshaped = tf.reshape(inputs, rot_shape)
  181. # x_reshaped.shape = (batch_size, n_heads, n_channels, sequence_length, key_dim // 2, 2)
  182. # Calculate the rotary position embedding
  183. rope = tf.stack(
  184. [
  185. rot_vals[..., 0] * x_reshaped[..., 0]
  186. - rot_vals[..., 1] * x_reshaped[..., 1],
  187. rot_vals[..., 1] * x_reshaped[..., 0]
  188. + rot_vals[..., 0] * x_reshaped[..., 1],
  189. ],
  190. axis=-1,
  191. )
  192. # rope.shape = (batch_size, n_heads, n_channels, sequence_length, key_dim // 2, 2)
  193. rope = tf.reshape(rope, tf.shape(inputs))
  194. # rope.shape = (batch_size, n_heads, n_channels, sequence_length, key_dim)
  195. return rope
  196. class ALiBiPositionEmbeddingLayer(tf.keras.layers.Layer):
  197. """
  198. Layer for generating Attention with Linear Biases (ALiBi) position embedding. Implemented
  199. based on "Train Short, Test Long: Attention with Linear Biases Enables Input Length Extrapolation"
  200. (Press et al., 2022) and adapted from the references below.
  201. References:
  202. - https://github.com/ofirpress/attention_with_linear_biases
  203. - https://nn.labml.ai/transformers/alibi/index.html
  204. Parameters
  205. ----------
  206. n_heads : int
  207. Number of attention heads.
  208. n_patches : int
  209. Number of patches to attend to.
  210. patch_length : int
  211. Patch length.
  212. unpatched_length : int
  213. Number of unpatched elements to attend to.
  214. """
  215. def __init__(
  216. self,
  217. n_heads: int,
  218. n_patches: int,
  219. patch_length: int,
  220. unpatched_length: int,
  221. **kwargs,
  222. ):
  223. super().__init__(**kwargs)
  224. self.n_heads = n_heads
  225. self.n_patches = n_patches
  226. self.patch_length = patch_length
  227. self.unpatched_length = unpatched_length
  228. def _get_slopes(self) -> tf.Tensor:
  229. # Compute the slopes for the linear biases (static, non-learned)
  230. h = tf.cast(self.n_heads, tf.float32)
  231. n = 2 ** tf.math.floor(tf.math.log(h) / tf.math.log(2.0)) # nearest power of 2
  232. ratio = tf.pow(2.0, -8.0 / n)
  233. m = tf.pow(
  234. ratio, tf.range(1, tf.cast(n, tf.int32) + 1, dtype=tf.float32)
  235. ) # head-specific slopes
  236. # When n_heads is not a power of 2, then we add additional slopes
  237. if n < h:
  238. ratio_hat = tf.pow(2.0, -4.0 / n)
  239. m_hat = tf.pow(ratio_hat, tf.range(1, 1 + 2 * (h - n), 2, dtype=tf.float32))
  240. m = tf.concat([m, m_hat], axis=0)
  241. return m
  242. def _get_relative_bias_matrix(self, mask: tf.Tensor) -> tf.Tensor:
  243. # Get head-specific slopes
  244. m = self._get_slopes()
  245. # Get a binary masking matrix
  246. mask = 1 - mask
  247. cum_mask = tf.cumsum(mask, axis=0) # cumulative sum of the mask
  248. # NOTE: In the mask matrix, columns are where you attend to and rows are the targets.
  249. # Compute the relative distance matrix for the patches
  250. patch_indices = tf.cast(tf.range(self.n_patches), dtype=tf.float32)
  251. avg_idx = (
  252. (self.patch_length * patch_indices + 1)
  253. + self.patch_length * (patch_indices + 1)
  254. ) / 2.0
  255. distance_patched = (
  256. -(cum_mask[:, : self.n_patches] + self.patch_length * (patch_indices + 1))
  257. + avg_idx
  258. )
  259. # NOTE: Considering a patch as a single receptive field, we calculate relative distance between the patch
  260. # and the current sequence index (i.e., the distance between the center of the patch and the current index).
  261. # This amounts to the average distance between the current index and all the indices in the patch.
  262. # Compute the relative distance matrix for the unpatched elements
  263. distance_unpatched = -cum_mask[
  264. :, self.n_patches : self.n_patches + self.unpatched_length
  265. ]
  266. # Concatenate two matrices
  267. distance = tf.concat([distance_patched, distance_unpatched], axis=1)
  268. # distance.shape = (latent_sequence_length, n_patches + unpatched_length)
  269. # Compute head-specific relative bias matrices
  270. masked_distance = distance * mask
  271. alibi_matrix = masked_distance[None, :, :] * m[:, None, None]
  272. # alibi_matrix.shape = (n_heads, latent_sequence_length, n_patches + unpatched_length)
  273. return alibi_matrix
  274. def call(self, inputs, mask: tf.Tensor, **kwargs):
  275. # Get the relative bias matrix
  276. alibi_matrix = self._get_relative_bias_matrix(mask)
  277. # alibi_matrix.shape = (n_heads, latent_sequence_length, n_patches + unpatched_length)
  278. # Reshape the relative bias matrix to match the attention matrix
  279. alibi_matrix = tf.reshape(
  280. alibi_matrix, (1, self.n_heads, 1, tf.shape(mask)[0], tf.shape(mask)[1])
  281. )
  282. # Add the relative bias matrix to the attention matrix
  283. attention_with_alibi = inputs + alibi_matrix
  284. # attention_with_alibi.shape = (batch_size, n_heads, n_channels, out_sequence_length, in_sequence_length)
  285. return attention_with_alibi
  286. class TokenWeightsLayer(tf.keras.layers.Layer):
  287. """
  288. Layer for computing token weights.
  289. Parameters
  290. ----------
  291. output_dim : int
  292. Dimension of the output.
  293. activation : str, optional
  294. Activation function to use.
  295. """
  296. def __init__(self, output_dim: int, activation: str = "linear", **kwargs):
  297. super().__init__(**kwargs)
  298. self.output_dim = output_dim
  299. self.dense_layer = tf.keras.layers.Dense(output_dim, activation=activation)
  300. self.activation_layer = tf.keras.layers.Activation(activation)
  301. self.norm_layer = tf.keras.layers.LayerNormalization()
  302. self.temperature = tf.Variable(0.0, trainable=False)
  303. def call(self, inputs, training=None, **kwargs):
  304. ell = self.activation_layer(self.dense_layer(inputs))
  305. ell = self.norm_layer(ell) / 0.1
  306. # Shape: (batch_size * n_channels, sequence_length, n_tokens)
  307. if training:
  308. # Sample from gumbel softmax parameterized by ell
  309. theta_sample = tf.argmax(ell, axis=2)
  310. theta_sample = tf.one_hot(theta_sample, self.output_dim)
  311. # Annealing
  312. theta_weight = tf.nn.softmax(ell, axis=2)
  313. # Shape: (batch_size * n_channels, sequence_length, n_tokens)
  314. token_weight = (
  315. self.temperature * theta_weight + (1 - self.temperature) * theta_sample
  316. )
  317. # Shape: (batch_size * n_channels, sequence_length, n_tokens)
  318. else:
  319. token_weight = tf.one_hot(tf.argmax(ell, axis=2), self.output_dim)
  320. return token_weight
  321. class PositionEmbedding(tf.keras.layers.Layer):
  322. """
  323. Layer for learning position embeddings.
  324. Parameters
  325. ----------
  326. sequence_length : int
  327. Sequence length.
  328. initializer : str, optional
  329. Initializer for the position embeddings.
  330. """
  331. def __init__(
  332. self,
  333. sequence_length: int,
  334. initializer: str = "glorot_uniform",
  335. **kwargs,
  336. ):
  337. super().__init__(**kwargs)
  338. self.sequence_length = sequence_length
  339. self.initializer = tf.keras.initializers.get(initializer)
  340. def build(self, inputs_shape):
  341. feature_size = inputs_shape[-1]
  342. self.position_embeddings = self.add_weight(
  343. name="embeddings",
  344. shape=[self.sequence_length, feature_size],
  345. initializer=self.initializer,
  346. trainable=True,
  347. )
  348. self.built = True
  349. def call(self, inputs, start_index=0):
  350. inputs_shape = tf.shape(inputs)
  351. feature_length = inputs_shape[-1]
  352. sequence_length = inputs_shape[-2]
  353. # trim to match the length of the input sequence, which might be less
  354. # than the sequence_length of the layer.
  355. position_embeddings = tf.convert_to_tensor(self.position_embeddings)
  356. position_embeddings = tf.slice(
  357. position_embeddings,
  358. (start_index, 0),
  359. (sequence_length, feature_length),
  360. )
  361. return tf.broadcast_to(position_embeddings, inputs_shape)
  362. class NormalizationLayer(tf.keras.layers.Layer):
  363. """
  364. Layer for performing normalization.
  365. Parameters
  366. ----------
  367. norm_type : str
  368. Type of normalization to perform.
  369. Options: "layer", "batch", "group".
  370. n_groups : int, optional
  371. Number of groups for group normalization.
  372. Required if norm_type is "group".
  373. """
  374. def __init__(self, norm_type: str = "layer", n_groups=None, **kwargs):
  375. super().__init__(**kwargs)
  376. if norm_type == "layer":
  377. self.norm_layer = tf.keras.layers.LayerNormalization()
  378. elif norm_type == "batch":
  379. self.norm_layer = tf.keras.layers.BatchNormalization()
  380. elif norm_type == "group":
  381. if n_groups is None:
  382. raise ValueError("n_groups must be specified for group normalization")
  383. self.norm_layer = tf.keras.layers.GroupNormalization(groups=n_groups)
  384. else:
  385. raise ValueError(f"Unknown normalization type: {norm_type}")
  386. def call(self, inputs, **kwargs):
  387. return self.norm_layer(inputs)
  388. class TimeAttentionLayer(tf.keras.layers.Layer):
  389. """
  390. Layer for performing time attention.
  391. Parameters
  392. ----------
  393. n_heads : int
  394. Number of heads.
  395. latent_sequence_length : int
  396. Sequence length of latent space.
  397. key_dim : int
  398. Key dimension.
  399. n_patches : int
  400. Number of patches to attend to.
  401. patch_length : int
  402. Patch length.
  403. unpatched_length : int
  404. Number of unpatched elements to attend to.
  405. transform_attention : str, optional
  406. Type of attention transformation to apply.
  407. """
  408. def __init__(
  409. self,
  410. n_heads: int,
  411. latent_sequence_length: int,
  412. key_dim: int,
  413. n_patches: int,
  414. patch_length: int,
  415. unpatched_length: int,
  416. transform_attention: str = None,
  417. **kwargs,
  418. ):
  419. super().__init__(**kwargs)
  420. self.n_heads = n_heads
  421. self.latent_sequence_length = latent_sequence_length
  422. self.key_dim = key_dim
  423. self.n_patches = n_patches
  424. self.patch_length = patch_length
  425. self.unpatched_length = unpatched_length
  426. self.transform_attention = transform_attention
  427. # Create a position embedding layer
  428. if transform_attention == "rope":
  429. self.rotary_position_embedding_layer = RotaryPositionEmbeddingLayer(
  430. embedding_dim=self.key_dim,
  431. max_sequence_length=tf.cast(
  432. tf.math.maximum(
  433. self.latent_sequence_length,
  434. self.n_patches + self.unpatched_length,
  435. ),
  436. dtype=tf.int32,
  437. ),
  438. )
  439. if transform_attention == "alibi":
  440. self.alibi_position_embedding_layer = ALiBiPositionEmbeddingLayer(
  441. n_heads=self.n_heads,
  442. n_patches=self.n_patches,
  443. patch_length=self.patch_length,
  444. unpatched_length=self.unpatched_length,
  445. )
  446. def call(self, inputs, mask=None, **kwargs):
  447. # q (Query): (batch_size, n_heads, out_sequence_length, out_n_channels, key_dim)
  448. # k (Key): (batch_size, n_heads, in_sequence_length, out_n_channels, key_dim)
  449. # v (Value): (batch_size, n_heads, out_sequence_length, out_n_channels, in_sequence_length, key_dim)
  450. q, k, v = inputs
  451. # Transpose inputs for time attention
  452. q = tf.transpose(q, perm=(0, 1, 3, 2, 4))
  453. # q: (batch_size, n_heads, out_n_channels, out_sequence_length, key_dim)
  454. k = tf.transpose(k, perm=(0, 1, 3, 2, 4))
  455. # k: (batch_size, n_heads, out_n_channels, in_sequence_length, key_dim)
  456. if self.transform_attention == "rope":
  457. # Apply rotary position embedding to the query and key vectors
  458. q = self.rotary_position_embedding_layer(q)
  459. k = self.rotary_position_embedding_layer(k)
  460. # Compute attention
  461. attention = tf.matmul(q, k, transpose_b=True) / tf.math.sqrt(
  462. tf.cast(self.key_dim, tf.float32)
  463. )
  464. # attention: (batch_size, n_heads, out_n_channels, out_sequence_length, in_sequence_length)
  465. if self.transform_attention == "alibi":
  466. # Apply ALiBi position embedding to the attention matrix
  467. attention = self.alibi_position_embedding_layer(attention, mask=mask)
  468. # Apply mask
  469. if mask is not None:
  470. attention += -1e9 * mask
  471. # Normalise attention with softmax
  472. attention = tf.nn.softmax(attention, axis=-1)
  473. if mask is not None:
  474. attention = attention * (1 - mask)
  475. # Apply attention to value
  476. attention = tf.expand_dims(
  477. tf.transpose(attention, perm=(0, 1, 3, 2, 4)), axis=-2
  478. )
  479. # attention: (batch_size, n_heads, out_sequence_length, out_n_channels, 1, in_sequence_length)
  480. output = tf.matmul(attention, v)
  481. # output: (batch_size, n_heads, out_sequence_length, out_n_channels, 1, key_dim)
  482. output = tf.squeeze(output, axis=-2)
  483. # output: (batch_size, n_heads, out_sequence_length, out_n_channels, key_dim)
  484. return output
  485. class ChannelAttention(tf.keras.layers.Layer):
  486. """
  487. Layer for performing channel attention.
  488. Parameters
  489. ----------
  490. key_dim : int
  491. Key dimension.
  492. """
  493. def __init__(self, key_dim: int, **kwargs):
  494. super().__init__(**kwargs)
  495. self.key_dim = key_dim
  496. def call(self, inputs, mask=None, **kwargs):
  497. # q (Query): (batch_size, n_heads, out_sequence_length, out_n_channels, key_dim)
  498. # k (Key): (batch_size, n_heads, out_sequence_length, in_n_channels, key_dim)
  499. # v (Value): (batch_size, n_heads, in_sequence_length, in_n_channels, key_dim)
  500. q, k, v = inputs
  501. batch_size = tf.shape(q)[0]
  502. n_heads = tf.shape(q)[1]
  503. out_sequence_length = tf.shape(q)[2]
  504. out_n_channels = tf.shape(q)[3]
  505. in_n_channels = tf.shape(k)[3]
  506. in_sequence_length = tf.shape(v)[2]
  507. # Compute attention
  508. attention = tf.matmul(q, k, transpose_b=True) / tf.math.sqrt(
  509. tf.cast(self.key_dim, tf.float32)
  510. )
  511. # attention: (batch_size, n_heads, out_sequence_length, out_n_channels, in_n_channels)
  512. # Apply mask
  513. if mask is not None:
  514. attention += -1e9 * mask
  515. # Normalise attention with softmax
  516. attention = tf.nn.softmax(attention, axis=-1)
  517. if mask is not None:
  518. attention = attention * (1 - mask)
  519. # Apply attention to value
  520. v = tf.transpose(v, perm=(0, 1, 3, 2, 4))
  521. # v: (batch_size, n_heads, in_n_channels, in_sequence_length, key_dim)
  522. v = tf.reshape(v, shape=(batch_size, n_heads, 1, in_n_channels, -1))
  523. # v: (batch_size, n_heads, 1, in_n_channels, in_sequence_length * key_dim)
  524. output = tf.matmul(attention, v)
  525. # output: (batch_size, n_heads, out_sequence_length, out_n_channels, in_sequence_length * key_dim)
  526. output = tf.reshape(
  527. output,
  528. shape=(
  529. batch_size,
  530. n_heads,
  531. out_sequence_length,
  532. out_n_channels,
  533. in_sequence_length,
  534. self.key_dim,
  535. ),
  536. )
  537. # output: (batch_size, n_heads, out_sequence_length, out_n_channels, in_sequence_length, key_dim)
  538. return output
  539. class PASSTALayer(tf.keras.layers.Layer):
  540. """
  541. The Perceiver AR Separable Space-Time self-Attention (PASSTA) layer.
  542. This layer performs space-time attention on the input tensor.
  543. Parameters
  544. ----------
  545. n_heads : int
  546. Number of heads.
  547. n_channels : int
  548. Number of channels.
  549. latent_sequence_length : int
  550. Sequence length of latent space.
  551. n_patches : int
  552. Number of patches to attend to.
  553. patch_length : int
  554. Patch length.
  555. unpatched_length : int
  556. Number of unpatched elements to attend to.
  557. key_dim : int
  558. Key dimension.
  559. pos_embedding_type : str
  560. Type of positional embedding to use.
  561. channel_attention_dropout : float
  562. Dropout rate for channel attention.
  563. Values greater than 1.0 means no channel attention.
  564. Values less than 0.0 means no dropout.
  565. within_channel_attention_dropout : float
  566. Dropout rate for within-channel attention.
  567. Values greater than 1.0 means no within-channel attention.
  568. Values less than 0.0 means no dropout.
  569. """
  570. def __init__(
  571. self,
  572. n_heads: int,
  573. n_channels: int,
  574. latent_sequence_length: int,
  575. n_patches: int,
  576. patch_length: int,
  577. unpatched_length: int,
  578. key_dim: int,
  579. pos_embedding_type: str,
  580. channel_attention_dropout: float,
  581. within_channel_attention_dropout: float,
  582. **kwargs,
  583. ):
  584. super().__init__(**kwargs)
  585. self.n_heads = n_heads
  586. self.n_channels = n_channels
  587. self.latent_sequence_length = latent_sequence_length
  588. self.key_dim = key_dim
  589. self.pos_embedding_type = pos_embedding_type
  590. self.n_patches = n_patches
  591. self.patch_length = patch_length
  592. self.unpatched_length = unpatched_length
  593. # Time attention layer
  594. self.time_attention_layer = TimeAttentionLayer(
  595. n_heads,
  596. latent_sequence_length,
  597. key_dim,
  598. n_patches,
  599. patch_length,
  600. unpatched_length,
  601. pos_embedding_type,
  602. )
  603. # Mask for time attention (This is fixed).
  604. self.time_attention_mask = self._compute_time_attention_mask()
  605. # Channel attention layer
  606. self.channel_attention_layer = ChannelAttention(key_dim)
  607. # Channel attention dropouts
  608. self.channel_attention_dropout = tf.Variable(
  609. channel_attention_dropout, trainable=False
  610. )
  611. self.within_channel_attention_dropout = tf.Variable(
  612. within_channel_attention_dropout, trainable=False
  613. )
  614. def _compute_time_attention_mask(self) -> tf.Tensor:
  615. """
  616. Compute the mask for time attention.
  617. Returns
  618. -------
  619. mask : tf.Tensor
  620. Mask for time attention.
  621. Shape: (latent_sequence_length, n_patches + unpatched_length).
  622. """
  623. mask = np.zeros(
  624. (self.latent_sequence_length, self.n_patches + self.unpatched_length)
  625. )
  626. # Patch masking
  627. for i in range(self.n_patches):
  628. m_indx = max(
  629. 0,
  630. (i + 1 - self.n_patches) * self.patch_length
  631. + self.latent_sequence_length
  632. - 1,
  633. )
  634. mask[:m_indx, i] = 1
  635. # Unpatched masking
  636. for i in range(self.n_patches, self.n_patches + self.unpatched_length):
  637. mask[
  638. : self.latent_sequence_length
  639. - (self.n_patches + self.unpatched_length)
  640. + i,
  641. i,
  642. ] = 1
  643. mask = tf.constant(mask, dtype=tf.float32)
  644. return mask
  645. def _compute_channel_attention_mask(self, training: bool) -> Union[tf.Tensor, None]:
  646. """
  647. Compute the mask for channel attention.
  648. Parameters
  649. ----------
  650. training : bool
  651. Whether the model is training.
  652. If False, no dropout is applied.
  653. Returns
  654. -------
  655. mask : tf.Tensor or None
  656. Mask for channel attention.
  657. Shape: (n_channels, n_channels).
  658. If None, full attention is applied.
  659. """
  660. if self.channel_attention_dropout < 1.0:
  661. # We apply channel attention dropout only during training
  662. if not training:
  663. return tf.zeros((self.n_channels, self.n_channels), dtype=tf.float32)
  664. else:
  665. uniform_sampler = tfp.distributions.Uniform()
  666. # Sample whether to apply channel attention
  667. if uniform_sampler.sample() < self.channel_attention_dropout:
  668. # Does not apply channel attention and mask all off-diagonal elements
  669. mask = 1 - np.eye(self.n_channels)
  670. return tf.constant(mask, dtype=tf.float32)
  671. # Sample whether to apply within-channel attention
  672. elif uniform_sampler.sample() < self.within_channel_attention_dropout:
  673. # Does not apply within-channel attention and mask all diagonal elements
  674. mask = np.eye(self.n_channels)
  675. return tf.constant(mask, dtype=tf.float32)
  676. # Apply channel attention and does not mask any elements, return None
  677. else:
  678. return tf.zeros(
  679. (self.n_channels, self.n_channels), dtype=tf.float32
  680. )
  681. else:
  682. # If channel_attention_dropout >= 1.0, no channel attention is applied
  683. # Mask all off-diagonal elements
  684. mask = 1 - np.eye(self.n_channels)
  685. return tf.constant(mask, dtype=tf.float32)
  686. def call(self, inputs, training=None, **kwargs):
  687. # ---------- Unpack Inputs ---------- #
  688. # q (Query): (batch_size, n_heads, latent_sequence_length, out_n_channels, key_dim)
  689. # k (Key): (batch_size, n_heads, n_patches + unpatched_length, out_n_channels, key_dim)
  690. # v (Value): (batch_size, n_heads, n_patches + unpatched_length, in_n_channels, key_dim)
  691. # c_q (Channel Query): (batch_size, n_heads, latent_sequence_length, out_n_channels, key_dim)
  692. # c_k (Channel Key): (batch_size, n_heads, latent_sequence_length, in_n_channels, key_dim)
  693. q, k, v, c_q, c_k = inputs
  694. # ---------- Channel Attention ---------- #
  695. # First sample channel attention mask
  696. channel_attention_mask = self._compute_channel_attention_mask(training=training)
  697. # Apply channel attention
  698. output = self.channel_attention_layer(
  699. [c_q, c_k, v], mask=channel_attention_mask, training=training, **kwargs
  700. )
  701. # ---------- Time Attention ---------- #
  702. output = self.time_attention_layer(
  703. [q, k, output], mask=self.time_attention_mask, training=training, **kwargs
  704. )
  705. return output
  706. class MultiHeadPASSTALayer(tf.keras.layers.Layer):
  707. """
  708. The Multi-head PASSTA layer.
  709. Parameters
  710. ----------
  711. n_heads : int
  712. Number of heads.
  713. model_dim : int
  714. Model dimension.
  715. n_channels : int
  716. Number of channels.
  717. sequence_length : int
  718. Sequence length.
  719. latent_sequence_length : int
  720. Latent sequence length.
  721. n_patches : int
  722. Number of patches to attend to.
  723. patch_length : int
  724. Patch length.
  725. unpatched_length : int
  726. Number of unpatched elements to attend to.
  727. pos_embedding_type : str
  728. Type of positional embedding to use.
  729. channel_attention_dropout : float
  730. Dropout rate for channel attention.
  731. Values greater than 1.0 means no channel attention.
  732. Values less than 0.0 means no dropout.
  733. within_channel_attention_dropout : float
  734. Dropout rate for within-channel attention.
  735. Values greater than 1.0 means no within-channel attention.
  736. Values less than 0.0 means no dropout
  737. """
  738. def __init__(
  739. self,
  740. n_heads: int,
  741. model_dim: int,
  742. n_channels: int,
  743. sequence_length: int,
  744. latent_sequence_length: int,
  745. n_patches: int,
  746. patch_length: int,
  747. unpatched_length: int,
  748. pos_embedding_type: str,
  749. channel_attention_dropout: float,
  750. within_channel_attention_dropout: float,
  751. **kwargs,
  752. ):
  753. super().__init__(**kwargs)
  754. self.n_heads = n_heads
  755. self.model_dim = model_dim
  756. self.key_dim = model_dim // n_heads
  757. self.n_channels = n_channels
  758. self.sequence_length = sequence_length
  759. self.latent_sequence_length = latent_sequence_length
  760. self.n_patches = n_patches
  761. self.patch_length = patch_length
  762. self.unpatched_length = unpatched_length
  763. self.pos_embedding_type = pos_embedding_type
  764. self.channel_attention_dropout = channel_attention_dropout
  765. self.within_channel_attention_dropout = within_channel_attention_dropout
  766. # Patch projection
  767. self.patch_projection = tf.keras.layers.Dense(n_heads)
  768. # Input projections
  769. self.time_patched_projection = tf.keras.layers.Dense(2 * self.model_dim)
  770. self.time_unpatched_projection = tf.keras.layers.Dense(2 * self.model_dim)
  771. self.time_query_projection = tf.keras.layers.Dense(self.model_dim)
  772. self.channel_projection = tf.keras.layers.Dense(2 * self.model_dim)
  773. # PASSTA layer for time and channel attention
  774. self.passta_layer = PASSTALayer(
  775. n_heads,
  776. n_channels,
  777. latent_sequence_length,
  778. n_patches,
  779. patch_length,
  780. unpatched_length,
  781. self.key_dim,
  782. pos_embedding_type,
  783. channel_attention_dropout,
  784. within_channel_attention_dropout,
  785. )
  786. # Output projection
  787. self.output_projection = tf.keras.layers.Dense(self.model_dim)
  788. def _patch_x(self, x: tf.Tensor) -> tf.Tensor:
  789. # x.shape: (batch_size, sequence_length, n_channels, model_dim)
  790. x = tf.transpose(x, perm=(0, 2, 3, 1))
  791. # x.shape: (batch_size, n_channels, model_dim, sequence_length)
  792. x = tf.reshape(
  793. x,
  794. (
  795. tf.shape(x)[0],
  796. self.n_channels,
  797. self.model_dim,
  798. self.n_patches,
  799. self.patch_length,
  800. ),
  801. )
  802. # x.shape: (batch_size, n_channels, model_dim, n_patches, patch_length)
  803. x = tf.transpose(x, perm=(0, 3, 1, 2, 4))
  804. # x.shape: (batch_size, n_patches, n_channels, model_dim, patch_length)
  805. return x
  806. def _perceiver_x(self, x: tf.Tensor) -> tf.Tensor:
  807. """Get the last latent_sequence_length elements of the sequence."""
  808. # x.shape: (batch_size, sequence_length, n_channels, model_dim)
  809. x = tf.slice(
  810. x,
  811. [0, self.sequence_length - self.latent_sequence_length, 0, 0],
  812. [-1, -1, -1, -1],
  813. )
  814. # x.shape: (batch_size, latent_sequence_length, n_channels, model_dim)
  815. return x
  816. def _unpatch_x(self, x: tf.Tensor) -> tf.Tensor:
  817. """Get the last unpatched_length elements of the sequence."""
  818. # x.shape: (batch_size, sequence_length, n_channels, model_dim)
  819. x = tf.slice(
  820. x,
  821. [0, self.sequence_length - self.unpatched_length, 0, 0],
  822. [-1, -1, -1, -1],
  823. )
  824. # x.shape: (batch_size, unpatched_length, n_channels, model_dim)
  825. return x
  826. def _split_heads(self, x: tf.Tensor) -> tf.Tensor:
  827. # x.shape: (batch_size, time_length, n_channels, model_dim)
  828. # Here time_length is either n_patches + unpatched_length or latent_sequence_length
  829. x = tf.reshape(
  830. x,
  831. (
  832. tf.shape(x)[0],
  833. tf.shape(x)[1],
  834. self.n_channels,
  835. self.n_heads,
  836. self.key_dim,
  837. ),
  838. )
  839. # x.shape: (batch_size, time_length, n_channels, n_heads, key_dim)
  840. x = tf.transpose(x, perm=(0, 3, 1, 2, 4))
  841. # x.shape: (batch_size, n_heads, time_length, n_channels, key_dim)
  842. return x
  843. def _combine_heads(self, x: tf.Tensor) -> tf.Tensor:
  844. # x.shape: (batch_size, n_heads, latent_sequence_length, n_channels, key_dim)
  845. x = tf.transpose(x, perm=(0, 2, 3, 1, 4))
  846. # x.shape: (batch_size, latent_sequence_length, n_channels, n_heads, key_dim)
  847. x = tf.reshape(
  848. x,
  849. (
  850. tf.shape(x)[0],
  851. self.latent_sequence_length,
  852. self.n_channels,
  853. self.model_dim,
  854. ),
  855. )
  856. # x.shape: (batch_size, latent_sequence_length, n_channels, model_dim)
  857. return x
  858. def call(self, inputs, training=None, **kwargs):
  859. x = inputs
  860. # x.shape: (batch_size, sequence_length, n_channels, model_dim)
  861. # ---------- Process inputs ---------- #
  862. # Input is processed into 3 parts:
  863. # 1. Patched input: patched_x
  864. # 2. Unpatched input: unpatched_x
  865. # 3. Perceiver input: perceiver_x
  866. patched_x = self._patch_x(x)
  867. # patched_x.shape: (batch_size, n_patches, n_channels, model_dim, patch_length)
  868. patched_x = tf.reshape(
  869. patched_x,
  870. (
  871. tf.shape(patched_x)[0],
  872. self.n_patches,
  873. self.n_channels,
  874. self.key_dim,
  875. self.n_heads,
  876. self.patch_length,
  877. ),
  878. )
  879. # patched_x.shape: (batch_size, n_patches, n_channels, key_dim, n_heads, patch_length)
  880. patched_x = tf.reshape(
  881. patched_x,
  882. (
  883. tf.shape(patched_x)[0],
  884. self.n_patches,
  885. self.n_channels,
  886. self.key_dim,
  887. self.n_heads * self.patch_length,
  888. ),
  889. )
  890. # patched_x.shape: (batch_size, n_patches, n_channels, key_dim, n_heads * patch_length)
  891. patched_x = self.patch_projection(patched_x)
  892. # patched_x.shape: (batch_size, n_patches, n_channels, key_dim, n_heads)
  893. patched_x = tf.reshape(
  894. patched_x,
  895. (
  896. tf.shape(patched_x)[0],
  897. self.n_patches,
  898. self.n_channels,
  899. self.model_dim,
  900. ),
  901. )
  902. # patched_x.shape: (batch_size, n_patches, n_channels, model_dim)
  903. perceiver_x = self._perceiver_x(x)
  904. # perceiver_x.shape: (batch_size, latent_sequence_length, n_channels, model_dim)
  905. unpatched_x = self._unpatch_x(x)
  906. # unpatched_x.shape: (batch_size, unpatched_length, n_channels, model_dim)
  907. # ---------- Project inputs to Q, K, V ---------- #
  908. q = self.time_query_projection(perceiver_x)
  909. # q.shape: (batch_size, latent_sequence_length, n_channels, model_dim)
  910. k_patched, v_patched = tf.split(
  911. self.time_patched_projection(patched_x), 2, axis=-1
  912. )
  913. # k_patched.shape: (batch_size, n_patches, n_channels, model_dim)
  914. # v_patched.shape: (batch_size, n_patches, n_channels, model_dim)
  915. k_unpatched, v_unpatched = tf.split(
  916. self.time_unpatched_projection(unpatched_x), 2, axis=-1
  917. )
  918. # k_unpatched.shape: (batch_size, unpatched_length, n_channels, model_dim)
  919. # v_unpatched.shape: (batch_size, unpatched_length, n_channels, model_dim)
  920. # Concatenate k and v for time attention
  921. k = tf.concat([k_patched, k_unpatched], axis=1)
  922. v = tf.concat([v_patched, v_unpatched], axis=1)
  923. # k.shape: (batch_size, n_patches + unpatched_length, n_channels, model_dim)
  924. # v.shape: (batch_size, n_patches + unpatched_length, n_channels, model_dim)
  925. c_q, c_k = tf.split(self.channel_projection(perceiver_x), 2, axis=-1)
  926. # c_q.shape: (batch_size, latent_sequence_length, n_channels, model_dim)
  927. # c_k.shape: (batch_size, latent_sequence_length, n_channels, model_dim)
  928. # ---------- Split heads ---------- #
  929. q = self._split_heads(q)
  930. # (batch_size, n_heads, latent_sequence_length, n_channels, key_dim)
  931. k = self._split_heads(k)
  932. # (batch_size, n_heads, n_patches + unpatched_length, n_channels, key_dim)
  933. v = self._split_heads(v)
  934. # (batch_size, n_heads, n_patches + unpatched_length, n_channels, key_dim)
  935. c_q = self._split_heads(c_q)
  936. # (batch_size, n_heads, latent_sequence_length, n_channels, key_dim)
  937. c_k = self._split_heads(c_k)
  938. # (batch_size, n_heads, latent_sequence_length, n_channels, key_dim)
  939. # ---------- PASSTA Layer ---------- #
  940. output = self.passta_layer([q, k, v, c_q, c_k], training=training, **kwargs)
  941. # output.shape: (batch_size, n_heads, latent_sequence_length, n_channels, key_dim)
  942. # ---------- Combine heads ---------- #
  943. output = self._combine_heads(output)
  944. # output.shape: (batch_size, latent_sequence_length, n_channels, model_dim)
  945. # ---------- Output projection ---------- #
  946. output = self.output_projection(output)
  947. # output.shape: (batch_size, latent_sequence_length, n_channels, model_dim)
  948. return output

layers.py at commit ffbeff3, under MIT · at the source

Overview

Authors: Rukuang Huang1,2, SungJun Cho1,3, Chetan Gohil1,2, Oiwi Parker Jones1,4, Mark Woolrich1,2
ORCID iDs: Chetan Gohil
  1. Oxford Centre for Integrative Neuroimaging (OxCIN), University of Oxford, Oxford, United Kingdom
  2. Department of Psychiatry, University of Oxford, Oxford, United Kingdom
  3. Nuffield Department of Clinical Neurosciences, University of Oxford, Oxford, United Kingdom
  4. Department of Engineering Science, University of Oxford, Oxford, United Kingdom
Institutions: University of Oxford (United Kingdom); Wellcome Centre for Integrative Neuroimaging (United Kingdom)
Journal: Imaging neuroscience (Cambridge, Mass.), volume 4, article IMAG.a.1301
Dates: received 10 January 2026; accepted 16 June 2026; published online 24 July 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1162/imag.a.1301 · PMID 42516171 · PMCID PMC13403652 · OpenAlex W4416055129
Open access: diamond, a free copy (OpenAlex)
Status: code verified
Categories: MEG (modality), methods / tools (subfield)
Methods: Spectral & time-frequency, Preprocessing, Connectivity, Smoothing, state filtering, decompositions, Machine learning
Keywords: electrophysiology, MEG, GPT, transformer, foundation model, tokenisation
Topic: EEG and Brain-Computer Interfaces (Cognitive Neuroscience, Neuroscience), according to OpenAlex
Funding: Wellcome Trust (203139/Z/16/Z, 106183/Z/14/Z, 203139/A/16/Z, 215573/Z/19/Z); National Institute for Health Research (NIHR) (NIHR203316)
Citations: not cited yet (Europe PMC); 71 references in the paper

Abstract

Modelling the complex spatio-temporal patterns of large-scale brain dynamics is crucial for neuroscience, but traditional methods fail to capture the rich structure in modalities such as magnetoencephalography (MEG). Recent advances in deep learning have enabled significant progress in other domains, such as language and vision, by using foundation models at scale. Here, we introduce MEG-GPT, a transformer-based foundation model that uses time-attention and next time-point prediction. To facilitate this, we also introduce a novel data-driven tokeniser for continuous MEG data, which preserves the high temporal resolution of continuous MEG signals without lossy transformations. We trained MEG-GPT on tokenised brain region time courses extracted from a large-scale MEG dataset (N = 612, eyes-closed rest, Cam-CAN data), and show that the learnt model can generate data with realistic spatio-spectral properties, including transient events and population variability. Critically, it performs well in downstream decoding tasks, improving downstream supervised prediction task, showing improved zero-shot generalisation across sessions (improving accuracy from 0.56 to 0.59) and subjects (improving accuracy from 0.45 to 0.49) compared with a PCA baseline method. Furthermore, we show the model can be efficiently fine-tuned on a smaller labelled dataset to boost performance in cross-subject decoding scenarios. This work establishes a powerful foundation model for electrophysiological data, paving the way for applications in computational neuroscience and neural decoding.

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 13 matches between paragraphs and lines of code.

OHBA-analysis/osl-foundation

License: MIT
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: ffbeff318811de3adea33d7dfeaa0af15a5b902b, 8 July 2026
Languages: Python (80)
Size: 103 files, 80 scripts
Software Heritage: not archived
Found in: “Data and Code Availability”
Holds: README, license file, environment (pyproject.toml)
Not found: CITATION.cff, tests, continuous integration, documentation
Tools: NumPy (45 files), Matplotlib (17 files), TensorFlow (15 files), pandas (6 files), seaborn (6 files), scikit-learn (4 files), MNE-Python (3 files), SciPy (2 files), NiBabel (1 file), Nilearn (1 file), UMAP (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
82 files

Zenodo 11099418

License: CC-BY-4.0
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Size: 9 files
Software Heritage: not checked
Found in: the references
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Availability: 1 check, the latest on 27 September 2026: the link answers (HTTP 200)
  • 27 September 2026: the link answers (HTTP 200)

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;
  • 80 scripts, each with its path and the digest of its content;
  • 13 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 and Code Availability

Data used are publicly available. For the Wakeman–Henson dataset, we refer the readers to the original paper (Wakeman & Henson, 2015). For the Cam-CAN dataset, we refer the readers to the original paper (Taylor et al., 2017).

Source code and scripts for reproducing results in the paper using TensorFlow are available on GitHub: https://github.com/OHBA-analysis/osl-foundation. Example code and tutorials for training a foundation model and applying it to new data are also provided.

A PyTorch implemention of the tokenizer is available here: https://github.com/OHBA-analysis/EphysTokenizer, and a PyTorch implementation of MEG-GPT is available here: https://github.com/OHBA-analysis/MEG-GPT.

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, 5 authors, 6 keywords, 2 funders, 64 references.

Cite

This paper

Huang, R., Cho, S., Gohil, C., Jones, O. P., & Woolrich, M. (2026). MEG-GPT: A transformer-based foundation model for magnetoencephalography data. Imaging neuroscience (Cambridge, Mass.), 4, IMAG.a.1301. https://doi.org/10.1162/imag.a.1301

BibTeX

@article{huang2026meg,
author = {Huang, Rukuang and Cho, SungJun and Gohil, Chetan and Jones, Oiwi Parker and Woolrich, Mark},
title = {{MEG-GPT: A transformer-based foundation model for magnetoencephalography data}},
journal = {Imaging neuroscience (Cambridge, Mass.)},
year = {2026},
month = jul,
volume = {4},
pages = {IMAG.a.1301},
publisher = {MIT Press},
issn = {2837-6056},
doi = {10.1162/imag.a.1301},
url = {https://doi.org/10.1162/imag.a.1301},
pmid = {42516171},
pmcid = {PMC13403652}
}

RIS

TY - JOUR
AU - Huang, Rukuang
AU - Cho, SungJun
AU - Gohil, Chetan
AU - Jones, Oiwi Parker
AU - Woolrich, Mark
TI - MEG-GPT: A transformer-based foundation model for magnetoencephalography data
T2 - Imaging neuroscience (Cambridge, Mass.)
J2 - Imaging Neurosci (Camb)
PY - 2026
DA - 2026/07/24
VL - 4
SP - IMAG.a.1301
SN - 2837-6056
PB - MIT Press
DO - 10.1162/imag.a.1301
UR - https://doi.org/10.1162/imag.a.1301
LA - en
ER -

CSL-JSON

{
"id": "10.1162/imag.a.1301",
"type": "article-journal",
"title": "MEG-GPT: A transformer-based foundation model for magnetoencephalography data",
"container-title": "Imaging neuroscience (Cambridge, Mass.)",
"author": [
{
"family": "Huang",
"given": "Rukuang"
},
{
"family": "Cho",
"given": "SungJun"
},
{
"family": "Gohil",
"given": "Chetan"
},
{
"family": "Jones",
"given": "Oiwi Parker"
},
{
"family": "Woolrich",
"given": "Mark"
}
],
"container-title-short": "Imaging Neurosci (Camb)",
"volume": "4",
"page": "IMAG.a.1301",
"DOI": "10.1162/imag.a.1301",
"PMID": "42516171",
"PMCID": "PMC13403652",
"ISSN": "2837-6056",
"publisher": "MIT Press",
"URL": "https://doi.org/10.1162/imag.a.1301",
"language": "en",
"issued": {
"date-parts": [
[
2026,
7,
24
]
]
}
}

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.1162/imag.a.1188 [code]
Modelling variability in functional brain networks using embeddings.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: NiBabel, seaborn, scikit-learn, 4 other tools, 15 references
[2] doi:10.1162/imag.a.1190 [code]
Canonical Hidden Markov Model Networks for studying M/EEG.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: MNE-Python, Nilearn, NiBabel, 5 other tools, MEG, 9 references, author Chetan Gohil
[3] doi:10.1162/imag.a.1237 [code]
Modelling discrete states and long-term dynamics in functional brain networks.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: MNE-Python, TensorFlow, seaborn, 5 other tools, MEG, 8 references, author Chetan Gohil
[4] doi:10.1002/hbm.70516 [code]
Effects of Age on Resting-State Cortical Networks.
Journal: Human brain mapping
In common: MNE-Python, Nilearn, NiBabel, 5 other tools, MEG, 8 references, author Chetan Gohil
[5] doi:10.1038/s41531-026-01372-1 [code]
Varying patterns of association between cortical large-scale networks and subthalamic nucleus activity in Parkinson's disease.
Journal: NPJ Parkinson's disease
In common: MNE-Python, Nilearn, NiBabel, 6 other tools, 8 references
[6] doi:10.1162/imag.a.1269 [code]
From early to contemporary normative modeling: Mapping individual differences in neurophysiological signals.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: MNE-Python, Nilearn, NiBabel, 6 other tools, MEG, 3 references
[7] doi:10.1093/braincomms/fcag236 [code]
Dynamic, state-dependent characteristics of cognitive fluctuations in Lewy body dementia: a magnetoencephalography study.
Journal: Brain communications
In common: MNE-Python, pandas, NumPy, MEG, 7 references
[8] doi:10.7554/elife.100605 [code]
Age-related changes in ‘cortical’ 1/f dynamics are linked to cardiac activity
Journal: —
In common: MNE-Python, seaborn, scikit-learn, 4 other tools, MEG, 5 references
[9] doi:10.1038/s41597-026-07350-9 [code]
An open multi-center MEG-EEG dataset for studying conscious visual perception.
Journal: Scientific data
In common: MNE-Python, Nilearn, NiBabel, 6 other tools, MEG, 2 references
[10] doi:10.1093/nc/niag029 [code]
A data-driven approach to identifying and evaluating connectivity-based neural correlates of conscious visual perception.
Journal: Neuroscience of consciousness
In common: MNE-Python, Nilearn, NiBabel, 6 other tools, MEG, 1 reference

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.

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.