radames/Text2Human-API
1
1import argparse2import logging3import os.path as osp4import random5 6import torch7 8from data.pose_attr_dataset import DeepFashionAttrPoseDataset9from models import create_model10from utils.logger import get_root_logger11from utils.options import dict2str, dict_to_nonedict, parse12from utils.util import make_exp_dirs, set_random_seed13 14 15def main():16 # options17 parser = argparse.ArgumentParser()18 parser.add_argument('-opt', type=str, help='Path to option YAML file.')19 args = parser.parse_args()20 opt = parse(args.opt, is_train=False)21 22 # mkdir and loggers23 make_exp_dirs(opt)24 log_file = osp.join(opt['path']['log'], f"test_{opt['name']}.log")25 logger = get_root_logger(26 logger_name='base', log_level=logging.INFO, log_file=log_file)27 logger.info(dict2str(opt))28 29 # convert to NoneDict, which returns None for missing keys30 opt = dict_to_nonedict(opt)31 32 # random seed33 seed = opt['manual_seed']34 if seed is None:35 seed = random.randint(1, 10000)36 logger.info(f'Random seed: {seed}')37 set_random_seed(seed)38 39 test_dataset = DeepFashionAttrPoseDataset(40 pose_dir=opt['pose_dir'],41 texture_ann_dir=opt['texture_ann_file'],42 shape_ann_path=opt['shape_ann_path'])43 test_loader = torch.utils.data.DataLoader(44 dataset=test_dataset, batch_size=4, shuffle=False)45 logger.info(f'Number of test set: {len(test_dataset)}.')46 47 model = create_model(opt)48 _ = model.inference(test_loader, opt['path']['results_root'])49 50 51if __name__ == '__main__':52 main()53 