import json import torchaudio import boto3 import tempfile import torch import numpy as np from tqdm import tqdm import sys from suno_utils.audio import Audio from suno_utils.gpt import chirp_v2_5 as chirp_v3 sys.path.append("/home/sara/neon/ditto-training") from ditto_v2.models.ditto import Ditto EXAMPLES = 1000 # None to run all examples in the jsonl DATA_PATH = "/app/suno/sara/metas_v0_cover.jsonl" S = torch.load("/app/suno/minz/models/ditto_cover.ckpt") SS = {k.replace("model.", ""): v for k, v in S["state_dict"].items()} ditto = Ditto(latent_dim=128, is_flash=False) ditto.load_state_dict(SS) ditto = ditto.eval() ditto = ditto.cuda() ditto.music_encoder = ditto.music_encoder.cuda() DITTO_EMBEDDING_DIM = 128 DITTO_SR = 24000 DITTO_LEN_S = 120 TASK = "cover_sim" DITTO_SIMILARITY_CUTOFF = 0.5 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): # print(a.shape, b.shape) 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 get_audio_emb(audio): num_chunks = 4 chunk_length_samples = 15 * 24000 available_samples = max(0, len(audio) - chunk_length_samples) start_positions = ( [int(i * available_samples / (num_chunks - 1)) for i in range(num_chunks)] if num_chunks > 1 and available_samples > 0 else [0] * num_chunks ) # Extract chunks chunks = [] for start_pos in start_positions: chunk = audio[start_pos : start_pos + chunk_length_samples] if len(chunk) < chunk_length_samples: chunk = np.pad(chunk, (0, chunk_length_samples - len(chunk))) chunks.append(chunk) # Stack chunks and get embeddings chunks_tensor = torch.tensor(np.stack(chunks), device="cuda") emb = ditto.music_to_latent(chunks_tensor, "cover_sim").detach().cpu().numpy() return emb # print(emb.shape) # return emb.mean(axis=0) 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) print(waveform.shape) waveform = torch.mean(waveform, dim=0) if sr != DITTO_SR: resampler = torchaudio.transforms.Resample(sr, DITTO_SR) waveform = resampler(waveform) return waveform if __name__ == "__main__": 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 = get_audio_emb(parent_waveform) for cover_id, cover_s3 in children_s3: waveform = load_webm_audio_from_s3(cover_s3) embedding = get_audio_emb(waveform) 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)