Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
80 changes: 80 additions & 0 deletions src/molearn/models/vae.py
Original file line number Diff line number Diff line change
@@ -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')
75 changes: 75 additions & 0 deletions src/molearn/trainers/openmm_physics_vae_trainer.py
Original file line number Diff line number Diff line change
@@ -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 <molearn.trainers.OpenMM_Physics_Trainer.train_step>` and adds an additional 'kld_loss' term.
Called from :func:`OpenMM_Physics_Trainer.train_epoch <molearn.trainers.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 <molearn.trainers.Trainer.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 <molearn.trainers.OpenMM_Physics_Trainer.valid_step>` and adds an additional 'kld_loss' term.

Differently to :func:`train_step <molearn.trainers.OpenMM_Physics_Trainer_VAE.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<molearn.trainer.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
Loading