import os import sys import time import datetime import random import torch import torch.nn as nn import torch.distributed as dist import torch.multiprocessing as mp from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data import DataLoader, DistributedSampler import numpy as np from sklearn import metrics as skm from einops import rearrange from ditto_v2.data_loaders.multi import MultiTaskDataset, MultiTaskBatchSampler from ditto_v2.models.ditto import Ditto sys.path.append("/home/minz/neon/sunoGPT/") from torch.distributed import destroy_process_group, init_process_group from utils.helpers import ( dist_barrier, load_checkpoint, load_old_state_dict, load_old_optimizer_state_dict, print_with_time, print_with_time_master, save_checkpoint, save_old_checkpoint, suppress_logging, verify_preload_model_args, ) ddp = int(os.environ.get("RANK", -1)) != -1 # is this a ddp run? if ddp: init_process_group(backend="nccl", timeout=datetime.timedelta(seconds=24 * 60 * 60)) ddp_rank = int(os.environ["RANK"]) # global gpu rank ddp_local_rank = int(os.environ["LOCAL_RANK"]) # gpu rank within node world_size = torch.distributed.get_world_size() # total number of gpus device = f"cuda:{ddp_local_rank}" torch.cuda.set_device(device) master_process = ddp_rank == 0 # this process will do logging, checkpointing etc. seed_offset = ddp_rank + 1 # each process gets a different seed print_with_time(f"ddp init, rank {ddp_rank}, local_rank {ddp_local_rank}") else: ddp_rank = 0 ddp_local_rank = 0 world_size = 1 # if not ddp, we are running on a single gpu, and one process master_process = True seed_offset = 1 n_gpus_per_node = torch.cuda.device_count() dist_barrier() print_with_time_master(f"ddp init: world size {world_size} ddp_rank {ddp_rank}.") def setup(rank, world_size): os.environ["MASTER_ADDR"] = "localhost" os.environ["MASTER_PORT"] = "12355" init_process_group("nccl", rank=rank, world_size=world_size) def cleanup(): destroy_process_group() def train(rank, world_size): setup(rank, world_size) # Set device for this process device = torch.device(f"cuda:{rank}") torch.cuda.set_device(device) # Create model and move it to GPU with DDP model = Ditto( latent_dim=128, model_path=None, is_flash=True ) # Initialize your Ditto model here model = model.to(device) model = DDP(model, device_ids=[rank]) # Define loss function and optimizer criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=learning_rate) # Create dataset and dataloader train_dataset = MultiTaskDataset(split="train") valid_dataset = MultiTaskDataset(split="valid") train_dataloader = data.DataLoader( dataset=train_dataset, batch_sampler=MultiTaskBatchSampler(train_datset, batch_size=16), ) sampler = DistributedSampler(dataset, num_replicas=world_size, rank=rank) dataloader = DataLoader(dataset, batch_size=batch_size, sampler=sampler) # Training loop for epoch in range(num_epochs): model.train() sampler.set_epoch(epoch) for batch in dataloader: # Move batch to device batch = {k: v.to(device) for k, v in batch.items()} # Forward pass outputs = model(batch) loss = criterion(outputs, batch["labels"]) # Backward pass and optimize optimizer.zero_grad() loss.backward() optimizer.step() # Print progress, save checkpoints, etc. if rank == 0: print(f"Epoch {epoch+1}/{num_epochs}, Loss: {loss.item()}") # Save checkpoint logic here cleanup() if __name__ == "__main__": world_size = torch.cuda.device_count() mp.spawn(train, args=(world_size,), nprocs=world_size, join=True) # 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 DittoModule(nn.Module): # 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 # 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(z_a.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}_median_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 forward(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] # 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] # inp2 = [line[:random_length2] for line in inp2] # outputs = self.model(inp1, inp2, task) # modality_ab = "m2t" # modality_ba = "t2m" # inp1_emb, inp2_emb, logit_scale = outputs # # get metrics # metrics = self.get_metrics(inp1_emb, inp2_emb, logit_scale) # return metrics, modality_ab, modality_ba # def train(rank, world_size, cfg): # dist.init_process_group("nccl", rank=rank, world_size=world_size) # torch.cuda.set_device(rank) # model = Ditto( # latent_dim=cfg.model.latent_dim, # model_path=cfg.model.model_path, # is_flash=True # ) # model = model.to(rank) # ditto_module = DittoModule( # model=model, # learning_rate=cfg.optim.learning_rate, # ) # ditto_module = DDP(ditto_module, device_ids=[rank]) # optimizer = torch.optim.AdamW( # [ # {"params": ditto_module.module.model.music_projection.parameters(), "lr": ditto_module.module.lr}, # {"params": ditto_module.module.model.text_projection.parameters(), "lr": ditto_module.module.lr}, # {"params": ditto_module.module.model.music_encoder.parameters(), "lr": ditto_module.module.lr / 10}, # {"params": ditto_module.module.model.text_encoder.parameters(), "lr": ditto_module.module.lr / 10}, # ], # lr=ditto_module.module.lr, # ) # train_dataset = MultiTaskDataset(split="train") # valid_dataset = MultiTaskDataset(split="valid") # train_sampler = DistributedSampler(train_dataset, num_replicas=world_size, rank=rank) # valid_sampler = DistributedSampler(valid_dataset, num_replicas=world_size, rank=rank) # train_dataloader = DataLoader( # dataset=train_dataset, # batch_sampler=MultiTaskBatchSampler(train_dataset, batch_size=cfg.data.batch_size), # num_workers=cfg.data.num_workers, # sampler=train_sampler, # ) # validation_dataloader = DataLoader( # dataset=valid_dataset, # batch_sampler=MultiTaskBatchSampler(valid_dataset, batch_size=cfg.data.batch_size), # num_workers=cfg.data.num_workers, # sampler=valid_sampler, # ) # for epoch in range(cfg.core.max_epochs): # ditto_module.train() # for batch in train_dataloader: # optimizer.zero_grad() # metrics, modality_ab, modality_ba = ditto_module(batch, "train") # loss = metrics["loss"] # loss.backward() # torch.nn.utils.clip_grad_norm_(ditto_module.parameters(), ditto_module.module.gradient_clip_val) # optimizer.step() # if rank == 0: # print(f"Epoch {epoch}, Loss: {loss.item()}") # print(f"mAP10-{modality_ab}_train: {metrics['a_to_b_mAP@10']}") # print(f"mAP10-{modality_ba}_train: {metrics['b_to_a_mAP@10']}") # ditto_module.eval() # val_loss = 0 # with torch.no_grad(): # for batch in validation_dataloader: # metrics, modality_ab, modality_ba = ditto_module(batch, "validation") # val_loss += metrics["loss"].item() # val_loss /= len(validation_dataloader) # if rank == 0: # print(f"Validation Loss: {val_loss}") # def main(cfg): # world_size = torch.cuda.device_count() # mp.spawn(train, args=(world_size, cfg), nprocs=world_size, join=True) # os.environ['MASTER_ADDR'] = 'localhost' # os.environ['MASTER_PORT'] = '12355' # mp.spawn(train, args=(world_size, cfg), nprocs=world_size, join=True) # if __name__ == "__main__": # import sys # import warnings # warnings.filterwarnings("ignore", category=FutureWarning) # import omegaconf # cfg = omegaconf.OmegaConf.load(sys.argv[1]) # main(cfg)