procgne/Plonk
0
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 