import os import sys sys.path.append("/home/minz/neon/ditto-training") import tqdm import random import glob import torch import numpy as np from torch.utils import data from suno_utils.audio import Audio from einops import rearrange from ditto_v2.data_loaders.self_sim import SelfSimDataset from ditto_v2.data_loaders.self_vox_sim import SelfVoxSimDataset from ditto_v2.data_loaders.self_lyric_sim import SelfLyricSimDataset from ditto_v2.data_loaders.artist_sim import ArtistSimDataset from ditto_v2.data_loaders.artist_vox_sim import ArtistVoxSimDataset from ditto_v2.data_loaders.album_sim import AlbumSimDataset from ditto_v2.data_loaders.genre_sim import GenreSimDataset from ditto_v2.data_loaders.lyric_sim import LyricSimDataset from ditto_v2.models.ditto import Ditto from suno_utils.utils.text import read_jsonl, read_json split = "valid" self_sim_input_length_s = 15.0 self_vox_sim_input_length_s = 15.0 self_lyric_sim_input_length = 300 artist_sim_input_length_s = 15.0 artist_vox_sim_input_length_s = 15.0 album_sim_input_length_s = 15.0 genre_sim_input_length_s = 15.0 genre_tag_input_length = 300 lyric_sim_input_length_s = 15.0 lyric_sim_input_text_length = 300 sample_rate = 24000 num_samples = 1000 num_self_sim_samples = -1 num_vox_sim_samples = -1 num_self_lyric_sim_samples = -1 num_artist_sim_samples = -1 num_artist_vox_sim_samples = -1 num_album_sim_samples = -1 num_genre_sim_samples = -1 num_lyric_sim_samples = -1 metadata = read_jsonl("/app/suno/data/v2_audio/metadata/metas_%s.jsonl" % split) task_indices = read_json("/app/suno/data/v2_audio/metadata/task_indices.json") # get model ditto = Ditto( latent_dim=128, model_path="/home/minz/logs/ditto_v2_local_8gpu/epoch=29.ckpt", is_flash=False, ) ditto = ditto.eval() ditto = ditto.cuda() # preprocess data ( self_sim_emb, artist_sim_emb, album_sim_emb, genre_sim_emb, self_vox_sim_emb, artist_vox_sim_emb, ) = [], [], [], [], [], [] class GeniusDataset(data.Dataset): def __init__(self, filelist, sample_rate=24000): self.filelist = filelist self.sample_rate = sample_rate def __len__(self): return len(self.filelist) def __getitem__(self, idx): filepath = self.filelist[idx] mix_tensor = self.get_tensor(filepath) vox_filepath = os.path.join( "/app/suno/data/v2_audio/genius_vox/", os.path.basename(filepath) ) vox_tensor = self.get_tensor(vox_filepath) return mix_tensor, vox_tensor def get_tensor(self, filepath): inp_wav = Audio.from_file(filepath, sample_rate=self.sample_rate) inp_array = inp_wav.array_float total_duration = len(inp_array) chunk_duration = self.sample_rate * 15 # Calculate 4 evenly spaced start points start_points = [ int(i * (total_duration - chunk_duration) / 3) for i in range(4) ] chunks = [inp_array[start : start + chunk_duration] for start in start_points] return torch.tensor(np.vstack(chunks)) def get_emb(filepath, task): inp = self.get_tensor(filepath) out = ditto.music_to_latent(inp.cuda(), task) out = out.mean(dim=0).detach().cpu().numpy() return out genius_filelist = glob.glob("/app/suno/data/v2_audio/genius_hq/*.wav") random.seed(134) random.shuffle(genius_filelist) genius_filelist = genius_filelist[:10000] dataset = GeniusDataset(genius_filelist) batch_size = 64 dataloader = data.DataLoader( dataset, batch_size=batch_size, num_workers=4, pin_memory=True ) for mix_batch, vox_batch in tqdm.tqdm(dataloader): mix_batch = mix_batch.cuda() vox_batch = vox_batch.cuda() mix_batch = rearrange(mix_batch, "b c t -> (b c) t") vox_batch = rearrange(vox_batch, "b c t -> (b c) t") # Process mix out = ditto.music_to_latent(mix_batch, "self_sim") out = rearrange(out, "(b c) e -> b c e", b=batch_size) out = out.mean(dim=1).detach().cpu().numpy() self_sim_emb.append(out) out = ditto.music_to_latent(mix_batch, "artist_sim") out = rearrange(out, "(b c) e -> b c e", b=batch_size) out = out.mean(dim=1).detach().cpu().numpy() artist_sim_emb.append(out) out = ditto.music_to_latent(mix_batch, "album_sim") out = rearrange(out, "(b c) e -> b c e", b=batch_size) out = out.mean(dim=1).detach().cpu().numpy() album_sim_emb.append(out) out = ditto.music_to_latent(mix_batch, "genre_sim") out = rearrange(out, "(b c) e -> b c e", b=batch_size) out = out.mean(dim=1).detach().cpu().numpy() genre_sim_emb.append(out) # Process vox out = ditto.music_to_latent(vox_batch, "self_vox_sim") out = rearrange(out, "(b c) e -> b c e", b=batch_size) out = out.mean(dim=1).detach().cpu().numpy() self_vox_sim_emb.append(out) out = ditto.music_to_latent(vox_batch, "artist_vox_sim") out = rearrange(out, "(b c) e -> b c e", b=batch_size) out = out.mean(dim=1).detach().cpu().numpy() artist_vox_sim_emb.append(out) # Concatenate results self_sim_emb = np.concatenate(self_sim_emb) artist_sim_emb = np.concatenate(artist_sim_emb) album_sim_emb = np.concatenate(album_sim_emb) genre_sim_emb = np.concatenate(genre_sim_emb) self_vox_sim_emb = np.concatenate(self_vox_sim_emb) artist_vox_sim_emb = np.concatenate(artist_vox_sim_emb) np.save(open("/home/minz/temp/genius_self_sim_emb.npy", "wb"), self_sim_emb) np.save(open("/home/minz/temp/genius_artist_sim_emb.npy", "wb"), artist_sim_emb) np.save(open("/home/minz/temp/genius_album_sim_emb.npy", "wb"), album_sim_emb) np.save(open("/home/minz/temp/genius_genre_sim_emb.npy", "wb"), genre_sim_emb) np.save(open("/home/minz/temp/genius_self_vox_sim_emb.npy", "wb"), self_vox_sim_emb) np.save(open("/home/minz/temp/genius_artist_vox_sim_emb.npy", "wb"), artist_vox_sim_emb)