Team Ai
Apppublic

OpenMotionLab/MotionGPT

sourceHugging Facemitupdated 1y agoView on Hugging Face
118likes
config.py218 linesDownload Raw Back to mGPT
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