| 4 | import torch |
| 5 | |
| 6 | class BaseOptions(): |
| 7 | def __init__(self): |
| 8 | self._parser = argparse.ArgumentParser() |
| 9 | self._initialized = False |
| 10 | |
| 11 | def initialize(self): |
| 12 | self._parser.add_argument('--data_dir', type=str, help='path to dataset') |
| 13 | self._parser.add_argument('--train_ids_file', type=str, default='train_ids.csv', help='file containing train ids') |
| 14 | self._parser.add_argument('--test_ids_file', type=str, default='test_ids.csv', help='file containing test ids') |
| 15 | self._parser.add_argument('--images_folder', type=str, default='imgs', help='images folder') |
| 16 | self._parser.add_argument('--aus_file', type=str, default='aus_openface.pkl', help='file containing samples aus') |
| 17 | |
| 18 | self._parser.add_argument('--load_epoch', type=int, default=-1, help='which epoch to load? set to -1 to use latest cached model') |
| 19 | self._parser.add_argument('--batch_size', type=int, default=4, help='input batch size') |
| 20 | self._parser.add_argument('--image_size', type=int, default=128, help='input image size') |
| 21 | self._parser.add_argument('--cond_nc', type=int, default=17, help='# of conditions') |
| 22 | self._parser.add_argument('--gpu_ids', type=str, default='0', help='gpu ids: e.g. 0 0,1,2, 0,2. use -1 for CPU') |
| 23 | self._parser.add_argument('--name', type=str, default='experiment_1', help='name of the experiment. It decides where to store samples and models') |
| 24 | self._parser.add_argument('--dataset_mode', type=str, default='aus', help='chooses dataset to be used') |
| 25 | self._parser.add_argument('--model', type=str, default='ganimation', help='model to run[au_net_model]') |
| 26 | self._parser.add_argument('--n_threads_test', default=1, type=int, help='# threads for loading data') |
| 27 | self._parser.add_argument('--checkpoints_dir', type=str, default='./checkpoints', help='models are saved here') |
| 28 | self._parser.add_argument('--serial_batches', action='store_true', help='if true, takes images in order to make batches, otherwise takes them randomly') |
| 29 | self._parser.add_argument('--do_saturate_mask', action="store_true", default=False, help='do use mask_fake for mask_cyc') |
| 30 | |
| 31 | |
| 32 | |
| 33 | |
| 34 | self._initialized = True |
| 35 | |
| 36 | def parse(self): |
| 37 | if not self._initialized: |
| 38 | self.initialize() |
| 39 | self._opt = self._parser.parse_args() |
| 40 | |
| 41 | # set is train or set |
| 42 | self._opt.is_train = self.is_train |
| 43 | |
| 44 | # set and check load_epoch |
| 45 | self._set_and_check_load_epoch() |
| 46 | |
| 47 | # get and set gpus |
| 48 | self._get_set_gpus() |
| 49 | |
| 50 | args = vars(self._opt) |
| 51 | |
| 52 | # print in terminal args |
| 53 | self._print(args) |
| 54 | |
| 55 | # save args to file |
| 56 | self._save(args) |
| 57 | |
| 58 | return self._opt |
| 59 | |
| 60 | def _set_and_check_load_epoch(self): |
| 61 | models_dir = os.path.join(self._opt.checkpoints_dir, self._opt.name) |
| 62 | if os.path.exists(models_dir): |
| 63 | if self._opt.load_epoch == -1: |
nothing calls this directly
no outgoing calls
no test coverage detected