All
This commit is contained in:
@@ -0,0 +1,108 @@
|
||||
import argparse
|
||||
import os
|
||||
from utils import util
|
||||
import torch
|
||||
|
||||
class BaseOptions():
|
||||
def __init__(self):
|
||||
self._parser = argparse.ArgumentParser()
|
||||
self._initialized = False
|
||||
|
||||
def initialize(self):
|
||||
self._parser.add_argument('--data_dir', type=str, help='path to dataset')
|
||||
self._parser.add_argument('--train_ids_file', type=str, default='train_ids.csv', help='file containing train ids')
|
||||
self._parser.add_argument('--test_ids_file', type=str, default='test_ids.csv', help='file containing test ids')
|
||||
self._parser.add_argument('--images_folder', type=str, default='imgs', help='images folder')
|
||||
self._parser.add_argument('--aus_file', type=str, default='aus_openface.pkl', help='file containing samples aus')
|
||||
|
||||
self._parser.add_argument('--load_epoch', type=int, default=-1, help='which epoch to load? set to -1 to use latest cached model')
|
||||
self._parser.add_argument('--batch_size', type=int, default=4, help='input batch size')
|
||||
self._parser.add_argument('--image_size', type=int, default=128, help='input image size')
|
||||
self._parser.add_argument('--cond_nc', type=int, default=17, help='# of conditions')
|
||||
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')
|
||||
self._parser.add_argument('--name', type=str, default='experiment_1', help='name of the experiment. It decides where to store samples and models')
|
||||
self._parser.add_argument('--dataset_mode', type=str, default='aus', help='chooses dataset to be used')
|
||||
self._parser.add_argument('--model', type=str, default='ganimation', help='model to run[au_net_model]')
|
||||
self._parser.add_argument('--n_threads_test', default=1, type=int, help='# threads for loading data')
|
||||
self._parser.add_argument('--checkpoints_dir', type=str, default='./checkpoints', help='models are saved here')
|
||||
self._parser.add_argument('--serial_batches', action='store_true', help='if true, takes images in order to make batches, otherwise takes them randomly')
|
||||
self._parser.add_argument('--do_saturate_mask', action="store_true", default=False, help='do use mask_fake for mask_cyc')
|
||||
|
||||
|
||||
|
||||
|
||||
self._initialized = True
|
||||
|
||||
def parse(self):
|
||||
if not self._initialized:
|
||||
self.initialize()
|
||||
self._opt = self._parser.parse_args()
|
||||
|
||||
# set is train or set
|
||||
self._opt.is_train = self.is_train
|
||||
|
||||
# set and check load_epoch
|
||||
self._set_and_check_load_epoch()
|
||||
|
||||
# get and set gpus
|
||||
self._get_set_gpus()
|
||||
|
||||
args = vars(self._opt)
|
||||
|
||||
# print in terminal args
|
||||
self._print(args)
|
||||
|
||||
# save args to file
|
||||
self._save(args)
|
||||
|
||||
return self._opt
|
||||
|
||||
def _set_and_check_load_epoch(self):
|
||||
models_dir = os.path.join(self._opt.checkpoints_dir, self._opt.name)
|
||||
if os.path.exists(models_dir):
|
||||
if self._opt.load_epoch == -1:
|
||||
load_epoch = 0
|
||||
for file in os.listdir(models_dir):
|
||||
if file.startswith("net_epoch_"):
|
||||
load_epoch = max(load_epoch, int(file.split('_')[2]))
|
||||
self._opt.load_epoch = load_epoch
|
||||
else:
|
||||
found = False
|
||||
for file in os.listdir(models_dir):
|
||||
if file.startswith("net_epoch_"):
|
||||
found = int(file.split('_')[2]) == self._opt.load_epoch
|
||||
if found: break
|
||||
assert found, 'Model for epoch %i not found' % self._opt.load_epoch
|
||||
else:
|
||||
assert self._opt.load_epoch < 1, 'Model for epoch %i not found' % self._opt.load_epoch
|
||||
self._opt.load_epoch = 0
|
||||
|
||||
def _get_set_gpus(self):
|
||||
# get gpu ids
|
||||
str_ids = self._opt.gpu_ids.split(',')
|
||||
self._opt.gpu_ids = []
|
||||
for str_id in str_ids:
|
||||
id = int(str_id)
|
||||
if id >= 0:
|
||||
self._opt.gpu_ids.append(id)
|
||||
|
||||
# set gpu ids
|
||||
if len(self._opt.gpu_ids) > 0:
|
||||
torch.cuda.set_device(self._opt.gpu_ids[0])
|
||||
|
||||
def _print(self, args):
|
||||
print('------------ Options -------------')
|
||||
for k, v in sorted(args.items()):
|
||||
print('%s: %s' % (str(k), str(v)))
|
||||
print('-------------- End ----------------')
|
||||
|
||||
def _save(self, args):
|
||||
expr_dir = os.path.join(self._opt.checkpoints_dir, self._opt.name)
|
||||
print(expr_dir)
|
||||
util.mkdirs(expr_dir)
|
||||
file_name = os.path.join(expr_dir, 'opt_%s.txt' % ('train' if self.is_train else 'test'))
|
||||
with open(file_name, 'wt') as opt_file:
|
||||
opt_file.write('------------ Options -------------\n')
|
||||
for k, v in sorted(args.items()):
|
||||
opt_file.write('%s: %s\n' % (str(k), str(v)))
|
||||
opt_file.write('-------------- End ----------------\n')
|
||||
@@ -0,0 +1,9 @@
|
||||
from .base_options import BaseOptions
|
||||
|
||||
|
||||
class TestOptions(BaseOptions):
|
||||
def initialize(self):
|
||||
BaseOptions.initialize(self)
|
||||
self._parser.add_argument('--input_path', type=str, help='path to image')
|
||||
self._parser.add_argument('--output_dir', type=str, default='./output', help='output path')
|
||||
self.is_train = False
|
||||
@@ -0,0 +1,31 @@
|
||||
from .base_options import BaseOptions
|
||||
|
||||
|
||||
class TrainOptions(BaseOptions):
|
||||
def initialize(self):
|
||||
BaseOptions.initialize(self)
|
||||
self._parser.add_argument('--n_threads_train', default=4, type=int, help='# threads for loading data')
|
||||
self._parser.add_argument('--num_iters_validate', default=1, type=int, help='# batches to use when validating')
|
||||
self._parser.add_argument('--print_freq_s', type=int, default=60, help='frequency of showing training results on console')
|
||||
self._parser.add_argument('--display_freq_s', type=int, default=300, help='frequency [s] of showing training results on screen')
|
||||
self._parser.add_argument('--save_latest_freq_s', type=int, default=3600, help='frequency of saving the latest results')
|
||||
|
||||
self._parser.add_argument('--nepochs_no_decay', type=int, default=20, help='# of epochs at starting learning rate')
|
||||
self._parser.add_argument('--nepochs_decay', type=int, default=10, help='# of epochs to linearly decay learning rate to zero')
|
||||
|
||||
self._parser.add_argument('--train_G_every_n_iterations', type=int, default=5, help='train G every n interations')
|
||||
self._parser.add_argument('--poses_g_sigma', type=float, default=0.06, help='initial learning rate for adam')
|
||||
self._parser.add_argument('--lr_G', type=float, default=0.0001, help='initial learning rate for G adam')
|
||||
self._parser.add_argument('--G_adam_b1', type=float, default=0.5, help='beta1 for G adam')
|
||||
self._parser.add_argument('--G_adam_b2', type=float, default=0.999, help='beta2 for G adam')
|
||||
self._parser.add_argument('--lr_D', type=float, default=0.0001, help='initial learning rate for D adam')
|
||||
self._parser.add_argument('--D_adam_b1', type=float, default=0.5, help='beta1 for D adam')
|
||||
self._parser.add_argument('--D_adam_b2', type=float, default=0.999, help='beta2 for D adam')
|
||||
self._parser.add_argument('--lambda_D_prob', type=float, default=1, help='lambda for real/fake discriminator loss')
|
||||
self._parser.add_argument('--lambda_D_cond', type=float, default=4000, help='lambda for condition discriminator loss')
|
||||
self._parser.add_argument('--lambda_cyc', type=float, default=10, help='lambda cycle loss')
|
||||
self._parser.add_argument('--lambda_mask', type=float, default=0.1, help='lambda mask loss')
|
||||
self._parser.add_argument('--lambda_D_gp', type=float, default=10, help='lambda gradient penalty loss')
|
||||
self._parser.add_argument('--lambda_mask_smooth', type=float, default=1e-5, help='lambda mask smooth loss')
|
||||
|
||||
self.is_train = True
|
||||
Reference in New Issue
Block a user