import os import tqdm import torch import numpy as np from einops import rearrange import torch.distributed as dist from torch.utils.data import DataLoader, DistributedSampler, Dataset from torch.nn.parallel import DistributedDataParallel as DDP from suno_utils.models.ditto.ditto import Ditto from suno_utils.audio import Audio from suno_utils.utils.text import read_jsonl import sys sys.path.append("/home/minz/glockenspiel/ditto-training") from ditto.data_loaders.cleaned_ytmsd import CleanedYTMSDDataset class CleanedYTMSDDataset(Dataset): def __init__( self, data_path="/app/suno/data/audio_mono_24khz/cleaned_ytmsd", split="train", input_length_s=30.0, sample_rate=24000, num_samples=-1, ): assert split in ["train", "valid"] self.data_path = data_path self.split = split self.input_length_s = input_length_s self.sample_rate = sample_rate self.num_samples = num_samples self.max_chunks = 8 # 30s * 8 chunks = 240 s # load files self.metadata = read_jsonl(os.path.join(data_path, "%s.jsonl" % split)) print("%d files are available for %s set" % (len(self.metadata), split)) def __getitem__(self, index): # read data metadata = self.metadata[index] # load audio audio_path = metadata["filepath"] audio = Audio.from_file(audio_path, sample_rate=self.sample_rate) # max 4 min wav = audio.get_segment(from_s=0.0, to_s=240).array_float num_chunks = len(wav) // self.input_length_s // self.sample_rate wav[int(num_chunks * self.input_length_s * self.sample_rate) :] = 0.0 # append short if len(wav) < int(self.sample_rate * self.max_chunks * self.input_length_s): pad = int(self.sample_rate * self.max_chunks * self.input_length_s) - len( wav ) wav = np.pad(wav, (0, pad), mode="constant", constant_values=0) if num_chunks == 0: audio_path = "invalid" # make stacks wav = wav.reshape(self.max_chunks, int(self.input_length_s * self.sample_rate)) return wav, num_chunks, audio_path def __len__(self): if self.num_samples > 0: return self.num_samples else: return len(self.metadata) # data parallel dist.init_process_group(backend="nccl") gpu_id = torch.distributed.get_rank() torch.cuda.set_device(gpu_id) # load model ditto = Ditto( music_encoder_name="musicfm_concat", latent_dim=128, model_path="/home/minz/logs/ditto_cleaned_ytmsd_128/step_370k.pt", is_flash=False, ) ditto = ditto.eval().bfloat16().cuda(gpu_id) ditto = DDP(ditto, device_ids=[gpu_id]) # train data train_ds = CleanedYTMSDDataset(split="train") train_sampler = DistributedSampler(train_ds, num_replicas=4, rank=gpu_id, shuffle=False) train_dl = DataLoader( dataset=train_ds, batch_size=4, shuffle=False, sampler=train_sampler, drop_last=False, num_workers=8, ) train_memmap_fn = "/app/suno/minz/ytmsd_train.dat" train_memmap_file_fn = "/app/suno/minz/ytmsd_train_files.dat" train_embeddings = np.memmap( train_memmap_fn, dtype=np.float32, mode="w+", shape=(len(train_ds), 128) ) train_filenames = np.memmap( train_memmap_file_fn, dtype=f"S{100}", mode="w+", shape=(len(train_ds),) ) # for multi gpu total_samples = len(train_ds) samples_per_gpu = total_samples // dist.get_world_size() offset = samples_per_gpu * dist.get_rank() for batch_idx, (wav, counter, fn) in enumerate(tqdm.tqdm(train_dl)): b = len(wav) wav = rearrange(wav, "b c t -> (b c) t") wav = wav.cuda().bfloat16() emb = ditto.module.music_to_latent(wav).cpu().detach().numpy() emb = rearrange(emb, "(b c) d -> b c d", b=b, c=8) emb = emb.sum(axis=1) / counter.numpy()[:, np.newaxis] start_idx = offset + batch_idx * train_dl.batch_size end_idx = start_idx + b train_embeddings[start_idx:end_idx] = emb train_filenames[start_idx:end_idx] = list(fn) if batch_idx % 10 == 0: train_embeddings.flush() train_filenames.flush() train_embeddings.flush() train_filenames.flush() # valid data valid_ds = CleanedYTMSDDataset(split="valid") valid_sampler = DistributedSampler(valid_ds, num_replicas=4, rank=gpu_id, shuffle=False) valid_dl = DataLoader( dataset=valid_ds, batch_size=4, shuffle=False, sampler=valid_sampler, drop_last=False, num_workers=8, ) valid_memmap_fn = "/app/suno/minz/ytmsd_valid.dat" valid_memmap_file_fn = "/app/suno/minz/ytmsd_valid_files.dat" valid_embeddings = np.memmap( valid_memmap_fn, dtype=np.float32, mode="w+", shape=(len(valid_ds), 128) ) valid_filenames = np.memmap( valid_memmap_file_fn, dtype=f"S{100}", mode="w+", shape=(len(valid_ds),) ) # for multi gpu total_samples = len(valid_ds) samples_per_gpu = total_samples // dist.get_world_size() offset = samples_per_gpu * dist.get_rank() for batch_idx, (wav, counter, fn) in enumerate(tqdm.tqdm(valid_dl)): b = len(wav) wav = rearrange(wav, "b c t -> (b c) t") wav = wav.cuda().bfloat16() emb = ditto.module.music_to_latent(wav).cpu().detach().numpy() emb = rearrange(emb, "(b c) d -> b c d", b=b, c=8) emb = emb.sum(axis=1) / counter.numpy()[:, np.newaxis] start_idx = offset + batch_idx * valid_dl.batch_size end_idx = start_idx + b valid_embeddings[start_idx:end_idx] = emb valid_filenames[start_idx:end_idx] = list(fn) if batch_idx % 10 == 0: valid_embeddings.flush() valid_filenames.flush() valid_embeddings.flush() valid_filenames.flush()