import json import torchaudio import boto3 import tempfile import torch import numpy as np from tqdm import tqdm from suno_utils.gpt import chirp_v2_5 as chirp_v3 from suno_utils.models.ditto_v2.ditto_v2 import Ditto EXAMPLES = 5000 # None to run all examples in the jsonl DATA_PATH = "/app/suno/sara/metas_v0_cover.jsonl" DITTO_S3_PATH = "s3://suno-data/minz/models/ditto_v2_epoch_57.pt" DITTO_EMBEDDING_DIM = 128 DITTO_SR = 24000 DITTO_LEN_S = 120 TASK = "self_sim" DITTO_SIMILARITY_CUTOFF = 0.75 def load_jsonl_cover(data_path): cover_sources = {} # parent id -> parent s3 cover_children = {} # parent id -> [children s3] PARENT = "parent_id" S3 = "s3_filepath" ID = "id" lines_read = 0 with open(data_path, "r", encoding="utf-8") as file: for line in tqdm(file): try: sample = json.loads(line) if PARENT in sample: parent = sample[PARENT] if parent not in cover_children: cover_children[parent] = [] cover_children[parent].append((sample[ID], sample[S3])) else: cover_sources[sample[ID]] = 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 cover_sources, cover_children 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) if sr != DITTO_SR: resampler = torchaudio.transforms.Resample(sr, DITTO_SR) waveform = resampler(waveform)[:, : (DITTO_SR * DITTO_LEN_S)] return waveform 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") cover_sources, cover_children = load_jsonl_cover(DATA_PATH) score_data = {} bad_covers = {} for parent_id, parent_s3 in tqdm(cover_sources.items()): children_s3 = cover_children[parent_id] score_data[parent_id] = [] parent_waveform = load_webm_audio_from_s3(parent_s3) parent_embed = ditto.music_to_latent(parent_waveform, task=TASK)[0].detach().cpu().numpy() for cover_id, cover_s3 in children_s3: waveform = load_webm_audio_from_s3(cover_s3) embedding = ditto.music_to_latent(waveform, task=TASK)[0].detach().cpu().numpy() score = cosine_similarity(parent_embed, embedding) if score > DITTO_SIMILARITY_CUTOFF: if parent_id not in bad_covers: bad_covers[parent_id] = [] bad_covers[parent_id].append(cover_id) score_data[parent_id].append(score) total_bad_covers = sum([len(covers) for covers in bad_covers.values()]) total_covers = sum([len(covers) for covers in score_data.values()]) print( f"Found {total_bad_covers} bad covers to filter, {(total_bad_covers / total_covers) * 100}% of all covers" ) np.savez(f"training_data_cover_{TASK}_{DITTO_LEN_S}.npz", **score_data) np.savez(f"bad_covers_{TASK}.npz", **bad_covers)