import json import numpy as np from tqdm import tqdm import glob TASK = "self_sim" # update this DATA_PATH = f"/app/suno/sara/ditto_v2_{TASK}/metas/*.jsonl" DITTO_PATH = f"/app/suno/sara/ditto_v2_{TASK}/*.npz" COVER_PATH = "/app/suno/sara/metas_v0_cover.jsonl" OUT_FILE = f"/app/suno/sara/ditto_v2_{TASK}_cover_scores.jsonl" def load_jsonl_cover(data_path): covers = {} # parent id -> [child ids] PARENT = "parent_id" ID = "id" 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 covers: covers[parent] = [] covers[parent].append(sample[ID]) except json.JSONDecodeError as e: print(f"Error decoding JSON: {e}") return covers 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_file_locs(ditto_path): ditto_locs = {} for file_path in tqdm(glob.glob(ditto_path)): ditto_data = np.load(file_path) for key in ditto_data: ditto_locs[key] = file_path return ditto_locs if __name__ == "__main__": print("Parsing cover data...") covers = load_jsonl_cover(COVER_PATH) print("Parsing ditto data...") ditto_locs = get_ditto_file_locs(DITTO_PATH) ditto_scores = [] misses = 0 entries = 0 print("Calculating scores...") for parent in tqdm(covers): child_tracks = covers[parent] if parent not in ditto_locs: misses += 1 continue parent_ditto = np.load(ditto_locs[parent])[parent] parent_mean = np.mean(parent_ditto, axis=0) for child in child_tracks: if child not in ditto_locs: misses += 1 continue child_ditto = np.load(ditto_locs[child])[child] child_mean = np.mean(child_ditto, axis=0) score = cosine_similarity(parent_mean, child_mean) ditto_scores.append({"parent_id": parent, "child_id": child, f"score_{TASK}": float(score)}) entries += 1 # if entries > 10000: # break num_pairs = len(ditto_scores) print(f"Done processing {num_pairs} pairs") print(f"{misses} missed ids when processing ditto covers, thats {misses/num_pairs * 100}%") print(f"Writing results to {OUT_FILE}...") try: with open(OUT_FILE, "w", encoding="utf-8") as f: for item in ditto_scores: json_line = json.dumps(item, ensure_ascii=False) f.write(json_line + "\n") except Exception as e: raise Exception(f"Error writing to JSONL file: {str(e)}") print("Done!")