CoolFace
Apppublic

blanchon/Metric3D

sourceHugging Facegpl-3.0updated 2y agoView on Hugging Face
0likes
test_scale_cano.py158 linesDownload Raw Back to tools
1import os2import os.path as osp3import cv24import time5import sys6CODE_SPACE=os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))7sys.path.append(CODE_SPACE)8import argparse9import mmcv10import torch11import torch.distributed as dist12import torch.multiprocessing as mp13 14try:15    from mmcv.utils import Config, DictAction16except:17    from mmengine import Config, DictAction18from datetime import timedelta19import random20import numpy as np21from mono.utils.logger import setup_logger22import glob23from mono.utils.comm import init_env24from mono.model.monodepth_model import get_configured_monodepth_model25from mono.utils.running import load_ckpt26from mono.utils.do_test import do_scalecano_test_with_custom_data27from mono.utils.mldb import load_data_info, reset_ckpt_path28from mono.utils.custom_data import load_from_annos, load_data29 30def parse_args():31    parser = argparse.ArgumentParser(description='Train a segmentor')32    parser.add_argument('config', help='train config file path')33    parser.add_argument('--show-dir', help='the dir to save logs and visualization results')34    parser.add_argument('--load-from', help='the checkpoint file to load weights from')35    parser.add_argument('--node_rank', type=int, default=0)36    parser.add_argument('--nnodes', type=int, default=1, help='number of nodes')37    parser.add_argument('--options', nargs='+', action=DictAction, help='custom options')38    parser.add_argument('--launcher', choices=['None', 'pytorch', 'slurm', 'mpi', 'ror'], default='slurm', help='job launcher')39    parser.add_argument('--test_data_path', default='None', type=str, help='the path of test data')40    args = parser.parse_args()41    return args42 43def main(args):44    os.chdir(CODE_SPACE)45    cfg = Config.fromfile(args.config)46    47    if args.options is not None:48        cfg.merge_from_dict(args.options)49        50    # show_dir is determined in this priority: CLI > segment in file > filename51    if args.show_dir is not None:52        # update configs according to CLI args if args.show_dir is not None53        cfg.show_dir = args.show_dir54    else:55        # use condig filename + timestamp as default show_dir if args.show_dir is None56        cfg.show_dir = osp.join('./show_dirs', 57                                osp.splitext(osp.basename(args.config))[0],58                                args.timestamp)59    60    # ckpt path61    if args.load_from is None:62        raise RuntimeError('Please set model path!')63    cfg.load_from = args.load_from64    65    # load data info66    data_info = {}67    load_data_info('data_info', data_info=data_info)68    cfg.mldb_info = data_info69    # update check point info70    reset_ckpt_path(cfg.model, data_info)71    72    # create show dir73    os.makedirs(osp.abspath(cfg.show_dir), exist_ok=True)74    75    # init the logger before other steps76    cfg.log_file = osp.join(cfg.show_dir, f'{args.timestamp}.log')77    logger = setup_logger(cfg.log_file)78    79    # log some basic info80    logger.info(f'Config:\n{cfg.pretty_text}')81    82    # init distributed env dirst, since logger depends on the dist info83    if args.launcher == 'None':84        cfg.distributed = False85    else:86        cfg.distributed = True87        init_env(args.launcher, cfg)88    logger.info(f'Distributed training: {cfg.distributed}')89    90    # dump config 91    cfg.dump(osp.join(cfg.show_dir, osp.basename(args.config)))92    test_data_path = args.test_data_path93    if not os.path.isabs(test_data_path):94        test_data_path = osp.join(CODE_SPACE, test_data_path)95 96    if 'json' in test_data_path:97        test_data = load_from_annos(test_data_path)98    else:99        test_data = load_data(args.test_data_path)100    101    if not cfg.distributed:102        main_worker(0, cfg, args.launcher, test_data)103    else:104        # distributed training105        if args.launcher == 'ror':106            local_rank = cfg.dist_params.local_rank107            main_worker(local_rank, cfg, args.launcher, test_data)108        else:109            mp.spawn(main_worker, nprocs=cfg.dist_params.num_gpus_per_node, args=(cfg, args.launcher, test_data))110        111def main_worker(local_rank: int, cfg: dict, launcher: str, test_data: list):112    if cfg.distributed:113        cfg.dist_params.global_rank = cfg.dist_params.node_rank * cfg.dist_params.num_gpus_per_node + local_rank114        cfg.dist_params.local_rank = local_rank115 116        if launcher == 'ror':117            init_torch_process_group(use_hvd=False)118        else:119            torch.cuda.set_device(local_rank)120            default_timeout = timedelta(minutes=30)121            dist.init_process_group(122                backend=cfg.dist_params.backend,123                init_method=cfg.dist_params.dist_url,124                world_size=cfg.dist_params.world_size,125                rank=cfg.dist_params.global_rank,126                timeout=default_timeout)127    128    logger = setup_logger(cfg.log_file)129    # build model130    model = get_configured_monodepth_model(cfg, )131    132    # config distributed training133    if cfg.distributed:134        model = torch.nn.parallel.DistributedDataParallel(model.cuda(),135                                                          device_ids=[local_rank],136                                                          output_device=local_rank,137                                                          find_unused_parameters=True)138    else:139        model = torch.nn.DataParallel(model).cuda()140        141    # load ckpt142    model, _,  _, _ = load_ckpt(cfg.load_from, model, strict_match=False)143    model.eval()144    145    do_scalecano_test_with_custom_data(146        model, 147        cfg,148        test_data,149        logger,150        cfg.distributed,151        local_rank152    )153    154if __name__ == '__main__':155    args = parse_args()156    timestamp = time.strftime('%Y%m%d_%H%M%S', time.localtime())157    args.timestamp = timestamp158    main(args)