CoolFace
Apppublic

YashwanthSC/Image-to-Mesh

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
train.py287 linesDownload Raw Back to root
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