CoolFace
Modelpublic

OneScience-Group/DiffDock

sourceHugging Facemitupdated 1mo agoView on Hugging Face
0likes31downloads
layers.py129 linesDownload Raw Back to models
1import torch2from torch import nn3 4ACTIVATIONS = {5    "relu": nn.ReLU,6    "silu": nn.SiLU,7}8 9 10def FCBlock(in_dim, hidden_dim, out_dim, layers, dropout, activation="relu"):11    activation = ACTIVATIONS[activation]12    assert layers >= 213    sequential = [nn.Linear(in_dim, hidden_dim), activation(), nn.Dropout(dropout)]14    for i in range(layers - 2):15        sequential += [nn.Linear(hidden_dim, hidden_dim), activation(), nn.Dropout(dropout)]16    sequential += [nn.Linear(hidden_dim, out_dim)]17    return nn.Sequential(*sequential)18 19 20class GaussianSmearing(torch.nn.Module):21    # used to embed the edge distances22    def __init__(self, start=0.0, stop=5.0, num_gaussians=50):23        super().__init__()24        offset = torch.linspace(start, stop, num_gaussians)25        self.coeff = -0.5 / (offset[1] - offset[0]).item() ** 226        self.register_buffer("offset", offset)27 28    def forward(self, dist):29        dist = dist.view(-1, 1) - self.offset.view(1, -1)30        return torch.exp(self.coeff * torch.pow(dist, 2))31 32 33class AtomEncoder(torch.nn.Module):34    def __init__(self, emb_dim, feature_dims, sigma_embed_dim, lm_embedding_dim=0):35        """36 37        Parameters38        ----------39        emb_dim40        feature_dims41            first element of feature_dims tuple is a list with the length of each categorical feature,42            and the second is the number of scalar features43        sigma_embed_dim44        lm_embedding_dim45        """46        super(AtomEncoder, self).__init__()47        self.atom_embedding_list = torch.nn.ModuleList()48        self.num_categorical_features = len(feature_dims[0])49        self.additional_features_dim = feature_dims[1] + sigma_embed_dim + lm_embedding_dim50        for i, dim in enumerate(feature_dims[0]):51            emb = torch.nn.Embedding(dim, emb_dim)52            torch.nn.init.xavier_uniform_(emb.weight.data)53            self.atom_embedding_list.append(emb)54 55        if self.additional_features_dim > 0:56            self.additional_features_embedder = torch.nn.Linear(57                self.additional_features_dim + emb_dim, emb_dim58            )59 60    def forward(self, x):61        x_embedding = 062        assert x.shape[1] == self.num_categorical_features + self.additional_features_dim63        for i in range(self.num_categorical_features):64            x_embedding += self.atom_embedding_list[i](x[:, i].long())65 66        if self.additional_features_dim > 0:67            x_embedding = self.additional_features_embedder(68                torch.cat([x_embedding, x[:, self.num_categorical_features :]], axis=1)69            )70        return x_embedding71 72 73class OldAtomEncoder(torch.nn.Module):74    def __init__(self, emb_dim, feature_dims, sigma_embed_dim, lm_embedding_type=None):75        """76 77        Parameters78        ----------79        emb_dim80        feature_dims81            first element of feature_dims tuple is a list with the length of each categorical feature,82            and the second is the number of scalar features83        sigma_embed_dim84        lm_embedding_type85        """86        super(OldAtomEncoder, self).__init__()87        self.atom_embedding_list = torch.nn.ModuleList()88        self.num_categorical_features = len(feature_dims[0])89        self.num_scalar_features = feature_dims[1] + sigma_embed_dim90        self.lm_embedding_type = lm_embedding_type91        for i, dim in enumerate(feature_dims[0]):92            emb = torch.nn.Embedding(dim, emb_dim)93            torch.nn.init.xavier_uniform_(emb.weight.data)94            self.atom_embedding_list.append(emb)95 96        if self.num_scalar_features > 0:97            self.linear = torch.nn.Linear(self.num_scalar_features, emb_dim)98        if self.lm_embedding_type is not None:99            if self.lm_embedding_type == "esm":100                self.lm_embedding_dim = 1280101            else:102                raise ValueError(103                    "LM Embedding type was not correctly determined. LM embedding type: ",104                    self.lm_embedding_type,105                )106            self.lm_embedding_layer = torch.nn.Linear(self.lm_embedding_dim + emb_dim, emb_dim)107 108    def forward(self, x):109        x_embedding = 0110        if self.lm_embedding_type is not None:111            assert (112                x.shape[1]113                == self.num_categorical_features + self.num_scalar_features + self.lm_embedding_dim114            )115        else:116            assert x.shape[1] == self.num_categorical_features + self.num_scalar_features117        for i in range(self.num_categorical_features):118            x_embedding += self.atom_embedding_list[i](x[:, i].long())119 120        if self.num_scalar_features > 0:121            x_embedding += self.linear(122                x[:, self.num_categorical_features : self.num_categorical_features + self.num_scalar_features]123            )124        if self.lm_embedding_type is not None:125            x_embedding = self.lm_embedding_layer(126                torch.cat([x_embedding, x[:, -self.lm_embedding_dim :]], axis=1)127            )128        return x_embedding129