diff --git a/src/molearn/models/vae.py b/src/molearn/models/vae.py new file mode 100644 index 0000000..69b5d81 --- /dev/null +++ b/src/molearn/models/vae.py @@ -0,0 +1,80 @@ +import torch +from torch import nn +import torch.nn.functional as F + +from molearn.models.foldingnet import AutoEncoder, GraphLayer, knn, index_points + + +class Encoder(nn.Module): + ''' + Graph based encoder + ''' + def __init__(self, latent_dimension=2, **kwargs): + super(Encoder, self).__init__() + self.latent_dimension = latent_dimension + self.conv1 = nn.Conv1d(12, 64, 1) + self.conv2 = nn.Conv1d(64, 64, 1) + self.conv3 = nn.Conv1d(64, 64, 1) + + self.bn1 = nn.BatchNorm1d(64) + self.bn2 = nn.BatchNorm1d(64) + self.bn3 = nn.BatchNorm1d(64) + + self.graph_layer1 = GraphLayer(in_channel=64, out_channel=128, k=16) + self.graph_layer2 = GraphLayer(in_channel=128, out_channel=1024, k=16) + + self.conv4 = nn.Conv1d(1024, 512, 1) + self.bn4 = nn.BatchNorm1d(512) + self.conv_mu = nn.Conv1d(512, latent_dimension,1) + self.conv_logvar = nn.Conv1d(512, latent_dimension,1) + + def forward(self, x): + b, c, n = x.size() + + # get the covariances, reshape and concatenate with x + knn_idx = knn(x, k=16) + knn_x = index_points(x.permute(0, 2, 1), knn_idx) # (B, N, 16, 3) + mean = torch.mean(knn_x, dim=2, keepdim=True) + knn_x = knn_x - mean + covariances = torch.matmul(knn_x.transpose(2, 3), knn_x).view(b, n, -1).permute(0, 2, 1) + x = torch.cat([x, covariances], dim=1) # (B, 12, N) + + # three layer MLP + x = F.relu(self.bn1(self.conv1(x))) + x = F.relu(self.bn2(self.conv2(x))) + x = F.relu(self.bn3(self.conv3(x))) + + # two consecutive graph layers + x = self.graph_layer1(x) + x = self.graph_layer2(x) + + x = self.bn4(self.conv4(x)) + + x = torch.max(x, dim=-1)[0].unsqueeze(-1) + + mu = self.conv_mu(x).squeeze(-1) + logvar = self.conv_logvar(x).squeeze(-1) + std = torch.exp(logvar * 0.5) + eps = torch.randn_like(std) + z = eps * std + mu + + return z.squeeze(-1), mu, logvar + + +class VAE(AutoEncoder): + """ + Variational autoencoder architecture derived from FoldingNet. + """ + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.encoder = Encoder(*args, **kwargs) + + def forward(self, x): + z, _, _ = self.encode(x) + x_rec = self.decode(z) + return x_rec + + +if __name__=='__main__': + print('Nothing to see here') diff --git a/src/molearn/trainers/openmm_physics_vae_trainer.py b/src/molearn/trainers/openmm_physics_vae_trainer.py new file mode 100644 index 0000000..83b1809 --- /dev/null +++ b/src/molearn/trainers/openmm_physics_vae_trainer.py @@ -0,0 +1,75 @@ +import os +import torch +from molearn.loss_functions import openmm_energy + +from molearn.trainers import OpenMM_Physics_Trainer + +class OpenMM_Physics_Trainer_VAE(OpenMM_Physics_Trainer): + """ + A modified OpenMM Physics Trainer suitable for training a Variational Autoencoder + """ + + def __init__(self, kld_weight=1e-4, physics_inter_weight=0, *args, **kwargs): + super().__init__(*args, **kwargs) + self.kld_weight = kld_weight + + def common_step(self, batch): + """ + Called from both train_step and valid_step. + Calculates the mean squared error loss for self.autoencoder. + Encoded and decoded frames are saved in self._internal under keys ``encoded`` and ``decoded`` respectively should you wish to use them elsewhere. + + :param torch.Tensor batch: Tensor of shape [Batch size, Number of Atoms, 3] A mini-batch of protein frames normalised. To recover original data multiple by ``self.std``. + :returns: Return calculated mse_loss and the kld_loss (KL(q_\phi(z|x) || p(z))) + :rtype: dict + """ + self._internal = {} + encoded, mu, logvar = self.autoencoder.encode(batch) + self._internal["encoded"] = encoded + decoded = self.autoencoder.decode(encoded)[:, : batch.size(1), :] + self._internal["decoded"] = decoded + geometric_loss = torch.nn.functional.mse_loss(batch , decoded) + kld_loss = torch.mean(-0.5 * torch.sum(1 + logvar - mu ** 2 - logvar.exp(), dim=1), dim=0) # Closed form KL between Gaussians + return dict(mse_loss=geometric_loss, kld_loss=kld_loss) + + def train_step(self, batch): + """ + This method overrides :func:`OpenMM_Physics_Trainer.train_step ` and adds an additional 'kld_loss' term. + Called from :func:`OpenMM_Physics_Trainer.train_epoch `. + + :param torch.Tensor batch: tensor shape [Batch size, Number of Atoms, 3]. A mini-batch of protein frames normalised. To recover original data multiple by ``self.std``. + :returns: Return loss. The dictionary must contain an entry with key ``'loss'`` that :func:`self.train_epoch ` will call ``result['loss'].backwards()`` to obtain gradients. + :rtype: dict + """ + + results = self.common_step(batch) + results.update(self.common_physics_step(batch, self._internal["encoded"])) + loss = results["mse_loss"] + self.kld_weight * results["kld_loss"] + self.physics_inter_weight * results["inter_physics_loss"] + results["loss"] = loss + return results + + def valid_step(self, batch): + """ + This method overrides :func:`OpenMM_Physics_Trainer.valid_step ` and adds an additional 'kld_loss' term. + + Differently to :func:`train_step ` this method sums the logs of mse_loss, kld_loss, and physics_loss ``final_loss = torch.log(results['mse_loss'])+kld_scale*torch.log(results["kld_loss"])+scale*torch.log(results['physics_loss'])`` + + Called from super class :func:`OpenMM_Physics_Trainer.valid_epoch` on every mini-batch. + + :param torch.Tensor batch: Tensor of shape [Batch size, Number of Atoms, 3]. A mini-batch of protein frames normalised. To recover original data multiple by ``self.std``. + :returns: Return loss. The dictionary must contain an entry with key ``'loss'`` that will be the score via which the best checkpoint is determined. + :rtype: dict + + """ + + results = self.common_step(batch) + results.update(self.common_physics_step(batch, self._internal["encoded"])) + physics_loss = self.physics_inter_weight * torch.log(results["inter_physics_loss"]) + kld_loss = self.kld_weight * torch.log(results["kld_loss"]) + final_loss = torch.log(results["mse_loss"]) + physics_loss + kld_loss + results["loss"] = final_loss + return results + + +if __name__ == "__main__": + pass