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 ditto.data_loaders.selector import get_dataset from ditto.models.ditto import Ditto from ditto.modules.lightning_module import DittoLitModule def main(cfg): model = Ditto( music_encoder_name=cfg.model.music_encoder_name, text_encoder_name=cfg.model.text_encoder_name, latent_dim=cfg.model.latent_dim, model_path=cfg.model.model_path, ) # lightning module lit_module = DittoLitModule( model=model, learning_rate=cfg.optim.learning_rate, ) # data loaders train_dataloader = data.DataLoader( dataset=get_dataset( cfg.data.train_dataset, split="train", num_samples=cfg.data.num_train_samples, ), batch_size=cfg.data.batch_size, shuffle=True, drop_last=True, num_workers=cfg.data.num_workers, ) validation_dataloader = data.DataLoader( dataset=get_dataset( cfg.data.valid_dataset, split="valid", num_samples=cfg.data.num_val_samples, ), batch_size=cfg.data.batch_size, shuffle=False, drop_last=True, num_workers=cfg.data.num_workers, ) # callbacks callbacks = [ ModelCheckpoint( save_last=True, save_top_k=cfg.core.save_top_k, monitor="loss_validation", mode="min", dirpath="/home/minz/logs/%s" % cfg.core.version, ) ] # 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, ckpt_path=cfg.core.ckpt_path, ) if __name__ == "__main__": cfg = omegaconf.OmegaConf.load(sys.argv[1]) main(cfg)