MCPcopy Create free account
hub / github.com/togethercomputer/OpenChatKit / main

Function main

training/dist_clm_train.py:214–355  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

212
213
214def main():
215 parser = argparse.ArgumentParser(description='Gpipe-GPT')
216 add_device_arguments(parser)
217 add_torch_distributed_arguments(parser)
218 add_model_arguments(parser)
219 add_task_arguments(parser)
220 add_training_hyper_parameter_arguments(parser)
221 add_mixed_precision_arguments(parser)
222 add_parallel_schema_arguments(parser)
223 parser.add_argument('--model-name', type=str, default='gpt2', metavar='S',
224 help='model name or path')
225 parser.add_argument('--tokenizer-name', type=str, default='gpt2', metavar='S',
226 help='tokenizer name or path')
227 parser.add_argument('--model-type', type=str, default='gpt2', metavar='S',
228 help='model name or path')
229 parser.add_argument('--checkpoint-path', type=str, default='model_checkpoints/gpt2')
230 parser.add_argument('--task-name', type=str, default='cot', metavar='S',
231 help='task name')
232 parser.add_argument('--warmup-steps', type=int, default=0, help='-')
233 parser.add_argument('--train-warmup-steps', type=int, default=0, help='-')
234 parser.add_argument('--total-steps', type=int, default=None, help='-')
235 parser.add_argument('--load-pretrained-model',
236 type=lambda x: x.lower()=='true', default=True, metavar='S',
237 help='load pretrained model or not.')
238 parser.add_argument('--load-checkpoint',
239 type=lambda x: x.lower()=='true', default=True, metavar='S',
240 help='load pretrained model or not.')
241 parser.add_argument('--seed', type=int, default=1, metavar='S',
242 help='random seed (default: 1)')
243 parser.add_argument('--profiling', type=str, default='no-profiling', metavar='S',
244 help='enable which profiling? default: tidy mode')
245 parser.add_argument('--trace-postfix', type=str, default='default', metavar='S',
246 help='postfix of the tracing file name.')
247 parser.add_argument('--evaluation-steps',
248 type=int, default=0, metavar='S',
249 help='every x steps, do evaluation. (0 means do not do evaluation)')
250 parser.add_argument('--evaluation-data',
251 type=str, default=None, help="path of eval data in jsonl")
252 parser.add_argument('--evaluation-num-batch',
253 type=int, default=None, help="for debug purpose, only eval the first several batch.")
254 parser.add_argument('--checkpoint-steps',
255 type=int, default=0, metavar='S',
256 help='every x steps, save checkpoint. (0 means do not save checkpoint)')
257 parser.add_argument('--net-interface',
258 type=str, default='lo', metavar='S',
259 help='net_interface')
260 parser.add_argument('--job-id',
261 type=str, default="0", metavar='S',
262 help='an uuid')
263 args = parser.parse_args()
264
265 torch.manual_seed(args.seed)
266 random.seed(args.seed)
267 np.random.seed(args.seed)
268
269 if args.use_cuda:
270 assert (torch.cuda.is_available())
271 device = torch.device('cuda', args.cuda_id)

Callers 1

dist_clm_train.pyFile · 0.70

Calls 15

build_tokenizerFunction · 0.90
get_train_data_loaderFunction · 0.90
get_eval_data_loaderFunction · 0.90
get_pp_moduleFunction · 0.90
add_device_argumentsFunction · 0.85
add_model_argumentsFunction · 0.85
add_task_argumentsFunction · 0.85
init_communicatorsFunction · 0.85

Tested by

no test coverage detected