Skip to content

Commit 1d35357

Browse files
committed
SphereNet: adding new node feature
1 parent 69c3e7e commit 1d35357

1 file changed

Lines changed: 22 additions & 6 deletions

File tree

‎dig/threedgraph/method/spherenet/spherenet.py‎

Lines changed: 22 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -51,17 +51,21 @@ def forward(self, x):
5151

5252

5353
class init(torch.nn.Module):
54-
def __init__(self, num_radial, hidden_channels, act=swish, use_node_features=True):
54+
def __init__(self, num_radial, hidden_channels, act=swish, use_node_features=True, use_extra_node_feature=False):
5555
super(init, self).__init__()
5656
self.act = act
5757
self.use_node_features = use_node_features
58+
self.use_extra_node_feature = use_extra_node_feature
5859
if self.use_node_features:
5960
self.emb = Embedding(95, hidden_channels)
6061
else: # option to use no node features and a learned embedding vector for each node instead
6162
self.node_embedding = nn.Parameter(torch.empty((hidden_channels,)))
6263
nn.init.normal_(self.node_embedding)
6364
self.lin_rbf_0 = Linear(num_radial, hidden_channels)
64-
self.lin = Linear(3 * hidden_channels, hidden_channels)
65+
if self.use_extra_node_feature:
66+
self.lin = Linear(5 * hidden_channels, hidden_channels)
67+
else:
68+
self.lin = Linear(3 * hidden_channels, hidden_channels)
6569
self.lin_rbf_1 = nn.Linear(num_radial, hidden_channels, bias=False)
6670
self.reset_parameters()
6771

@@ -72,12 +76,14 @@ def reset_parameters(self):
7276
self.lin.reset_parameters()
7377
glorot_orthogonal(self.lin_rbf_1.weight, scale=2.0)
7478

75-
def forward(self, x, emb, i, j):
79+
def forward(self, x, node_feature, emb, i, j):
7680
rbf,_,_ = emb
7781
if self.use_node_features:
7882
x = self.emb(x)
7983
else:
8084
x = self.node_embedding[None, :].expand(x.shape[0], -1)
85+
if node_feature != None and self.use_extra_node_feature:
86+
x = torch.cat((x, node_feature), 1)
8187
rbf0 = self.act(self.lin_rbf_0(rbf))
8288
e1 = self.act(self.lin(torch.cat([x[i], x[j], rbf0], dim=-1)))
8389
e2 = self.lin_rbf_1(rbf) * e1
@@ -250,13 +256,17 @@ def __init__(
250256
basis_emb_size_dist=8, basis_emb_size_angle=8, basis_emb_size_torsion=8, out_emb_channels=256,
251257
num_spherical=7, num_radial=6, envelope_exponent=5,
252258
num_before_skip=1, num_after_skip=2, num_output_layers=3,
253-
act=swish, output_init='GlorotOrthogonal', use_node_features=True):
259+
act=swish, output_init='GlorotOrthogonal', use_node_features=True, use_extra_node_feature=False, extra_node_feature_dim=1):
254260
super(SphereNet, self).__init__()
255261

256262
self.cutoff = cutoff
257263
self.energy_and_force = energy_and_force
264+
self.use_extra_node_feature = use_extra_node_feature
265+
266+
if use_extra_node_feature:
267+
self.extra_emb = Linear(extra_node_feature_dim, hidden_channels)
258268

259-
self.init_e = init(num_radial, hidden_channels, act, use_node_features=use_node_features)
269+
self.init_e = init(num_radial, hidden_channels, act, use_node_features=use_node_features, use_extra_node_feature=use_extra_node_feature)
260270
self.init_v = update_v(hidden_channels, out_emb_channels, out_channels, num_output_layers, act, output_init)
261271
self.init_u = update_u()
262272
self.emb = emb(num_spherical, num_radial, self.cutoff, envelope_exponent)
@@ -272,6 +282,8 @@ def __init__(
272282
self.reset_parameters()
273283

274284
def reset_parameters(self):
285+
if self.use_extra_node_feature:
286+
self.extra_emb.reset_parameters()
275287
self.init_e.reset_parameters()
276288
self.init_v.reset_parameters()
277289
self.emb.reset_parameters()
@@ -283,6 +295,10 @@ def reset_parameters(self):
283295

284296
def forward(self, batch_data):
285297
z, pos, batch = batch_data.z, batch_data.pos, batch_data.batch
298+
if self.use_extra_node_feature and batch_data.node_feature != None:
299+
extra_node_feature = self.extra_emb(batch_data.node_feature)
300+
else:
301+
extra_node_feature = None
286302
if self.energy_and_force:
287303
pos.requires_grad_()
288304
edge_index = radius_graph(pos, r=self.cutoff, batch=batch)
@@ -292,7 +308,7 @@ def forward(self, batch_data):
292308
emb = self.emb(dist, angle, torsion, idx_kj)
293309

294310
#Initialize edge, node, graph features
295-
e = self.init_e(z, emb, i, j)
311+
e = self.init_e(z, extra_node_feature, emb, i, j)
296312
v = self.init_v(e, i)
297313
u = self.init_u(torch.zeros_like(scatter(v, batch, dim=0)), v, batch) #scatter(v, batch, dim=0)
298314

0 commit comments

Comments
 (0)