matcha.torch.models.pretraining.gin_pretraining
GIN pretraining model for supervised multi-task learning on graphs.
Jointly predicts user-provided atom-level and molecule-level labels via a shared GIN encoder with separate prediction heads.
Classes
GIN encoder variant that returns node embeddings for pretraining. |
|
GIN model for pretraining with node-level and graph-level predictions. |
Module Contents
- class matcha.torch.models.pretraining.gin_pretraining.GINPretrainingEncoder(num_layers: int, atom_input_dim: int, bond_input_dim: int, atom_hidden_dim: int, activation: str, aggregation: 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.HyperparametersMixinGIN 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 message passing layers
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
activation (str) – Activation function name
aggregation (str) – Aggregation function for message passing
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
- norms
- norm_type
- 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.gin_pretraining.GINPretraining(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 = 300, enc_aggregation: str = 'vpa', enc_jk: str = 'concat', enc_norm: str | None = 'graph', 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 = 20, 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.BaseGraphPretrainingModelGIN model for pretraining with node-level and graph-level predictions.
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
Both prediction targets are user-provided labels — the model does not reconstruct its own input atom features. Typical node-level targets include computed atom properties (partial charges, SASA, electronegativity, etc.), while graph-level targets are molecular descriptors (logP, MW, etc.).
Example usage:
model = GINPretraining( num_node_targets=3, # e.g., partial charge, SASA, electronegativity num_graph_targets=2, # e.g., logP, molecular weight ) trainer = L.Trainer(max_epochs=50) trainer.fit(model=model, train_dataloaders=train_dataloader)
- Parameters:
num_node_targets (int) – Number of per-atom label dimensions to predict
num_graph_targets (int) – Number of per-molecule label dimensions to predict
enc_num_layers (int) – Number of message passing layers, 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 300
enc_aggregation (str) – Message aggregation strategy, defaults to ‘vpa’
enc_jk (str) – Jumping knowledge strategy, defaults to ‘concat’
enc_norm (str | None) – Normalization type, defaults to ‘graph’
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