MCPcopy Create free account
hub / github.com/IBM/Project_CodeNet / init_model

Function init_model

model-experiments/gnn-based-experiments/src/main.py:295–317  ·  view source on GitHub ↗
(args, node_encoder, numclass=275, edge_attr_dim=2)

Source from the content-addressed store, hash-verified

293
294
295def init_model(args, node_encoder, numclass=275, edge_attr_dim=2):
296 # this was only relevant for regression version
297 n, m = 1, 1
298 if args.gnn == 'gin':
299 model = GNN(num_vocab=n, max_seq_len=m, node_encoder=node_encoder,
300 num_layer=args.num_layer, gnn_type='gin', emb_dim=args.emb_dim, drop_ratio=args.drop_ratio,
301 virtual_node=False, num_class=numclass, edge_attr_dim=edge_attr_dim)
302 elif args.gnn == 'gin-virtual':
303 model = GNN(num_vocab=n, max_seq_len=m, node_encoder=node_encoder,
304 num_layer=args.num_layer, gnn_type='gin', emb_dim=args.emb_dim, drop_ratio=args.drop_ratio,
305 virtual_node=True, num_class=numclass, edge_attr_dim=edge_attr_dim)
306 elif args.gnn == 'gcn':
307 model = GNN(num_vocab=n, max_seq_len=m, node_encoder=node_encoder,
308 num_layer=args.num_layer, gnn_type='gcn', emb_dim=args.emb_dim, drop_ratio=args.drop_ratio,
309 virtual_node=False, num_class=numclass, edge_attr_dim=edge_attr_dim)
310 elif args.gnn == 'gcn-virtual':
311 model = GNN(num_vocab=n, max_seq_len=m, node_encoder=node_encoder,
312 num_layer=args.num_layer, gnn_type='gcn', emb_dim=args.emb_dim, drop_ratio=args.drop_ratio,
313 virtual_node=True, num_class=numclass, edge_attr_dim=edge_attr_dim)
314 else:
315 raise ValueError('Invalid GNN type')
316
317 return model
318
319
320if __name__ == "__main__":

Callers 1

mainFunction · 0.85

Calls 1

GNNClass · 0.90

Tested by

no test coverage detected