| 27 | from lib.model.loss import * |
| 28 | |
| 29 | def parse_args(): |
| 30 | parser = argparse.ArgumentParser() |
| 31 | parser.add_argument("--config", type=str, default="configs/pretrain.yaml", help="Path to the config file.") |
| 32 | parser.add_argument('-c', '--checkpoint', default='checkpoint', type=str, metavar='PATH', help='checkpoint directory') |
| 33 | parser.add_argument('-p', '--pretrained', default='checkpoint', type=str, metavar='PATH', help='pretrained checkpoint directory') |
| 34 | parser.add_argument('-r', '--resume', default='', type=str, metavar='FILENAME', help='checkpoint to resume (file name)') |
| 35 | parser.add_argument('-e', '--evaluate', default='', type=str, metavar='FILENAME', help='checkpoint to evaluate (file name)') |
| 36 | parser.add_argument('-ms', '--selection', default='latest_epoch.bin', type=str, metavar='FILENAME', help='checkpoint to finetune (file name)') |
| 37 | parser.add_argument('-sd', '--seed', default=0, type=int, help='random seed') |
| 38 | opts = parser.parse_args() |
| 39 | return opts |
| 40 | |
| 41 | def set_random_seed(seed): |
| 42 | random.seed(seed) |