import os import sys import json import torch import tqdm import argparse import librosa import numpy as np from tqdm.contrib.concurrent import thread_map from suno_utils.utils.text import read_json, read_jsonl from suno_utils.audio import Audio from suno_utils.gpt.chirp_v2_5 import ( preload_semantic_models, preload_codec_models, _get_model_if_needed, GenerationConfig, generate, semantic_encode, codec_encode, codec_decode_stream_to_full_audio, codec_decode ) sys.path.append("/home/minz/neon/ditto-cover") from ditto.models.ditto import Ditto def get_emb(wav): # resample wav = librosa.resample(wav[:48000*360], orig_sr=48000, target_sr=24000) # make chunks chunks = librosa.util.frame(wav, frame_length=24000 * 15, hop_length=24000 * 15) chunks = chunks.transpose(1, 0) # get embedding inp = torch.tensor(chunks).unsqueeze(1).cuda() emb = ditto.music_to_latent(inp) return emb.mean(dim=0).detach().cpu().numpy() def get_cover_similarity(ix): try: seed = keys[ix] print('processing ', seed) seed_codec = data[int(seed)][:, 1:] seed_codec = seed_codec[:np.where(seed_codec == 2048)[0][0]] seed_audio = codec_decode(torch.tensor(seed_codec).long()) seed_emb = get_emb(seed_audio.array_float.mean(axis=0)) except Exception as e: print("error in seed", e) return None, None sims = [] covers = info[seed] for cover in covers: try: cover_codec = data[int(cover)][:, 1:] cover_codec = cover_codec[:np.where(cover_codec == 2048)[0][0]] cover_audio = codec_decode(torch.tensor(cover_codec).long()) cover_emb = get_emb(cover_audio.array_float.mean(axis=0)) sim = np.dot(seed_emb, cover_emb) / (np.linalg.norm(seed_emb) * np.linalg.norm(cover_emb)) sims.append({"similarity": sim, "seed": seed, "cover": cover}) except Exception as e: print("error in cover", e) continue print('completed ', seed) return seed, sims if __name__ == "__main__": parser = argparse.ArgumentParser() parser.add_argument("--trial_index", type=int, default=0) args = parser.parse_args() process_start_index = args.trial_index * 12000 # load models print("loading models...") _ = preload_codec_models("/app/suno/models/chirp_v2/dac_2c_25x12.pt") ditto = Ditto(model_path="/home/minz/logs/ditto_cover/epoch40.ckpt", is_flash=False) ditto = ditto.eval() ditto = ditto.cuda() print("done!") # read data print("loading data...") SPLIT = "tr" # tr or val OUTPUT_PATH = "/app/suno/minz/cover_sim/" + SPLIT info = read_json("/app/suno/data/chirp_v4/multi/info_%s.json" % SPLIT)["discogs_covers"]["idx_map"] keys = list(info.keys()) data = np.memmap("/app/suno/data/chirp_v4/multi/data_%s.bin" % SPLIT, dtype=np.uint16, mode="r") data = data.reshape(-1, 6016, 13) # metadata = read_jsonl("/app/suno/data/chirp_v4/multi/metas_%s.jsonl" % SPLIT) print("done!") print("get similarity...") num_threads = 8 batch_size = 100 for i in tqdm.tqdm(range(process_start_index, process_start_index + 12000, batch_size)): start_index = i end_index = min(start_index + batch_size, len(info)) if start_index > len(info): continue output_json_path = os.path.join(OUTPUT_PATH, f"batch_{start_index}.json") with open(output_json_path, "w") as fp: fp.write("") analyzed_outputs = thread_map( get_cover_similarity, range(start_index, end_index), max_workers=num_threads, chunksize=1, disable=True ) out = {seed: sims for seed, sims in analyzed_outputs if sims is not None} with open(output_json_path, "w") as fp: json.dump(out, fp, default=float) print("done!")