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")
