CoolFace
Apppublic

procgne/Plonk

sourceHugging Facemitupdated 10mo agoView on Hugging Face
0likes
train_random.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 RandomGeolocalizer16 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 = RandomGeolocalizer.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 = RandomGeolocalizer(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