OpenMotionLab/MotionGPT
118
1import importlib2from argparse import ArgumentParser3from omegaconf import OmegaConf4from os.path import join as pjoin5import os6import glob7 8 9def get_module_config(cfg, filepath="./configs"):10 """11 Load yaml config files from subfolders12 """13 14 yamls = glob.glob(pjoin(filepath, '*', '*.yaml'))15 yamls = [y.replace(filepath, '') for y in yamls]16 for yaml in yamls:17 nodes = yaml.replace('.yaml', '').replace('/', '.')18 nodes = nodes[1:] if nodes[0] == '.' else nodes19 OmegaConf.update(cfg, nodes, OmegaConf.load('./configs' + yaml))20 21 return cfg22 23 24def get_obj_from_str(string, reload=False):25 """26 Get object from string27 """28 29 module, cls = string.rsplit(".", 1)30 if reload:31 module_imp = importlib.import_module(module)32 importlib.reload(module_imp)33 return getattr(importlib.import_module(module, package=None), cls)34 35 36def instantiate_from_config(config):37 """38 Instantiate object from config39 """40 if not "target" in config:41 raise KeyError("Expected key `target` to instantiate.")42 return get_obj_from_str(config["target"])(**config.get("params", dict()))43 44 45def resume_config(cfg: OmegaConf):46 """47 Resume model and wandb48 """49 50 if cfg.TRAIN.RESUME:51 resume = cfg.TRAIN.RESUME52 if os.path.exists(resume):53 # Checkpoints54 cfg.TRAIN.PRETRAINED = pjoin(resume, "checkpoints", "last.ckpt")55 # Wandb56 wandb_files = os.listdir(pjoin(resume, "wandb", "latest-run"))57 wandb_run = [item for item in wandb_files if "run-" in item][0]58 cfg.LOGGER.WANDB.params.id = wandb_run.replace("run-","").replace(".wandb", "")59 else:60 raise ValueError("Resume path is not right.")61 62 return cfg63 64def parse_args(phase="train"):65 """66 Parse arguments and load config files67 """68 69 parser = ArgumentParser()70 group = parser.add_argument_group("Training options")71 72 # Assets73 group.add_argument(74 "--cfg_assets",75 type=str,76 required=False,77 default="./configs/assets.yaml",78 help="config file for asset paths",79 )80 81 # Default config82 if phase in ["train", "test"]:83 cfg_defualt = "./configs/default.yaml"84 elif phase == "render":85 cfg_defualt = "./configs/render.yaml"86 elif phase == "webui":87 cfg_defualt = "./configs/webui.yaml"88 89 group.add_argument(90 "--cfg",91 type=str,92 required=False,93 default=cfg_defualt,94 help="config file",95 )96 97 # Parse for each phase98 if phase in ["train", "test"]:99 group.add_argument("--batch_size",100 type=int,101 required=False,102 help="training batch size")103 group.add_argument("--num_nodes",104 type=int,105 required=False,106 help="number of nodes")107 group.add_argument("--device",108 type=int,109 nargs="+",110 required=False,111 help="training device")112 group.add_argument("--task",113 type=str,114 required=False,115 help="evaluation task type")116 group.add_argument("--nodebug",117 action="store_true",118 required=False,119 help="debug or not")120 121 122 if phase == "demo":123 group.add_argument(124 "--example",125 type=str,126 required=False,127 help="input text and lengths with txt format",128 )129 group.add_argument(130 "--out_dir",131 type=str,132 required=False,133 help="output dir",134 )135 group.add_argument("--task",136 type=str,137 required=False,138 help="evaluation task type")139 140 if phase == "render":141 group.add_argument("--npy",142 type=str,143 required=False,144 default=None,145 help="npy motion files")146 group.add_argument("--dir",147 type=str,148 required=False,149 default=None,150 help="npy motion folder")151 group.add_argument("--fps",152 type=int,153 required=False,154 default=30,155 help="render fps")156 group.add_argument(157 "--mode",158 type=str,159 required=False,160 default="sequence",161 help="render target: video, sequence, frame",162 )163 164 params = parser.parse_args()165 166 # Load yaml config files167 OmegaConf.register_new_resolver("eval", eval)168 cfg_assets = OmegaConf.load(params.cfg_assets)169 cfg_base = OmegaConf.load(pjoin(cfg_assets.CONFIG_FOLDER, 'default.yaml'))170 cfg_exp = OmegaConf.merge(cfg_base, OmegaConf.load(params.cfg))171 if not cfg_exp.FULL_CONFIG:172 cfg_exp = get_module_config(cfg_exp, cfg_assets.CONFIG_FOLDER)173 cfg = OmegaConf.merge(cfg_exp, cfg_assets)174 175 # Update config with arguments176 if phase in ["train", "test"]:177 cfg.TRAIN.BATCH_SIZE = params.batch_size if params.batch_size else cfg.TRAIN.BATCH_SIZE178 cfg.DEVICE = params.device if params.device else cfg.DEVICE179 cfg.NUM_NODES = params.num_nodes if params.num_nodes else cfg.NUM_NODES180 cfg.model.params.task = params.task if params.task else cfg.model.params.task181 cfg.DEBUG = not params.nodebug if params.nodebug is not None else cfg.DEBUG182 183 # Force no debug in test184 if phase == "test":185 cfg.DEBUG = False186 cfg.DEVICE = [0]187 print("Force no debugging and one gpu when testing")188 189 if phase == "demo":190 cfg.DEMO.RENDER = params.render191 cfg.DEMO.FRAME_RATE = params.frame_rate192 cfg.DEMO.EXAMPLE = params.example193 cfg.DEMO.TASK = params.task194 cfg.TEST.FOLDER = params.out_dir if params.out_dir else cfg.TEST.FOLDER195 os.makedirs(cfg.TEST.FOLDER, exist_ok=True)196 197 if phase == "render":198 if params.npy:199 cfg.RENDER.NPY = params.npy200 cfg.RENDER.INPUT_MODE = "npy"201 if params.dir:202 cfg.RENDER.DIR = params.dir203 cfg.RENDER.INPUT_MODE = "dir"204 if params.fps:205 cfg.RENDER.FPS = float(params.fps)206 cfg.RENDER.MODE = params.mode207 208 # Debug mode209 if cfg.DEBUG:210 cfg.NAME = "debug--" + cfg.NAME211 cfg.LOGGER.WANDB.params.offline = True212 cfg.LOGGER.VAL_EVERY_STEPS = 1213 214 # Resume config215 cfg = resume_config(cfg)216 217 return cfg218 