import os os.environ["TOKENIZERS_PARALLELISM"] = "false" import sys import warnings warnings.filterwarnings("ignore", category=FutureWarning) 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_v2.data_loaders.multi import MultiTaskDataset, MultiTaskBatchSampler from ditto_v2.models.ditto import Ditto from ditto_v2.modules.lightning_module import DittoLitModule def main(cfg): model = Ditto( latent_dim=cfg.model.latent_dim, model_path=cfg.model.model_path, is_flash=True ) # lightning module lit_module = DittoLitModule( model=model, learning_rate=cfg.optim.learning_rate, ) # data loaders train_dataset = MultiTaskDataset(split="train") valid_dataset = MultiTaskDataset(split="valid") train_dataloader = data.DataLoader( dataset=train_dataset, batch_sampler=MultiTaskBatchSampler( train_dataset, batch_size=cfg.data.batch_size ), num_workers=cfg.data.num_workers, ) validation_dataloader = data.DataLoader( dataset=valid_dataset, batch_sampler=MultiTaskBatchSampler( valid_dataset, batch_size=cfg.data.batch_size ), 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, use_distributed_sampler=False, ) 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)