| 1 | import argparse |
| 2 | |
| 3 | def parse_args(*args): |
| 4 | parser = argparse.ArgumentParser(description='MOSS-RLHF Reward Model @Fudan NLP Group') |
| 5 | # training settings |
| 6 | parser.add_argument('--seed', type=int, default=42, help='seed') |
| 7 | parser.add_argument('--lr', type=float, default=5e-6, help='learning rate of reward model') |
| 8 | parser.add_argument('--batch_size', type=int, default=8, help='training batch size for single GPU') |
| 9 | parser.add_argument('--gradient_checkpoint', action='store_true', help='deepspeed') |
| 10 | parser.add_argument('--reward_lm_loss_factor', type=float, default=0., help='calculate lm loss on rm model') |
| 11 | parser.add_argument('--warmup_steps', type=int, default=500, help='warmup steps') |
| 12 | parser.add_argument('--train_steps', type=int, default=10000, help='train steps') |
| 13 | parser.add_argument('--fp32_loss', action='store_true', help='use fp32 to calculate cross-entropy loss, enable when numeric stability problem occurs') |
| 14 | parser.add_argument('--save_per_step', type=int, default=200, help='save ckpt and save validation tensorboard per steps') |
| 15 | parser.add_argument('--print_interval', type=int, default=5, help='print training state and save training tensorboard per steps') |
| 16 | parser.add_argument('--validation_metric', type=str, default='loss', help='metric to select the best model') |
| 17 | |
| 18 | # Optimizer , Scheduler and Dataloader |
| 19 | parser.add_argument('--beta1', type=float, default=0.9, help='adam') |
| 20 | parser.add_argument('--beta2', type=float, default=0.95, help='adam') |
| 21 | parser.add_argument('--eps', type=float, default=1e-6, help='optimizer') |
| 22 | parser.add_argument('--num_prefetch', type=int, default=32, help='dataloader') |
| 23 | parser.add_argument('--num_workers', type=int, default=1, help='dataloader') |
| 24 | parser.add_argument('--weight_decay', type=float, default=0., help='l2 weight decay') |
| 25 | |
| 26 | # Path |
| 27 | parser.add_argument('--data_path', type=str, default='./data', help='dataset for training and validation') |
| 28 | parser.add_argument('--init_checkpoint_model', type=str, default=None, help='checkpoint used to initialize the model, used for fine-tuning') |
| 29 | parser.add_argument('--logdir', type=str, default=None, help='path to save tensorboard logs') |
| 30 | parser.add_argument('--model_save_path', type=str, default='./outputs/', help='checkpoint path, used for save model and training') |
| 31 | parser.add_argument('--hf_model_name_or_path', type=str, default='meta-llama/Llama-2-7b-hf', help='Hugging model name used to load tokenizer, configs and pretained models') |
| 32 | |
| 33 | # LLM settings |
| 34 | parser.add_argument('--context_truncate', type=int, default=2048, help='max length for history') |
| 35 | parser.add_argument('--delimiter', type=str, default='\n', help='delimiter to seperate dialog history') |
| 36 | |
| 37 | |
| 38 | args = parser.parse_args() |
| 39 | return args |
| 40 | |
| 41 | |