File size: 4,805 Bytes
4947683 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 | from typing import Any, Dict
import hydra
import numpy as np
import omegaconf
import torch
import pytorch_lightning as pl
import torch.nn as nn
from torch.nn import functional as F
from torch_scatter import scatter
from tqdm import tqdm
from diffcsp.common.utils import PROJECT_ROOT
from diffcsp.common.data_utils import (
EPSILON, cart_to_frac_coords, mard, lengths_angles_to_volume,
frac_to_cart_coords, min_distance_sqr_pbc)
MAX_ATOMIC_NUM = 100
def build_mlp(in_dim, hidden_dim, fc_num_layers, out_dim):
mods = [nn.Linear(in_dim, hidden_dim), nn.ReLU()]
for i in range(fc_num_layers-1):
mods += [nn.Linear(hidden_dim, hidden_dim), nn.ReLU()]
mods += [nn.Linear(hidden_dim, out_dim)]
return nn.Sequential(*mods)
class BaseModule(pl.LightningModule):
def __init__(self, *args, **kwargs) -> None:
super().__init__()
# populate self.hparams with args and kwargs automagically!
self.save_hyperparameters()
if hasattr(self.hparams, "model"):
self._hparams = self.hparams.model
def configure_optimizers(self):
opt = hydra.utils.instantiate(
self.hparams.optim.optimizer, params=self.parameters(), _convert_="partial"
)
if not self.hparams.optim.use_lr_scheduler:
return [opt]
scheduler = hydra.utils.instantiate(
self.hparams.optim.lr_scheduler, optimizer=opt
)
return {"optimizer": opt, "lr_scheduler": scheduler, "monitor": "val_loss"}
class CrystGNN_Supervise(BaseModule):
"""
GNN model for fitting the supervised objectives for crystals.
"""
def __init__(self, *args, **kwargs) -> None:
super().__init__(*args, **kwargs)
self.encoder = hydra.utils.instantiate(self.hparams.encoder)
def forward(self, batch) -> Dict[str, torch.Tensor]:
preds = self.encoder(batch) # shape (N, 1)
return preds
def training_step(self, batch: Any, batch_idx: int) -> torch.Tensor:
preds = self(batch)
loss = F.mse_loss(preds, batch.y)
self.log_dict(
{'train_loss': loss},
on_step=True,
on_epoch=True,
prog_bar=True,
)
return loss
def validation_step(self, batch: Any, batch_idx: int) -> torch.Tensor:
preds = self(batch)
log_dict, loss = self.compute_stats(batch, preds, prefix='val')
self.log_dict(
log_dict,
on_step=False,
on_epoch=True,
prog_bar=True,
)
return loss
def test_step(self, batch: Any, batch_idx: int) -> torch.Tensor:
preds = self(batch)
log_dict, loss = self.compute_stats(batch, preds, prefix='test')
self.log_dict(
log_dict,
)
return loss
def compute_stats(self, batch, preds, prefix):
loss = F.mse_loss(preds, batch.y)
self.scaler.match_device(preds)
scaled_preds = self.scaler.inverse_transform(preds)
scaled_y = self.scaler.inverse_transform(batch.y)
mae = torch.mean(torch.abs(scaled_preds - scaled_y))
log_dict = {
f'{prefix}_loss': loss,
f'{prefix}_mae': mae,
}
if self.hparams.data.prop == 'scaled_lattice':
pred_lengths = scaled_preds[:, :3]
pred_angles = scaled_preds[:, 3:]
if self.hparams.data.lattice_scale_method == 'scale_length':
pred_lengths = pred_lengths * \
batch.num_atoms.view(-1, 1).float()**(1/3)
lengths_mae = torch.mean(torch.abs(pred_lengths - batch.lengths))
angles_mae = torch.mean(torch.abs(pred_angles - batch.angles))
lengths_mard = mard(batch.lengths, pred_lengths)
angles_mard = mard(batch.angles, pred_angles)
pred_volumes = lengths_angles_to_volume(pred_lengths, pred_angles)
true_volumes = lengths_angles_to_volume(
batch.lengths, batch.angles)
volumes_mard = mard(true_volumes, pred_volumes)
log_dict.update({
f'{prefix}_lengths_mae': lengths_mae,
f'{prefix}_angles_mae': angles_mae,
f'{prefix}_lengths_mard': lengths_mard,
f'{prefix}_angles_mard': angles_mard,
f'{prefix}_volumes_mard': volumes_mard,
})
return log_dict, loss
@hydra.main(config_path=str(PROJECT_ROOT / "conf"), config_name="default", version_base="1.1")
def main(cfg: omegaconf.DictConfig):
model: pl.LightningModule = hydra.utils.instantiate(
cfg.model,
optim=cfg.optim,
data=cfg.data,
logging=cfg.logging,
_recursive_=False,
)
return model
if __name__ == "__main__":
main()
|