Spaces:
Paused
Paused
| from __future__ import annotations | |
| from typing import Iterable | |
| from lightning import pytorch as pl | |
| import torch | |
| from torch import nn, Tensor, optim | |
| from chemprop.data import TrainingBatch, BatchMolGraph | |
| from chemprop.nn.metrics import Metric | |
| from chemprop.nn import MessagePassing, Aggregation, Predictor, LossFunction | |
| from chemprop.schedulers import NoamLR | |
| from chemprop.nn.transforms import ScaleTransform | |
| class MPNN(pl.LightningModule): | |
| r"""An :class:`MPNN` is a sequence of message passing layers, an aggregation routine, and a | |
| predictor routine. | |
| The first two modules calculate learned fingerprints from an input molecule | |
| reaction graph, and the final module takes these learned fingerprints as input to calculate a | |
| final prediction. I.e., the following operation: | |
| .. math:: | |
| \mathtt{MPNN}(\mathcal{G}) = | |
| \mathtt{predictor}(\mathtt{agg}(\mathtt{message\_passing}(\mathcal{G}))) | |
| The full model is trained end-to-end. | |
| Parameters | |
| ---------- | |
| message_passing : MessagePassing | |
| the message passing block to use to calculate learned fingerprints | |
| agg : Aggregation | |
| the aggregation operation to use during molecule-level predictor | |
| predictor : Predictor | |
| the function to use to calculate the final prediction | |
| batch_norm : bool, default=True | |
| if `True`, apply batch normalization to the output of the aggregation operation | |
| metrics : Iterable[Metric] | None, default=None | |
| the metrics to use to evaluate the model during training and evaluation | |
| warmup_epochs : int, default=2 | |
| the number of epochs to use for the learning rate warmup | |
| init_lr : int, default=1e-4 | |
| the initial learning rate | |
| max_lr : float, default=1e-3 | |
| the maximum learning rate | |
| final_lr : float, default=1e-4 | |
| the final learning rate | |
| Raises | |
| ------ | |
| ValueError | |
| if the output dimension of the message passing block does not match the input dimension of | |
| the predictor function | |
| """ | |
| def __init__( | |
| self, | |
| message_passing: MessagePassing, | |
| agg: Aggregation, | |
| predictor: Predictor, | |
| batch_norm: bool = True, | |
| metrics: Iterable[Metric] | None = None, | |
| warmup_epochs: int = 2, | |
| init_lr: float = 1e-4, | |
| max_lr: float = 1e-3, | |
| final_lr: float = 1e-4, | |
| X_d_transform: ScaleTransform | None = None, | |
| ): | |
| super().__init__() | |
| self.save_hyperparameters(ignore=["message_passing", "agg", "predictor"]) | |
| self.hparams.update( | |
| { | |
| "message_passing": message_passing.hparams, | |
| "agg": agg.hparams, | |
| "predictor": predictor.hparams, | |
| } | |
| ) | |
| self.message_passing = message_passing | |
| self.agg = agg | |
| self.bn = nn.BatchNorm1d(self.message_passing.output_dim) if batch_norm else nn.Identity() | |
| self.predictor = predictor | |
| self.X_d_transform = X_d_transform if X_d_transform is not None else nn.Identity() | |
| self.metrics = ( | |
| [*metrics, self.criterion] | |
| if metrics | |
| else [self.predictor._T_default_metric(), self.criterion] | |
| ) | |
| self.warmup_epochs = warmup_epochs | |
| self.init_lr = init_lr | |
| self.max_lr = max_lr | |
| self.final_lr = final_lr | |
| def output_dim(self) -> int: | |
| return self.predictor.output_dim | |
| def n_tasks(self) -> int: | |
| return self.predictor.n_tasks | |
| def n_targets(self) -> int: | |
| return self.predictor.n_targets | |
| def criterion(self) -> LossFunction: | |
| return self.predictor.criterion | |
| def fingerprint( | |
| self, bmg: BatchMolGraph, V_d: Tensor | None = None, X_d: Tensor | None = None | |
| ) -> Tensor: | |
| """the learned fingerprints for the input molecules""" | |
| H_v = self.message_passing(bmg, V_d) | |
| H = self.agg(H_v, bmg.batch) | |
| H = self.bn(H) | |
| return H if X_d is None else torch.cat((H, self.X_d_transform(X_d)), 1) | |
| def encoding( | |
| self, bmg: BatchMolGraph, V_d: Tensor | None = None, X_d: Tensor | None = None, i: int = -1 | |
| ) -> Tensor: | |
| """Calculate the :attr:`i`-th hidden representation""" | |
| return self.predictor.encode(self.fingerprint(bmg, V_d, X_d), i) | |
| def forward( | |
| self, bmg: BatchMolGraph, V_d: Tensor | None = None, X_d: Tensor | None = None | |
| ) -> Tensor: | |
| """Generate predictions for the input molecules/reactions""" | |
| return self.predictor(self.fingerprint(bmg, V_d, X_d)) | |
| def training_step(self, batch: TrainingBatch, batch_idx): | |
| bmg, V_d, X_d, targets, weights, lt_mask, gt_mask = batch | |
| mask = targets.isfinite() | |
| targets = targets.nan_to_num(nan=0.0) | |
| Z = self.fingerprint(bmg, V_d, X_d) | |
| preds = self.predictor.train_step(Z) | |
| l = self.criterion(preds, targets, mask, weights, lt_mask, gt_mask) | |
| self.log("train_loss", l, prog_bar=True) | |
| return l | |
| def on_validation_model_eval(self) -> None: | |
| self.eval() | |
| self.predictor.output_transform.train() | |
| def validation_step(self, batch: TrainingBatch, batch_idx: int = 0): | |
| losses = self._evaluate_batch(batch) | |
| metric2loss = {f"val/{m.alias}": l for m, l in zip(self.metrics, losses)} | |
| self.log_dict(metric2loss, batch_size=len(batch[0])) | |
| self.log("val_loss", losses[0], batch_size=len(batch[0]), prog_bar=True) | |
| def test_step(self, batch: TrainingBatch, batch_idx: int = 0): | |
| losses = self._evaluate_batch(batch) | |
| metric2loss = {f"batch_averaged_test/{m.alias}": l for m, l in zip(self.metrics, losses)} | |
| self.log_dict(metric2loss, batch_size=len(batch[0])) | |
| def _evaluate_batch(self, batch) -> list[Tensor]: | |
| bmg, V_d, X_d, targets, _, lt_mask, gt_mask = batch | |
| mask = targets.isfinite() | |
| targets = targets.nan_to_num(nan=0.0) | |
| preds = self(bmg, V_d, X_d) | |
| return [ | |
| metric(preds, targets, mask, None, lt_mask, gt_mask) for metric in self.metrics[:-1] | |
| ] | |
| def predict_step(self, batch: TrainingBatch, batch_idx: int, dataloader_idx: int = 0) -> Tensor: | |
| """Return the predictions of the input batch | |
| Parameters | |
| ---------- | |
| batch : TrainingBatch | |
| the input batch | |
| Returns | |
| ------- | |
| Tensor | |
| a tensor of varying shape depending on the task type: | |
| * regression/binary classification: ``n x (t * s)``, where ``n`` is the number of input | |
| molecules/reactions, ``t`` is the number of tasks, and ``s`` is the number of targets | |
| per task. The final dimension is flattened, so that the targets for each task are | |
| grouped. I.e., the first ``t`` elements are the first target for each task, the second | |
| ``t`` elements the second target, etc. | |
| * multiclass classification: ``n x t x c``, where ``c`` is the number of classes | |
| """ | |
| bmg, X_vd, X_d, *_ = batch | |
| return self(bmg, X_vd, X_d) | |
| def configure_optimizers(self): | |
| opt = optim.Adam(self.parameters(), self.init_lr) | |
| lr_sched = NoamLR( | |
| opt, | |
| self.warmup_epochs, | |
| self.trainer.max_epochs, | |
| self.trainer.estimated_stepping_batches // self.trainer.max_epochs, | |
| self.init_lr, | |
| self.max_lr, | |
| self.final_lr, | |
| ) | |
| lr_sched_config = { | |
| "scheduler": lr_sched, | |
| "interval": "step" if isinstance(lr_sched, NoamLR) else "batch", | |
| } | |
| return {"optimizer": opt, "lr_scheduler": lr_sched_config} | |
| def load_submodules(cls, checkpoint_path, **kwargs): | |
| hparams = torch.load(checkpoint_path)["hyper_parameters"] | |
| kwargs |= { | |
| key: hparams[key].pop("cls")(**hparams[key]) | |
| for key in ("message_passing", "agg", "predictor") | |
| if key not in kwargs | |
| } | |
| return kwargs | |
| def load_from_checkpoint( | |
| cls, checkpoint_path, map_location=None, hparams_file=None, strict=True, **kwargs | |
| ) -> MPNN: | |
| kwargs = cls.load_submodules(checkpoint_path, **kwargs) | |
| return super().load_from_checkpoint( | |
| checkpoint_path, map_location, hparams_file, strict, **kwargs | |
| ) | |
| def load_from_file(cls, model_path, map_location=None, strict=True) -> MPNN: | |
| d = torch.load(model_path, map_location=map_location) | |
| try: | |
| hparams = d["hyper_parameters"] | |
| state_dict = d["state_dict"] | |
| except KeyError: | |
| raise KeyError(f"Could not find hyper parameters and/or state dict in {model_path}. ") | |
| for key in ["message_passing", "agg", "predictor"]: | |
| hparam_kwargs = hparams[key] | |
| hparam_cls = hparam_kwargs.pop("cls") | |
| hparams[key] = hparam_cls(**hparam_kwargs) | |
| model = cls(**hparams) | |
| model.load_state_dict(state_dict, strict=strict) | |
| return model | |