{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "from tqdm import tqdm\n",
    "import glob\n",
    "from suno_utils.utils.text import read_jsonl, write_json, read_json, write_jsonl\n",
    "import os\n",
    "import math"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Get Similarity Scores"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "TASK = \"self_sim\"  # update this\n",
    "DITTO_PATH = f\"/app/suno/sara/cover_filter/ditto_v2_{TASK}_raw/*.npz\"\n",
    "COVER_PATH = \"/app/suno/sara/cover_filter/filtered_cover_name_view_detailed.jsonl\"\n",
    "DISCOGS_PATH = \"/app/suno/sara/metas_v0_discogs.jsonl\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def cosine_similarity(a, b):\n",
    "    dot_product = np.dot(a, b)\n",
    "    magnitude_a = np.sqrt(np.dot(a, a))\n",
    "    magnitude_b = np.sqrt(np.dot(b, b))\n",
    "    return dot_product / (magnitude_a * magnitude_b)\n",
    "\n",
    "\n",
    "def get_ditto_scores(ditto_path):\n",
    "    ditto_scores = {}\n",
    "    idx = 0\n",
    "    for file_path in tqdm(glob.glob(ditto_path)):\n",
    "        ditto_data = np.load(file_path)\n",
    "        for key in ditto_data:\n",
    "            ditto_mean = np.mean(ditto_data[key], axis=0)\n",
    "            ditto_scores[key] = ditto_mean\n",
    "        idx += 1\n",
    "        if idx > 1000:\n",
    "            break\n",
    "\n",
    "    return ditto_scores"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "ditto_scores = get_ditto_scores(DITTO_PATH)\n",
    "print(len(ditto_scores))\n",
    "# write_json(ditto_scores, \"/app/suno/sara/cover_filter/raw_parent_to_child_ditto.json\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "discogs = read_jsonl(DISCOGS_PATH, progress=True)\n",
    "discogs_ids = {}\n",
    "for data in discogs:\n",
    "    discogs_ids[data[\"id\"]] = 0  # data\n",
    "del discogs"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "test = read_jsonl(COVER_PATH, progress=True)\n",
    "missing_from_discogs = 0\n",
    "for idx, val in tqdm(enumerate(test)):\n",
    "    if \"parent_id\" not in val:\n",
    "        if val[\"id\"] not in discogs_ids:\n",
    "            missing_from_discogs += 1\n",
    "        else:\n",
    "            test[idx] = discogs_ids[val[\"id\"]]\n",
    "print(missing_from_discogs)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "write_jsonl(test, \"/app/suno/sara/cover_filter/filtered_cover_name_view_detailed.jsonl\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from joblib import Parallel, delayed\n",
    "\n",
    "\n",
    "def process_batch(batch, ditto_scores):\n",
    "    skipped = 0\n",
    "    found = 0\n",
    "    ditto_score_map = {}\n",
    "    for row in batch:\n",
    "        if \"parent_id\" in row:\n",
    "            parent = row[\"parent_id\"]\n",
    "            child = row[\"id\"]\n",
    "\n",
    "            if parent not in ditto_scores or child not in ditto_scores:\n",
    "                skipped += 1\n",
    "            else:\n",
    "                parent_ditto = ditto_scores[parent]\n",
    "                child_ditto = ditto_scores[child]\n",
    "                similarity_score = cosine_similarity(parent_ditto, child_ditto)\n",
    "                if parent not in ditto_score_map:\n",
    "                    ditto_score_map[parent] = {}\n",
    "                ditto_score_map[parent][child] = str(similarity_score)\n",
    "                found += 1\n",
    "\n",
    "    return ditto_score_map, skipped, found"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "test = read_jsonl(\n",
    "    \"/app/suno/sara/cover_filter/filtered_cover_name_view_detailed.jsonl\",\n",
    "    progress=True,\n",
    "    max_lines=1000000,\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "batch_size = 10000\n",
    "n_jobs = 32\n",
    "batches = [test[i : i + batch_size] for i in range(0, len(test), batch_size)]\n",
    "\n",
    "# Show progress with tqdm and use joblib for parallelization\n",
    "results = Parallel(n_jobs=n_jobs)(\n",
    "    delayed(process_batch)(batch, ditto_scores)\n",
    "    for batch in tqdm(batches, desc=\"Processing\")\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Combine results\n",
    "final_ditto_score_map = {}\n",
    "total_skipped = 0\n",
    "total_found = 0\n",
    "\n",
    "for ditto_map, skipped, found in results:\n",
    "    # Merge dictionaries\n",
    "    for parent, children in ditto_map.items():\n",
    "        if parent not in final_ditto_score_map:\n",
    "            final_ditto_score_map[parent] = {}\n",
    "        final_ditto_score_map[parent].update(children)\n",
    "\n",
    "    total_skipped += skipped\n",
    "    total_found += found\n",
    "\n",
    "\n",
    "print(total_found, total_skipped)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "chunksize = 200\n",
    "start_idx = 0\n",
    "end_idx = 1000000\n",
    "n_chunks = int(np.ceil((end_idx - start_idx) / chunksize))\n",
    "njobs = 5\n",
    "n_chunks"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "for idx_chunk in tqdm(\n",
    "    funcy.chunks(chunksize, list(range(start_idx, end_idx))), total=n_chunks\n",
    "):\n",
    "    n_jobs = np.min([njobs, chunksize, len(idx_chunk)])\n",
    "    # collect relevant parts of meta file to avoid copying all to subprocesses\n",
    "    tmp_uid_chunks = Parallel(n_jobs=n_jobs, prefer=\"threads\")(\n",
    "        delayed(process_data)(\n",
    "            f\"s3://suno-data/sara/cover_filter/{SEMANTIC_EMBED_DIR}/\"\n",
    "            + f\"metas/part_{idx}.jsonl\"\n",
    "        )\n",
    "        for idx in idx_chunk\n",
    "    )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "n_jobs = 16  # Use all available cores except one\n",
    "batch_size = 100000  # Process this many rows per batch\n",
    "\n",
    "# Calculate number of batches\n",
    "num_batches = math.ceil(len(test) / batch_size)\n",
    "\n",
    "print(f\"Processing {len(test)} rows using {n_jobs} cores\")\n",
    "print(f\"Data split into {num_batches} batches, processing {batch_size} rows per batch\")\n",
    "\n",
    "# Initialize result containers\n",
    "ditto_score_map = {}\n",
    "total_skipped = 0\n",
    "total_found = 0\n",
    "\n",
    "# Create batch indices for better management\n",
    "batches = [test[i : (i + batch_size)] for i in range(0, len(test), batch_size)]\n",
    "\n",
    "# Process in smaller chunks to show progress\n",
    "chunk_size = 10  # Process this many batches at a time\n",
    "for chunk_start in range(0, len(batch_indices), chunk_size):\n",
    "    chunk_end = min(chunk_start + chunk_size, len(batch_indices))\n",
    "    current_chunk = batch_indices[chunk_start:chunk_end]\n",
    "\n",
    "    print(\n",
    "        f\"Processing chunk {chunk_start//chunk_size + 1}/{math.ceil(len(batch_indices)/chunk_size)}\"\n",
    "    )\n",
    "\n",
    "    # Process this chunk of batches in parallel with timeout\n",
    "    results = Parallel(n_jobs=n_jobs, verbose=0, timeout=600)(\n",
    "        delayed(process_batch)(i, batch, ditto_scores)\n",
    "        for i, batch in tqdm(current_chunk, desc=f\"Batches {chunk_start}-{chunk_end-1}\")\n",
    "    )\n",
    "\n",
    "    # Merge results from this chunk\n",
    "    for local_map, local_skipped, local_found in results:\n",
    "        # Merge the dictionaries\n",
    "        for parent, children in local_map.items():\n",
    "            if parent not in ditto_score_map:\n",
    "                ditto_score_map[parent] = {}\n",
    "            ditto_score_map[parent].update(children)\n",
    "\n",
    "        total_skipped += local_skipped\n",
    "        total_found += local_found\n",
    "\n",
    "    # Print progress after each chunk\n",
    "    completion = min(100, (chunk_end / len(batch_indices)) * 100)\n",
    "    print(\n",
    "        f\"Progress: {completion:.1f}% - Processed: {total_found}, Skipped: {total_skipped}\"\n",
    "    )\n",
    "\n",
    "print(f\"Completed! Processed {total_found} pairs, skipped {total_skipped} pairs\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "test = read_jsonl(\"/app/suno/sara/cover_filter/filtered_cover_name_view_detailed.jsonl\")\n",
    "# ditto_scores = read_json(\"/app/suno/sara/cover_filter/raw_parent_to_child_ditto.json\")\n",
    "\n",
    "ditto_score_map = {}\n",
    "parent_to_data = {}\n",
    "parent_to_covers = {}\n",
    "skipped = 0\n",
    "for row in tqdm(test):\n",
    "    if \"parent_id\" not in row:\n",
    "        parent_to_data[row[\"id\"]] = row\n",
    "        parent_to_covers[row[\"id\"]] = []\n",
    "for row in tqdm(test):\n",
    "    if \"parent_id\" in row:\n",
    "        parent = row[\"parent_id\"]\n",
    "        child = row[\"id\"]\n",
    "\n",
    "        if parent not in ditto_scores or child not in ditto_scores:\n",
    "            skipped += 1\n",
    "            continue\n",
    "        similarity_score = cosine_similarity(ditto_scores[parent], ditto_scores[child])\n",
    "        if parent not in ditto_score_map:\n",
    "            ditto_score_map[parent] = {}\n",
    "        ditto_score_map[parent][child] = str(similarity_score)\n",
    "        if similarity_score > 0.4 and similarity_score < 0.8:\n",
    "            parent_to_covers[parent].append(row)\n",
    "\n",
    "final_data = []\n",
    "sources = 0\n",
    "num_covers = 0\n",
    "for parent, covers in parent_to_covers.items():\n",
    "    if len(covers) > 0:\n",
    "        final_data.append(parent_to_data[parent])\n",
    "        sources += 1\n",
    "        trimmed_covers = covers[:10]\n",
    "        for cover in trimmed_covers:\n",
    "            final_data.append(cover)\n",
    "            num_covers += 1\n",
    "\n",
    "print(len(final_data), sources, num_covers, skipped)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "for row in test:\n",
    "    if \"parent_id\" not in row:\n",
    "        parent_to_data[row[\"id\"]]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "idx_map_professional = {}\n",
    "idx_map_amateur = {}\n",
    "parent_to_idx = {}\n",
    "skipped_parents = 0\n",
    "found_parents = 0\n",
    "found_children = 0\n",
    "skipped_children = 0\n",
    "missing_from_discogs = 0\n",
    "total_prof = 0\n",
    "total_am = 0\n",
    "for idx, val in tqdm(enumerate(test)):\n",
    "    if \"parent_id\" not in val:\n",
    "        if val[\"id\"] in ditto_scores:\n",
    "            # parent\n",
    "            # if val[\"id\"] not in discogs_ids:\n",
    "            #    missing_from_discogs += 1\n",
    "            idx_map_professional[val[\"id\"]] = {}\n",
    "            # idx_map_amateur[val[\"id\"]] = {}\n",
    "            found_parents += 1\n",
    "        else:\n",
    "            skipped_parents += 1\n",
    "    else:\n",
    "        parent = val[\"parent_id\"]\n",
    "        child = val[\"id\"]\n",
    "        if child not in ditto_scores:\n",
    "            skipped_children += 1\n",
    "        if parent in ditto_scores and child in ditto_scores:\n",
    "            found_children += 1\n",
    "            parent_ditto = ditto_scores[parent]\n",
    "            child_ditto = ditto_scores[child]\n",
    "            similarity_score = cosine_similarity(parent_ditto, child_ditto)\n",
    "\n",
    "            idx_map_professional[parent][child] = str(similarity_score)\n",
    "            # if child in discogs_ids:\n",
    "            #    idx_map_professional[parent][child] = str(similarity_score)\n",
    "            #    total_prof += 1\n",
    "            # else:\n",
    "            #    idx_map_amateur[parent][child] = str(similarity_score)\n",
    "            #    total_am += 1\n",
    "\n",
    "\n",
    "print(len(test))\n",
    "print(missing_from_discogs, total_prof, total_am)\n",
    "print(\n",
    "    f\"Skipped {skipped_parents} parents and {skipped_children} children. Found {found_parents} parents, {found_children} children\"\n",
    ")\n",
    "print(f\"{found_parents + found_children}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "write_json(\n",
    "    idx_map_professional,\n",
    "    \"/app/suno/sara/cover_filter/parent_to_cover_self_sim_scores_raw.json\",\n",
    ")\n",
    "# write_json(idx_map_professional, \"parent_to_cover_self_sim_scores_v0_prof.json\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Make Index Maps"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "SPLIT = \"val\"  # \"tr\" or \"val\"\n",
    "SOURCE_PATH = \"/app/suno/sara/cover_filter_v2/\"\n",
    "SUFFIX = \"\"\n",
    "OUT_PATH = \"/app/suno/sara/cover_filter_v2/\"\n",
    "OUT_SUFFIX = \"_prof\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "meta_data = read_jsonl(os.path.join(SOURCE_PATH, f\"metas_{SPLIT}.jsonl\"), progress=True)\n",
    "parent_to_cover = read_json(\n",
    "    \"/app/suno/sara/cover_filter/parent_to_cover_self_sim_scores_raw.json\"\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "idx_vals = {}\n",
    "parent_to_idx = {}\n",
    "covers = 0\n",
    "skipped = 0\n",
    "filtered = 0\n",
    "\n",
    "for idx, val in tqdm(enumerate(meta_data)):\n",
    "    if \"parent_id\" not in val:\n",
    "        parent = str(val[\"id\"])\n",
    "        idx_vals[idx] = []\n",
    "        parent_to_idx[parent] = idx\n",
    "\n",
    "for idx, val in tqdm(enumerate(meta_data)):\n",
    "    if \"parent_id\" in val:\n",
    "        parent = str(val[\"id\"])\n",
    "        child = val[\"id\"]\n",
    "        parent = val[\"parent_id\"]\n",
    "        if (\n",
    "            parent in parent_to_cover\n",
    "            and child in parent_to_cover[parent]\n",
    "            and parent in parent_to_idx\n",
    "        ):\n",
    "            similarity_score = float(parent_to_cover[parent][child])\n",
    "            if (\n",
    "                similarity_score > 0.0\n",
    "                and similarity_score < 1.0\n",
    "                and child in discogs_ids\n",
    "            ):\n",
    "                parent_idx = parent_to_idx[parent]\n",
    "                if len(idx_vals[parent_idx]) < 50:\n",
    "                    idx_vals[parent_idx].append(idx)\n",
    "                    covers += 1\n",
    "            else:\n",
    "                filtered += 1\n",
    "        else:\n",
    "            # print(parent in parent_to_cover, parent in parent_to_idx)\n",
    "            skipped += 1\n",
    "\n",
    "final_idx_vals = {}\n",
    "for k, v in idx_vals.items():\n",
    "    if len(v) > 0:\n",
    "        final_idx_vals[k] = v\n",
    "\n",
    "print(f\"Found {covers}, skipped {skipped}, filtered {filtered}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "info_data = {\"covers\": {\"task\": \"covers\", \"idx_map\": final_idx_vals}}\n",
    "write_json(info_data, os.path.join(OUT_PATH, f\"info_{SPLIT}{OUT_SUFFIX}.json\"))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "info_data[\"covers\"]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "suno_clean",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "codemirror_mode": {
    "name": "ipython",
    "version": 3
   },
   "file_extension": ".py",
   "mimetype": "text/x-python",
   "name": "python",
   "nbconvert_exporter": "python",
   "pygments_lexer": "ipython3",
   "version": "3.10.15"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
