blanchon/Metric3D
0
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) 