CoolFace
Apppublic

procgne/Plonk

sourceHugging Facemitupdated 10mo agoView on Hugging Face
0likes
train.py147 linesDownload Raw Back to root
1import os2import hydra3import wandb4from os.path import isfile, join5from shutil import copyfile6 7import torch8 9from omegaconf import OmegaConf10from hydra.core.hydra_config import HydraConfig11from hydra.utils import instantiate12from pytorch_lightning.callbacks import LearningRateMonitor13from lightning_fabric.utilities.rank_zero import _get_rank14from callbacks import EMACallback, FixNANinGrad, IncreaseDataEpoch15from models.module import DiffGeolocalizer16 17torch.set_float32_matmul_precision("high")  # TODO do we need that?18 19# Registering the "eval" resolver allows for advanced config20# interpolation with arithmetic operations in hydra:21# https://omegaconf.readthedocs.io/en/2.3_branch/how_to_guides.html22OmegaConf.register_new_resolver("eval", eval)23 24 25def wandb_init(cfg):26    directory = cfg.checkpoints.dirpath27    if isfile(join(directory, "wandb_id.txt")) and cfg.logger_suffix == "":28        with open(join(directory, "wandb_id.txt"), "r") as f:29            wandb_id = f.readline()30    else:31        rank = _get_rank()32        wandb_id = wandb.util.generate_id()33        print(f"Generated wandb id: {wandb_id}")34        if rank == 0 or rank is None:35            with open(join(directory, "wandb_id.txt"), "w") as f:36                f.write(str(wandb_id))37 38    return wandb_id39 40 41def load_model(cfg, dict_config, wandb_id, callbacks):42    directory = cfg.checkpoints.dirpath43    if isfile(join(directory, "last.ckpt")):44        checkpoint_path = join(directory, "last.ckpt")45        logger = instantiate(cfg.logger, id=wandb_id, resume="allow")46        model = DiffGeolocalizer.load_from_checkpoint(checkpoint_path, cfg=cfg.model)47        ckpt_path = join(directory, "last.ckpt")48        print(f"Loading form checkpoint ... {ckpt_path}")49    else:50        ckpt_path = None51        logger = instantiate(cfg.logger, id=wandb_id, resume="allow")52        log_dict = {"model": dict_config["model"], "dataset": dict_config["dataset"]}53        logger._wandb_init.update({"config": log_dict})54        model = DiffGeolocalizer(cfg.model)55 56    trainer, strategy = cfg.trainer, cfg.trainer.strategy57    # from pytorch_lightning.profilers import PyTorchProfiler58 59    trainer = instantiate(60        trainer,61        strategy=strategy,62        logger=logger,63        callbacks=callbacks,64        # profiler=PyTorchProfiler(65        #     dirpath="logs",66        #     schedule=torch.profiler.schedule(wait=1, warmup=3, active=3, repeat=1),67        #     on_trace_ready=torch.profiler.tensorboard_trace_handler("./logs"),68        #     record_shapes=True,69        #     with_stack=True,70        #     with_flops=True,71        #     with_modules=True,72        # ),73    )74    return trainer, model, ckpt_path75 76 77def project_init(cfg):78    print("Working directory set to {}".format(os.getcwd()))79    directory = cfg.checkpoints.dirpath80    os.makedirs(directory, exist_ok=True)81    copyfile(".hydra/config.yaml", join(directory, "config.yaml"))82 83 84def callback_init(cfg):85    checkpoint_callback = instantiate(cfg.checkpoints)86    progress_bar = instantiate(cfg.progress_bar)87    lr_monitor = LearningRateMonitor()88    ema_callback = EMACallback(89        "network",90        "ema_network",91        decay=cfg.model.ema_decay,92        start_ema_step=cfg.model.start_ema_step,93        init_ema_random=False,94    )95    fix_nan_callback = FixNANinGrad(96        monitor=["train/loss"],97    )98    increase_data_epoch_callback = IncreaseDataEpoch()99    callbacks = [100        checkpoint_callback,101        progress_bar,102        lr_monitor,103        ema_callback,104        fix_nan_callback,105        increase_data_epoch_callback,106    ]107    return callbacks108 109 110def init_datamodule(cfg):111    datamodule = instantiate(cfg.datamodule)112    return datamodule113 114 115def hydra_boilerplate(cfg):116    dict_config = OmegaConf.to_container(cfg, resolve=True)117    callbacks = callback_init(cfg)118    datamodule = init_datamodule(cfg)119    project_init(cfg)120    wandb_id = wandb_init(cfg)121    trainer, model, ckpt_path = load_model(cfg, dict_config, wandb_id, callbacks)122    return trainer, model, datamodule, ckpt_path123 124 125@hydra.main(config_path="configs", config_name="config", version_base=None)126def main(cfg):127    if "stage" in cfg and cfg.stage == "debug":128        import lovely_tensors as lt129 130        lt.monkey_patch()131    trainer, model, datamodule, ckpt_path = hydra_boilerplate(cfg)132    model.datamodule = datamodule133    # model = torch.compile(model)134    if cfg.mode == "train":135        trainer.fit(model, datamodule=datamodule, ckpt_path=ckpt_path)136    elif cfg.mode == "eval":137        trainer.test(model, datamodule=datamodule)138    elif cfg.mode == "traineval":139        cfg.mode = "train"140        trainer.fit(model, datamodule=datamodule, ckpt_path=ckpt_path)141        cfg.mode = "test"142        trainer.test(model, datamodule=datamodule)143 144 145if __name__ == "__main__":146    main()147