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 model for pretraining with node-level and graph-level predictions. |
Module Contents
- 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_eps: float = 0.0, enc_train_eps: bool = False, 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
enc_eps (float) – Initial value of the
epsterm inGINEConv, defaults to 0.0 (matches PyG default and preserves the pre-unification pretraining behaviour).enc_train_eps (bool) – Whether to learn
epsas a parameter, defaults to False (matches PyG default and preserves the pre-unification pretraining behaviour).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