YashwanthSC/Image-to-Mesh
0
1import os, sys2import argparse3import shutil4import subprocess5from omegaconf import OmegaConf6 7from pytorch_lightning import seed_everything8from pytorch_lightning.trainer import Trainer9from pytorch_lightning.strategies import DDPStrategy10from pytorch_lightning.callbacks import Callback11from pytorch_lightning.utilities import rank_zero_only, rank_zero_warn12 13from src.utils.train_util import instantiate_from_config14 15 16@rank_zero_only17def rank_zero_print(*args):18 print(*args)19 20 21def get_parser(**parser_kwargs):22 def str2bool(v):23 if isinstance(v, bool):24 return v25 if v.lower() in ("yes", "true", "t", "y", "1"):26 return True27 elif v.lower() in ("no", "false", "f", "n", "0"):28 return False29 else:30 raise argparse.ArgumentTypeError("Boolean value expected.")31 32 parser = argparse.ArgumentParser(**parser_kwargs)33 parser.add_argument(34 "-r",35 "--resume",36 type=str,37 default=None,38 help="resume from checkpoint",39 )40 parser.add_argument(41 "--resume_weights_only",42 action="store_true",43 help="only resume model weights",44 )45 parser.add_argument(46 "-b",47 "--base",48 type=str,49 default="base_config.yaml",50 help="path to base configs",51 )52 parser.add_argument(53 "-n",54 "--name",55 type=str,56 default="",57 help="experiment name",58 )59 parser.add_argument(60 "--num_nodes",61 type=int,62 default=1,63 help="number of nodes to use",64 )65 parser.add_argument(66 "--gpus",67 type=str,68 default="0,",69 help="gpu ids to use",70 )71 parser.add_argument(72 "-s",73 "--seed",74 type=int,75 default=42,76 help="seed for seed_everything",77 )78 parser.add_argument(79 "-l",80 "--logdir",81 type=str,82 default="logs",83 help="directory for logging data",84 )85 return parser86 87 88class SetupCallback(Callback):89 def __init__(self, resume, logdir, ckptdir, cfgdir, config):90 super().__init__()91 self.resume = resume92 self.logdir = logdir93 self.ckptdir = ckptdir94 self.cfgdir = cfgdir95 self.config = config96 97 def on_fit_start(self, trainer, pl_module):98 if trainer.global_rank == 0:99 # Create logdirs and save configs100 os.makedirs(self.logdir, exist_ok=True)101 os.makedirs(self.ckptdir, exist_ok=True)102 os.makedirs(self.cfgdir, exist_ok=True)103 104 rank_zero_print("Project config")105 rank_zero_print(OmegaConf.to_yaml(self.config))106 OmegaConf.save(self.config,107 os.path.join(self.cfgdir, "project.yaml"))108 109 110class CodeSnapshot(Callback):111 """112 Modified from https://github.com/threestudio-project/threestudio/blob/main/threestudio/utils/callbacks.py#L60113 """114 def __init__(self, savedir):115 self.savedir = savedir116 117 def get_file_list(self):118 return [119 b.decode()120 for b in set(121 subprocess.check_output(122 'git ls-files -- ":!:configs/*"', shell=True123 ).splitlines()124 )125 | set( # hard code, TODO: use config to exclude folders or files126 subprocess.check_output(127 "git ls-files --others --exclude-standard", shell=True128 ).splitlines()129 )130 ]131 132 @rank_zero_only133 def save_code_snapshot(self):134 os.makedirs(self.savedir, exist_ok=True)135 for f in self.get_file_list():136 if not os.path.exists(f) or os.path.isdir(f):137 continue138 os.makedirs(os.path.join(self.savedir, os.path.dirname(f)), exist_ok=True)139 shutil.copyfile(f, os.path.join(self.savedir, f))140 141 def on_fit_start(self, trainer, pl_module):142 try:143 self.save_code_snapshot()144 except:145 rank_zero_warn(146 "Code snapshot is not saved. Please make sure you have git installed and are in a git repository."147 )148 149 150if __name__ == "__main__":151 # add cwd for convenience and to make classes in this file available when152 # running as `python main.py`153 sys.path.append(os.getcwd())154 155 parser = get_parser()156 opt, unknown = parser.parse_known_args()157 158 cfg_fname = os.path.split(opt.base)[-1]159 cfg_name = os.path.splitext(cfg_fname)[0]160 exp_name = "-" + opt.name if opt.name != "" else ""161 logdir = os.path.join(opt.logdir, cfg_name+exp_name)162 163 ckptdir = os.path.join(logdir, "checkpoints")164 cfgdir = os.path.join(logdir, "configs")165 codedir = os.path.join(logdir, "code")166 seed_everything(opt.seed)167 168 # init configs169 config = OmegaConf.load(opt.base)170 lightning_config = config.lightning171 trainer_config = lightning_config.trainer172 173 trainer_config["accelerator"] = "gpu"174 rank_zero_print(f"Running on GPUs {opt.gpus}")175 ngpu = len(opt.gpus.strip(",").split(','))176 trainer_config['devices'] = ngpu177 178 trainer_opt = argparse.Namespace(**trainer_config)179 lightning_config.trainer = trainer_config180 181 # model182 model = instantiate_from_config(config.model)183 if opt.resume and opt.resume_weights_only:184 model = model.__class__.load_from_checkpoint(opt.resume, **config.model.params)185 186 model.logdir = logdir187 188 # trainer and callbacks189 trainer_kwargs = dict()190 191 # logger192 default_logger_cfg = {193 "target": "pytorch_lightning.loggers.TensorBoardLogger",194 "params": {195 "name": "tensorboard",196 "save_dir": logdir, 197 "version": "0",198 }199 }200 logger_cfg = OmegaConf.merge(default_logger_cfg)201 trainer_kwargs["logger"] = instantiate_from_config(logger_cfg)202 203 # model checkpoint204 default_modelckpt_cfg = {205 "target": "pytorch_lightning.callbacks.ModelCheckpoint",206 "params": {207 "dirpath": ckptdir,208 "filename": "{step:08}",209 "verbose": True,210 "save_last": True,211 "every_n_train_steps": 5000,212 "save_top_k": -1, # save all checkpoints213 }214 }215 216 if "modelcheckpoint" in lightning_config:217 modelckpt_cfg = lightning_config.modelcheckpoint218 else:219 modelckpt_cfg = OmegaConf.create()220 modelckpt_cfg = OmegaConf.merge(default_modelckpt_cfg, modelckpt_cfg)221 222 # callbacks223 default_callbacks_cfg = {224 "setup_callback": {225 "target": "train.SetupCallback",226 "params": {227 "resume": opt.resume,228 "logdir": logdir,229 "ckptdir": ckptdir,230 "cfgdir": cfgdir,231 "config": config,232 }233 },234 "learning_rate_logger": {235 "target": "pytorch_lightning.callbacks.LearningRateMonitor",236 "params": {237 "logging_interval": "step",238 }239 },240 "code_snapshot": {241 "target": "train.CodeSnapshot",242 "params": {243 "savedir": codedir,244 }245 },246 }247 default_callbacks_cfg["checkpoint_callback"] = modelckpt_cfg248 249 if "callbacks" in lightning_config:250 callbacks_cfg = lightning_config.callbacks251 else:252 callbacks_cfg = OmegaConf.create()253 callbacks_cfg = OmegaConf.merge(default_callbacks_cfg, callbacks_cfg)254 255 trainer_kwargs["callbacks"] = [256 instantiate_from_config(callbacks_cfg[k]) for k in callbacks_cfg]257 258 trainer_kwargs['precision'] = '32-true'259 trainer_kwargs["strategy"] = DDPStrategy(find_unused_parameters=True)260 261 # trainer262 trainer = Trainer(**trainer_config, **trainer_kwargs, num_nodes=opt.num_nodes)263 trainer.logdir = logdir264 265 # data266 data = instantiate_from_config(config.data)267 data.prepare_data()268 data.setup("fit")269 270 # configure learning rate271 base_lr = config.model.base_learning_rate272 if 'accumulate_grad_batches' in lightning_config.trainer:273 accumulate_grad_batches = lightning_config.trainer.accumulate_grad_batches274 else:275 accumulate_grad_batches = 1276 rank_zero_print(f"accumulate_grad_batches = {accumulate_grad_batches}")277 lightning_config.trainer.accumulate_grad_batches = accumulate_grad_batches278 model.learning_rate = base_lr279 rank_zero_print("++++ NOT USING LR SCALING ++++")280 rank_zero_print(f"Setting learning rate to {model.learning_rate:.2e}")281 282 # run training loop283 if opt.resume and not opt.resume_weights_only:284 trainer.fit(model, data, ckpt_path=opt.resume)285 else:286 trainer.fit(model, data)287 