()
| 188 | |
| 189 | |
| 190 | def 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) |
no test coverage detected