import os import sys import json import torch import argparse import librosa import pandas as pd import numpy as np from tqdm.contrib.concurrent import thread_map from datetime import datetime from suno_utils.gpt.chirp_v2_5 import ( preload_codec_models, codec_decode, ) # sys.path.append("/home/minz/neon/ditto-cover") # from ditto.models.ditto import Ditto sys.path.append("/home/minz/neon/ditto-training") from ditto_v2.models.ditto import Ditto def get_emb(wav, task="cover"): # resample wav = librosa.resample(wav[: 48000 * 360], orig_sr=48000, target_sr=24000) expected_len = 24000 * 15 if len(wav) < expected_len: wav = np.pad(wav, (0, expected_len - len(wav)), mode="constant") # make chunks chunks = librosa.util.frame(wav, frame_length=expected_len, hop_length=expected_len) chunks = chunks.transpose(1, 0) # get embedding inp = torch.tensor(chunks).unsqueeze(1).cuda() # # self_sim: audio, random aug, two random aug for same song # artist_sim: audio, multiple artists, multiple songs, 2 for each artist, # alubm_sim: audio, same album # artist_vox_sim: two song per artist, source sep vocal, sim # genre_sim: genre_text to audio if task == "cover": emb = ditto.music_to_latent(inp, "self_sim") elif task == "artist": emb = ditto.music_to_latent(inp, "artist_sim") return emb.mean(dim=0).detach().cpu().numpy() def get_similarity(input_job): s3_id, task = input_job assert task in ["cover", "artist"] try: NPZ_DIR = "/app/suno/data/dpo/30b_npz" local_path = ( f"{NPZ_DIR}/{s3_id + ('_gen_cycle' if 'cycle' in NPZ_DIR else '')}.npz" ) if not os.path.exists(local_path): raise ValueError() try: temp_npz = np.load(local_path) if "v4.0_raw" in temp_npz: arr = temp_npz["v4.0_raw"] elif "v3.5_raw" in temp_npz: if "cycle" not in NPZ_DIR: print(f"weird, {local_path}, with only v3.5") arr = temp_npz["v3.5_raw"] elif "v3.0_raw" in temp_npz: if "cycle" not in NPZ_DIR: print(f"weird, {local_path}, with only v3.0") arr = temp_npz["v3.0_raw"] else: raise ValueError() except Exception as e: print(local_path) raise eval assert arr.shape[1] == 13 if task == "cover": prompt_arr = temp_npz["cover_arr"] elif task == "artist": prompt_arr = temp_npz["artist_arr"] assert prompt_arr.shape[1] == 13 try: seed_codec = prompt_arr[:, 1:] seed_audio = codec_decode(torch.tensor(seed_codec).long()) seed_emb = get_emb(seed_audio.array_float.mean(axis=0), task=task) except Exception as e: print("error in seed embedding --> ", e) return s3_id, 2 # seed_codec = prompt_arr[:, 1:] # seed_audio = codec_decode(torch.tensor(seed_codec).long()) # seed_emb = get_emb(seed_audio.array_float.mean(axis=0)) codec_arr = arr[:, 1:] audio = codec_decode(torch.tensor(codec_arr).long()) audio_emb = get_emb(audio.array_float.mean(axis=0)) output_npz_path = f"/app/suno/data/dpo/ditto/{s3_id}_ditto_emb.npz" np.savez(output_npz_path, audio_emb=audio_emb, seed_emb=seed_emb) sim = np.dot(seed_emb, audio_emb) return s3_id, float(sim) except Exception as e: print("error in processing --> ", e) return s3_id, 2 if __name__ == "__main__": print(f"working with GPU:{os.environ['CUDA_VISIBLE_DEVICES']} ") parser = argparse.ArgumentParser() parser.add_argument("--task_index", type=int, default=0) parser.add_argument("--max_index", type=int, default=4) args = parser.parse_args() # Ensure task_index is valid if args.task_index < 0 or args.task_index >= args.max_index: raise ValueError(f"task_index must be between 0 and {args.max_index - 1}") # load models print("loading models...") _ = preload_codec_models("/app/suno/models/chirp_v2/dac_2c_25x12.pt", device="cuda") ditto = Ditto( latent_dim=128, model_path="/home/minz/logs/ditto_v2_local_8gpu_cont/ditto_v2_epoch_57.pt", is_flash=False, ) ditto = ditto.eval() ditto = ditto.cuda() print("Finished loading models!") # read data print("loading data...") input_df = pd.read_pickle( "/home/tony/Data/Preference/30b_v6/interesting_clips_v4_h_t_6_20250115_full.pkl" ) cover_ids = sorted(list(input_df[input_df["task"] == "cover"]["s3_id"].unique())) artist_ids = sorted( list(input_df[input_df["task"] == "artist_consistency"]["s3_id"].unique()) ) print(f"Total cover_ids: {len(cover_ids)}") print(f"Total artist_ids: {len(artist_ids)}") input_tasks = [] for cover_id in cover_ids: input_tasks.append((cover_id, "cover")) for artist_id in artist_ids: input_tasks.append((artist_id, "artist")) # Split the data into chunks chunk_size = len(input_tasks) // args.max_index start_index = args.task_index * chunk_size end_index = ( start_index + chunk_size if args.task_index < args.max_index - 1 else len(input_tasks) ) chunk_input_ids = input_tasks[start_index:end_index] print( f"Processing chunk {args.task_index + 1}/{args.max_index} with {len(chunk_input_ids)} ids" ) print("done!") print("get similarity...") num_threads = 10 today_str = datetime.now().strftime("%Y%m%d") output_json_path = os.path.join( f"/home/tony/Data/Preference/30b_v6/similarity_{today_str}_chunk{args.task_index + 1}of{args.max_index}.json" ) analyzed_outputs = thread_map( get_similarity, chunk_input_ids, max_workers=num_threads, chunksize=1, ) out = {input_id: sim for input_id, sim in analyzed_outputs if sim is not None} with open(output_json_path, "w") as fp: json.dump(out, fp) print(f"done! Output saved to {output_json_path}")