import json import torchaudio import boto3 import tempfile import torch import numpy as np import random from tqdm import tqdm import librosa from suno_utils.gpt import chirp_v2_5 as chirp_v3 from suno_utils.models.ditto_v2.ditto_v2 import Ditto from suno_utils.tasks import ss_vad from suno_utils.tasks.dac_vae_100hz_peaq import preload_models as preload_vae_models DATA_PATH = "/app/suno/sara/metas_v0_artist.jsonl" DITTO_S3_PATH = "s3://suno-data/minz/models/ditto_v2_epoch_57.pt" VAE_S3_PATH = "s3://suno-data/christian/25hz_vae_peaq_kl_0.005.pth" EXAMPLES = 5000 # to run all DITTO_EMBEDDING_DIM = 128 SEP_SAMPLE_RATE = 44100 DITTO_SR = 24000 DITTO_LEN_S = 120 TASK = "artist_vox_sim" ITEMS_PER_ARTIST = 10 DITTO_SIMILARITY_CUTOFF = 0.4 SILENCE_CUTOFF = 30 def load_jsonl_artist(data_path): artists = {} # artist id -> [s3 paths] S3 = "s3_filepath" ARTIST = "artists" LYRICS = "text" lines_read = 0 with open(DATA_PATH, "r", encoding="utf-8") as file: for line in file: try: sample = json.loads(line) if LYRICS not in sample or len(sample[LYRICS]) < 50: continue if ARTIST not in sample: continue artist_ids = sample[ARTIST] for artist_id in artist_ids: if artist_id not in artists: artists[artist_id] = [] artists[artist_id].append(sample[S3]) except json.JSONDecodeError as e: print(f"Error decoding JSON: {e}") lines_read += 1 if EXAMPLES is not None and lines_read > EXAMPLES: break return artists def cosine_similarity(a, b): dot_product = np.dot(a, b) magnitude_a = np.sqrt(np.dot(a, a)) magnitude_b = np.sqrt(np.dot(b, b)) return dot_product / (magnitude_a * magnitude_b) def load_webm_audio_from_s3(s3_path): s3 = boto3.client("s3") prefix = "s3://" path = s3_path[len(prefix) :] bucket, *key_parts = path.split("/") key = "/".join(key_parts) with tempfile.NamedTemporaryFile(suffix=".webm") as temp_file: s3.download_file(bucket, key, temp_file.name) waveform, sr = torchaudio.load(temp_file.name) waveform = torch.mean(waveform, dim=0).unsqueeze(0)[:, : sr * 120] if sr != SEP_SAMPLE_RATE: resampler = torchaudio.transforms.Resample(sr, SEP_SAMPLE_RATE) waveform = resampler(waveform) return waveform def get_vocal_embed(s3_id, task, trim=True): waveform = load_webm_audio_from_s3(s3_id) vocals = ss_vad.encode(waveform) if trim: vocals, _ = librosa.effects.trim(vocals, top_db=SILENCE_CUTOFF) vocals = torch.from_numpy(vocals) resampler = torchaudio.transforms.Resample(SEP_SAMPLE_RATE, DITTO_SR) vocals_embed = ditto.music_to_latent(resampler(vocals), task=task)[0].detach().cpu().numpy() return vocals, vocals_embed if __name__ == "__main__": ditto_path = chirp_v3._get_model_if_needed(DITTO_S3_PATH) ditto = Ditto( latent_dim=DITTO_EMBEDDING_DIM, model_path=ditto_path, is_flash=False, is_serving=True, ) ditto = ditto.eval().to("cuda") ss_vad_config_path = chirp_v3._get_model_if_needed(ss_vad.YAML_PATH) ss_vad_model_path = chirp_v3._get_model_if_needed(ss_vad.MODEL_PATH) ss_vad.preload_models(checkpoint_filepath=ss_vad_model_path, config_path=ss_vad_config_path) preload_vae_models(VAE_S3_PATH) print("Finish loading models") artists = load_jsonl_artist(DATA_PATH) score_data = {} sus_personas = {} for parent_id, parent_s3 in tqdm(artists.items()): num_items = len(parent_s3) if num_items < 2: continue random_idx = random.sample(range(0, num_items), min(ITEMS_PER_ARTIST * 2, num_items)) for idx in range(0, len(random_idx) - 1, 2): artist_parent = parent_s3[random_idx[idx]] parent_vocals, parent_embed = get_vocal_embed(artist_parent, task=TASK) artist_child = parent_s3[random_idx[idx + 1]] child_vocals, child_embed = get_vocal_embed(artist_child, task=TASK) score = cosine_similarity(parent_embed, child_embed) if parent_id not in score_data: score_data[parent_id] = [] score_data[parent_id].append(score) if score < DITTO_SIMILARITY_CUTOFF: if parent_id not in sus_personas: sus_personas[parent_id] = [] sus_personas[parent_id].append((artist_parent, artist_child)) total_bad_personas = sum([len(personas) for personas in sus_personas.values()]) total_personas = sum([len(personas) for personas in score_data.values()]) print( f"Found {total_bad_personas} bad personas to filter, {(total_bad_personas / total_personas) * 100}% of all persona pairs" ) np.savez(f"training_data_artist_{TASK}_{DITTO_LEN_S}_{SILENCE_CUTOFF}.npz", **score_data) np.savez(f"bad_persona_pairs_{TASK}.npz", **sus_personas)