| 293 | |
| 294 | |
| 295 | def 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 | |
| 320 | if __name__ == "__main__": |