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

Function main

training/dist_prefixlm_train.py:190–338  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

188
189
190def main():
191 parser = argparse.ArgumentParser(description='Gpipe-GPT')
192 add_device_arguments(parser)
193 add_torch_distributed_arguments(parser)
194 add_model_arguments(parser)
195 add_task_arguments(parser)
196 add_training_hyper_parameter_arguments(parser)
197 add_mixed_precision_arguments(parser)
198 add_parallel_schema_arguments(parser)
199 parser.add_argument('--model-name', type=str, default='gpt2', metavar='S',
200 help='model name or path')
201 parser.add_argument('--tokenizer-name', type=str, default='gpt2', metavar='S',
202 help='tokenizer name or path')
203 parser.add_argument('--model-type', type=str, default='gpt2', metavar='S',
204 help='model name or path')
205 parser.add_argument('--checkpoint-path', type=str, default='model_checkpoints/gpt2')
206 parser.add_argument('--task-name', type=str, default='cot', metavar='S',
207 help='task name')
208 parser.add_argument('--warmup-steps', type=int, default=0, help='-')
209 parser.add_argument('--train-warmup-steps', type=int, default=0, help='-')
210 parser.add_argument('--total-steps', type=int, default=None, help='-')
211 parser.add_argument('--load-pretrained-model',
212 type=lambda x: x.lower()=='true', default=True, metavar='S',
213 help='load pretrained model or not.')
214 parser.add_argument('--load-checkpoint',
215 type=lambda x: x.lower()=='true', default=True, metavar='S',
216 help='load pretrained model or not.')
217 parser.add_argument('--seed', type=int, default=1, metavar='S',
218 help='random seed (default: 1)')
219 parser.add_argument('--profiling', type=str, default='no-profiling', metavar='S',
220 help='enable which profiling? default: tidy mode')
221 parser.add_argument('--trace-postfix', type=str, default='default', metavar='S',
222 help='postfix of the tracing file name.')
223 parser.add_argument('--evaluation-steps',
224 type=int, default=0, metavar='S',
225 help='every x steps, do evaluation. (0 means do not do evaluation)')
226 parser.add_argument('--evaluation-data',
227 type=str, default=None, help="path of eval data in jsonl")
228 parser.add_argument('--evaluation-num-batch',
229 type=int, default=None, help="for debug purpose, only eval the first several batch.")
230 parser.add_argument('--checkpoint-steps',
231 type=int, default=0, metavar='S',
232 help='every x steps, save checkpoint. (0 means do not save checkpoint)')
233 parser.add_argument('--net-interface',
234 type=str, default='lo', metavar='S',
235 help='net_interface')
236 parser.add_argument('--job-id',
237 type=str, default="0", metavar='S',
238 help='an uuid')
239 args = parser.parse_args()
240
241 torch.manual_seed(args.seed)
242 random.seed(args.seed)
243 np.random.seed(args.seed)
244
245 if args.use_cuda:
246 assert (torch.cuda.is_available())
247 device = torch.device('cuda', args.cuda_id)

Callers 1

Calls 15

build_tokenizerFunction · 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
get_data_parallel_commFunction · 0.85

Tested by

no test coverage detected