import numpy as np from tqdm import tqdm import glob from suno_utils.utils.text import read_jsonl from joblib import Parallel, delayed from suno_utils.utils.text import write_json """ This script assumes you've already run modal encode and have the ditto embeddings saved to DITTO_PATH It saves a dictionary of format {parent_id: {child_id: ditto_cosine_sim}} for looking up ditto similarity This is meant to be run before make_filtered_idx.py or make_filtered_meta.py """ # update these as needed TASK = "self_sim" DITTO_PATH = f"/app/suno/sara/cover_filter/ditto_v2_{TASK}_raw/*.npz" COVER_PATH = "/app/suno/sara/cover_filter/filtered_cover_name_view_detailed.jsonl" 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 get_ditto_scores(ditto_path): ditto_scores = {} idx = 0 for file_path in tqdm(glob.glob(ditto_path)): ditto_data = np.load(file_path) for key in ditto_data: ditto_mean = np.mean(ditto_data[key], axis=0) ditto_scores[key] = ditto_mean idx += 1 return ditto_scores def process_batch(batch, ditto_scores): skipped = 0 found = 0 ditto_score_map = {} for row in batch: if "parent_id" in row: parent = row["parent_id"] child = row["id"] if parent not in ditto_scores or child not in ditto_scores: skipped += 1 else: parent_ditto = ditto_scores[parent] child_ditto = ditto_scores[child] similarity_score = cosine_similarity(parent_ditto, child_ditto) if parent not in ditto_score_map: ditto_score_map[parent] = {} ditto_score_map[parent][child] = str(similarity_score) found += 1 return ditto_score_map, skipped, found print("Loading ditto scores...") ditto_scores = get_ditto_scores(DITTO_PATH) print(f"Found {len(ditto_scores)} ditto embeddings") print("Loading cover data...") test = read_jsonl(COVER_PATH, progress=True) batch_size = 10000 n_jobs = 32 batches = [test[i : i + batch_size] for i in range(0, len(test), batch_size)] # Show progress with tqdm and use joblib for parallelization results = Parallel(n_jobs=n_jobs, prefer="threads")( delayed(process_batch)(batch, ditto_scores) for batch in tqdm(batches, desc="Scoring cover similarity") ) # Combine results print("Combining results...") final_ditto_score_map = {} total_skipped = 0 total_found = 0 for ditto_map, skipped, found in results: # Merge dictionaries for parent, children in ditto_map.items(): if parent not in final_ditto_score_map: final_ditto_score_map[parent] = {} final_ditto_score_map[parent].update(children) total_skipped += skipped total_found += found print( f"Found {total_found} scores, skipped {total_skipped}, {total_found * 100.0 / (total_found + total_skipped)}% hits" ) print("Writing to file...") write_json(final_ditto_score_map, "/app/suno/sara/cover_filter/raw_ditto_mappings_v0.json")