{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 1,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "import torch\n",
    "import pandas as pd\n",
    "from suno_utils.utils.text import read_jsonl, write_jsonl"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# let's load the original metas so that we can also remove examples that do not have lyrics and tags \n",
    "dataset_dir = \"/app/suno/data/diffusion_v5/v0\"\n",
    "train_metas_filepath = os.path.join(dataset_dir, \"metas_tr.jsonl\")\n",
    "metas = read_jsonl(train_metas_filepath, progress=False)\n",
    "print(len(metas))\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "metas[100]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# special map for genius\n",
    "old_genius_metas_filepath = \"/home/christian/code/christian/metadata/genius_hq_metas.jsonl\"\n",
    "old_genius_metas = read_jsonl(old_genius_metas_filepath, progress=False)\n",
    "print(len(old_genius_metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 64,
   "metadata": {},
   "outputs": [],
   "source": [
    "genius_old_to_new_id_map = {m[\"id\"]: m[\"original_id\"] for m in old_genius_metas}\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from tqdm import tqdm\n",
    "from collections import defaultdict\n",
    "\n",
    "valid_ids_metadata = defaultdict(set)\n",
    "subset_valid_counts_metadata = {}\n",
    "subset_total_counts_metadata = {}\n",
    "\n",
    "subset_meta_id_to_text_lang = {}\n",
    "\n",
    "for meta in tqdm(metas):\n",
    "    # check for full text, text_aligned, tags\n",
    "    full_text = meta.get(\"text\", None)\n",
    "    text_aligned = meta.get(\"text_aligned\", None)\n",
    "    tags = meta.get(\"tags\", None)\n",
    "    dataset = meta.get(\"dataset\", None)\n",
    "\n",
    "    text_lang = meta.get(\"lang\", None)\n",
    "    if text_lang is None:\n",
    "        text_lang = meta.get(\"text_lang\", None)\n",
    "\n",
    "    # Track total counts per dataset\n",
    "    subset_total_counts_metadata[dataset] = subset_total_counts_metadata.get(dataset, 0) + 1\n",
    "    \n",
    "    if dataset == \"discogs_subset\":\n",
    "        dataset = \"discogs\"\n",
    "    else:\n",
    "        continue\n",
    "\n",
    "    if dataset not in subset_meta_id_to_text_lang:\n",
    "        subset_meta_id_to_text_lang[dataset] = {}\n",
    "\n",
    "    if dataset == \"imslp\":\n",
    "        valid_ids_metadata[dataset].add(meta[\"id\"])\n",
    "    else:\n",
    "        if tags is not None:\n",
    "            num_tags = len(tags)\n",
    "            if full_text is not None and num_tags > 3:\n",
    "                valid_ids_metadata[dataset].add(meta[\"id\"])\n",
    "                subset_meta_id_to_text_lang[dataset][meta[\"id\"]] = text_lang\n",
    "\n",
    "    # Update valid counts\n",
    "    for ds in valid_ids_metadata:\n",
    "        subset_valid_counts_metadata[ds] = len(valid_ids_metadata[ds])\n",
    "\n",
    "# Print total valid IDs across all datasets\n",
    "total_valid_ids = sum(len(ids) for ids in valid_ids_metadata.values())\n",
    "print(f\"Total valid IDs: {total_valid_ids}\")\n",
    "\n",
    "# print the counts for each dataset\n",
    "for ds in valid_ids_metadata:\n",
    "    print(f\"{ds}: {len(valid_ids_metadata[ds])}\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "for ds in valid_ids_metadata:\n",
    "    print(f\"{ds}: {len(valid_ids_metadata[ds])}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 3,
   "metadata": {},
   "outputs": [],
   "source": [
    "# for diffusion data cut we want to remove examples with high cer and low quality (we can probably use the audio features)\n",
    "# do do the data cut, we will need a way to filter out examples. i think to start we can just save the dataset and the ids of the examples that we want to keep"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# load alignments\n",
    "# load some alignemnts info (h5 alignments)\n",
    "genius_alignments_filepath = (\n",
    "    \"/home/tony/Work/tony/hoot/tmp/genius_hq_alignments_h4_t30_v1.jsonl\"\n",
    ")\n",
    "discogs_alignments_filepath = (\n",
    "    \"/home/tony/Work/tony/hoot/tmp/discogs_hq_alignments_h4_t30_v1.jsonl\"\n",
    ")\n",
    "deezer_alignments_filepath = (\n",
    "    \"/home/tony/Work/tony/hoot/tmp/deezer_hq_alignments_h4_t30_v1.jsonl\"\n",
    ")\n",
    "\n",
    "genius_alignments = read_jsonl(genius_alignments_filepath, progress=False)\n",
    "print(len(genius_alignments))\n",
    "discogs_alignments = read_jsonl(discogs_alignments_filepath, progress=False)\n",
    "print(len(discogs_alignments))\n",
    "deezer_alignments = read_jsonl(deezer_alignments_filepath, progress=False)\n",
    "print(len(deezer_alignments))\n",
    "\n",
    "# make alignments into a dictionary \n",
    "genius_alignments_map = {k: {\"texts\": v, \"cer\": cer} for k, v, cer in genius_alignments}\n",
    "discogs_alignments_map = {k: {\"texts\": v, \"cer\": cer} for k, v, cer in discogs_alignments}\n",
    "deezer_alignments_map = {k: {\"texts\": v, \"cer\": cer} for k, v, cer in deezer_alignments}"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 10,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Define the mapping between subsets and their corresponding alignment maps\n",
    "subset_to_alignment_map = {\n",
    "    \"discogs\": discogs_alignments_map,\n",
    "    #\"genius\": genius_alignments_map, \n",
    "    #\"deezer\": deezer_alignments_map\n",
    "}\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# import counter\n",
    "from collections import Counter\n",
    "from tqdm import tqdm\n",
    "\n",
    "counter = Counter()\n",
    "cer_by_lang = {}\n",
    "\n",
    "# Process each subset\n",
    "for subset, alignment_map in subset_to_alignment_map.items():\n",
    "    counter[subset] = 0\n",
    "    \n",
    "    for meta_id, alignment_info in alignment_map.items():\n",
    "        # Here we would update the features with alignment info\n",
    "        # Since we don't have direct access to modify the dataframe in this way,\n",
    "        # we'll just count the matches\n",
    "        counter[subset] += 1\n",
    "    \n",
    "\n",
    "        # Track CER by language if available\n",
    "        lang = subset_meta_id_to_text_lang[subset].get(meta_id, None)\n",
    "        \n",
    "        if lang not in cer_by_lang:\n",
    "            cer_by_lang[lang] = []\n",
    "        cer_by_lang[lang].append(alignment_info[\"cer\"])\n",
    "\n",
    "# Print the counts\n",
    "for subset in counter:\n",
    "    print(f\"{subset}: {counter[subset]}\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "\n",
    "# compute a language specific cutoff for cer\n",
    "cer_cutoff_by_lang = {}\n",
    "\n",
    "for lang, cers in cer_by_lang.items():\n",
    "    print(f\"{lang}: {np.mean(cers)} {np.percentile(cers, 90)}\")\n",
    "    cer_cutoff_by_lang[lang] = np.percentile(cers, 90)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# now apply the cutoff to the alignments\n",
    "valid_ids_cer_by_subset = {}\n",
    "subset_valid_counts_cer = {}\n",
    "subset_total_counts_cer = {}\n",
    "\n",
    "for subset, alignment_map in subset_to_alignment_map.items():\n",
    "    valid_ids_cer_by_subset[subset] = set()\n",
    "    subset_total_counts_cer[subset] = len(alignment_map)\n",
    "    subset_valid_counts_cer[subset] = 0\n",
    "\n",
    "    for meta_id, alignment_info in alignment_map.items():\n",
    "        cer = alignment_info.get('cer', None)\n",
    "        lang = subset_meta_id_to_text_lang[subset].get(meta_id, None)\n",
    "        if subset == \"genius\": # we need to map old id to new id\n",
    "            meta_id = genius_old_to_new_id_map[meta_id]\n",
    "\n",
    "        if cer < cer_cutoff_by_lang[lang]:\n",
    "            valid_ids_cer_by_subset[subset].add(meta_id)\n",
    "            subset_valid_counts_cer[subset] += 1\n",
    "\n",
    "# Calculate total valid examples across all subsets\n",
    "total_valid = sum(len(valid_ids) for valid_ids in valid_ids_cer_by_subset.values())\n",
    "print(f\"Total valid examples after CER filtering: {total_valid}\")\n",
    "\n",
    "# since imslp have no cer, we need to add all the ids to the valid ids\n",
    "valid_ids_cer_by_subset[\"imslp\"] = set(valid_ids_metadata[\"imslp\"])\n",
    "subset_valid_counts_cer[\"imslp\"] = len(valid_ids_metadata[\"imslp\"])\n",
    "\n",
    "# Print statistics for each subset\n",
    "for subset in subset_total_counts_cer:\n",
    "    total = subset_total_counts_cer[subset]\n",
    "    kept = subset_valid_counts_cer[subset]\n",
    "    percent = (kept / total) * 100 if total > 0 else 0\n",
    "    print(f\"{subset}: kept {kept}/{total} ({percent:.2f}%)\")\n",
    "\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "for ds in valid_ids_cer_by_subset:\n",
    "    print(f\"{ds}: {len(valid_ids_cer_by_subset[ds])}\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 15,
   "metadata": {},
   "outputs": [],
   "source": [
    "feature_bounds = {\n",
    "    \"loudness\": [-32, -4],\n",
    "    \"spectral_centroid\": [1750, 5000],\n",
    "    \"spectral_flatness\": [0.02, 0.22],\n",
    "    \"crest_factor\": [1.0, 3],\n",
    "    \"bass\" : [0.1, 0.5],\n",
    "    \"mid\" : [0.4, 1.0],\n",
    "    \"high\" : [0.15, 1.25],\n",
    "    \"stereo_width\" : [0.1, 0.3]\n",
    "}\n",
    "\n",
    "\n",
    "# load audio features\n",
    "discogs_audio_features = pd.read_csv(\"/home/christian/code/christian/metadata/v4/discogs_subset_audio_production_features_v2.csv\")\n",
    "imslp_audio_features = pd.read_csv(\"/home/christian/code/christian/metadata/v4/imslp_audio_production_features_v2.csv\")\n",
    "genius_audio_features = pd.read_csv(\"/home/christian/code/christian/metadata/v4/genius_audio_production_features_v2.csv\")\n",
    "\n",
    "subset_to_feature_map = {\n",
    "    \"discogs\": {m[\"id\"]: m for m in discogs_audio_features.to_dict(orient=\"records\")},\n",
    "    #\"imslp\": {m[\"id\"]: m for m in imslp_audio_features.to_dict(orient=\"records\")},\n",
    "    #\"genius\": {m[\"id\"]: m for m in genius_audio_features.to_dict(orient=\"records\")},\n",
    "}"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# now lets iterate over the audio features and remove the examples that are outside the bounds\n",
    "# in this case we will create a set for each subset to store the ids of the examples that are within the bounds\n",
    "\n",
    "valid_ids_audio_features_by_subset = {}\n",
    "subset_valid_counts_audio_features = {}\n",
    "subset_total_counts_audio_features = {}\n",
    "\n",
    "for subset, feature_map in subset_to_feature_map.items():\n",
    "    valid_ids_audio_features_by_subset[subset] = set()\n",
    "    subset_total_counts_audio_features[subset] = len(feature_map)\n",
    "    subset_valid_counts_audio_features[subset] = 0\n",
    "    \n",
    "    for meta_id, features in tqdm(feature_map.items(), total=len(feature_map)):\n",
    "        is_valid = True\n",
    "        for feature_name, bounds in feature_bounds.items():\n",
    "            if feature_name in features:\n",
    "                feature_value = features[feature_name]\n",
    "                if feature_value < bounds[0] or feature_value > bounds[1]:\n",
    "                    is_valid = False\n",
    "                    break\n",
    "        \n",
    "        if is_valid:\n",
    "            valid_ids_audio_features_by_subset[subset].add(meta_id)\n",
    "            subset_valid_counts_audio_features[subset] += 1\n",
    "\n",
    "# Calculate total valid examples across all subsets\n",
    "total_valid = sum(len(valid_ids) for valid_ids in valid_ids_audio_features_by_subset.values())\n",
    "print(f\"Total valid examples after audio feature filtering: {total_valid}\")\n",
    "\n",
    "# Print retention statistics for each subset\n",
    "for subset in subset_valid_counts_audio_features:\n",
    "    retention_percent = (subset_valid_counts_audio_features[subset] / subset_total_counts_audio_features[subset]) * 100\n",
    "    print(f\"{subset}: {subset_valid_counts_audio_features[subset]} / {subset_total_counts_audio_features[subset]} ({retention_percent:.2f}%)\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "print(\"\\naudio features\")\n",
    "for ds in valid_ids_audio_features_by_subset:\n",
    "    print(f\"{ds}: {len(valid_ids_audio_features_by_subset[ds])}\")\n",
    "\n",
    "print(\"\\ncer\")\n",
    "for ds in valid_ids_cer_by_subset:\n",
    "    print(f\"{ds}: {len(valid_ids_cer_by_subset[ds])}\")\n",
    "\n",
    "print(\"\\nmetadata\")\n",
    "for ds in valid_ids_metadata:\n",
    "    print(f\"{ds}: {len(valid_ids_metadata[ds])}\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(list(valid_ids_audio_features_by_subset[\"genius\"])[:10])\n",
    "print(list(valid_ids_cer_by_subset[\"genius\"])[:10])\n",
    "print(list(valid_ids_metadata[\"genius\"])[:10])\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Let's merge the valid IDs for each subset separately\n",
    "valid_ids_merged_by_subset = {}\n",
    "total_valid_merged = 0\n",
    "\n",
    "#for subset in [\"discogs\", \"imslp\", \"genius\"]:\n",
    "for subset in [\"discogs\"]:\n",
    "    # Get the valid IDs from each filtering step for this subset\n",
    "    audio_feature_valid = valid_ids_audio_features_by_subset.get(subset, set())\n",
    "    cer_valid = valid_ids_cer_by_subset.get(subset, set())\n",
    "    metadata_valid = valid_ids_metadata.get(subset, set())\n",
    "\n",
    "    # If we have all three filtering steps\n",
    "    merged = audio_feature_valid.intersection(cer_valid)#.intersection(metadata_valid)\n",
    "    \n",
    "    valid_ids_merged_by_subset[subset] = merged\n",
    "    total_valid_merged += len(merged)\n",
    "    \n",
    "    # Print statistics for this subset\n",
    "    print(f\"{subset}: {len(merged)} valid examples after merging all filters\")\n",
    "\n",
    "print(f\"Total valid examples after merging across all subsets: {total_valid_merged}\")\n",
    "\n",
    "\n",
    "\n",
    "import json\n",
    "# now we can save the valid ids for each subset\n",
    "#with open(\"/app/suno/data/diffusion_v5/v0/info_tr_v1.json\", \"w\") as f:\n",
    "#    json_data = {subset: list(ids) for subset, ids in valid_ids_merged_by_subset.items()}\n",
    "#    json.dump(json_data, f)\n",
    "\n",
    "json_data = {\n",
    "    \"discogs\": list(valid_ids_merged_by_subset[\"discogs\"])\n",
    "}\n",
    "with open(\"/home/christian/code/christian/metadata/v4/discogs_subset_filtered_ids.json\", \"w\") as f:\n",
    "    json.dump(json_data, f)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "json_data.keys()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "suno_gpt45",
   "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.14"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
