MEG-GPT: A transformer-based foundation model for magnetoencephalography data.
The 13 matches
- [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] § 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] § 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] § 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] § 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] § 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] § 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] § 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] § 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] § 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] § 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] § 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] § 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
- from typing import Union
- import numpy as np
- import tensorflow as tf
- import tensorflow_probability as tfp
- def rnn_layer(
- rnn_type: str, rnn_n_units: int, return_sequences: bool = True
- ) -> tf.keras.layers.Layer:
- """
- Create an RNN layer.
- Parameters
- ----------
- rnn_type : str
- Type of RNN layer. Options are 'gru', 'lstm'.
- rnn_n_units : int
- Number of units in the RNN layer.
- return_sequences : bool, optional
- Whether to return sequences.
- Returns
- -------
- rnn_layer : tf.keras.layers.Layer
- RNN layer.
- """
- if rnn_type == "gru":
- return tf.keras.layers.GRU(
- rnn_n_units, return_sequences=return_sequences, stateful=False
- )
- elif rnn_type == "lstm":
- return tf.keras.layers.LSTM(
- rnn_n_units, return_sequences=return_sequences, stateful=False
- )
- else:
- raise ValueError(f"Unknown RNN type: {rnn_type}")
- class IdentityLayer(tf.keras.layers.Layer):
- """
- Identity layer.
- This layer directly returns the input tensor.
- """
- def call(self, inputs, **kwargs):
- return inputs
- class MSELossLayer(tf.keras.layers.Layer):
- """
- Layer for computing the mean squared error loss.
- This is a wrapper around tf.keras.losses.MeanSquaredError.
- """
- def __init__(self, **kwargs):
- super().__init__(**kwargs)
- self.loss_fn = lambda y_true, y_pred: tf.reduce_mean(tf.square(y_true - y_pred))
- def call(self, y_true, y_pred, **kwargs):
- loss = self.loss_fn(y_true, y_pred)
- self.add_loss(loss)
- return tf.expand_dims(loss, -1)
- class SinusoidalPositionalEncodingLayer(tf.keras.layers.Layer):
- """
- Layer for generating sinusoidal positional encoding for position embeddings.
- Implemented as in "Attention is All You Need" (Vaswani et al., 2017), and partly
- adpated from the `huggingface/transformers` library and TensorFlow
- (https://www.tensorflow.org/text/tutorials/transformer).
- Parameters
- ----------
- sequence_length : int
- Length of the sequence.
- """
- def __init__(self, sequence_length: int, **kwargs):
- super().__init__(**kwargs)
- self.sequence_length = sequence_length
- def call(self, inputs, **kwargs):
- # Validation
- embedding_dim = tf.shape(inputs)[-1]
- tf.debugging.assert_equal(
- tf.math.floormod(embedding_dim, 2),
- tf.constant(0, dtype=tf.int32),
- message="embedding dimension must be even for positional encoding.",
- )
- # Precompute scaling factor
- embedding_dim_float = tf.cast(embedding_dim, tf.float32)
- denominator = tf.math.pow(
- 10000.0,
- 2 * (tf.range(embedding_dim_float) // 2) / embedding_dim_float,
- )
- # Get position indices
- position_indices = tf.range(self.sequence_length, dtype=tf.float32)[
- :, tf.newaxis
- ]
- # position_indices.shape = (sequence_length, 1)
- # Compute positional encoding
- angle_rads = position_indices / denominator
- # angle_rads.shape = (sequence_length, embedding_dim)
- pos_encoding = tf.TensorArray(dtype=tf.float32, size=embedding_dim)
- for i in range(embedding_dim):
- if i % 2 == 0:
- pos_encoding = pos_encoding.write(i, tf.sin(angle_rads[:, i]))
- else:
- pos_encoding = pos_encoding.write(i, tf.cos(angle_rads[:, i]))
- pos_encoding = tf.transpose(pos_encoding.stack())
- # pos_encoding.shape = (sequence_length, embedding_dim)
- # Reshape the positions to match shapes of other embeddings
- positions = tf.expand_dims(
- tf.expand_dims(pos_encoding, axis=0), axis=0
- ) # positions.shape = (1, 1, sequence_length, embedding_dim)
- return positions
- class RotaryPositionEmbeddingLayer(tf.keras.layers.Layer):
- """
- Layer for generating Rotary Position Embedding (RoPE). Implemented as in
- "RoFormer: Enhanced Transformer with Rotary Position Embedding" (Su et al., 2022), and
- adpated from the PyTorch and Keras 3 libraries (see below for the reference).
- References:
- - https://pytorch.org/torchtune/0.2/_modules/torchtune/modules/position_embeddings.html#RotaryPositionalEmbeddings
- - https://github.com/keras-team/keras-hub/blob/v0.17.0/keras_hub/src/layers/modeling/rotary_embedding.py
- Parameters
- ----------
- embedding_dim : int
- Dimension of the rotary position embedding. When using multi-head attention blocks,
- the dimension is usually `embedding_dim // n_heads`.
- max_sequence_length : int
- Maximum expected length of the sequence. If the model dimensions of the key and
- query vectors differ, the sequence length should be the maximum of the two.
- base : int, optional
- The base value for computing the rotation angles.
- """
- def __init__(
- self, embedding_dim: int, max_sequence_length: int, base: int = 10000, **kwargs
- ):
- super().__init__(**kwargs)
- # Validation
- if embedding_dim % 2 != 0:
- raise ValueError(
- "embedding dimension must be even for positional encoding."
- )
- self.embedding_dim = embedding_dim
- self.max_sequence_length = max_sequence_length
- self.base = base
- self.rotation_values = self._compute_rotation_matrix()
- def _compute_rotation_matrix(self) -> tf.Tensor:
- # Compute pre-defined theta angles
- theta = 1.0 / (
- self.base
- ** (
- tf.range(0, self.embedding_dim, 2, dtype=tf.float32)[
- : (self.embedding_dim // 2)
- ]
- / self.embedding_dim
- )
- )
- # Create position indices
- seq_idx = tf.range(self.max_sequence_length, dtype=tf.float32)
- # NOTE: Having 0 as a starting index is okay, since we don't have to rotate the first position.
- # Compute outer product of position index and theta
- idx_theta = tf.einsum("i,j->ij", seq_idx, theta)
- # idx_theta.shape = (sequence_length, embedding_dim // 2)
- # Compute elements of the rotation matrix
- rotation_values = tf.stack([tf.cos(idx_theta), tf.sin(idx_theta)], axis=-1)
- # rotation_values.shape = (sequence_length, embedding_dim // 2, 2)
- return rotation_values
- def call(self, inputs, **kwargs):
- # Get input sequence length
- sequence_length = tf.shape(inputs)[3]
- # inputs.shape = (batch_size, n_heads, n_channels, sequence_length, key_dim)
- # Validation
- tf.debugging.assert_greater_equal(
- self.max_sequence_length,
- sequence_length,
- message="input sequence length should not exceed max_sequence_length.",
- )
- tf.debugging.assert_equal(
- tf.shape(inputs)[-1],
- self.embedding_dim,
- message="the model dimension of the input tensor must match the embedding dimension.",
- )
- # Reshape the rotation values to be compabitle with the input shape
- rot_vals = tf.reshape(
- self.rotation_values[:sequence_length],
- (1, 1, 1, sequence_length, self.embedding_dim // 2, 2),
- )
- # NOTE: We clip the rotation values to the sequence length, in case sequence_length is
- # less than max_sequence_length as the dimension for key and query differs.
- # Match the input shape to the rope shape
- rot_shape = tf.concat(
- [tf.shape(inputs)[:-1], [self.embedding_dim // 2, 2]], axis=0
- )
- x_reshaped = tf.reshape(inputs, rot_shape)
- # x_reshaped.shape = (batch_size, n_heads, n_channels, sequence_length, key_dim // 2, 2)
- # Calculate the rotary position embedding
- rope = tf.stack(
- [
- rot_vals[..., 0] * x_reshaped[..., 0]
- - rot_vals[..., 1] * x_reshaped[..., 1],
- rot_vals[..., 1] * x_reshaped[..., 0]
- + rot_vals[..., 0] * x_reshaped[..., 1],
- ],
- axis=-1,
- )
- # rope.shape = (batch_size, n_heads, n_channels, sequence_length, key_dim // 2, 2)
- rope = tf.reshape(rope, tf.shape(inputs))
- # rope.shape = (batch_size, n_heads, n_channels, sequence_length, key_dim)
- return rope
- class ALiBiPositionEmbeddingLayer(tf.keras.layers.Layer):
- """
- Layer for generating Attention with Linear Biases (ALiBi) position embedding. Implemented
- based on "Train Short, Test Long: Attention with Linear Biases Enables Input Length Extrapolation"
- (Press et al., 2022) and adapted from the references below.
- References:
- - https://github.com/ofirpress/attention_with_linear_biases
- - https://nn.labml.ai/transformers/alibi/index.html
- Parameters
- ----------
- n_heads : int
- Number of attention heads.
- n_patches : int
- Number of patches to attend to.
- patch_length : int
- Patch length.
- unpatched_length : int
- Number of unpatched elements to attend to.
- """
- def __init__(
- self,
- n_heads: int,
- n_patches: int,
- patch_length: int,
- unpatched_length: int,
- **kwargs,
- ):
- super().__init__(**kwargs)
- self.n_heads = n_heads
- self.n_patches = n_patches
- self.patch_length = patch_length
- self.unpatched_length = unpatched_length
- def _get_slopes(self) -> tf.Tensor:
- # Compute the slopes for the linear biases (static, non-learned)
- h = tf.cast(self.n_heads, tf.float32)
- n = 2 ** tf.math.floor(tf.math.log(h) / tf.math.log(2.0)) # nearest power of 2
- ratio = tf.pow(2.0, -8.0 / n)
- m = tf.pow(
- ratio, tf.range(1, tf.cast(n, tf.int32) + 1, dtype=tf.float32)
- ) # head-specific slopes
- # When n_heads is not a power of 2, then we add additional slopes
- if n < h:
- ratio_hat = tf.pow(2.0, -4.0 / n)
- m_hat = tf.pow(ratio_hat, tf.range(1, 1 + 2 * (h - n), 2, dtype=tf.float32))
- m = tf.concat([m, m_hat], axis=0)
- return m
- def _get_relative_bias_matrix(self, mask: tf.Tensor) -> tf.Tensor:
- # Get head-specific slopes
- m = self._get_slopes()
- # Get a binary masking matrix
- mask = 1 - mask
- cum_mask = tf.cumsum(mask, axis=0) # cumulative sum of the mask
- # NOTE: In the mask matrix, columns are where you attend to and rows are the targets.
- # Compute the relative distance matrix for the patches
- patch_indices = tf.cast(tf.range(self.n_patches), dtype=tf.float32)
- avg_idx = (
- (self.patch_length * patch_indices + 1)
- + self.patch_length * (patch_indices + 1)
- ) / 2.0
- distance_patched = (
- -(cum_mask[:, : self.n_patches] + self.patch_length * (patch_indices + 1))
- + avg_idx
- )
- # NOTE: Considering a patch as a single receptive field, we calculate relative distance between the patch
- # and the current sequence index (i.e., the distance between the center of the patch and the current index).
- # This amounts to the average distance between the current index and all the indices in the patch.
- # Compute the relative distance matrix for the unpatched elements
- distance_unpatched = -cum_mask[
- :, self.n_patches : self.n_patches + self.unpatched_length
- ]
- # Concatenate two matrices
- distance = tf.concat([distance_patched, distance_unpatched], axis=1)
- # distance.shape = (latent_sequence_length, n_patches + unpatched_length)
- # Compute head-specific relative bias matrices
- masked_distance = distance * mask
- alibi_matrix = masked_distance[None, :, :] * m[:, None, None]
- # alibi_matrix.shape = (n_heads, latent_sequence_length, n_patches + unpatched_length)
- return alibi_matrix
- def call(self, inputs, mask: tf.Tensor, **kwargs):
- # Get the relative bias matrix
- alibi_matrix = self._get_relative_bias_matrix(mask)
- # alibi_matrix.shape = (n_heads, latent_sequence_length, n_patches + unpatched_length)
- # Reshape the relative bias matrix to match the attention matrix
- alibi_matrix = tf.reshape(
- alibi_matrix, (1, self.n_heads, 1, tf.shape(mask)[0], tf.shape(mask)[1])
- )
- # Add the relative bias matrix to the attention matrix
- attention_with_alibi = inputs + alibi_matrix
- # attention_with_alibi.shape = (batch_size, n_heads, n_channels, out_sequence_length, in_sequence_length)
- return attention_with_alibi
- class TokenWeightsLayer(tf.keras.layers.Layer):
- """
- Layer for computing token weights.
- Parameters
- ----------
- output_dim : int
- Dimension of the output.
- activation : str, optional
- Activation function to use.
- """
- def __init__(self, output_dim: int, activation: str = "linear", **kwargs):
- super().__init__(**kwargs)
- self.output_dim = output_dim
- self.dense_layer = tf.keras.layers.Dense(output_dim, activation=activation)
- self.activation_layer = tf.keras.layers.Activation(activation)
- self.norm_layer = tf.keras.layers.LayerNormalization()
- self.temperature = tf.Variable(0.0, trainable=False)
- def call(self, inputs, training=None, **kwargs):
- ell = self.activation_layer(self.dense_layer(inputs))
- ell = self.norm_layer(ell) / 0.1
- # Shape: (batch_size * n_channels, sequence_length, n_tokens)
- if training:
- # Sample from gumbel softmax parameterized by ell
- theta_sample = tf.argmax(ell, axis=2)
- theta_sample = tf.one_hot(theta_sample, self.output_dim)
- # Annealing
- theta_weight = tf.nn.softmax(ell, axis=2)
- # Shape: (batch_size * n_channels, sequence_length, n_tokens)
- token_weight = (
- self.temperature * theta_weight + (1 - self.temperature) * theta_sample
- )
- # Shape: (batch_size * n_channels, sequence_length, n_tokens)
- else:
- token_weight = tf.one_hot(tf.argmax(ell, axis=2), self.output_dim)
- return token_weight
- class PositionEmbedding(tf.keras.layers.Layer):
- """
- Layer for learning position embeddings.
- Parameters
- ----------
- sequence_length : int
- Sequence length.
- initializer : str, optional
- Initializer for the position embeddings.
- """
- def __init__(
- self,
- sequence_length: int,
- initializer: str = "glorot_uniform",
- **kwargs,
- ):
- super().__init__(**kwargs)
- self.sequence_length = sequence_length
- self.initializer = tf.keras.initializers.get(initializer)
- def build(self, inputs_shape):
- feature_size = inputs_shape[-1]
- self.position_embeddings = self.add_weight(
- name="embeddings",
- shape=[self.sequence_length, feature_size],
- initializer=self.initializer,
- trainable=True,
- )
- self.built = True
- def call(self, inputs, start_index=0):
- inputs_shape = tf.shape(inputs)
- feature_length = inputs_shape[-1]
- sequence_length = inputs_shape[-2]
- # trim to match the length of the input sequence, which might be less
- # than the sequence_length of the layer.
- position_embeddings = tf.convert_to_tensor(self.position_embeddings)
- position_embeddings = tf.slice(
- position_embeddings,
- (start_index, 0),
- (sequence_length, feature_length),
- )
- return tf.broadcast_to(position_embeddings, inputs_shape)
- class NormalizationLayer(tf.keras.layers.Layer):
- """
- Layer for performing normalization.
- Parameters
- ----------
- norm_type : str
- Type of normalization to perform.
- Options: "layer", "batch", "group".
- n_groups : int, optional
- Number of groups for group normalization.
- Required if norm_type is "group".
- """
- def __init__(self, norm_type: str = "layer", n_groups=None, **kwargs):
- super().__init__(**kwargs)
- if norm_type == "layer":
- self.norm_layer = tf.keras.layers.LayerNormalization()
- elif norm_type == "batch":
- self.norm_layer = tf.keras.layers.BatchNormalization()
- elif norm_type == "group":
- if n_groups is None:
- raise ValueError("n_groups must be specified for group normalization")
- self.norm_layer = tf.keras.layers.GroupNormalization(groups=n_groups)
- else:
- raise ValueError(f"Unknown normalization type: {norm_type}")
- def call(self, inputs, **kwargs):
- return self.norm_layer(inputs)
- class TimeAttentionLayer(tf.keras.layers.Layer):
- """
- Layer for performing time attention.
- Parameters
- ----------
- n_heads : int
- Number of heads.
- latent_sequence_length : int
- Sequence length of latent space.
- key_dim : int
- Key dimension.
- n_patches : int
- Number of patches to attend to.
- patch_length : int
- Patch length.
- unpatched_length : int
- Number of unpatched elements to attend to.
- transform_attention : str, optional
- Type of attention transformation to apply.
- """
- def __init__(
- self,
- n_heads: int,
- latent_sequence_length: int,
- key_dim: int,
- n_patches: int,
- patch_length: int,
- unpatched_length: int,
- transform_attention: str = None,
- **kwargs,
- ):
- super().__init__(**kwargs)
- self.n_heads = n_heads
- self.latent_sequence_length = latent_sequence_length
- self.key_dim = key_dim
- self.n_patches = n_patches
- self.patch_length = patch_length
- self.unpatched_length = unpatched_length
- self.transform_attention = transform_attention
- # Create a position embedding layer
- if transform_attention == "rope":
- self.rotary_position_embedding_layer = RotaryPositionEmbeddingLayer(
- embedding_dim=self.key_dim,
- max_sequence_length=tf.cast(
- tf.math.maximum(
- self.latent_sequence_length,
- self.n_patches + self.unpatched_length,
- ),
- dtype=tf.int32,
- ),
- )
- if transform_attention == "alibi":
- self.alibi_position_embedding_layer = ALiBiPositionEmbeddingLayer(
- n_heads=self.n_heads,
- n_patches=self.n_patches,
- patch_length=self.patch_length,
- unpatched_length=self.unpatched_length,
- )
- def call(self, inputs, mask=None, **kwargs):
- # q (Query): (batch_size, n_heads, out_sequence_length, out_n_channels, key_dim)
- # k (Key): (batch_size, n_heads, in_sequence_length, out_n_channels, key_dim)
- # v (Value): (batch_size, n_heads, out_sequence_length, out_n_channels, in_sequence_length, key_dim)
- q, k, v = inputs
- # Transpose inputs for time attention
- q = tf.transpose(q, perm=(0, 1, 3, 2, 4))
- # q: (batch_size, n_heads, out_n_channels, out_sequence_length, key_dim)
- k = tf.transpose(k, perm=(0, 1, 3, 2, 4))
- # k: (batch_size, n_heads, out_n_channels, in_sequence_length, key_dim)
- if self.transform_attention == "rope":
- # Apply rotary position embedding to the query and key vectors
- q = self.rotary_position_embedding_layer(q)
- k = self.rotary_position_embedding_layer(k)
- # Compute attention
- attention = tf.matmul(q, k, transpose_b=True) / tf.math.sqrt(
- tf.cast(self.key_dim, tf.float32)
- )
- # attention: (batch_size, n_heads, out_n_channels, out_sequence_length, in_sequence_length)
- if self.transform_attention == "alibi":
- # Apply ALiBi position embedding to the attention matrix
- attention = self.alibi_position_embedding_layer(attention, mask=mask)
- # Apply mask
- if mask is not None:
- attention += -1e9 * mask
- # Normalise attention with softmax
- attention = tf.nn.softmax(attention, axis=-1)
- if mask is not None:
- attention = attention * (1 - mask)
- # Apply attention to value
- attention = tf.expand_dims(
- tf.transpose(attention, perm=(0, 1, 3, 2, 4)), axis=-2
- )
- # attention: (batch_size, n_heads, out_sequence_length, out_n_channels, 1, in_sequence_length)
- output = tf.matmul(attention, v)
- # output: (batch_size, n_heads, out_sequence_length, out_n_channels, 1, key_dim)
- output = tf.squeeze(output, axis=-2)
- # output: (batch_size, n_heads, out_sequence_length, out_n_channels, key_dim)
- return output
- class ChannelAttention(tf.keras.layers.Layer):
- """
- Layer for performing channel attention.
- Parameters
- ----------
- key_dim : int
- Key dimension.
- """
- def __init__(self, key_dim: int, **kwargs):
- super().__init__(**kwargs)
- self.key_dim = key_dim
- def call(self, inputs, mask=None, **kwargs):
- # q (Query): (batch_size, n_heads, out_sequence_length, out_n_channels, key_dim)
- # k (Key): (batch_size, n_heads, out_sequence_length, in_n_channels, key_dim)
- # v (Value): (batch_size, n_heads, in_sequence_length, in_n_channels, key_dim)
- q, k, v = inputs
- batch_size = tf.shape(q)[0]
- n_heads = tf.shape(q)[1]
- out_sequence_length = tf.shape(q)[2]
- out_n_channels = tf.shape(q)[3]
- in_n_channels = tf.shape(k)[3]
- in_sequence_length = tf.shape(v)[2]
- # Compute attention
- attention = tf.matmul(q, k, transpose_b=True) / tf.math.sqrt(
- tf.cast(self.key_dim, tf.float32)
- )
- # attention: (batch_size, n_heads, out_sequence_length, out_n_channels, in_n_channels)
- # Apply mask
- if mask is not None:
- attention += -1e9 * mask
- # Normalise attention with softmax
- attention = tf.nn.softmax(attention, axis=-1)
- if mask is not None:
- attention = attention * (1 - mask)
- # Apply attention to value
- v = tf.transpose(v, perm=(0, 1, 3, 2, 4))
- # v: (batch_size, n_heads, in_n_channels, in_sequence_length, key_dim)
- v = tf.reshape(v, shape=(batch_size, n_heads, 1, in_n_channels, -1))
- # v: (batch_size, n_heads, 1, in_n_channels, in_sequence_length * key_dim)
- output = tf.matmul(attention, v)
- # output: (batch_size, n_heads, out_sequence_length, out_n_channels, in_sequence_length * key_dim)
- output = tf.reshape(
- output,
- shape=(
- batch_size,
- n_heads,
- out_sequence_length,
- out_n_channels,
- in_sequence_length,
- self.key_dim,
- ),
- )
- # output: (batch_size, n_heads, out_sequence_length, out_n_channels, in_sequence_length, key_dim)
- return output
- class PASSTALayer(tf.keras.layers.Layer):
- """
- The Perceiver AR Separable Space-Time self-Attention (PASSTA) layer.
- This layer performs space-time attention on the input tensor.
- Parameters
- ----------
- n_heads : int
- Number of heads.
- n_channels : int
- Number of channels.
- latent_sequence_length : int
- Sequence length of latent space.
- n_patches : int
- Number of patches to attend to.
- patch_length : int
- Patch length.
- unpatched_length : int
- Number of unpatched elements to attend to.
- key_dim : int
- Key dimension.
- pos_embedding_type : str
- Type of positional embedding to use.
- channel_attention_dropout : float
- Dropout rate for channel attention.
- Values greater than 1.0 means no channel attention.
- Values less than 0.0 means no dropout.
- within_channel_attention_dropout : float
- Dropout rate for within-channel attention.
- Values greater than 1.0 means no within-channel attention.
- Values less than 0.0 means no dropout.
- """
- def __init__(
- self,
- n_heads: int,
- n_channels: int,
- latent_sequence_length: int,
- n_patches: int,
- patch_length: int,
- unpatched_length: int,
- key_dim: int,
- pos_embedding_type: str,
- channel_attention_dropout: float,
- within_channel_attention_dropout: float,
- **kwargs,
- ):
- super().__init__(**kwargs)
- self.n_heads = n_heads
- self.n_channels = n_channels
- self.latent_sequence_length = latent_sequence_length
- self.key_dim = key_dim
- self.pos_embedding_type = pos_embedding_type
- self.n_patches = n_patches
- self.patch_length = patch_length
- self.unpatched_length = unpatched_length
- # Time attention layer
- self.time_attention_layer = TimeAttentionLayer(
- n_heads,
- latent_sequence_length,
- key_dim,
- n_patches,
- patch_length,
- unpatched_length,
- pos_embedding_type,
- )
- # Mask for time attention (This is fixed).
- self.time_attention_mask = self._compute_time_attention_mask()
- # Channel attention layer
- self.channel_attention_layer = ChannelAttention(key_dim)
- # Channel attention dropouts
- self.channel_attention_dropout = tf.Variable(
- channel_attention_dropout, trainable=False
- )
- self.within_channel_attention_dropout = tf.Variable(
- within_channel_attention_dropout, trainable=False
- )
- def _compute_time_attention_mask(self) -> tf.Tensor:
- """
- Compute the mask for time attention.
- Returns
- -------
- mask : tf.Tensor
- Mask for time attention.
- Shape: (latent_sequence_length, n_patches + unpatched_length).
- """
- mask = np.zeros(
- (self.latent_sequence_length, self.n_patches + self.unpatched_length)
- )
- # Patch masking
- for i in range(self.n_patches):
- m_indx = max(
- 0,
- (i + 1 - self.n_patches) * self.patch_length
- + self.latent_sequence_length
- - 1,
- )
- mask[:m_indx, i] = 1
- # Unpatched masking
- for i in range(self.n_patches, self.n_patches + self.unpatched_length):
- mask[
- : self.latent_sequence_length
- - (self.n_patches + self.unpatched_length)
- + i,
- i,
- ] = 1
- mask = tf.constant(mask, dtype=tf.float32)
- return mask
- def _compute_channel_attention_mask(self, training: bool) -> Union[tf.Tensor, None]:
- """
- Compute the mask for channel attention.
- Parameters
- ----------
- training : bool
- Whether the model is training.
- If False, no dropout is applied.
- Returns
- -------
- mask : tf.Tensor or None
- Mask for channel attention.
- Shape: (n_channels, n_channels).
- If None, full attention is applied.
- """
- if self.channel_attention_dropout < 1.0:
- # We apply channel attention dropout only during training
- if not training:
- return tf.zeros((self.n_channels, self.n_channels), dtype=tf.float32)
- else:
- uniform_sampler = tfp.distributions.Uniform()
- # Sample whether to apply channel attention
- if uniform_sampler.sample() < self.channel_attention_dropout:
- # Does not apply channel attention and mask all off-diagonal elements
- mask = 1 - np.eye(self.n_channels)
- return tf.constant(mask, dtype=tf.float32)
- # Sample whether to apply within-channel attention
- elif uniform_sampler.sample() < self.within_channel_attention_dropout:
- # Does not apply within-channel attention and mask all diagonal elements
- mask = np.eye(self.n_channels)
- return tf.constant(mask, dtype=tf.float32)
- # Apply channel attention and does not mask any elements, return None
- else:
- return tf.zeros(
- (self.n_channels, self.n_channels), dtype=tf.float32
- )
- else:
- # If channel_attention_dropout >= 1.0, no channel attention is applied
- # Mask all off-diagonal elements
- mask = 1 - np.eye(self.n_channels)
- return tf.constant(mask, dtype=tf.float32)
- def call(self, inputs, training=None, **kwargs):
- # ---------- Unpack Inputs ---------- #
- # q (Query): (batch_size, n_heads, latent_sequence_length, out_n_channels, key_dim)
- # k (Key): (batch_size, n_heads, n_patches + unpatched_length, out_n_channels, key_dim)
- # v (Value): (batch_size, n_heads, n_patches + unpatched_length, in_n_channels, key_dim)
- # c_q (Channel Query): (batch_size, n_heads, latent_sequence_length, out_n_channels, key_dim)
- # c_k (Channel Key): (batch_size, n_heads, latent_sequence_length, in_n_channels, key_dim)
- q, k, v, c_q, c_k = inputs
- # ---------- Channel Attention ---------- #
- # First sample channel attention mask
- channel_attention_mask = self._compute_channel_attention_mask(training=training)
- # Apply channel attention
- output = self.channel_attention_layer(
- [c_q, c_k, v], mask=channel_attention_mask, training=training, **kwargs
- )
- # ---------- Time Attention ---------- #
- output = self.time_attention_layer(
- [q, k, output], mask=self.time_attention_mask, training=training, **kwargs
- )
- return output
- class MultiHeadPASSTALayer(tf.keras.layers.Layer):
- """
- The Multi-head PASSTA layer.
- Parameters
- ----------
- n_heads : int
- Number of heads.
- model_dim : int
- Model dimension.
- n_channels : int
- Number of channels.
- sequence_length : int
- Sequence length.
- latent_sequence_length : int
- Latent sequence length.
- n_patches : int
- Number of patches to attend to.
- patch_length : int
- Patch length.
- unpatched_length : int
- Number of unpatched elements to attend to.
- pos_embedding_type : str
- Type of positional embedding to use.
- channel_attention_dropout : float
- Dropout rate for channel attention.
- Values greater than 1.0 means no channel attention.
- Values less than 0.0 means no dropout.
- within_channel_attention_dropout : float
- Dropout rate for within-channel attention.
- Values greater than 1.0 means no within-channel attention.
- Values less than 0.0 means no dropout
- """
- def __init__(
- self,
- n_heads: int,
- model_dim: int,
- n_channels: int,
- sequence_length: int,
- latent_sequence_length: int,
- n_patches: int,
- patch_length: int,
- unpatched_length: int,
- pos_embedding_type: str,
- channel_attention_dropout: float,
- within_channel_attention_dropout: float,
- **kwargs,
- ):
- super().__init__(**kwargs)
- self.n_heads = n_heads
- self.model_dim = model_dim
- self.key_dim = model_dim // n_heads
- self.n_channels = n_channels
- self.sequence_length = sequence_length
- self.latent_sequence_length = latent_sequence_length
- self.n_patches = n_patches
- self.patch_length = patch_length
- self.unpatched_length = unpatched_length
- self.pos_embedding_type = pos_embedding_type
- self.channel_attention_dropout = channel_attention_dropout
- self.within_channel_attention_dropout = within_channel_attention_dropout
- # Patch projection
- self.patch_projection = tf.keras.layers.Dense(n_heads)
- # Input projections
- self.time_patched_projection = tf.keras.layers.Dense(2 * self.model_dim)
- self.time_unpatched_projection = tf.keras.layers.Dense(2 * self.model_dim)
- self.time_query_projection = tf.keras.layers.Dense(self.model_dim)
- self.channel_projection = tf.keras.layers.Dense(2 * self.model_dim)
- # PASSTA layer for time and channel attention
- self.passta_layer = PASSTALayer(
- n_heads,
- n_channels,
- latent_sequence_length,
- n_patches,
- patch_length,
- unpatched_length,
- self.key_dim,
- pos_embedding_type,
- channel_attention_dropout,
- within_channel_attention_dropout,
- )
- # Output projection
- self.output_projection = tf.keras.layers.Dense(self.model_dim)
- def _patch_x(self, x: tf.Tensor) -> tf.Tensor:
- # x.shape: (batch_size, sequence_length, n_channels, model_dim)
- x = tf.transpose(x, perm=(0, 2, 3, 1))
- # x.shape: (batch_size, n_channels, model_dim, sequence_length)
- x = tf.reshape(
- x,
- (
- tf.shape(x)[0],
- self.n_channels,
- self.model_dim,
- self.n_patches,
- self.patch_length,
- ),
- )
- # x.shape: (batch_size, n_channels, model_dim, n_patches, patch_length)
- x = tf.transpose(x, perm=(0, 3, 1, 2, 4))
- # x.shape: (batch_size, n_patches, n_channels, model_dim, patch_length)
- return x
- def _perceiver_x(self, x: tf.Tensor) -> tf.Tensor:
- """Get the last latent_sequence_length elements of the sequence."""
- # x.shape: (batch_size, sequence_length, n_channels, model_dim)
- x = tf.slice(
- x,
- [0, self.sequence_length - self.latent_sequence_length, 0, 0],
- [-1, -1, -1, -1],
- )
- # x.shape: (batch_size, latent_sequence_length, n_channels, model_dim)
- return x
- def _unpatch_x(self, x: tf.Tensor) -> tf.Tensor:
- """Get the last unpatched_length elements of the sequence."""
- # x.shape: (batch_size, sequence_length, n_channels, model_dim)
- x = tf.slice(
- x,
- [0, self.sequence_length - self.unpatched_length, 0, 0],
- [-1, -1, -1, -1],
- )
- # x.shape: (batch_size, unpatched_length, n_channels, model_dim)
- return x
- def _split_heads(self, x: tf.Tensor) -> tf.Tensor:
- # x.shape: (batch_size, time_length, n_channels, model_dim)
- # Here time_length is either n_patches + unpatched_length or latent_sequence_length
- x = tf.reshape(
- x,
- (
- tf.shape(x)[0],
- tf.shape(x)[1],
- self.n_channels,
- self.n_heads,
- self.key_dim,
- ),
- )
- # x.shape: (batch_size, time_length, n_channels, n_heads, key_dim)
- x = tf.transpose(x, perm=(0, 3, 1, 2, 4))
- # x.shape: (batch_size, n_heads, time_length, n_channels, key_dim)
- return x
- def _combine_heads(self, x: tf.Tensor) -> tf.Tensor:
- # x.shape: (batch_size, n_heads, latent_sequence_length, n_channels, key_dim)
- x = tf.transpose(x, perm=(0, 2, 3, 1, 4))
- # x.shape: (batch_size, latent_sequence_length, n_channels, n_heads, key_dim)
- x = tf.reshape(
- x,
- (
- tf.shape(x)[0],
- self.latent_sequence_length,
- self.n_channels,
- self.model_dim,
- ),
- )
- # x.shape: (batch_size, latent_sequence_length, n_channels, model_dim)
- return x
- def call(self, inputs, training=None, **kwargs):
- x = inputs
- # x.shape: (batch_size, sequence_length, n_channels, model_dim)
- # ---------- Process inputs ---------- #
- # Input is processed into 3 parts:
- # 1. Patched input: patched_x
- # 2. Unpatched input: unpatched_x
- # 3. Perceiver input: perceiver_x
- patched_x = self._patch_x(x)
- # patched_x.shape: (batch_size, n_patches, n_channels, model_dim, patch_length)
- patched_x = tf.reshape(
- patched_x,
- (
- tf.shape(patched_x)[0],
- self.n_patches,
- self.n_channels,
- self.key_dim,
- self.n_heads,
- self.patch_length,
- ),
- )
- # patched_x.shape: (batch_size, n_patches, n_channels, key_dim, n_heads, patch_length)
- patched_x = tf.reshape(
- patched_x,
- (
- tf.shape(patched_x)[0],
- self.n_patches,
- self.n_channels,
- self.key_dim,
- self.n_heads * self.patch_length,
- ),
- )
- # patched_x.shape: (batch_size, n_patches, n_channels, key_dim, n_heads * patch_length)
- patched_x = self.patch_projection(patched_x)
- # patched_x.shape: (batch_size, n_patches, n_channels, key_dim, n_heads)
- patched_x = tf.reshape(
- patched_x,
- (
- tf.shape(patched_x)[0],
- self.n_patches,
- self.n_channels,
- self.model_dim,
- ),
- )
- # patched_x.shape: (batch_size, n_patches, n_channels, model_dim)
- perceiver_x = self._perceiver_x(x)
- # perceiver_x.shape: (batch_size, latent_sequence_length, n_channels, model_dim)
- unpatched_x = self._unpatch_x(x)
- # unpatched_x.shape: (batch_size, unpatched_length, n_channels, model_dim)
- # ---------- Project inputs to Q, K, V ---------- #
- q = self.time_query_projection(perceiver_x)
- # q.shape: (batch_size, latent_sequence_length, n_channels, model_dim)
- k_patched, v_patched = tf.split(
- self.time_patched_projection(patched_x), 2, axis=-1
- )
- # k_patched.shape: (batch_size, n_patches, n_channels, model_dim)
- # v_patched.shape: (batch_size, n_patches, n_channels, model_dim)
- k_unpatched, v_unpatched = tf.split(
- self.time_unpatched_projection(unpatched_x), 2, axis=-1
- )
- # k_unpatched.shape: (batch_size, unpatched_length, n_channels, model_dim)
- # v_unpatched.shape: (batch_size, unpatched_length, n_channels, model_dim)
- # Concatenate k and v for time attention
- k = tf.concat([k_patched, k_unpatched], axis=1)
- v = tf.concat([v_patched, v_unpatched], axis=1)
- # k.shape: (batch_size, n_patches + unpatched_length, n_channels, model_dim)
- # v.shape: (batch_size, n_patches + unpatched_length, n_channels, model_dim)
- c_q, c_k = tf.split(self.channel_projection(perceiver_x), 2, axis=-1)
- # c_q.shape: (batch_size, latent_sequence_length, n_channels, model_dim)
- # c_k.shape: (batch_size, latent_sequence_length, n_channels, model_dim)
- # ---------- Split heads ---------- #
- q = self._split_heads(q)
- # (batch_size, n_heads, latent_sequence_length, n_channels, key_dim)
- k = self._split_heads(k)
- # (batch_size, n_heads, n_patches + unpatched_length, n_channels, key_dim)
- v = self._split_heads(v)
- # (batch_size, n_heads, n_patches + unpatched_length, n_channels, key_dim)
- c_q = self._split_heads(c_q)
- # (batch_size, n_heads, latent_sequence_length, n_channels, key_dim)
- c_k = self._split_heads(c_k)
- # (batch_size, n_heads, latent_sequence_length, n_channels, key_dim)
- # ---------- PASSTA Layer ---------- #
- output = self.passta_layer([q, k, v, c_q, c_k], training=training, **kwargs)
- # output.shape: (batch_size, n_heads, latent_sequence_length, n_channels, key_dim)
- # ---------- Combine heads ---------- #
- output = self._combine_heads(output)
- # output.shape: (batch_size, latent_sequence_length, n_channels, model_dim)
- # ---------- Output projection ---------- #
- output = self.output_projection(output)
- # output.shape: (batch_size, latent_sequence_length, n_channels, model_dim)
- return output
layers.py at commit ffbeff3, under MIT · at the source
Overview
- Oxford Centre for Integrative Neuroimaging (OxCIN), University of Oxford, Oxford, United Kingdom
- Department of Psychiatry, University of Oxford, Oxford, United Kingdom
- Nuffield Department of Clinical Neurosciences, University of Oxford, Oxford, United Kingdom
- Department of Engineering Science, University of Oxford, Oxford, United Kingdom
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
ffbeff318811de3adea33d7dfeaa0af15a5b902b, 8 July 2026Availability: 1 check, the latest on 27 September 2026: the link answers
- 27 September 2026: the link answers
82 files
- examples/
paper/ — Python, 315 lines, 2 matchesanalysis/ 1_tokeniser_reconstructs .py - examples/
paper/ — Python, 211 linesanalysis/ 2_spatial_spectral.py - examples/
paper/ — Python, 723 lines, 1 matchanalysis/ 3_subject_fingerprints.p y - examples/
paper/ — Python, 556 lines, 2 matchesanalysis/ 4_bursting_behaviour.py - examples/
paper/ — Python, 363 linesanalysis/ plot_embeddings.py - examples/
paper/ — Python, 95 linesdecoding/ 1_get_baseline_features. py - examples/
paper/ — Python, 38 linesdecoding/ 2_tokenize_data.py - examples/
paper/ — Python, 142 linesdecoding/ 3_get_zero_shot_features .py - examples/
paper/ — Python, 28 linesdecoding/ 4_prepare_tfrecords.py - examples/
paper/ — Python, 49 linesdecoding/ 5_fine_tuning.py - examples/
paper/ — Python, 142 linesdecoding/ 6_get_fine_tune_features .py - examples/
paper/ — Python, 232 lines, 2 matchesdecoding/ 7_decoding_performances. py - examples/
paper/ — Python, 78 linesgenerate_data/ 1_meg-gpt.py - examples/
paper/ — Python, 59 linesgenerate_data/ 2_ar_model.py - examples/
paper/ — Python, 32 linestrain_models/ 1_train_tokenizer.py - examples/
paper/ — Python, 43 linestrain_models/ 2_tokenize_data.py - examples/
paper/ — Python, 35 linestrain_models/ 3_prepare_data.py - examples/
paper/ — Python, 33 linestrain_models/ 4_train_meg-gpt.py - examples/
paper/ — Python, 186 linestrain_models/ 5_train_ar.py - examples/
random_tokens/ — Python, 47 linesephys_gpt.py - examples/
real_data/ — Python, 48 linescamcan_rest/ 1_train_tokenizer.py - examples/
real_data/ — Python, 51 linescamcan_rest/ 2_tokenize_data.py - examples/
real_data/ — Python, 71 linescamcan_rest/ gradual_space_attention/ 1_train_generator.py - examples/
real_data/ — Python, 34 linescamcan_rest/ gradual_space_attention/ 2_generate_data.py - examples/
real_data/ — Python, 138 linescamcan_rest/ gradual_space_attention/ 3_plot_results.py - examples/
real_data/ — Python, 62 linescamcan_rest/ gradual_space_attention/ 4_train_hmm.py - examples/
real_data/ — Python, 9 linescamcan_rest/ gradual_space_attention/ plot_history.py - examples/
real_data/ — Python, 27 linescamcan_rest/ gradual_space_attention/ submit.py - examples/
real_data/ — Python, 32 linescamcan_rest/ with_space_attention/ 1_train_generator.py - examples/
real_data/ — Python, 34 linescamcan_rest/ with_space_attention/ 2_generate_data.py - examples/
real_data/ — Python, 138 linescamcan_rest/ with_space_attention/ 3_plot_results.py - examples/
real_data/ — Python, 61 linescamcan_rest/ with_space_attention/ 4_train_hmm.py - examples/
real_data/ — Python, 9 linescamcan_rest/ with_space_attention/ plot_history.py - examples/
real_data/ — Python, 27 linescamcan_rest/ with_space_attention/ submit.py - examples/
real_data/ — Python, 32 linescamcan_rest/ without_space_attention/ 1_train_generator.py - examples/
real_data/ — Python, 34 linescamcan_rest/ without_space_attention/ 2_generate_data.py - examples/
real_data/ — Python, 138 linescamcan_rest/ without_space_attention/ 3_plot_results.py - examples/
real_data/ — Python, 175 linescamcan_rest/ without_space_attention/ 4_bursts.py - examples/
real_data/ — Python, 9 linescamcan_rest/ without_space_attention/ plot_history.py - examples/
real_data/ — Python, 27 linescamcan_rest/ without_space_attention/ submit.py - examples/
simulation/ — Python, 43 linessession_variability/ 1_simulate_data.py - examples/
simulation/ — Python, 76 linessession_variability/ 2_train_tokenizer.py - examples/
simulation/ — Python, 79 linessession_variability/ 3_train_generator.py - examples/
simulation/ — Python, 108 linessession_variability/ 4_plot_results.py - examples/
simulation/ — Python, 40 linessimple_simulation/ 1_simulate_data.py - examples/
simulation/ — Python, 77 linessimple_simulation/ 2_train_tokenizer.py - examples/
simulation/ — Python, 46 linessimple_simulation/ 3_tokenize_data.py - examples/
simulation/ — Python, 38 linessimple_simulation/ gradual_space_attention/ 1_train_generator.py - examples/
simulation/ — Python, 22 linessimple_simulation/ gradual_space_attention/ 2_generate_data.py - examples/
simulation/ — Python, 165 linessimple_simulation/ gradual_space_attention/ 3_plot_results.py - examples/
simulation/ — Python, 9 linessimple_simulation/ gradual_space_attention/ plot_history.py - examples/
simulation/ — Python, 34 linessimple_simulation/ mode_labels/ 1_train_generator.py - examples/
simulation/ — Python, 52 linessimple_simulation/ mode_labels/ 2_generate_data.py - examples/
simulation/ — Python, 85 linessimple_simulation/ mode_labels/ 3_plot_results.py - examples/
simulation/ — Python, 38 linessimple_simulation/ with_space_attention/ 1_train_generator.py - examples/
simulation/ — Python, 22 linessimple_simulation/ with_space_attention/ 2_generate_data.py - examples/
simulation/ — Python, 165 linessimple_simulation/ with_space_attention/ 3_plot_results.py - examples/
simulation/ — Python, 9 linessimple_simulation/ with_space_attention/ plot_history.py - examples/
simulation/ — Python, 38 linessimple_simulation/ without_space_attention/ 1_train_generator.py - examples/
simulation/ — Python, 22 linessimple_simulation/ without_space_attention/ 2_generate_data.py - examples/
simulation/ — Python, 165 linessimple_simulation/ without_space_attention/ 3_plot_results.py - examples/
simulation/ — Python, 9 linessimple_simulation/ without_space_attention/ plot_history.py - osl_foundation/
__init__.py — Python, 23 lines - osl_foundation/
config/ — Python, 286 lines__init__.py - osl_foundation/
config/ — Python, 86 linesbase.py - osl_foundation/
config/ — Python, 214 linesgenerator_config.py - osl_foundation/
config/ — Python, 70 linestokenizer_config.py - osl_foundation/
inference/ — Python, 1 line__init__.py - osl_foundation/
inference/ — Python, 170 linescallbacks.py - osl_foundation/
inference/ — Python, 1,134 lines, 3 matcheslayers.py - osl_foundation/
models/ — Python, 90 lines__init__.py - osl_foundation/
models/ — Python, 383 linesbase.py - osl_foundation/
models/ — Python, 995 lines, 2 matchesmeg_gpt.py - osl_foundation/
models/ — Python, 1,659 linestokenizers.py - osl_foundation/
simulation/ — Python, 382 linesbursts.py - osl_foundation/
utils/ — Python, 1 line__init__.py - osl_foundation/
utils/ — Python, 11 linesmisc.py - osl_foundation/
utils/ — Python, 383 linesplotting.py - osl_foundation/
utils/ — Python, 133 lines, 1 matchsampling.py - osl_foundation/
utils/ — Python, 64 linestesting.py - LICENSE — License, 21 lines
- README.md — Text, 65 lines
Zenodo 11099418
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://
A PyTorch implemention of the tokenizer is available here: 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, 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://
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/
url = {https://
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/
VL - 4
SP - IMAG.a.1301
SN - 2837-6056
PB - MIT Press
DO - 10.1162/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1162/
"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":
"volume": "4",
"page": "IMAG.a.1301",
"DOI": "10.1162/
"PMID": "42516171",
"PMCID": "PMC13403652",
"ISSN": "2837-6056",
"publisher": "MIT Press",
"URL": "https://
"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 mappingIn 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 diseaseIn 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 communicationsIn 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 dataIn 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 consciousnessIn 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.
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, 80 scripts, and 13 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:5888fd44e9a4580f…
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.
