| 31 | from torch.utils.data import DataLoader |
| 32 | |
| 33 | def parse_args(): |
| 34 | parser = argparse.ArgumentParser() |
| 35 | parser.add_argument("--config", type=str, default="configs/pretrain.yaml", help="Path to the config file.") |
| 36 | parser.add_argument('-c', '--checkpoint', default='checkpoint', type=str, metavar='PATH', help='checkpoint directory') |
| 37 | parser.add_argument('-p', '--pretrained', default='checkpoint', type=str, metavar='PATH', help='pretrained checkpoint directory') |
| 38 | parser.add_argument('-r', '--resume', default='', type=str, metavar='FILENAME', help='checkpoint to resume (file name)') |
| 39 | parser.add_argument('-e', '--evaluate', default='', type=str, metavar='FILENAME', help='checkpoint to evaluate (file name)') |
| 40 | parser.add_argument('-freq', '--print_freq', default=100) |
| 41 | parser.add_argument('-ms', '--selection', default='latest_epoch.bin', type=str, metavar='FILENAME', help='checkpoint to finetune (file name)') |
| 42 | parser.add_argument('-sd', '--seed', default=0, type=int, help='random seed') |
| 43 | opts = parser.parse_args() |
| 44 | return opts |
| 45 | |
| 46 | def set_random_seed(seed): |
| 47 | random.seed(seed) |