matcha.torch.models.pretraining.gps_pretraining

GPS pretraining model for self-supervised learning on graphs.

Classes

GPSPretrainingEncoder

GPS encoder variant that returns node embeddings for pretraining.

GPSPretraining

GPS model for self-supervised pretraining with node and graph level predictions.

Module Contents

class matcha.torch.models.pretraining.gps_pretraining.GPSPretrainingEncoder(num_layers: int, atom_input_dim: int, bond_input_dim: int, atom_hidden_dim: int, num_heads: int, expansion_k: int, distance_k: int | None, activation: str, dropout: float, norm: str | None, jk: str, readout: str, laplacian_k: int, rwse_k: int, elstatic_k: int, distmat_k: int, rrwp_k: int)[source]

Bases: matcha.torch.encoders.base_graph_encoder.BaseGraphEncoder, lightning.pytorch.core.mixins.HyperparametersMixin

GPS encoder variant that returns node embeddings for pretraining.

This encoder outputs node-level embeddings that can be used for both node-level and graph-level pretraining tasks.

Parameters:
  • num_layers (int) – Number of GPS blocks

  • atom_input_dim (int) – Number of input atom features

  • bond_input_dim (int) – Number of input bond features

  • atom_hidden_dim (int) – Hidden dimension for atom features

  • num_heads (int) – Number of attention heads

  • expansion_k (int) – FFN expansion factor

  • distance_k (int | None) – Maximum distance for spatial encoding

  • activation (str) – Activation function name

  • dropout (float) – Dropout rate

  • norm (str | None) – Normalization type

  • jk (str) – Jumping knowledge strategy

  • readout (str) – Readout function for graph-level aggregation

  • laplacian_k (int) – Laplacian positional encoding dimension

  • rwse_k (int) – Random walk structural encoding dimension

  • elstatic_k (int) – Electrostatic encoding dimension

  • distmat_k (int) – Distance matrix encoding dimension

  • rrwp_k (int) – Relative random walk probability dimension

layers
atom_projection
bond_projection
forward(graph: torch_geometric.data.Batch) torch.Tensor[source]

Forward pass returning graph-level embeddings.

Parameters:

graph – Batched PyG graph

Returns:

Graph-level embeddings [batch_size, hidden_dim]

forward_nodes(graph: torch_geometric.data.Batch) tuple[torch.Tensor, torch_geometric.data.Batch][source]

Forward pass returning node-level embeddings (JK-merged).

Parameters:

graph – Batched PyG graph

Returns:

Tuple of (node_embeddings, processed_graph)

forward_nodes_per_layer(graph: torch_geometric.data.Batch) tuple[list[torch.Tensor], torch_geometric.data.Batch][source]

Forward pass returning per-layer node embeddings.

Parameters:

graph – Batched PyG graph

Returns:

Tuple of (per_layer_embeddings, processed_graph)

class matcha.torch.models.pretraining.gps_pretraining.GPSPretraining(num_node_targets: int = 1, num_graph_targets: int = 1, enc_num_layers: int = 6, enc_atom_input_dim: int = ATOM_FEAT_DIM, enc_bond_input_dim: int = BOND_FEAT_DIM, enc_atom_hidden_dim: int = 512, enc_num_heads: int = 8, enc_expansion_k: int = 1, enc_distance_k: int | None = None, enc_jk: str = 'last', enc_norm: str | None = 'layer', enc_readout: str = 'vpa', enc_activation: str = 'swish', enc_dropout: float = 0.2, enc_laplacian_k: int = 10, enc_rwse_k: int = 20, enc_elstatic_k: int = 0, enc_distmat_k: int = 0, enc_rrwp_k: int = 0, enc_node_encoder_depth: int | None = None, node_head_dims: list[int] | None = None, graph_head_dims: list[int] | None = None, node_task_head_dims: list[int] | None = None, graph_task_head_dims: list[int] | None = None, pred_activation: str = 'swish', pred_dropout: float = 0.2, loss_fn: str = 'mse', loss_args: dict = {}, optimizer: str = 'adamw', optimizer_args: dict = {'lr': 0.0001}, scheduler: str = 'cosine_annealing', scheduler_args: dict = {'min_lr': 1e-06, 'total_steps': 50}, node_loss_weight: float = 0.5, graph_loss_weight: float = 0.5, per_task_log_every_n_steps: int = 1)[source]

Bases: matcha.torch.models.pretraining.base_graph_pretraining.BaseGraphPretrainingModel

GPS model for self-supervised pretraining with node and graph level predictions.

GPS (General, Powerful, Scalable) combines local message passing with global self-attention. This model performs a single encoder forward pass and produces both: - Node-level predictions via a dedicated node MLP head - Graph-level predictions via readout + graph MLP head

Example usage:

model = GPSPretraining(
    num_node_targets=44,   # e.g., predict atom types
    num_graph_targets=1,   # e.g., predict molecular property
)

trainer = L.Trainer(max_epochs=50)
trainer.fit(model=model, train_dataloaders=train_dataloader)
Parameters:
  • num_node_targets (int) – Number of node-level prediction targets

  • num_graph_targets (int) – Number of graph-level prediction targets

  • enc_num_layers (int) – Number of GPS blocks, defaults to 6

  • enc_atom_input_dim (int) – Input atom feature dimension, defaults to ATOM_FEAT_DIM

  • enc_bond_input_dim (int) – Input bond feature dimension, defaults to BOND_FEAT_DIM

  • enc_atom_hidden_dim (int) – Hidden dimension for atoms, defaults to 512

  • enc_num_heads (int) – Number of attention heads, defaults to 8

  • enc_expansion_k (int) – FFN expansion factor, defaults to 1

  • enc_distance_k (int | None) – Maximum distance for spatial encoding

  • enc_jk (str) – Jumping knowledge strategy, defaults to ‘last’

  • enc_norm (str | None) – Normalization type, defaults to ‘layer’

  • enc_readout (str) – Readout function, defaults to ‘vpa’

  • enc_activation (str) – Activation function, defaults to ‘swish’

  • enc_dropout (float) – Dropout rate, defaults to 0.2

  • node_head_dims (list[int]) – Hidden dims for shared node prediction head

  • graph_head_dims (list[int]) – Hidden dims for shared graph prediction head

  • node_task_head_dims (list[int]) – Per-task hidden dims for node head

  • graph_task_head_dims (list[int]) – Per-task hidden dims for graph head

  • pred_activation (str) – Activation for prediction heads

  • pred_dropout (float) – Dropout for prediction heads

  • loss_fn (str) – Loss function name

  • loss_args (dict) – Loss function arguments

  • optimizer (str) – Optimizer name

  • optimizer_args (dict) – Optimizer arguments

  • enc_node_encoder_depth (int | None) – Number of encoder layers used by the node prediction head. When set, the node MLP receives embeddings from only the first enc_node_encoder_depth layers while the graph MLP uses all enc_num_layers layers. Defaults to None (both heads use all layers).

  • scheduler (str) – Scheduler name

  • scheduler_args (dict) – Scheduler arguments

  • node_loss_weight (float) – Constant weight for node-level loss

  • graph_loss_weight (float) – Constant weight for graph-level loss

node_head
graph_head