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

GINPretraining

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.BaseGraphPretrainingModel

GIN 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 eps term in GINEConv, defaults to 0.0 (matches PyG default and preserves the pre-unification pretraining behaviour).

  • enc_train_eps (bool) – Whether to learn eps as 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_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