import sys import omegaconf import lightning as L from torch.utils import data from lightning.pytorch.callbacks import ModelCheckpoint from lightning.pytorch.loggers import WandbLogger from benchmark.models.model import BenchmarkModel from benchmark.modules.selector import get_lightning_module from benchmark.data_loaders.selector import get_dataset def main(cfg): # model model = BenchmarkModel( frontend_name=cfg.model.frontend, backend_name=cfg.model.backend, latent_dim=cfg.model.latent_dim, output_dim=cfg.model.output_dim, layer_ix=cfg.model.layer_ix, is_flash=cfg.model.is_flash, output_size=cfg.model.output_size, ) # lightning module lit_module = get_lightning_module(cfg.core.task)( model=model, dataset=cfg.data.dataset, learning_rate=cfg.optim.learning_rate ) # data loaders train_dataloader = data.DataLoader( dataset=get_dataset(cfg.data.dataset)(split="train"), batch_size=cfg.data.batch_size, shuffle=True, drop_last=False, num_workers=cfg.data.num_workers ) validation_dataloader = data.DataLoader( dataset=get_dataset(cfg.data.dataset)(split="valid"), batch_size=cfg.data.val_batch_size, shuffle=False, drop_last=False, num_workers=cfg.data.num_workers ) # callbacks callbacks = [ ModelCheckpoint( save_last=True, save_top_k=cfg.core.save_top_k, monitor="valid_loss", mode="min", dirpath="/home/minz/logs/benchmark/%s_%s" % (cfg.core.task, cfg.data.dataset), ) ] # logger logger = WandbLogger(name=cfg.core.version, save_dir="/app/suno/minz/wandb_logs", log_model="all") # trainer trainer = L.Trainer( accelerator="gpu", devices=cfg.core.devices, num_nodes=cfg.core.num_nodes, strategy="deepspeed", precision=cfg.core.precision, limit_train_batches=cfg.data.limit_train, profiler="simple", # "simple" or "advanced" callbacks=callbacks, max_epochs=cfg.core.max_epochs, logger=logger, ) trainer.fit(lit_module, train_dataloader, validation_dataloader) if __name__ == "__main__": cfg = omegaconf.OmegaConf.load(sys.argv[1]) main(cfg)