Download models/property_pred/prop_egnn.py from OneScience-Group/TargetDiff: direct link, hf CLI and curl.
- Browser
- Download file 3.48 kB
-
https://huggingface.co/OneScience-Group/TargetDiff/resolve/main/models/property_pred/prop_egnn.py
- Command line
-
hf download hf://OneScience-Group/TargetDiff/models/property_pred/prop_egnn.py
-
curl -L -o prop_egnn.py https://huggingface.co/OneScience-Group/TargetDiff/resolve/main/models/property_pred/prop_egnn.py
3.48 kB
| import torch | |
| import torch.nn as nn | |
| from torch_scatter import scatter_sum | |
| from torch_geometric.nn import radius_graph, knn_graph | |
| from ..common import GaussianSmearing, MLP | |
| class EnBaseLayer(nn.Module): | |
| def __init__(self, hidden_dim, edge_feat_dim, num_r_gaussian, update_x=True, act_fn='relu', norm=False): | |
| super().__init__() | |
| self.r_min = 0. | |
| self.r_max = 10. ** 2 | |
| self.hidden_dim = hidden_dim | |
| self.num_r_gaussian = num_r_gaussian | |
| self.edge_feat_dim = edge_feat_dim | |
| self.update_x = update_x | |
| self.act_fn = act_fn | |
| self.norm = norm | |
| if num_r_gaussian > 1: | |
| self.r_expansion = GaussianSmearing(self.r_min, self.r_max, num_gaussians=num_r_gaussian, fixed_offset=False) | |
| self.edge_mlp = MLP(2 * hidden_dim + edge_feat_dim + num_r_gaussian, hidden_dim, hidden_dim, | |
| num_layer=2, norm=norm, act_fn=act_fn, act_last=True) | |
| self.edge_inf = nn.Sequential(nn.Linear(hidden_dim, 1), nn.Sigmoid()) | |
| if self.update_x: | |
| self.x_mlp = MLP(hidden_dim, 1, hidden_dim, num_layer=2, norm=norm, act_fn=act_fn) | |
| self.node_mlp = MLP(2 * hidden_dim, hidden_dim, hidden_dim, num_layer=2, norm=norm, act_fn=act_fn) | |
| def forward(self, h, edge_index, edge_attr): | |
| dst, src = edge_index | |
| hi, hj = h[dst], h[src] | |
| # \phi_e in Eq(3) | |
| mij = self.edge_mlp(torch.cat([edge_attr, hi, hj], -1)) | |
| eij = self.edge_inf(mij) | |
| mi = scatter_sum(mij * eij, dst, dim=0, dim_size=h.shape[0]) | |
| # h update in Eq(6) | |
| # h = h + self.node_mlp(torch.cat([mi, h], -1)) | |
| output = self.node_mlp(torch.cat([mi, h], -1)) | |
| # if self.update_x: | |
| # # x update in Eq(4) | |
| # xi, xj = x[dst], x[src] | |
| # delta_x = scatter_sum((xi - xj) * self.x_mlp(mij), dst, dim=0) | |
| # x = x + delta_x | |
| return output | |
| class EnEquiEncoder(nn.Module): | |
| def __init__(self, num_layers, hidden_dim, edge_feat_dim, num_r_gaussian, k=32, cutoff=10.0, | |
| update_x=True, act_fn='relu', norm=False): | |
| super().__init__() | |
| # Build the network | |
| self.num_layers = num_layers | |
| self.hidden_dim = hidden_dim | |
| self.edge_feat_dim = edge_feat_dim | |
| self.num_r_gaussian = num_r_gaussian | |
| self.update_x = update_x | |
| self.act_fn = act_fn | |
| self.norm = norm | |
| self.k = k | |
| self.cutoff = cutoff | |
| self.distance_expansion = GaussianSmearing(stop=cutoff, num_gaussians=num_r_gaussian, fixed_offset=False) | |
| self.net = self._build_network() | |
| def _build_network(self): | |
| # Equivariant layers | |
| layers = [] | |
| for l_idx in range(self.num_layers): | |
| layer = EnBaseLayer(self.hidden_dim, self.edge_feat_dim, self.num_r_gaussian, | |
| update_x=self.update_x, act_fn=self.act_fn, norm=self.norm) | |
| layers.append(layer) | |
| return nn.ModuleList(layers) | |
| def forward(self, node_attr, pos, batch): | |
| # edge_index = radius_graph(pos, self.cutoff, batch=batch, loop=False) | |
| edge_index = knn_graph(pos, k=self.k, batch=batch, flow='target_to_source') | |
| edge_length = torch.norm(pos[edge_index[0]] - pos[edge_index[1]], dim=1) | |
| edge_attr = self.distance_expansion(edge_length) | |
| h = node_attr | |
| for interaction in self.net: | |
| h = h + interaction(h, edge_index, edge_attr) | |
| return h | |