| 85 | return eval_results |
| 86 | |
| 87 | def parse_args(): |
| 88 | parser = argparse.ArgumentParser(description='Critique T5') |
| 89 | |
| 90 | parser.add_argument('--training-file', dest='training_file', required=False, help='Path to training file', |
| 91 | default=None) |
| 92 | parser.add_argument('--noisy-file', dest='noisy_file', required=False, help='Path to noisy file', |
| 93 | default=None) |
| 94 | parser.add_argument('--validation-file', dest='validation_file', required=False, help='Path to validation file') |
| 95 | parser.add_argument('--language-model', default='t5-base', help='Can be either some huggingface model or a ' |
| 96 | 'path to a model. If the path is in GCS we ' |
| 97 | 'download it first.') |
| 98 | parser.add_argument('--model-dir', dest='model_dir', required=True, |
| 99 | help='the folder/google bucket in which the model will be stored or loaded from.') |
| 100 | parser.add_argument('--critique_model-dir', dest='critique_model_dir', required=True, |
| 101 | help='the folder/google bucket in which the model will be stored or loaded from.') |
| 102 | parser.add_argument('--epochs', default=20, |
| 103 | help='number of epochs to train') |
| 104 | parser.add_argument('--batch-size', default=1, |
| 105 | help='batch size') |
| 106 | parser.add_argument('--val-batch-size', default=1, |
| 107 | help='validation batch size') |
| 108 | parser.add_argument('--number_turn', default=4, |
| 109 | help='learning rate') |
| 110 | parser.add_argument('--lr', default=0.0001, |
| 111 | help='learning rate') |
| 112 | parser.add_argument('--gradient-accumulation', default=1) |
| 113 | parser.add_argument('--local_rank', default=-1) |
| 114 | args = parser.parse_args() |
| 115 | |
| 116 | return args |
| 117 | |
| 118 | |
| 119 | def main(): |