MCPcopy Create free account
hub / github.com/albertpumarola/GANimation / BaseOptions

Class BaseOptions

options/base_options.py:6–108  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

4import torch
5
6class 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:

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected