{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "51b22da9",
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.audio import Audio\n",
    "import numpy as np\n",
    "import os\n",
    "import pandas as pd\n",
    "\n",
    "#base_dir = \"/app2/suno/data/christian/outputs/v3-base-data-ctx-t1/\"\n",
    "base_dir = \"/app2/suno/data/christian/outputs/v3-base-data-ctx-rs-t2/\"\n",
    "\n",
    "# get all files in each directory\n",
    "\n",
    "#model_name = \"16n_2b_flow_distill_bs1_N5_c5e5_g1e6_alt_dmd_cfg_2_residual_sft_220k\"\n",
    "#model_name = \"v3_flow_distill_s3177_lm_t1_0_7_cut_history_4x_1E6_beta100_n8_bt2_acc4_2k_last\"\n",
    "#model_name = \"v3_flow_distill_v1_t18_1E6_beta100_n4_bt2_acc2_4k_last\"\n",
    "model_name = \"4n_25hz_2b_flow_5e5_sft_t8_500k\"\n",
    "\n",
    "def process_dir(dirpath, N=10):\n",
    "    \"\"\"\n",
    "    Generalized to load any number of audio files (0, 1, ..., N) for a given dirpath.\n",
    "    Returns a list of dicts, each containing metadata and upsampled_vae for each index found.\n",
    "    \"\"\"\n",
    "    results = []\n",
    "    for idx in range(N):\n",
    "        metadata_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_{model_name}_{idx}__metadata.npz\")\n",
    "        upsampled_vae_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_{model_name}_{idx}_upsampled_vae.npz\")\n",
    "        if not (os.path.exists(metadata_filepath) and os.path.exists(upsampled_vae_filepath)):\n",
    "            # Stop if either file is missing for this index\n",
    "            break\n",
    "        # Load metadata\n",
    "        metadata_npz = np.load(metadata_filepath, allow_pickle=True)\n",
    "        metadata_dict = {key: metadata_npz[key].tolist() for key in metadata_npz.keys()}\n",
    "        # Load upsampled vae\n",
    "        #upsampled_vae = np.load(upsampled_vae_filepath)\n",
    "        results.append({\n",
    "            \"metadata\": metadata_dict,\n",
    "            #\"upsampled_vae\": upsampled_vae,\n",
    "            \"upsampled_vae_filepath\": upsampled_vae_filepath,\n",
    "            \"index\": idx,\n",
    "        })\n",
    "    return results"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "0fba14ad",
   "metadata": {},
   "outputs": [],
   "source": [
    "# load labels\n",
    "import pandas as pd\n",
    "\n",
    "filepath = \"/home/christian/code/christian/metadata/labelmaker/t4/dpo_annotations_export_0_7.csv\"\n",
    "#filepath = f\"{base_dir}ear_scores.csv\"\n",
    "df = pd.read_csv(filepath)\n",
    "\n",
    "# Compute the delta between 0 and 1 by taking the chosen slot minus the non-chosen slot\n",
    "# (i.e., score = score of chosen slot - score of non-chosen slot)\n",
    "def compute_delta(row):\n",
    "    chosen = int(row[\"chosen_slot\"])\n",
    "    not_chosen = 1 - chosen\n",
    "    delta = row[str(chosen)] - row[str(not_chosen)]\n",
    "    return delta\n",
    "\n",
    "#f[\"delta\"] = df.apply(compute_delta, axis=1)\n",
    "#/df.head()\n",
    "\n",
    "\n",
    "# convert this to a dict lookup that goes from clip_id to the label\n",
    "label_lookup = {}\n",
    "for index, row in df.iterrows():\n",
    "    clip_id = row[\"clip_id\"]\n",
    "    label = row[\"chosen_slot\"]\n",
    "    label_lookup[clip_id] = {\n",
    "        \"label\": label,\n",
    "        \"delta\": row.get(\"delta\", None),\n",
    "        \"agreement\": row.get(\"agreement\", None),\n",
    "    }\n",
    "\n",
    "print(len(label_lookup))\n",
    "\n",
    "# also read the loudness values from the csv file\n",
    "loudness_df = pd.read_csv(f\"{base_dir}/loudness_scores.csv\")\n",
    "# then add the loudness delta to the label lookup\n",
    "for index, row in loudness_df.iterrows():\n",
    "    dirname = row['id']\n",
    "    if dirname not in label_lookup:\n",
    "        continue\n",
    "    # get the positive index \n",
    "    pos_idx = str(label_lookup[dirname]['label'])\n",
    "    neg_idx = \"0\" if pos_idx == \"1\" else \"1\"\n",
    "    pos_loudness = row[str(pos_idx)]\n",
    "    neg_loudness = row[str(neg_idx)]\n",
    "    label_lookup[dirname]['pos_loudness'] = pos_loudness\n",
    "    label_lookup[dirname]['neg_loudness'] = neg_loudness\n",
    "    label_lookup[dirname]['loudness_delta'] = pos_loudness - neg_loudness\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "6747e8e7",
   "metadata": {},
   "outputs": [],
   "source": [
    "loundess_deltas = [label_lookup[dirname]['loudness_delta'] for dirname in label_lookup]\n",
    "print(np.mean(loundess_deltas))\n",
    "print(np.percentile(loundess_deltas, 5))\n",
    "print(np.percentile(loundess_deltas, 95))\n",
    "\n",
    "import matplotlib.pyplot as plt\n",
    "plt.hist(loundess_deltas, bins=100)\n",
    "plt.axvline(np.percentile(loundess_deltas, 5), color=\"red\")\n",
    "plt.axvline(np.percentile(loundess_deltas, 95), color=\"red\")\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "479442de",
   "metadata": {},
   "outputs": [],
   "source": [
    "# create a label lookup by setting the label to 1 always\n",
    "# get all dirs in the base_dir\n",
    "dirs = os.listdir(base_dir)\n",
    "# create a label lookup by setting the label to 1 always\n",
    "label_lookup = {dir: {\"label\": 1, \"delta\": None, \"agreement\": None} for dir in dirs}\n",
    "print(len(label_lookup))\n",
    "\n",
    "# also read the loudness values from the csv file\n",
    "if os.path.exists(f\"{base_dir}loudness_scores.csv\"):\n",
    "    loudness_df = pd.read_csv(f\"{base_dir}loudness_scores.csv\")\n",
    "    print(loudness_df.head())\n",
    "    # then add the loudness delta to the label lookup\n",
    "    for index, row in loudness_df.iterrows():\n",
    "        dirname = row['id']\n",
    "        label_lookup[dirname]['loudness_delta'] = row['loudness_delta']\n",
    "        label_lookup[dirname]['pos_loudness'] = row['1']\n",
    "        label_lookup[dirname]['neg_loudness'] = row['0']\n",
    "\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a08a70ba",
   "metadata": {},
   "outputs": [],
   "source": [
    "# select a random directory from the label lookup\n",
    "import random\n",
    "random_dir = random.choice(list(label_lookup.keys()))\n",
    "print(random_dir)\n",
    "\n",
    "\n",
    "label_lookup[random_dir]\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "e3739be5",
   "metadata": {},
   "outputs": [],
   "source": [
    "# check a directory \n",
    "\n",
    "dirpath = \"aa5e0ee2-1004-482a-8062-8345515a7f30\"\n",
    "semantic_codes_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_semantic.npz\")\n",
    "\n",
    "pos_vae_latents_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_{model_name}_{0}_upsampled_vae.npz\")\n",
    "neg_vae_latents_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_{model_name}_{1}_upsampled_vae.npz\")\n",
    "history_vae_latents_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_history_vae.npz\")\n",
    "\n",
    "semantic_codes = np.load(semantic_codes_filepath)[\"semantic_codes\"]\n",
    "pos_vae_latents = np.load(pos_vae_latents_filepath)[\"vae_latents\"]\n",
    "neg_vae_latents = np.load(neg_vae_latents_filepath)[\"vae_latents\"]\n",
    "history_vae_latents = np.load(history_vae_latents_filepath)[\"vae_latents\"]\n",
    "\n",
    "\n",
    "# now get the last 750 tokens of the history vae latents\n",
    "history_vae_latents = history_vae_latents[-750:]\n",
    "\n",
    "# concat with the pos and neg vae latents\n",
    "pos_vae_full = np.concatenate([history_vae_latents, pos_vae_latents[:750]], axis=0)\n",
    "neg_vae_full = np.concatenate([history_vae_latents, neg_vae_latents[:750]], axis=0)\n",
    "\n",
    "# decode the audio\n",
    "pos_audio = codec_decode(pos_vae_full).play()\n",
    "neg_audio = codec_decode(neg_vae_full).play()\n",
    "\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2dfb4a79",
   "metadata": {},
   "outputs": [],
   "source": [
    "# create metas # in this case using ear score\n",
    "from tqdm import tqdm\n",
    "from joblib import Parallel, delayed\n",
    "\n",
    "metas = []\n",
    "with_history = 0\n",
    "use_hoot = False\n",
    "\n",
    "# get all directories in base_dir\n",
    "dirs = os.listdir(base_dir)\n",
    "print(len(dirs))\n",
    "\n",
    "def process_single_dir(dirpath):\n",
    "    try:\n",
    "        results = process_dir(dirpath)\n",
    "    except Exception as e:\n",
    "        print(f\"Error processing {dirpath}: {e}\")\n",
    "        return None, 0\n",
    "\n",
    "    if len(results) < 2:\n",
    "        return None, 0\n",
    "\n",
    "    tags = results[0][\"metadata\"][\"tags\"]\n",
    "    text = results[0][\"metadata\"][\"text\"]\n",
    "    semantic_codes_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_semantic.npz\")\n",
    "\n",
    "    if use_hoot:\n",
    "        # get the result with the lower and higher cer\n",
    "        # find the index of the result with the lowest cer\n",
    "        # Handle the case where hoot_cer might be None\n",
    "        hoot_cers = [result[\"metadata\"].get(\"hoot_cer\", None) for result in results]\n",
    "        # Replace None with np.inf for min, -np.inf for max so they are always last\n",
    "        hoot_cers_for_min = [c if c is not None else np.inf for c in hoot_cers]\n",
    "        hoot_cers_for_max = [c if c is not None else -np.inf for c in hoot_cers]\n",
    "        min_idx = np.argmin(hoot_cers_for_min)\n",
    "        max_idx = np.argmax(hoot_cers_for_max)\n",
    "        pos_metadata = results[min_idx][\"metadata\"]\n",
    "        neg_metadata = results[max_idx][\"metadata\"]\n",
    "        pos_vae_latents_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_{model_name}_{min_idx}_upsampled_vae.npz\")\n",
    "        neg_vae_latents_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_{model_name}_{max_idx}_upsampled_vae.npz\")\n",
    "    else:\n",
    "        label_dict = label_lookup.get(dirpath, None)\n",
    "        if label_dict is None:\n",
    "            return None, 0\n",
    "        if label_dict[\"label\"] == 0:\n",
    "            pos_metadata = results[0][\"metadata\"]\n",
    "            neg_metadata = results[1][\"metadata\"]\n",
    "\n",
    "            pos_vae_latents_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_{model_name}_{0}_upsampled_vae.npz\")\n",
    "            neg_vae_latents_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_{model_name}_{1}_upsampled_vae.npz\")\n",
    "        else:\n",
    "            pos_metadata = results[1][\"metadata\"]\n",
    "            neg_metadata = results[0][\"metadata\"]\n",
    "            pos_vae_latents_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_{model_name}_{1}_upsampled_vae.npz\")\n",
    "            neg_vae_latents_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_{model_name}_{0}_upsampled_vae.npz\")\n",
    "\n",
    "\n",
    "        pos_loudness = label_dict.get(\"pos_loudness\", None)\n",
    "        neg_loudness = label_dict.get(\"neg_loudness\", None)\n",
    "        pos_metadata[\"loudness\"] = pos_loudness\n",
    "        neg_metadata[\"loudness\"] = neg_loudness\n",
    "\n",
    "    # check if we have history vae latents\n",
    "    history_vae_latents_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_history_vae.npz\")\n",
    "    if os.path.exists(history_vae_latents_filepath):\n",
    "        # history_vae_latents = np.load(history_vae_latents_filepath)[\"vae_latents\"]\n",
    "        with_history_flag = 1\n",
    "    else:\n",
    "        # history_vae_latents = None\n",
    "        history_vae_latents_filepath = None\n",
    "        with_history_flag = 0\n",
    "\n",
    "    history_vae_latents_filepath = None # for now we don't have history vae latents\n",
    "\n",
    "    meta = {\n",
    "        \"id\": dirpath,\n",
    "        \"id_x\": dirpath,\n",
    "        \"tags\": str(tags),\n",
    "        \"text\": str(text),\n",
    "        \"pos_metadata\": pos_metadata,\n",
    "        \"neg_metadata\": neg_metadata,\n",
    "        \"pos_vae_latents_filepath\": pos_vae_latents_filepath,\n",
    "        \"neg_vae_latents_filepath\": neg_vae_latents_filepath,\n",
    "        \"history_vae_latents_filepath\": history_vae_latents_filepath,\n",
    "        \"semantic_codes_filepath\": semantic_codes_filepath,\n",
    "        \"delta\": label_dict[\"delta\"],\n",
    "        \"agreement\": label_dict.get(\"agreement\", None),\n",
    "    }\n",
    "    return meta, with_history_flag\n",
    "\n",
    "# we should still use joblib but we should use backend loky i think\n",
    "results = Parallel(n_jobs=-1, backend=\"loky\")(\n",
    "    delayed(process_single_dir)(dirpath) for dirpath in tqdm(dirs)\n",
    ")\n",
    "\n",
    "# Unpack results\n",
    "metas = []\n",
    "with_history = 0\n",
    "for meta, with_history_flag in results:\n",
    "    if meta is not None:\n",
    "        metas.append(meta)\n",
    "        with_history += with_history_flag\n",
    "\n",
    "print(f\"Processed {len(metas)} valid, with history {with_history}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "bac3bc54",
   "metadata": {},
   "outputs": [],
   "source": [
    "import matplotlib.pyplot as plt\n",
    "# make a plot of the pos vs negative shimmer score\n",
    "shimmer_score_deltas = []\n",
    "ear_score_deltas = []\n",
    "ear_v3_score_deltas = []\n",
    "step_deltas = []\n",
    "hoot_cer_deltas = []\n",
    "stereo_score_deltas = []\n",
    "stereo_deltas = []\n",
    "loudness_deltas = []\n",
    "\n",
    "for meta in filtered_metas:\n",
    "    shimmer_score_deltas.append(meta[\"pos_metadata\"][\"shimmer_score\"] - meta[\"neg_metadata\"][\"shimmer_score\"])\n",
    "    ear_score_deltas.append(meta[\"pos_metadata\"][\"ear_score\"] - meta[\"neg_metadata\"][\"ear_score\"])\n",
    "    #ear_v3_score_deltas.append(meta[\"pos_metadata\"][\"ear_v3_score\"] - meta[\"neg_metadata\"][\"ear_v3_score\"])\n",
    "    step_deltas.append(meta[\"pos_metadata\"][\"diffusion\"][\"steps\"] - meta[\"neg_metadata\"][\"diffusion\"][\"steps\"])\n",
    "    \n",
    "    pos_stereo_width = meta[\"pos_metadata\"][\"stereo_width\"]\n",
    "    neg_stereo_width = meta[\"neg_metadata\"][\"stereo_width\"]\n",
    "    stereo_deltas.append(pos_stereo_width - neg_stereo_width)\n",
    "    ## compute stereo score as distance from 0.15 \n",
    "    #pos_stereo_score = 1.0 - np.abs(0.15 - pos_stereo_width)\n",
    "    #neg_stereo_score = 1.0 - np.abs(0.15 - neg_stereo_width)\n",
    "    #meta[\"pos_metadata\"][\"stereo_score\"] = pos_stereo_score\n",
    "    #meta[\"neg_metadata\"][\"stereo_score\"] = neg_stereo_score\n",
    "    #stereo_score_deltas.append(pos_stereo_score - neg_stereo_score)\n",
    "\n",
    "    pos_loudness = meta[\"pos_metadata\"][\"loudness\"]\n",
    "    neg_loudness = meta[\"neg_metadata\"][\"loudness\"]\n",
    "\n",
    "    # Check for nan or inf and handle accordingly\n",
    "    if (\n",
    "        pos_loudness is None or neg_loudness is None\n",
    "        or np.isnan(pos_loudness) or np.isnan(neg_loudness)\n",
    "        or np.isinf(pos_loudness) or np.isinf(neg_loudness)\n",
    "    ):\n",
    "        pass\n",
    "    else:\n",
    "        loudness_deltas.append(pos_loudness - neg_loudness)\n",
    "\n",
    "    pos_hoot_cer = meta[\"pos_metadata\"][\"hoot_cer\"]\n",
    "    neg_hoot_cer = meta[\"neg_metadata\"][\"hoot_cer\"]\n",
    "    if pos_hoot_cer is not None and neg_hoot_cer is not None:\n",
    "        hoot_cer_deltas.append(pos_hoot_cer - neg_hoot_cer)\n",
    "\n",
    "# Combine all plots into one figure with subplots and add text for mean, std, min, max\n",
    "print(len(loudness_deltas))\n",
    "import numpy as np\n",
    "\n",
    "# Prepare data and labels\n",
    "plot_data = [\n",
    "    {\n",
    "        \"data\": shimmer_score_deltas,\n",
    "        \"title\": f\"Shimmer score deltas for {model_name}\",\n",
    "        \"xlabel\": \"Shimmer Score Delta\"\n",
    "    },\n",
    "    {\n",
    "        \"data\": ear_score_deltas,\n",
    "        \"title\": f\"Ear score deltas for {model_name}\",\n",
    "        \"xlabel\": \"Ear Score Delta\"\n",
    "    },\n",
    "    {\n",
    "        \"data\": stereo_deltas,\n",
    "        \"title\": f\"Stereo deltas for {model_name}\",\n",
    "        \"xlabel\": \"Stereo Delta\"\n",
    "    },\n",
    "    {\n",
    "        \"data\": step_deltas,\n",
    "        \"title\": f\"Step deltas for {model_name}\",\n",
    "        \"xlabel\": \"Step Delta\"\n",
    "    },\n",
    "    {\n",
    "        \"data\": hoot_cer_deltas,\n",
    "        \"title\": f\"Hoot cer deltas for {model_name}\",\n",
    "        \"xlabel\": \"Hoot CER Delta\"\n",
    "    },\n",
    "    {\n",
    "        \"data\": loudness_deltas,\n",
    "        \"title\": f\"Loudness deltas for {model_name}\",\n",
    "        \"xlabel\": \"Loudness Delta\"\n",
    "    },\n",
    "\n",
    "]\n",
    "\n",
    "# Optionally add ear_v3_score_deltas if present\n",
    "if len(ear_v3_score_deltas) > 0:\n",
    "    plot_data.insert(3, {  # Insert after stereo_deltas\n",
    "        \"data\": ear_v3_score_deltas,\n",
    "        \"title\": f\"Ear v3 score deltas for {model_name}\",\n",
    "        \"xlabel\": \"Ear v3 Score Delta\"\n",
    "    })\n",
    "\n",
    "n_plots = len(plot_data)\n",
    "n_cols = 2\n",
    "n_rows = (n_plots + 1) // n_cols\n",
    "\n",
    "fig, axs = plt.subplots(n_rows, n_cols, figsize=(8, 2.5 * n_rows))\n",
    "axs = axs.flatten()\n",
    "\n",
    "for i, plot_info in enumerate(plot_data):\n",
    "    if len(plot_info[\"data\"]) == 0:\n",
    "        continue\n",
    "    data = np.array(plot_info[\"data\"])\n",
    "    ax = axs[i]\n",
    "    ax.hist(data, bins=50, color='skyblue', edgecolor='black')\n",
    "    p5 = np.percentile(data, 5)\n",
    "    p95 = np.percentile(data, 95)\n",
    "    mean = np.mean(data)\n",
    "    std = np.std(data)\n",
    "    dmin = np.min(data)\n",
    "    dmax = np.max(data)\n",
    "    ax.axvline(p5, color=\"red\", linestyle=\"--\", label=\"5th/95th percentile\")\n",
    "    ax.axvline(p95, color=\"red\", linestyle=\"--\")\n",
    "    ax.set_title(plot_info[\"title\"], fontsize=10)\n",
    "    ax.set_xlabel(plot_info[\"xlabel\"])\n",
    "    ax.set_ylabel(\"Count\")\n",
    "    # Add text box with stats\n",
    "    stats_text = (\n",
    "        f\"mean={mean:.3f}\\n\"\n",
    "        f\"std={std:.3f}\\n\"\n",
    "        f\"min={dmin:.3f}\\n\"\n",
    "        f\"max={dmax:.3f}\"\n",
    "    )\n",
    "    # Place the text in the upper right of the plot\n",
    "    ax.text(\n",
    "        0.98, 0.98, stats_text,\n",
    "        transform=ax.transAxes,\n",
    "        fontsize=8,\n",
    "        verticalalignment='top',\n",
    "        horizontalalignment='right',\n",
    "        bbox=dict(boxstyle=\"round,pad=0.3\", facecolor=\"white\", alpha=0.7)\n",
    "    )\n",
    "\n",
    "    # Optionally print the percentiles for reference\n",
    "    print(f\"{plot_info['title']}: 5th={p5:.3f}, 95th={p95:.3f}\")\n",
    "\n",
    "# Hide any unused subplots\n",
    "for j in range(i+1, len(axs)):\n",
    "    fig.delaxes(axs[j])\n",
    "\n",
    "plt.tight_layout()\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "93276b84",
   "metadata": {},
   "outputs": [],
   "source": [
    "# for shimmer, score lower is better, so we want the positive to be lower, hence delta is negative is better\n",
    "# for shimmer we want to cut off delta that is very positive\n",
    "# for ear, score higher is better, so we want the positive to be higher, hence delta is positive is better\n",
    "# for ear we want to cut off delta that is very negative\n",
    "\n",
    "filtered_metas = []\n",
    "\n",
    "for meta in metas:\n",
    "    # for relative filter\n",
    "    shimmer_score_delta = meta[\"pos_metadata\"][\"shimmer_score\"] - meta[\"neg_metadata\"][\"shimmer_score\"]\n",
    "    ear_score_delta = meta[\"pos_metadata\"][\"ear_score\"] - meta[\"neg_metadata\"][\"ear_score\"]\n",
    "    #ear_v3_score_delta = meta[\"pos_metadata\"][\"ear_v3_score\"] - meta[\"neg_metadata\"][\"ear_v3_score\"]\n",
    "    stereo_delta = meta[\"pos_metadata\"][\"stereo_width\"] - meta[\"neg_metadata\"][\"stereo_width\"]\n",
    "    #loudness_delta = meta[\"pos_metadata\"][\"loudness\"] - meta[\"neg_metadata\"][\"loudness\"]\n",
    "\n",
    "    # for absolute filter\n",
    "    pos_shimmer_score = meta[\"pos_metadata\"][\"shimmer_score\"]\n",
    "    pos_ear_score = meta[\"pos_metadata\"][\"ear_score\"]\n",
    "    pos_hoot_cer = meta[\"pos_metadata\"][\"hoot_cer\"]\n",
    "    neg_hoot_cer = meta[\"neg_metadata\"][\"hoot_cer\"]\n",
    "    if pos_hoot_cer is None:\n",
    "        pos_hoot_cer = 0.0\n",
    "    if neg_hoot_cer is None:\n",
    "        neg_hoot_cer = 0.0\n",
    "    hoot_cer_delta = pos_hoot_cer - neg_hoot_cer\n",
    "\n",
    "    if shimmer_score_delta < -1.0 and ear_score_delta > 1.0 \\\n",
    "        and pos_shimmer_score < 12 and pos_ear_score > 14 \\\n",
    "        and pos_hoot_cer < 0.8 and hoot_cer_delta <= -0.025:\n",
    "        filtered_metas.append(meta)\n",
    "\n",
    "    #if shimmer_score_delta < 2.0 and pos_shimmer_score < 8 \\\n",
    "    #    and loudness_delta > -0.5 and loudness_delta < 0.5 \\\n",
    "    #    and pos_ear_score > 10:\n",
    "    #    filtered_metas.append(meta)\n",
    "\n",
    "print(len(metas))\n",
    "print(len(filtered_metas))\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "e21ecb75",
   "metadata": {},
   "outputs": [],
   "source": [
    "# count the number of chunks with history vae latents\n",
    "history_count = 0\n",
    "for meta in filtered_metas:\n",
    "    if meta[\"history_vae_latents_filepath\"] is not None:\n",
    "        history_count += 1\n",
    "\n",
    "print(len(filtered_metas))\n",
    "print(history_count)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "13f9e183",
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "new_metas = []\n",
    "for meta in filtered_metas:\n",
    "    new_meta = meta.copy()\n",
    "    new_metas.append(new_meta)\n",
    "    if meta[\"history_vae_latents_filepath\"] is not None:\n",
    "        for i in range(8):\n",
    "            new_metas.append(new_meta)\n",
    "\n",
    "print(len(new_metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "8a176e3b",
   "metadata": {},
   "outputs": [],
   "source": [
    "filtered_metas = []\n",
    "\n",
    "for meta in metas:\n",
    "    # ensure the hoot_cer delta is greater than 0.02\n",
    "    pos_hoot_cer = meta[\"pos_metadata\"][\"hoot_cer\"]\n",
    "    neg_hoot_cer = meta[\"neg_metadata\"][\"hoot_cer\"]\n",
    "    if pos_hoot_cer is not None and neg_hoot_cer is not None:\n",
    "        hoot_cer_delta = pos_hoot_cer - neg_hoot_cer\n",
    "        if abs(hoot_cer_delta) > 0.02:\n",
    "            filtered_metas.append(meta)\n",
    "\n",
    "print(len(filtered_metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "7fe113c2",
   "metadata": {},
   "outputs": [],
   "source": [
    "filtered_metas = new_metas"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "0d48d16c",
   "metadata": {},
   "outputs": [],
   "source": [
    "metas_tr = filtered_metas[:int(len(filtered_metas) * 0.98)]\n",
    "metas_val = filtered_metas[int(len(filtered_metas) * 0.98):]\n",
    "\n",
    "print(len(metas_tr), len(metas_val))\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "e485be1a",
   "metadata": {},
   "outputs": [],
   "source": [
    "# to create a memmap for sft we will select the higheset scoring upsample_id for each base_s3_id\n",
    "# we also need to grab the correct vae latents and semantic codes and text prompt\n",
    "from tqdm import tqdm\n",
    "from suno_utils.utils.text import write_jsonl\n",
    "import gc\n",
    "import sys\n",
    "import shutil\n",
    "\n",
    "# aws s3 sync s3://suno-data/christian/outputs/corrupt/genius_t6_sampled_10k/ /app/suno/data/diff_syn_dpo/genius_t6_sampled_10k_corrupted/npz\n",
    "\n",
    "# t0 hoot cer 0.02\n",
    "\n",
    "memmap_name = \"d35\"\n",
    "memmap_dir = \"/app2/suno/data/christian/outputs/memmaps/\"\n",
    "SEMANTIC_RATE_HZ = 25\n",
    "CHUNK_SIZE_S = 30\n",
    "CHUNK_SIZE = int(CHUNK_SIZE_S * SEMANTIC_RATE_HZ)\n",
    "OUT_DATA_DIR = f\"{memmap_dir}/{memmap_name}\"\n",
    "\n",
    "#BASE_S3_DIR = \"s3://suno-data/christian/outputs/v2-infill-data-v1\"\n",
    "BASE_LOCAL_DIR = f\"{base_dir}/\"\n",
    "\n",
    "if not os.path.exists(OUT_DATA_DIR):\n",
    "    os.makedirs(OUT_DATA_DIR, exist_ok=True)\n",
    "else:\n",
    "    #shutil.rmtree(OUT_DATA_DIR)\n",
    "    os.makedirs(OUT_DATA_DIR, exist_ok=True)\n",
    "\n",
    "for dset_type in [\"val\", \"tr\"]:\n",
    "\n",
    "    if dset_type == \"tr\":\n",
    "        metas = metas_tr\n",
    "    else:\n",
    "        metas = metas_val\n",
    "\n",
    "    new_metas = []\n",
    "\n",
    "    out_mm_semantic_filepath = os.path.join(OUT_DATA_DIR, f\"data_semantic_{dset_type}.bin\")\n",
    "    out_mm_vae_filepath = os.path.join(OUT_DATA_DIR, f\"data_vae_{dset_type}.bin\")\n",
    "    out_metas_filepath = os.path.join(OUT_DATA_DIR, f\"metas_{dset_type}.jsonl\")\n",
    "\n",
    "    n_offs_v = 0\n",
    "    n_offs_s = 0\n",
    "    to_write_len_v = 0\n",
    "    to_write_len_s = 0\n",
    "    total_hours = 0  # Counter for total hours of audio\n",
    "\n",
    "    out_mm_semantic = np.memmap(\n",
    "        out_mm_semantic_filepath, dtype=np.uint16, mode=\"w+\", shape=(1,)\n",
    "    )\n",
    "\n",
    "    out_mm_vae = np.memmap(\n",
    "        out_mm_vae_filepath, dtype=np.float16, mode=\"w+\", shape=(1,)\n",
    "    )\n",
    "\n",
    "    # clear the metas file\n",
    "    with open(out_metas_filepath, \"w\") as f:\n",
    "        f.write(\"\")\n",
    "\n",
    "    # Create a tqdm progress bar with hours counter\n",
    "    pbar = tqdm(metas)\n",
    "    pbar.set_description(\"Hours: 0.00\")\n",
    "\n",
    "    for idx, (meta) in enumerate(pbar):\n",
    "\n",
    "        # load semantic codes from disk\n",
    "        semantic_codes_filepath = meta[\"semantic_codes_filepath\"]\n",
    "        pos_vae_latents_filepath = meta[\"pos_vae_latents_filepath\"]\n",
    "        neg_vae_latents_filepath = meta[\"neg_vae_latents_filepath\"]\n",
    "        history_vae_latents_filepath = meta[\"history_vae_latents_filepath\"]\n",
    "\n",
    "        try:\n",
    "            semantic_data = np.load(semantic_codes_filepath)[\"semantic_codes\"]\n",
    "        except Exception as e:\n",
    "            raise ValueError(f\"Error loading {semantic_codes_filepath}: {e}\")\n",
    "\n",
    "        try:\n",
    "            vae_data_pos = np.load(pos_vae_latents_filepath)[\"vae_latents\"].astype(np.float16)\n",
    "        except Exception as e:\n",
    "            raise ValueError(f\"Error loading {pos_vae_latents_filepath}: {e}\")\n",
    "\n",
    "        try:\n",
    "            vae_data_neg = np.load(neg_vae_latents_filepath)[\"vae_latents\"].astype(np.float16)\n",
    "        except Exception as e:\n",
    "            raise ValueError(f\"Error loading {neg_vae_latents_filepath}: {e}\")\n",
    "\n",
    "        if meta[\"history_vae_latents_filepath\"] is not None:\n",
    "            try:\n",
    "                history_vae_latents = np.load(history_vae_latents_filepath)[\"vae_latents\"].astype(np.float16)\n",
    "            except Exception as e:\n",
    "                raise ValueError(f\"Error loading {history_vae_latents_filepath}: {e}\")\n",
    "\n",
    "        # load the metadata for positive and negative\n",
    "        #print(semantic_data.shape, vae_data_pos.shape, vae_data_neg.shape)        \n",
    "        #break\n",
    "\n",
    "        #num_chunks = 2 if meta[\"history_vae_latents_filepath\"] is not None else 1\n",
    "        num_chunks = semantic_data.shape[0] // CHUNK_SIZE\n",
    "        #assert num_chunks == 1 # for this data\n",
    "\n",
    "        to_write_len_s = semantic_data[:750].size * num_chunks * 2\n",
    "        to_write_len_v = vae_data_pos[:750, :].size * num_chunks * 2\n",
    "        \n",
    "        if to_write_len_s == 0:\n",
    "            continue\n",
    "\n",
    "        if to_write_len_v == 0:\n",
    "            continue\n",
    "        \n",
    "        out_mm_semantic = np.memmap(\n",
    "            out_mm_semantic_filepath,\n",
    "            dtype=np.uint16,\n",
    "            mode=\"r+\",\n",
    "            shape=(n_offs_s + to_write_len_s,),\n",
    "        )\n",
    "\n",
    "        out_mm_vae = np.memmap(\n",
    "            out_mm_vae_filepath,\n",
    "            dtype=np.float16,\n",
    "            mode=\"r+\",\n",
    "            shape=(n_offs_v + to_write_len_v,),\n",
    "        )\n",
    "        # Add to total hours counter\n",
    "        audio_duration_hours = (num_chunks * CHUNK_SIZE_S) / 3600\n",
    "        total_hours += audio_duration_hours\n",
    "        \n",
    "        # Update progress bar description with current total hours\n",
    "        pbar.set_description(f\"Hours: {total_hours:.2f}\")\n",
    "\n",
    "        # we will write the positive and negative as interleaved chunks\n",
    "        # if we have history vae latents, we willy write them first \n",
    "        i = 0 \n",
    "        if meta[\"history_vae_latents_filepath\"] is not None:\n",
    "            for n in range(2):\n",
    "                vae_chunk = history_vae_latents[-750:]# last 750 tokens\n",
    "                assert vae_chunk.shape[0] == CHUNK_SIZE\n",
    "                out_mm_vae[n_offs_v : n_offs_v + vae_chunk.size] = vae_chunk.reshape(\n",
    "                    -1,\n",
    "                )\n",
    "                n_offs_v += vae_chunk.size\n",
    "\n",
    "                # also write semantic data, but this is going to be dummy data (do all pad, 4000)\n",
    "                semantic_chunk = np.zeros(CHUNK_SIZE, dtype=np.uint16)\n",
    "                out_mm_semantic[n_offs_s : n_offs_s + semantic_chunk.size] = semantic_chunk.reshape(\n",
    "                    -1,\n",
    "                )\n",
    "                n_offs_s += semantic_chunk.size\n",
    "\n",
    "                # create a new meta\n",
    "                new_meta = {\n",
    "                    \"id\": meta[\"id\"],\n",
    "                    'id_x': meta[\"id\"],\n",
    "                    \"start_s\": float(i*CHUNK_SIZE_S),\n",
    "                    \"end_s\": float((i+1)*CHUNK_SIZE_S),\n",
    "                    \"original_duration_s\": semantic_data.shape[0] / SEMANTIC_RATE_HZ,\n",
    "                    \"n_vae_tokens\": CHUNK_SIZE,\n",
    "                    \"n_semantic_tokens\": CHUNK_SIZE,\n",
    "                    \"text\" : meta[\"text\"],\n",
    "                    \"tags\" : meta[\"tags\"],\n",
    "                    \"offset_rows\" : 1,\n",
    "                    \"trainable\" : False,\n",
    "                    #\"agreement\": meta[\"agreement\"],\n",
    "                }\n",
    "                new_metas.append(new_meta)\n",
    "\n",
    "            i += 1\n",
    "\n",
    "        for chunk_idx in np.arange(start=i, stop=num_chunks, step=1):\n",
    "            start_s = chunk_idx*CHUNK_SIZE_S\n",
    "            end_s = (chunk_idx+1)*CHUNK_SIZE_S  \n",
    "            for pair_name, vae_data in [(\"neg\", vae_data_neg), (\"pos\", vae_data_pos)]:\n",
    "                #print(chunk_idx, start_s, end_s)\n",
    "                # create a new meta\n",
    "                new_meta = {\n",
    "                    \"id\": meta[\"id\"],\n",
    "                    \"id_x\": meta[\"id\"],\n",
    "                    \"start_s\": float(start_s),\n",
    "                    \"end_s\": float(end_s),\n",
    "                    \"original_duration_s\": semantic_data.shape[0] / SEMANTIC_RATE_HZ,\n",
    "                    \"n_vae_tokens\": CHUNK_SIZE,\n",
    "                    \"n_semantic_tokens\": CHUNK_SIZE,\n",
    "                    \"text\" : meta[\"text\"],\n",
    "                    \"tags\" : meta[\"tags\"],\n",
    "                    \"trainable\" : True,\n",
    "                    \"offset_rows\" : 2 if chunk_idx > 0 else 1,\n",
    "                    \"pair_name\" : pair_name,\n",
    "                    #\"agreement\": meta[\"agreement\"],\n",
    "                }\n",
    "                new_metas.append(new_meta)\n",
    "\n",
    "                #chunk_start = 0 * CHUNK_SIZE\n",
    "                chunk_start = chunk_idx * CHUNK_SIZE\n",
    "                chunk_end = chunk_start + CHUNK_SIZE\n",
    "                semantic_chunk = semantic_data[chunk_start:chunk_end]\n",
    "\n",
    "                out_mm_semantic[n_offs_s : n_offs_s + semantic_chunk.size] = semantic_chunk.reshape(\n",
    "                    -1,\n",
    "                )\n",
    "                n_offs_s += semantic_chunk.size\n",
    "\n",
    "                vae_chunk = vae_data[chunk_start:chunk_end, :]\n",
    "                out_mm_vae[n_offs_v : n_offs_v + vae_chunk.size] = vae_chunk.reshape(\n",
    "                    -1,\n",
    "                )\n",
    "                n_offs_v += vae_chunk.size\n",
    "\n",
    "    print(f\"Total hours of audio added: {total_hours:.4f} for {dset_type} set\")\n",
    "    print(f\"len(new_metas): {len(new_metas)}\")\n",
    "    write_jsonl(\n",
    "        new_metas,\n",
    "        os.path.join(out_metas_filepath),\n",
    "        do_append=True\n",
    "    )\n",
    "\n",
    "    out_mm_semantic.flush()\n",
    "    out_mm_vae.flush()\n",
    "    del out_mm_semantic, out_mm_vae, f\n",
    "    gc.collect()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3bbf813d",
   "metadata": {},
   "outputs": [],
   "source": [
    "print(base_dir)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "81a04d7e",
   "metadata": {},
   "source": [
    "# Check memmaps"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "04e415d5",
   "metadata": {},
   "outputs": [],
   "source": [
    "# test some of the local data\n",
    "CODEC_FILEPATH = \"s3://suno-data/minz/models/dac_vae_tuned_25hz.pth\"\n",
    "\n",
    "from suno_utils.tasks.dac_vae_fixed_25hz import (\n",
    "    preload_models as preload_codec_models,\n",
    "    decode as codec_decode,\n",
    "    encode as codec_encode,\n",
    "    decode_stream_to_full_audio,\n",
    ")\n",
    "\n",
    "_ = preload_codec_models(CODEC_FILEPATH)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "b6026917",
   "metadata": {},
   "outputs": [],
   "source": [
    "# test loading the memmaps \n",
    "from suno_utils.utils.text import read_jsonl\n",
    "import numpy as np\n",
    "\n",
    "VAE_DIM = 128\n",
    "VAE_N_MEMMAP_TOKENS = 750\n",
    "\n",
    "SEMANTIC_N_CODEBOOKS = 1\n",
    "SEMANTIC_N_MEMMAP_TOKENS = 750\n",
    "\n",
    "#dataset_dir = \"/app2/suno/data/christian/outputs/memmaps/d34\"\n",
    "dataset_dir = \"/app2/suno/data/dpo/diff3_carp_t1_v1/\"\n",
    "\n",
    "val_metas = read_jsonl(f\"{dataset_dir}/metas_val.jsonl\", progress=True)\n",
    "vae_memmap_filepath = f\"{dataset_dir}/data_vae_val.bin\"\n",
    "semantic_memmap_filepath = f\"{dataset_dir}/data_semantic_val.bin\"\n",
    "\n",
    "# load memmaps\n",
    "vae_memmap = np.memmap(vae_memmap_filepath, dtype=np.float16, mode=\"r\")\n",
    "semantic_memmap = np.memmap(semantic_memmap_filepath, dtype=np.uint16, mode=\"r\")\n",
    "\n",
    "# reshape memmaps\n",
    "vae_data = vae_memmap.reshape(-1, VAE_N_MEMMAP_TOKENS, VAE_DIM)\n",
    "semantic_data = semantic_memmap.reshape(-1, SEMANTIC_N_MEMMAP_TOKENS, SEMANTIC_N_CODEBOOKS)\n",
    "\n",
    "print(vae_data.shape, semantic_data.shape, len(val_metas))\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "366c727b",
   "metadata": {},
   "outputs": [],
   "source": [
    "# get the pos and neg vae latents \n",
    "\n",
    "neg_ctx_vae = None\n",
    "pos_ctx_vae = None\n",
    "\n",
    "neg_idx = 10\n",
    "pos_idx = neg_idx + 1\n",
    "\n",
    "neg_meta = val_metas[neg_idx]\n",
    "pos_meta = val_metas[pos_idx]\n",
    "\n",
    "if not neg_meta.get(\"trainable\", True):\n",
    "    print(\"neg is not trainable\")\n",
    "    neg_idx += 2\n",
    "    pos_idx = neg_idx + 1\n",
    "\n",
    "neg_offset_rows = neg_meta.get(\"offset_rows\", 1)\n",
    "if neg_offset_rows > 1:\n",
    "    neg_ctx_meta = val_metas[neg_idx - neg_offset_rows]\n",
    "    print(\"neg ctx\", neg_ctx_meta)\n",
    "    neg_ctx_vae = vae_data[neg_idx - neg_offset_rows]\n",
    "\n",
    "pos_offset_rows = pos_meta.get(\"offset_rows\", 1)\n",
    "if pos_offset_rows > 1:\n",
    "    pos_ctx_meta = val_metas[pos_idx - pos_offset_rows]\n",
    "    print(\"pos ctx\", pos_ctx_meta)\n",
    "    pos_ctx_vae = vae_data[pos_idx - pos_offset_rows]\n",
    "\n",
    "vae_neg = vae_data[neg_idx]\n",
    "vae_pos = vae_data[pos_idx]\n",
    "\n",
    "# decode \n",
    "print(\"neg\")\n",
    "print(neg_meta.get(\"start_s\", 0), neg_meta.get(\"end_s\", 0), neg_meta.get(\"trainable\", True))\n",
    "if neg_ctx_vae is not None:\n",
    "    full_vae = np.concatenate([neg_ctx_vae, vae_neg], axis=0)\n",
    "    neg_audio_ctx = codec_decode(full_vae).play()\n",
    "else:\n",
    "    neg_audio = codec_decode(vae_neg).play()\n",
    "\n",
    "\n",
    "print(\"pos\")\n",
    "print(pos_meta.get(\"start_s\", 0), pos_meta.get(\"end_s\", 0), pos_meta.get(\"trainable\", True))\n",
    "if pos_ctx_vae is not None:\n",
    "    full_vae = np.concatenate([neg_ctx_vae, vae_neg], axis=0)\n",
    "    pos_audio_ctx = codec_decode(full_vae).play()\n",
    "else:\n",
    "    pos_audio = codec_decode(vae_pos).play()\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3e3f2dfe",
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "suno_diff",
   "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.12.9"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
