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

Function main

model-experiments/gnn-based-experiments/src/main.py:88–292  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

86
87
88def main():
89 parser = argparse.ArgumentParser(description='GNN baselines on Project Codenet data with Pytorch Geometrics')
90 parser.add_argument('--device', type=int, default=0,
91 help='which gpu to use if any (default: 0)')
92 parser.add_argument('--gnn', type=str, default="gcn",
93 help='GNN gin, gin-virtual, or gcn, or gcn-virtual (default: gcn), ...')
94 parser.add_argument('--drop_ratio', type=float, default=0,
95 help='dropout ratio (default: 0)')
96 parser.add_argument('--num_layer', type=int, default=5,
97 help='number of GNN message passing layers (default: 5)')
98 parser.add_argument('--emb_dim', type=int, default=300,
99 help='dimensionality of hidden units in GNNs (default: 300)')
100 parser.add_argument('--feat_nums', type=str, default="",
101 help='comma separated string containing number of categories per feature '
102 '(default: "", meaning "let it be computed")')
103 parser.add_argument('--batch_size', type=int, default=80,
104 help='input batch size for training (default: 128)')
105 parser.add_argument('--epochs', type=int, default=1000,
106 help='number of epochs to train (default: 1000)')
107 parser.add_argument('--num_workers', type=int, default=0,
108 help='number of workers (default: 0)')
109 parser.add_argument('--dataset', type=str, default="small", choices={"small","Java250", "Python800", "C++1000", "C++1400"},
110 help='dataset name (default: python1k)')
111
112 parser.add_argument('--filename', type=str, default="test",
113 help='filename to output result (default: test)')
114
115 parser.add_argument('--dir_data', type=str, default=os.path.join(PATH, 'data'),
116 help='directory where data should be stored (default: REPOSITORY/data)')
117 parser.add_argument('--dir_results', type=str, default=os.path.join(PATH, 'results'),
118 help='results directory (default: REPOSITORY/results)')
119 parser.add_argument('--dir_save', default=os.path.join(PATH, 'saved_models'),
120 help='directory to save checkpoints in (default: REPOSITORY/saved_models)')
121 parser.add_argument('--checkpointing', default=0, type=int, choices=[0, 1],
122 help='if you want to use checkpointing (1) or not (0) (default: 0)')
123 parser.add_argument('--checkpoint', default="",
124 help='path of checkpoint if any')
125 parser.add_argument('--runs', default=10, type=int,
126 help='number of runs (default: 10)')
127 parser.add_argument('--clip', default=0, type=float,
128 help='clipping value if gradient clipping should be uses (default: 0) ')
129 parser.add_argument('--lr', default=1e-3, type=float,
130 help='learning rate (default: 1e-3)')
131 parser.add_argument('--patience', default=20, type=float,
132 help='patience (default: 20)')
133 ###
134
135 args = parser.parse_args()
136 device = torch.device("cuda:" + str(args.device)) if torch.cuda.is_available() else torch.device("cpu")
137
138 os.makedirs(args.dir_results, exist_ok=True)
139 os.makedirs(args.dir_save, exist_ok=True)
140
141 save_args(args, os.path.join(args.dir_results, "a_" + args.filename + '.csv'))
142
143 train_file = os.path.join(args.dir_results, 'd_' + args.filename + '.csv')
144 if not os.path.exists(train_file):
145 with open(train_file, 'w') as f:

Callers 1

main.pyFile · 0.70

Calls 15

get_idx_splitMethod · 0.95
augment_edgeFunction · 0.90
ASTNodeEncoderClass · 0.90
DataLoaderClass · 0.90
DataParallelClass · 0.90
save_argsFunction · 0.85
load_checkpoint_resultsFunction · 0.85
init_modelFunction · 0.85
load_checkpointFunction · 0.85
evalFunction · 0.85
create_checkpointFunction · 0.85

Tested by

no test coverage detected