import time import random import torch import lightning as L import torch.nn.functional as F import numpy as np from torch import nn from sklearn import metrics as skm from einops import rearrange def print_model_params(model): total_params = 0 trainable_params = 0 for _, parameter in model.named_parameters(): params = parameter.numel() total_params += params if parameter.requires_grad: trainable_params += params print(f"Total Params: {total_params:,}, Trainable_params: {trainable_params:,}") class DittoLitModule(L.LightningModule): def __init__( self, model, learning_rate=1e-4, gradient_clip_val=1.0, ): super().__init__() self.lr = learning_rate self.model = model print_model_params(model) self.loss_function = nn.CrossEntropyLoss() self.gradient_clip_val = gradient_clip_val self.save_hyperparameters(ignore=["model"]) ( self.inference_music_embeddings, self.inference_text_embeddings, self.inference_ids, self.start_ss, self.lyrics, ) = [], [], [], [], [] self.counter = 0 def get_metrics(self, z_a, z_b, logit_scale): metrics = {} # scaled dot product similarity logits_per_a = logit_scale * z_a @ z_b.t() logits_per_b = logits_per_a.t() labels = torch.arange(z_a.shape[0]).long().to(self.device) # get cross entropy loss loss = ( self.loss_function(logits_per_a, labels) + self.loss_function(logits_per_b, labels) ) / 2 metrics["loss"] = loss # get accuracy metrics["acc_a"] = skm.accuracy_score( labels.detach().cpu().numpy(), logits_per_a.argmax(dim=1).detach().cpu().numpy(), ) metrics["acc_b"] = skm.accuracy_score( labels.detach().cpu().numpy(), logits_per_b.argmax(dim=1).detach().cpu().numpy(), ) # get ranking metrics logits = { "a_to_b": logits_per_a.detach().cpu(), "b_to_a": logits_per_b.detach().cpu(), } ground_truth = torch.arange(len(z_b)).view(-1, 1) for name, logit in logits.items(): ranking = torch.argsort(logit, descending=True) preds = torch.where(ranking == ground_truth)[1] preds = preds.detach().cpu().numpy() metrics[f"{name}_mean_rank"] = preds.mean() + 1 metrics[f"{name}_mdeidan_rank"] = np.floor(np.median(preds)) + 1 for k in [1, 5, 10]: metrics[f"{name}_R@{k}"] = np.mean(preds < k) metrics[f"{name}_mAP@10"] = np.mean( np.where(preds < 10, 1 / (preds + 1), 0.0) ) return metrics def random_masking(self, x, mask_prob=0.125, mask_hop_s=0.5): """random masking of 500ms with given probability""" b, t = x.shape len_masking_raw = int(24000 * mask_hop_s) # get random mask indices start_indices = torch.rand(b, t // len_masking_raw) < mask_prob time_domain_masked_indices = torch.nonzero( start_indices.repeat_interleave(len_masking_raw, dim=1) ) # mask with random values masking_noise = ( torch.randn(time_domain_masked_indices.shape[0], dtype=x.dtype) * 0.1 ) # 0 mean 0.1 std x[tuple(time_domain_masked_indices.t())] = masking_noise.to(self.device) return x def sequence_masking(self, x, max_ratio=0.6): b, t = x.shape mask_len = random.randint(1, int(t * max_ratio)) masking_noise = torch.randn(b, mask_len, dtype=x.dtype) * 0.1 x[:, -mask_len:] = masking_noise.to(self.device) return x def step(self, batch, stage): # get batch inps, task = batch task = task[0] inp1, inp2 = inps # forward based on task if task in [ "self_sim", "self_vox_sim", "artist_sim", "artist_vox_sim", "album_sim", ]: if stage == "train": random_length1 = random.randint(24000 * 5, 24000 * 15) random_length2 = random.randint(24000 * 5, 24000 * 15) inp1 = inp1[:, :random_length1] inp2 = inp2[:, :random_length2] inp1 = inp1.to(self.device) inp2 = inp2.to(self.device) outputs = self.model(inp1, inp2, task) modality_ab = "m2m" modality_ba = "m2m" elif task in ["self_lyric_sim"]: if stage == "train": random_length1 = random.randint(50, 313) random_length2 = random.randint(50, 313) inp1 = [line[:random_length1] for line in inp1] inp2 = [line[:random_length2] for line in inp2] outputs = self.model(inp1, inp2, task) modality_ab = "t2t" modality_ba = "t2t" elif task in ["genre_sim", "lyric_sim"]: if stage == "train": random_length1 = random.randint(24000 * 5, 24000 * 15) random_length2 = random.randint(50, 313) inp1 = inp1[:, :random_length1] inp1 = inp1.to(self.device) inp2 = [line[:random_length2] for line in inp2] outputs = self.model(inp1, inp2, task) modality_ab = "m2t" modality_ba = "t2m" # gather multi-gpu outputs if self.trainer.world_size > 1: gathered_outputs = self.all_gather(outputs, sync_grads=(stage == "train")) inp1_emb = rearrange(gathered_outputs[0], "n b c -> (n b) c") inp2_emb = rearrange(gathered_outputs[1], "n b c -> (n b) c") else: inp1_emb, inp2_emb = outputs[0], outputs[1] logit_scale = outputs[2] # get metrics metrics = self.get_metrics(inp1_emb, inp2_emb, logit_scale) # log metrics self.log( "loss_%s_%s" % (stage, task), metrics["loss"], prog_bar=True, sync_dist=True, batch_size=len(inp1_emb), ) self.log( "loss_%s" % stage, metrics["loss"], prog_bar=True, sync_dist=True, batch_size=len(inp1_emb), ) self.log( "mAP10-%s_%s_%s" % (modality_ab, stage, task), metrics["a_to_b_mAP@10"], prog_bar=True, sync_dist=True, batch_size=len(inp1_emb), ) self.log( "mAP10-%s_%s_%s" % (modality_ba, stage, task), metrics["b_to_a_mAP@10"], prog_bar=True, sync_dist=True, batch_size=len(inp1_emb), ) return metrics def training_step(self, batch, batch_idx): try: metrics = self.step(batch, "train") except Exception as e: print(f"Error in training: {str(e)}") return metrics["loss"] def validation_step(self, batch, batch_idx): metrics = self.step(batch, "validation") return metrics["loss"] def configure_optimizers(self): optimizer = torch.optim.AdamW( [ {"params": self.model.music_projection.parameters(), "lr": self.lr}, {"params": self.model.text_projection.parameters(), "lr": self.lr}, {"params": self.model.music_encoder.parameters(), "lr": self.lr / 10}, {"params": self.model.text_encoder.parameters(), "lr": self.lr / 10}, ], lr=self.lr, ) return {"optimizer": optimizer, "gradient_clip_val": self.gradient_clip_val} # 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)