MCPcopy Create free account
hub / github.com/MotrixLab/FineMoGen / main

Function main

tools/train.py:60–142  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

58
59
60def main():
61 args = parse_args()
62
63 cfg = Config.fromfile(args.config)
64 if args.options is not None:
65 cfg.merge_from_dict(args.options)
66 # set cudnn_benchmark
67 if cfg.get('cudnn_benchmark', False):
68 torch.backends.cudnn.benchmark = True
69
70 # work_dir is determined in this priority: CLI > segment in file > filename
71 if args.work_dir is not None:
72 # update configs according to CLI args if args.work_dir is not None
73 cfg.work_dir = args.work_dir
74 elif cfg.get('work_dir', None) is None:
75 # use config filename as default work_dir if cfg.work_dir is None
76 cfg.work_dir = osp.join('./work_dirs',
77 osp.splitext(osp.basename(args.config))[0])
78 if args.resume_from is not None:
79 cfg.resume_from = args.resume_from
80 if args.gpu_ids is not None:
81 cfg.gpu_ids = args.gpu_ids
82 else:
83 cfg.gpu_ids = range(1) if args.gpus is None else range(args.gpus)
84
85 # init distributed env first, since logger depends on the dist info.
86 if args.launcher == 'none':
87 distributed = False
88 else:
89 distributed = True
90 init_dist(args.launcher, **cfg.dist_params)
91 _, world_size = get_dist_info()
92 cfg.gpu_ids = range(world_size)
93
94 # create work_dir
95 mmcv.mkdir_or_exist(osp.abspath(cfg.work_dir))
96 # dump config
97 cfg.dump(osp.join(cfg.work_dir, osp.basename(args.config)))
98 # init the logger before other steps
99 timestamp = time.strftime('%Y%m%d_%H%M%S', time.localtime())
100 log_file = osp.join(cfg.work_dir, f'{timestamp}.log')
101 logger = get_root_logger(log_file=log_file, log_level=cfg.log_level)
102
103 # init the meta dict to record some important information such as
104 # environment info and seed, which will be logged
105 meta = dict()
106 # log env info
107 env_info_dict = collect_env()
108 env_info = '\n'.join([(f'{k}: {v}') for k, v in env_info_dict.items()])
109 dash_line = '-' * 60 + '\n'
110 logger.info('Environment info:\n' + dash_line + env_info + '\n' +
111 dash_line)
112 meta['env_info'] = env_info
113
114 # log some basic info
115 logger.info(f'Distributed training: {distributed}')
116 logger.info(f'Config:\n{cfg.pretty_text}')
117

Callers 1

train.pyFile · 0.70

Calls 7

get_root_loggerFunction · 0.90
collect_envFunction · 0.90
set_random_seedFunction · 0.90
build_architectureFunction · 0.90
build_datasetFunction · 0.90
train_modelFunction · 0.90
parse_argsFunction · 0.70

Tested by

no test coverage detected