@@ -51,17 +51,21 @@ def forward(self, x):
5151
5252
5353class 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