Wthinker's picture
Publish AIDD open-source resources
4947683 verified
Raw History Blame Contribute Delete
4.81 kB
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()