matcha.torch.models.pretraining.gps_pretraining
GPS pretraining model for self-supervised learning on graphs.
Classes
GPS encoder variant that returns node embeddings for pretraining. |
|
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.HyperparametersMixinGPS 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.BaseGraphPretrainingModelGPS 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_depthlayers while the graph MLP uses allenc_num_layerslayers. 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