{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "6ce269a3",
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "import json\n",
    "import pandas as pd\n",
    "import numpy as np\n",
    "from tqdm import tqdm\n",
    "from suno_utils.utils.text import write_jsonl"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "149692c2",
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "import matplotlib.pyplot as plt\n",
    "\n",
    "def plot_feature_distributions(\n",
    "    features,\n",
    "    datasets,\n",
    "    *,\n",
    "    bins=50,\n",
    "    cols=2,\n",
    "    feature_titles=None,\n",
    "    xlabels_delta=None,\n",
    "    xlabels_pos=None\n",
    "):\n",
    "    \"\"\"\n",
    "    Plot per-feature distributions for delta metrics and (pos) metrics,\n",
    "    overlaying multiple datasets with shared bin edges.\n",
    "\n",
    "    Parameters\n",
    "    ----------\n",
    "    features : list[str]\n",
    "        Feature names to plot (keys into *_dicts below).\n",
    "    datasets : list[tuple]\n",
    "        A list of (label, deltas_dict, pos_dict).\n",
    "        Example:\n",
    "            [\n",
    "                (\"metas\", all_deltas, all_pos_features),\n",
    "                (\"filtered_metas\", filtered_deltas, filtered_pos_features),\n",
    "            ]\n",
    "    bins : int\n",
    "        Number of histogram bins.\n",
    "    cols : int\n",
    "        Number of subplot columns.\n",
    "    \"\"\"\n",
    "    if feature_titles is None:\n",
    "        feature_titles = {\n",
    "            \"shimmer_score\": \"Shimmer score\",\n",
    "            \"lufs_db\": \"Loudness\",\n",
    "            \"ear_score\": \"Ear score\",\n",
    "            \"ear_v3_score\": \"Ear v3 score\",\n",
    "            \"hoot_cer\": \"Hoot cer\",\n",
    "            \"stereo_width\": \"Stereo width\",\n",
    "            \"total_delta\": \"Total frequency delta\",\n",
    "            \"mse\": \"Spec. Reference MSE\",\n",
    "        }\n",
    "\n",
    "    if xlabels_delta is None:\n",
    "        xlabels_delta = {f: f\"{feature_titles.get(f, f)} Delta\" for f in features}\n",
    "    if xlabels_pos is None:\n",
    "        xlabels_pos = {f: f\"{feature_titles.get(f, f)} (pos)\" for f in features}\n",
    "\n",
    "    def _make_grid(n_items, cols, height_per_row=2.6, width=8):\n",
    "        rows = (n_items + cols - 1) // cols\n",
    "        fig, axs = plt.subplots(rows, cols, figsize=(width, height_per_row * rows))\n",
    "        axs = axs.flatten() if isinstance(axs, np.ndarray) else np.array([axs])\n",
    "        return fig, axs, rows\n",
    "\n",
    "    # --- Safe helper for shared bins ---\n",
    "    def _shared_bins(datasets, feature, is_pos=False, bins=50):\n",
    "        arrays = []\n",
    "        for _, deltas_dict, pos_dict in datasets:\n",
    "            arr = pos_dict.get(feature, []) if is_pos else deltas_dict.get(feature, [])\n",
    "            arr = np.asarray(arr, dtype=float)\n",
    "            arr = arr[np.isfinite(arr)]\n",
    "            if arr.size:\n",
    "                arrays.append(arr)\n",
    "\n",
    "        if not arrays:\n",
    "            return None  # No data at all for this feature\n",
    "\n",
    "        combined = np.concatenate(arrays)\n",
    "        if combined.size == 0:\n",
    "            return None\n",
    "\n",
    "        lo, hi = np.min(combined), np.max(combined)\n",
    "        if not np.isfinite(lo) or not np.isfinite(hi) or lo == hi:\n",
    "            # Handle degenerate or NaN-only data\n",
    "            span = 1.0 if not np.isfinite(lo) or not np.isfinite(hi) else max(1e-6, abs(lo) * 1e-6 + 1e-6)\n",
    "            return np.linspace(lo - span/2, hi + span/2, bins + 1)\n",
    "\n",
    "        return np.histogram_bin_edges(combined, bins=bins)\n",
    "\n",
    "    # ---------- DELTAS ----------\n",
    "    fig_delta, axs_delta, _ = _make_grid(len(features), cols)\n",
    "    for i, f in enumerate(features):\n",
    "        ax = axs_delta[i]\n",
    "        title = f\"{feature_titles.get(f, f)} deltas\"\n",
    "        ax.set_title(title, fontsize=10)\n",
    "        ax.set_xlabel(xlabels_delta.get(f, f\"{f} Delta\"))\n",
    "        ax.set_ylabel(\"Count\")\n",
    "\n",
    "        edges = _shared_bins(datasets, f, is_pos=False, bins=bins)\n",
    "        if edges is None:\n",
    "            ax.set_visible(False)\n",
    "            continue\n",
    "\n",
    "        any_plotted = False\n",
    "        for j, (label, deltas_dict, _) in enumerate(datasets):\n",
    "            data = np.asarray(deltas_dict.get(f, []), dtype=float)\n",
    "            data = data[~np.isnan(data)]\n",
    "            if data.size == 0:\n",
    "                continue\n",
    "\n",
    "            ax.hist(data, bins=edges, alpha=0.45, edgecolor=\"black\", label=label)\n",
    "            p5, p95 = np.percentile(data, [5, 95])\n",
    "            #ax.axvline(p5, linestyle=\"--\")\n",
    "            #ax.axvline(p95, linestyle=\"--\")\n",
    "\n",
    "            mean, std = data.mean(), data.std()\n",
    "            dmin, dmax = data.min(), data.max()\n",
    "            print(f\"[DELTAS] {title} | {label}: mean={mean:.3f}, std={std:.3f}, \"\n",
    "                  f\"min={dmin:.3f}, max={dmax:.3f}, p5={p5:.3f}, p95={p95:.3f}\")\n",
    "\n",
    "            if j == 0:\n",
    "                stats_text = (\n",
    "                    f\"{label}\\nmean={mean:.3f}\\nstd={std:.3f}\\nmin={dmin:.3f}\\nmax={dmax:.3f}\"\n",
    "                )\n",
    "                ax.text(\n",
    "                    0.98, 0.98, stats_text, transform=ax.transAxes,\n",
    "                    fontsize=8, va=\"top\", ha=\"right\",\n",
    "                    bbox=dict(boxstyle=\"round,pad=0.3\", facecolor=\"white\", alpha=0.7)\n",
    "                )\n",
    "            any_plotted = True\n",
    "\n",
    "        ax.grid(True)  # Add grid to this subplot\n",
    "\n",
    "        if any_plotted:\n",
    "            ax.legend(fontsize=8)\n",
    "        else:\n",
    "            ax.set_visible(False)\n",
    "\n",
    "    for k in range(len(features), len(axs_delta)):\n",
    "        fig_delta.delaxes(axs_delta[k])\n",
    "    fig_delta.tight_layout()\n",
    "\n",
    "    # ---------- POS ----------\n",
    "    fig_pos, axs_pos, _ = _make_grid(len(features), cols)\n",
    "    for i, f in enumerate(features):\n",
    "        ax = axs_pos[i]\n",
    "        title = f\"{feature_titles.get(f, f)} (pos)\"\n",
    "        ax.set_title(title, fontsize=10)\n",
    "        ax.set_xlabel(xlabels_pos.get(f, f\"{f} (pos)\"))\n",
    "        ax.set_ylabel(\"Count\")\n",
    "\n",
    "        edges = _shared_bins(datasets, f, is_pos=True, bins=bins)\n",
    "        if edges is None:\n",
    "            ax.set_visible(False)\n",
    "            continue\n",
    "\n",
    "        any_plotted = False\n",
    "        for j, (label, _, pos_dict) in enumerate(datasets):\n",
    "            data = np.asarray(pos_dict.get(f, []), dtype=float)\n",
    "            data = data[~np.isnan(data)]\n",
    "            if data.size == 0:\n",
    "                continue\n",
    "\n",
    "            ax.hist(data, bins=edges, alpha=0.45, edgecolor=\"black\", label=label)\n",
    "            p5, p95 = np.percentile(data, [5, 95])\n",
    "            #ax.axvline(p5, linestyle=\"--\")\n",
    "            #ax.axvline(p95, linestyle=\"--\")\n",
    "\n",
    "            mean, std = data.mean(), data.std()\n",
    "            dmin, dmax = data.min(), data.max()\n",
    "            print(f\"[POS] {title} | {label}: mean={mean:.3f}, std={std:.3f}, \"\n",
    "                  f\"min={dmin:.3f}, max={dmax:.3f}, p5={p5:.3f}, p95={p95:.3f}\")\n",
    "\n",
    "            if j == 0:\n",
    "                stats_text = (\n",
    "                    f\"{label}\\nmean={mean:.3f}\\nstd={std:.3f}\\nmin={dmin:.3f}\\nmax={dmax:.3f}\"\n",
    "                )\n",
    "                ax.text(\n",
    "                    0.98, 0.98, stats_text, transform=ax.transAxes,\n",
    "                    fontsize=8, va=\"top\", ha=\"right\",\n",
    "                    bbox=dict(boxstyle=\"round,pad=0.3\", facecolor=\"white\", alpha=0.7)\n",
    "                )\n",
    "            any_plotted = True\n",
    "\n",
    "        ax.grid(True)  # Add grid to this subplot\n",
    "\n",
    "        if any_plotted:\n",
    "            ax.legend(fontsize=8)\n",
    "        else:\n",
    "            ax.set_visible(False)\n",
    "\n",
    "    \n",
    "\n",
    "    for k in range(len(features), len(axs_pos)):\n",
    "        fig_pos.delaxes(axs_pos[k])\n",
    "    fig_pos.tight_layout()\n",
    "    plt.show(block=False)\n",
    "\n",
    "    return fig_delta, fig_pos\n"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "018eba9b",
   "metadata": {},
   "source": [
    "# From generated data (modal)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "34ef7cec",
   "metadata": {},
   "outputs": [],
   "source": [
    "#model_name = \"4n_25hz_2b_flow_5e5_sft_t8_500k\"\n",
    "model_name = \"16n_25hz_v45_infill_shared_flow_resume_1_75m\"\n",
    "base_dir = \"/app2/suno/data/christian/outputs/v3-base-data-ctx-rs-t5/\"\n",
    "\n",
    "dirnames = os.listdir(base_dir)\n",
    "# filter to only include dirs\n",
    "dirnames = [d for d in dirnames if os.path.isdir(os.path.join(base_dir, d))]\n",
    "print(len(dirnames))\n",
    "\n",
    "# load the extra metadata\n",
    "extra_metadata_filepath = os.path.join(base_dir, \"extra_metadata.npz\")\n",
    "if os.path.exists(extra_metadata_filepath):\n",
    "    print(f\"Loading extra metadata from {extra_metadata_filepath}\")\n",
    "    extra_metadata = np.load(extra_metadata_filepath, allow_pickle=True)\n",
    "    print(len(extra_metadata))\n",
    "else:\n",
    "    extra_metadata = {}\n",
    "\n",
    "# check for ear score in base_dir \n",
    "#ear_score_filename = \"ear_scores_s4862.csv\"\n",
    "#ear_score_filepath = os.path.join(base_dir, ear_score_filename)\n",
    "#if os.path.exists(ear_score_filepath):\n",
    "#    ear_scores = pd.read_csv(ear_score_filepath)\n",
    "\n",
    "# create a ear_scores dict with id mapped to 0_mean and 1_mean\n",
    "#ear_scores_dict = ear_scores.set_index(\"index\").to_dict(orient=\"records\")\n",
    "#ear_scores_dict = {d[\"id\"]: {\"0_mean\": d[\"0_mean\"], \"1_mean\": d[\"1_mean\"]} for d in ear_scores_dict}\n",
    "\n",
    "label_lookup = {}\n",
    "for dirname in tqdm(dirnames):\n",
    "    label = 1        \n",
    "    label_lookup[dirname] = {\n",
    "        \"label\": label, \n",
    "    }\n",
    "\n",
    "def process_dir(base_dir, dirname):\n",
    "    results = {}\n",
    "    semantic_codes_filepath = os.path.join(base_dir, dirname, f\"{dirname}_semantic.npz\")\n",
    "    for n in range(2):\n",
    "        upsampled_vae_filepath = os.path.join(base_dir, dirname, f\"{dirname}_{model_name}_{n}_upsampled_vae.npz\")\n",
    "        metadata_filepath = os.path.join(base_dir, dirname, f\"{dirname}_{model_name}_{n}__metadata.npz\")\n",
    "        try:\n",
    "            with np.load(metadata_filepath, allow_pickle=True) as data:\n",
    "                metadata_npz = dict(data)\n",
    "        except FileNotFoundError:\n",
    "            metadata_npz = None\n",
    "\n",
    "        if metadata_npz is None:\n",
    "            continue\n",
    "        metadata_dict = {key: metadata_npz[key].tolist() for key in metadata_npz.keys()}\n",
    "\n",
    "        results[n] = {\n",
    "            \"upsampled_vae_filepath\": upsampled_vae_filepath,\n",
    "            \"metadata\": metadata_dict,\n",
    "        }\n",
    "    return results, semantic_codes_filepath\n",
    "\n",
    "# get all the dirs in the base_dir\n",
    "dirs = os.listdir(base_dir)\n",
    "# filter to things that are only a dir\n",
    "dirs = [d for d in dirs if os.path.isdir(os.path.join(base_dir, d))]\n",
    "\n",
    "metas = []\n",
    "for dirname in tqdm(dirs):\n",
    "    # get the label from the label_lookup\n",
    "    label = label_lookup.get(dirname, None)\n",
    "    if label is None:\n",
    "        continue\n",
    "\n",
    "    results, semantic_codes_filepath = process_dir(base_dir, dirname)\n",
    "\n",
    "    # Improved check: ensure both 0 and 1 keys exist and their metadata is not None\n",
    "    if (\n",
    "        0 not in results or\n",
    "        1 not in results or\n",
    "        results[0].get(\"metadata\") is None or\n",
    "        results[1].get(\"metadata\") is None\n",
    "    ):\n",
    "        continue\n",
    "\n",
    "    if label[\"label\"] == 0:\n",
    "        pos_idx, neg_idx = 0, 1\n",
    "    else:\n",
    "        pos_idx, neg_idx = 1, 0\n",
    "\n",
    "    pos_metadata = results[pos_idx][\"metadata\"]\n",
    "    neg_metadata = results[neg_idx][\"metadata\"]\n",
    "\n",
    "    # check if we have extra metadata\n",
    "    if dirname in extra_metadata:\n",
    "        extra_metadata_for_dirname = extra_metadata[dirname].item()\n",
    "        pos_extra_metadata = extra_metadata_for_dirname[pos_idx]\n",
    "        neg_extra_metadata = extra_metadata_for_dirname[neg_idx]\n",
    "        pos_metadata.update(pos_extra_metadata)\n",
    "        neg_metadata.update(neg_extra_metadata)\n",
    "\n",
    "    text = str(pos_metadata[\"text\"])\n",
    "    tags = str(pos_metadata[\"tags\"])\n",
    "\n",
    "    meta = {\n",
    "        \"id\": dirname,\n",
    "        \"text\": str(pos_metadata[\"text\"]),\n",
    "        \"tags\": str(pos_metadata[\"tags\"]),\n",
    "        \"pos_vae_latents_filepath\": results[pos_idx][\"upsampled_vae_filepath\"],\n",
    "        \"neg_vae_latents_filepath\": results[neg_idx][\"upsampled_vae_filepath\"],\n",
    "        \"pos_metadata\": pos_metadata,\n",
    "        \"neg_metadata\": neg_metadata,\n",
    "        \"semantic_codes_filepath\": semantic_codes_filepath,\n",
    "    }\n",
    "    metas.append(meta)\n",
    "\n",
    "print(len(metas))\n"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "221d0fb8",
   "metadata": {},
   "source": [
    "# From Prod data"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "85daa2d7",
   "metadata": {},
   "outputs": [],
   "source": [
    "base_dir = \"/app2/suno/data/christian/outputs/dorado_t1\"\n",
    "os.makedirs(base_dir, exist_ok=True)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3433e0a6",
   "metadata": {},
   "outputs": [],
   "source": [
    "dataframe_filepath = \"/home/tony/Data/Preference/dorado_t1/interesting_clips_dorado_t1_20251003.pkl\"\n",
    "root_dir = \"/app2/suno/data/dpo/dorado_t1\"\n",
    "\n",
    "# load each dataframe\n",
    "# select the last 5% of rows as validation use the rest as training\n",
    "# then create two dataframes, one for training and one for validation\n",
    "# the odd index in the dataframe is the neative and even is the positive\n",
    "\n",
    "df = pd.read_pickle(dataframe_filepath)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "484ecab0",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "61836e36",
   "metadata": {},
   "outputs": [],
   "source": [
    "dataframes = [\n",
    "    #\"/home/tony/Data/Preference/up_v2_d3/interesting_clips_ahi_d3_20250608.pkl\",\n",
    "    #\"/home/tony/Data/Preference/up_v2_d4/interesting_clips_ahi_d4_20250714.pkl\",\n",
    "    #\"/home/tony/Data/Preference/up_v2_d5/interesting_clips_ahi_d5_20250824.pkl\"\n",
    "    #\"/home/tony/Data/Preference/up_v2_d5/fully_merged_up_v2_d5.pkl\"\n",
    "]\n",
    "\n",
    "root_dirs = [\n",
    "    #\"/app2/suno/data/dpo/diff2_v2_d3\",\n",
    "    #\"/app2/suno/data/dpo/diff2_v2_d4\",\n",
    "    #\"/app2/suno/data/dpo/diff2_v2_d5\"\n",
    "]\n",
    "\n",
    "#dataframe_filepath = \"/home/tony/Data/Preference/up_v2_d5/fully_merged_up_v2_d5.pkl\"\n",
    "#root_dir = \"/app2/suno/data/dpo/diff2_v2_d5\"\n",
    "\n",
    "dataframe_filepath = \"/home/tony/Data/Preference/dorado_t1/interesting_clips_dorado_t1_20251003.pkl\"\n",
    "root_dir = \"/app2/suno/data/dpo/dorado_t1\"\n",
    "\n",
    "# load each dataframe\n",
    "# select the last 5% of rows as validation use the rest as training\n",
    "# then create two dataframes, one for training and one for validation\n",
    "# the odd index in the dataframe is the neative and even is the positive\n",
    "\n",
    "df = pd.read_pickle(dataframe_filepath)\n",
    "\n",
    "# merge in the audio features and metrics\n",
    "with open(\"/home/tony/Data/Preference/dorado_t1/full_pair_quality.json\", \"r\") as file:\n",
    "    full_pair_quality = json.load(file)\n",
    "print(\"Total pair quality scores:\", len(full_pair_quality))\n",
    "\n",
    "# iterate over the dataframe and crate the pair metas\n",
    "# make indices an array of odd indices\n",
    "subset_metas = []\n",
    "indices = np.arange(len(df))\n",
    "indices = indices[indices % 2 == 1]\n",
    "# iterate over the indices and create the pair metas\n",
    "for index in tqdm(indices):\n",
    "    request_id = df.iloc[index][\"request_id\"]\n",
    "\n",
    "    pos_item_id = df.iloc[index][\"id\"]\n",
    "    neg_item_id = df.iloc[index - 1][\"id\"]\n",
    "\n",
    "    pos_norm_play_frac = df.iloc[index][\"norm_play_frac\"]\n",
    "    neg_norm_play_frac = df.iloc[index - 1][\"norm_play_frac\"]\n",
    "\n",
    "    pos_edited_clip_id = df.iloc[index][\"edited_clip_id\"]\n",
    "    neg_edited_clip_id = df.iloc[index - 1][\"edited_clip_id\"]\n",
    "    if pos_edited_clip_id != neg_edited_clip_id:\n",
    "        print(pos_edited_clip_id, neg_edited_clip_id, \"clip ids mismatch\")\n",
    "        continue\n",
    "\n",
    "    pos_duration_s = df.iloc[index][\"duration\"]\n",
    "    neg_duration_s = df.iloc[index - 1][\"duration\"]\n",
    "\n",
    "    pos_variation = df.iloc[index][\"metadata\"].get(\"remaster_sliders\", {}).get(\"variation_category\", None)\n",
    "    neg_variation = df.iloc[index - 1][\"metadata\"].get(\"remaster_sliders\", {}).get(\"variation_category\", None)\n",
    "\n",
    "    if pos_duration_s < 30.0 or neg_duration_s < 30.0:\n",
    "        continue\n",
    "\n",
    "    pos_metadata = full_pair_quality.get(request_id, {}).get(pos_item_id, None)\n",
    "    neg_metadata = full_pair_quality.get(request_id, {}).get(neg_item_id, None)\n",
    "\n",
    "    if pos_metadata is None or neg_metadata is None:\n",
    "        continue\n",
    "\n",
    "    # in the metadata, convert \"abs_loudness_factor\" to \"lufs_db\"\n",
    "    # convert \"ear-v3_quality_scores to \"ear_score\" and then take the mean\n",
    "    pos_metadata[\"lufs_db\"] = pos_metadata.get(\"abs_loudness_factor\", None)\n",
    "    pos_metadata[\"ear_score\"] = np.mean(pos_metadata.get(\"ear_v2_quality_scores\", None))\n",
    "    pos_metadata[\"norm_play_frac\"] = pos_norm_play_frac\n",
    "    neg_metadata[\"lufs_db\"] = neg_metadata.get(\"abs_loudness_factor\", None)\n",
    "    neg_metadata[\"ear_score\"] = np.mean(neg_metadata[\"ear_v2_quality_scores\"])\n",
    "    neg_metadata[\"norm_play_frac\"] = neg_norm_play_frac\n",
    "\n",
    "    pos_vae_latents_filepath = os.path.join(root_dir, f\"{pos_item_id}_vae.npz\")\n",
    "    neg_vae_latents_filepath = os.path.join(root_dir, f\"{neg_item_id}_vae.npz\")\n",
    "\n",
    "    if os.path.exists(pos_vae_latents_filepath) \\\n",
    "        and os.path.exists(neg_vae_latents_filepath) \\\n",
    "        and os.path.exists(os.path.join(root_dir, f\"{pos_edited_clip_id}.npz\")):\n",
    "        subset_metas.append({\n",
    "            \"id\": request_id,\n",
    "            \"pos_id\": pos_item_id,\n",
    "            \"neg_id\": neg_item_id,\n",
    "            \"text\": df.iloc[index][\"prompt_text\"],\n",
    "            \"tags\": df.iloc[index][\"metadata\"][\"tags\"],\n",
    "            \"pos_vae_latents_filepath\": pos_vae_latents_filepath,\n",
    "            \"neg_vae_latents_filepath\": neg_vae_latents_filepath,\n",
    "            \"pos_duration_s\": pos_duration_s,\n",
    "            \"neg_duration_s\": neg_duration_s,\n",
    "            \"pos_metadata\": pos_metadata,\n",
    "            \"neg_metadata\": neg_metadata,\n",
    "            \"pos_variation\": pos_variation,\n",
    "            \"neg_variation\": neg_variation,\n",
    "            \"semantic_codes_filepath\": os.path.join(root_dir, f\"{pos_edited_clip_id}.npz\")\n",
    "        })"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "63187046",
   "metadata": {},
   "outputs": [],
   "source": [
    "# bar chart of the variation categories\n",
    "variation_counts = {}\n",
    "for meta in subset_metas:\n",
    "    variation = meta[\"pos_variation\"]\n",
    "    if variation is None:\n",
    "        variation = \"normal\"\n",
    "    variation_counts[variation] = variation_counts.get(variation, 0) + 1\n",
    "\n",
    "plt.bar(variation_counts.keys(), variation_counts.values())"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "b5b32667",
   "metadata": {},
   "outputs": [],
   "source": [
    "filtered_subset_metas = []\n",
    "for meta in subset_metas:\n",
    "    if meta[\"pos_variation\"] == \"normal\" and meta[\"neg_variation\"] == \"normal\" or meta[\"pos_variation\"] is None and meta[\"neg_variation\"] is None:\n",
    "        filtered_subset_metas.append(meta)\n",
    "\n",
    "print(len(filtered_subset_metas))\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "93f2843b",
   "metadata": {},
   "outputs": [],
   "source": [
    "metas = filtered_subset_metas"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "dbac1d5e",
   "metadata": {},
   "source": [
    "# From LabelMaker"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "57ea84a1",
   "metadata": {},
   "outputs": [],
   "source": [
    "model_name = \"4n_25hz_2b_flow_5e5_sft_t8_500k\"\n",
    "base_dir = \"/app2/suno/data/christian/outputs/v3-base-data-ctx-t2/\"\n",
    "\n",
    "dirnames = os.listdir(base_dir)\n",
    "# filter to only include dirs\n",
    "dirnames = [d for d in dirnames if os.path.isdir(os.path.join(base_dir, d))]\n",
    "print(len(dirnames))\n",
    "\n",
    "ratings_filepath = \"/home/christian/code/christian/metadata/labelmaker/t5/dpo_annotations_export_0_7.csv\"\n",
    "ratings_df = pd.read_csv(ratings_filepath)\n",
    "print(ratings_df.head())"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "f2dd1976",
   "metadata": {},
   "outputs": [],
   "source": [
    "# now collect all the clips that have ratings for\n",
    "metas = []\n",
    "\n",
    "for index, row in tqdm(ratings_df.iterrows()):\n",
    "    clip_id = row[\"clip_id\"]\n",
    "    agreement = row[\"agreement\"]\n",
    "    num_ratings = row[\"num_ratings\"]\n",
    "    chosen_file_index = row[\"chosen_file_index\"]    \n",
    "    unchosen_file_index = row[\"unchosen_file_index\"]\n",
    "    # get the directory with the audio files\n",
    "    dirname = os.path.join(base_dir, f\"{clip_id}\")\n",
    "    pos_vae_latents_filepath = os.path.join(dirname, f\"{clip_id}_{model_name}_{chosen_file_index}_upsampled_vae.npz\")\n",
    "    pos_metadata_filepath = os.path.join(dirname, f\"{clip_id}_{model_name}_{chosen_file_index}__metadata.npz\")\n",
    "    neg_vae_latents_filepath = os.path.join(dirname, f\"{clip_id}_{model_name}_{unchosen_file_index}_upsampled_vae.npz\")\n",
    "    neg_metadata_filepath = os.path.join(dirname, f\"{clip_id}_{model_name}_{unchosen_file_index}__metadata.npz\")\n",
    "    semantic_codes_filepath = os.path.join(dirname, f\"{clip_id}_semantic.npz\")\n",
    "\n",
    "    # check if the files exist\n",
    "    if not os.path.exists(pos_vae_latents_filepath) \\\n",
    "        or not os.path.exists(neg_vae_latents_filepath) \\\n",
    "        or not os.path.exists(semantic_codes_filepath):\n",
    "        continue\n",
    "\n",
    "    history_vae_latents_filepath = os.path.join(dirname, f\"{clip_id}_{model_name}_history_vae.npz\")\n",
    "    history_vae_latents_filepath = history_vae_latents_filepath if os.path.exists(history_vae_latents_filepath) else None\n",
    "\n",
    "    # open the metadata files and get the text and tags\n",
    "    with np.load(pos_metadata_filepath, allow_pickle=True) as data:\n",
    "        pos_metadata = {key: data[key].tolist() for key in data.keys()}\n",
    "    with np.load(neg_metadata_filepath, allow_pickle=True) as data:\n",
    "        neg_metadata = {key: data[key].tolist() for key in data.keys()}\n",
    "    \n",
    "    text = str(pos_metadata[\"text\"])\n",
    "    tags = str(pos_metadata[\"tags\"])\n",
    "\n",
    "    metas.append({\n",
    "        \"id\": clip_id,\n",
    "        \"text\": text,\n",
    "        \"tags\": tags,\n",
    "        \"pos_vae_latents_filepath\": pos_vae_latents_filepath,\n",
    "        \"neg_vae_latents_filepath\": neg_vae_latents_filepath,\n",
    "        \"semantic_codes_filepath\": semantic_codes_filepath,\n",
    "        \"history_vae_latents_filepath\": history_vae_latents_filepath,\n",
    "        \"pos_metadata\": pos_metadata,\n",
    "        \"neg_metadata\": neg_metadata,\n",
    "        \"agreement\": agreement,\n",
    "        \"num_ratings\": num_ratings,\n",
    "        \"chosen_file_index\": chosen_file_index,\n",
    "        \"unchosen_file_index\": unchosen_file_index,\n",
    "    })\n",
    "\n",
    "print(len(metas))"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "b4e26477",
   "metadata": {},
   "source": [
    "# Filtering"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "7a00508a",
   "metadata": {},
   "outputs": [],
   "source": [
    "# apppy filtering here\n",
    "\n",
    "# Collect all feature values from pos_metadata and neg_metadata, and compute deltas (pos - neg)\n",
    "features = [\"shimmer_score\", \"lufs_db\", \"ear_score\", \"ear_v3_score\", \"hoot_cer\", \"stereo_width\", \"stereo_width_delta\", \"total_delta\", \"mse\"]\n",
    "#features = [\"shimmer_score\", \"lufs_db\", \"ear_score\", \"ear_v3_score\", \"hoot_cer\", \"stereo_width\", \"norm_play_frac\"]\n",
    "\n",
    "\n",
    "all_pos_features = {f: [] for f in features}\n",
    "all_neg_features = {f: [] for f in features}\n",
    "all_deltas = {f: [] for f in features}\n",
    "all_delta_percents = {f: [] for f in features}  # delta %: (pos - neg) / abs(neg) * 100\n",
    "\n",
    "for meta in metas:\n",
    "    pos_metadata = meta[\"pos_metadata\"]\n",
    "    neg_metadata = meta[\"neg_metadata\"]\n",
    "    for f in features:\n",
    "        pos_val = pos_metadata.get(f, None)\n",
    "        neg_val = neg_metadata.get(f, None)\n",
    "        if pos_val is not None and neg_val is not None:\n",
    "            pos_val = float(pos_val)\n",
    "            neg_val = float(neg_val)\n",
    "            all_pos_features[f].append(pos_val)\n",
    "            all_neg_features[f].append(neg_val)\n",
    "            all_deltas[f].append(pos_val - neg_val)\n",
    "            # Compute delta percent, handle neg_val == 0\n",
    "            if neg_val != 0:\n",
    "                delta_percent = ((pos_val - neg_val) / abs(neg_val)) * 100\n",
    "            else:\n",
    "                delta_percent = float('nan')\n",
    "            all_delta_percents[f].append(delta_percent)\n",
    "\n",
    "plot_feature_distributions(\n",
    "    features,\n",
    "    datasets=[\n",
    "        (\"metas\", all_deltas, all_pos_features),\n",
    "    ],\n",
    "    bins=50,\n",
    "    cols=2,\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "fc91449a",
   "metadata": {},
   "outputs": [],
   "source": [
    "# now apply the filtering here\n",
    "\n",
    "# start with the deltas\n",
    "delta_ranges = {\n",
    "    \"shimmer_score\": (-8.0, 3.0),\n",
    "    \"lufs_db\": (-3.0, 0.5),\n",
    "    \"ear_score\": (0.0, 9.0),\n",
    "    \"ear_v3_score\": (-2.5, 4.0),\n",
    "    \"hoot_cer\": (-1, 0.05),\n",
    "    \"stereo_width\": (-0.18, 0.18),\n",
    "    \"total_delta\": (-150, -15),\n",
    "    \"stereo_width_delta\": (-0.18, 0.18),\n",
    "    \"mse\": (-40, 0),\n",
    "}\n",
    "\n",
    "pos_ranges = {\n",
    "    \"shimmer_score\": (0.0, 4.0),\n",
    "    \"lufs_db\": (-24, -12),\n",
    "    \"ear_score\": (15, 25),\n",
    "    \"ear_v3_score\": (-5, 5),\n",
    "    \"hoot_cer\": (0, 1.0),\n",
    "    \"stereo_width\": (0.05, 0.3),\n",
    "    \"total_delta\": (0, 60),\n",
    "    \"stereo_width_delta\": (-0.15, 0.15),\n",
    "    \"mse\": (0, 20),\n",
    "    #\"norm_play_frac\": (1.0, 100.0),\n",
    "}\n",
    "\n",
    "cut_metadata = {\n",
    "    \"delta_ranges\": delta_ranges,\n",
    "    \"pos_ranges\": pos_ranges,\n",
    "}\n",
    "\n",
    "features = [\"lufs_db\"] #\n",
    "\n",
    "# Filtering logic with per-feature rejection counts:\n",
    "from collections import defaultdict\n",
    "\n",
    "filtered_metas = []\n",
    "rejection_counts = defaultdict(int)\n",
    "total_rejections = 0\n",
    "\n",
    "for meta in metas:\n",
    "    pos_metadata = meta.get(\"pos_metadata\", {})\n",
    "    neg_metadata = meta.get(\"neg_metadata\", {})\n",
    "    passed = True\n",
    "    for f in features:\n",
    "        pos_val = pos_metadata.get(f, None)\n",
    "        neg_val = neg_metadata.get(f, None)\n",
    "        # If either value is missing, skip this feature (do not check, let it pass)\n",
    "        if pos_val is None or neg_val is None:\n",
    "            continue\n",
    "        # Check delta ranges\n",
    "        if f in delta_ranges:\n",
    "            try:\n",
    "                delta = float(pos_val) - float(neg_val)\n",
    "            except Exception:\n",
    "                rejection_counts[f + \"_delta_cast\"] += 1\n",
    "                total_rejections += 1\n",
    "                passed = False\n",
    "                break\n",
    "            if not (delta_ranges[f][0] <= delta <= delta_ranges[f][1]):\n",
    "                rejection_counts[f + \"_delta\"] += 1\n",
    "                total_rejections += 1\n",
    "                passed = False\n",
    "                break\n",
    "        # Check positive sample ranges\n",
    "        if f in pos_ranges:\n",
    "            try:\n",
    "                pos_val_float = float(pos_val)\n",
    "            except Exception:\n",
    "                rejection_counts[f + \"_pos_cast\"] += 1\n",
    "                total_rejections += 1\n",
    "                passed = False\n",
    "                break\n",
    "            if not (pos_ranges[f][0] <= pos_val_float <= pos_ranges[f][1]):\n",
    "                rejection_counts[f + \"_pos\"] += 1\n",
    "                total_rejections += 1\n",
    "                passed = False\n",
    "                break\n",
    "    if passed:\n",
    "        new_meta = meta.copy()\n",
    "        # remove the pos_metadata and neg_metadata\n",
    "        #if \"pos_metadata\" in new_meta:\n",
    "        #    del new_meta[\"pos_metadata\"]\n",
    "        #if \"neg_metadata\" in new_meta:\n",
    "        #     del new_meta[\"neg_metadata\"]\n",
    "        filtered_metas.append(new_meta)\n",
    "\n",
    "print(f\"Remaining {len(filtered_metas)} metas from {len(metas)} ({100 * len(filtered_metas) / len(metas):.2f}%)\")\n",
    "print(f\"Total rejections: {total_rejections}\\n\")\n",
    "print(\"Rejection counts by feature:\")\n",
    "for k, v in sorted(rejection_counts.items(), key=lambda item: item[1], reverse=True):\n",
    "    print(f\"  {k:20} {v:>8} ({100 * v / total_rejections:.2f}%)\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "83802ad8",
   "metadata": {},
   "outputs": [],
   "source": [
    "# Collect all feature values from pos_metadata and neg_metadata, and compute deltas (pos - neg)\n",
    "features = [\"shimmer_score\", \"lufs_db\", \"ear_score\", \"ear_v3_score\", \"hoot_cer\", \"stereo_width\", \"total_delta\", \"mse\"]\n",
    "\n",
    "filtered_pos_features = {f: [] for f in features}\n",
    "filtered_neg_features = {f: [] for f in features}\n",
    "filtered_deltas = {f: [] for f in features}\n",
    "\n",
    "for meta in filtered_metas:\n",
    "    pos_metadata = meta[\"pos_metadata\"]\n",
    "    neg_metadata = meta[\"neg_metadata\"]\n",
    "    for f in features:\n",
    "        pos_val = pos_metadata.get(f, None)\n",
    "        neg_val = neg_metadata.get(f, None)\n",
    "        if pos_val is not None and neg_val is not None:\n",
    "            pos_val = float(pos_val)\n",
    "            neg_val = float(neg_val)\n",
    "            filtered_pos_features[f].append(pos_val)\n",
    "            filtered_neg_features[f].append(neg_val)\n",
    "            filtered_deltas[f].append(pos_val - neg_val)\n",
    "\n",
    "plot_feature_distributions(\n",
    "    features,\n",
    "    datasets=[\n",
    "        (\"metas\", all_deltas, all_pos_features),\n",
    "        (\"filtered_metas\", filtered_deltas, filtered_pos_features),\n",
    "    ],\n",
    "    bins=50,\n",
    "    cols=2,\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a60d0274",
   "metadata": {},
   "outputs": [],
   "source": [
    "# save the metas\n",
    "# shuffle filtered_metas before splitting\n",
    "import random\n",
    "random.shuffle(filtered_metas)\n",
    "\n",
    "# split metas into train and test\n",
    "split_idx = int(len(filtered_metas) * 0.95)\n",
    "train_metas = filtered_metas[:split_idx]\n",
    "val_metas = filtered_metas[split_idx:]\n",
    "\n",
    "# repeat the val metas 4 times\n",
    "val_metas = val_metas * 1\n",
    "\n",
    "print(len(train_metas), len(val_metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "f9518ab3",
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "version = \"t31\"\n",
    "tr_output_filepath = os.path.join(base_dir, f\"metas_tr_{version}.jsonl\")\n",
    "val_output_filepath = os.path.join(base_dir, f\"metas_val_{version}.jsonl\")\n",
    "\n",
    "# remove the pos_metadata and neg_metadata from the metas\n",
    "for meta in train_metas:\n",
    "    if \"pos_metadata\" in meta:\n",
    "        del meta[\"pos_metadata\"]\n",
    "    if \"neg_metadata\" in meta:\n",
    "        del meta[\"neg_metadata\"]\n",
    "\n",
    "for meta in val_metas:\n",
    "    if \"pos_metadata\" in meta:\n",
    "        del meta[\"pos_metadata\"]\n",
    "    if \"neg_metadata\" in meta:\n",
    "        del meta[\"neg_metadata\"]\n",
    "\n",
    "write_jsonl(train_metas, tr_output_filepath)\n",
    "write_jsonl(val_metas, val_output_filepath)\n",
    "\n",
    "# save this as a json file in the base_dir, with version string in filename\n",
    "with open(os.path.join(base_dir, f\"cut_metadata_{version}.json\"), \"w\") as f:\n",
    "    json.dump(cut_metadata, f)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "c50c7019",
   "metadata": {},
   "outputs": [],
   "source": [
    "base_dir"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "0cacabd4",
   "metadata": {},
   "source": [
    "# Extra junk "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a6ad38c7",
   "metadata": {},
   "outputs": [],
   "source": [
    "import sys\n",
    "sys.path.append(\"/home/christian/code/neon/sunoDiff\")\n",
    "from dataset import PairFileDataset\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "bdb113e8",
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "import torch\n",
    "\n",
    "dataset_dir = \"/app2/suno/data/christian/outputs/v3-base-data-ctx-t2/\"\n",
    "metas_filename = \"metas_tr_t8.jsonl\"\n",
    "\n",
    "\n",
    "dataset = PairFileDataset(\n",
    "    dataset_dir=dataset_dir,\n",
    "    metas_filename=metas_filename,\n",
    "    chunk_size=750,\n",
    "    cond_text_len=2560,\n",
    "    vae_scale_factor=0.4,\n",
    "    noise_ctx=1.0,\n",
    ")\n",
    "\n",
    "# make datalodaer\n",
    "dataloader = torch.utils.data.DataLoader(dataset, batch_size=4, shuffle=True, num_workers=4)\n",
    "\n",
    "import time\n",
    "from tqdm import tqdm\n",
    "# time how long it takes to load a batch\n",
    "# average over 100 batches\n",
    "start_time = time.time()\n",
    "count = 0\n",
    "for i, batch in tqdm(enumerate(dataloader)):\n",
    "    count += 1\n",
    "    if count > 100:\n",
    "        break\n",
    "\n",
    "print(f\"Time per batch: {(time.time() - start_time) / count} seconds\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "508e5c93",
   "metadata": {},
   "outputs": [],
   "source": [
    "# collec the spectrum backfilled data and save as a csv file\n",
    "import os\n",
    "import numpy as np\n",
    "model_name = \"4n_25hz_2b_flow_5e5_sft_t8_500k\"\n",
    "base_dir = \"/app2/suno/data/christian/outputs/v3-base-data-ctx-rs-t3/\"\n",
    "\n",
    "dirnames = os.listdir(base_dir)\n",
    "# filter to only include dirs\n",
    "dirnames = [d for d in dirnames if os.path.isdir(os.path.join(base_dir, d))]\n",
    "print(len(dirnames))\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "773fe905",
   "metadata": {},
   "outputs": [],
   "source": [
    "from tqdm import tqdm\n",
    "results = {}\n",
    "\n",
    "for dirname in tqdm(dirnames):\n",
    "    # look for files called \"*_0_spectrum.npz\" and \"*_1_spectrum.npz\"\n",
    "    import glob\n",
    "\n",
    "    spectrum_files_0 = glob.glob(os.path.join(base_dir, dirname, \"*_0_spectrum.npz\"))\n",
    "    spectrum_files_1 = glob.glob(os.path.join(base_dir, dirname, \"*_1_spectrum.npz\"))\n",
    "\n",
    "    if len(spectrum_files_0) == 0 or len(spectrum_files_1) == 0:\n",
    "        continue\n",
    "\n",
    "    # load the spectrum files\n",
    "    spectrum_0 = np.load(spectrum_files_0[0])\n",
    "    spectrum_1 = np.load(spectrum_files_1[0])\n",
    "\n",
    "\n",
    "    results[dirname] = {\n",
    "        0 : {f: spectrum_0[f] for f in spectrum_0.keys()},\n",
    "        1 : {f: spectrum_1[f] for f in spectrum_1.keys()},\n",
    "    }"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "bf67ecb1",
   "metadata": {},
   "outputs": [],
   "source": [
    "extra_metadata[\"e128df5a-962c-4f5e-911c-61696bd7af74\"]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a4941f85",
   "metadata": {},
   "outputs": [],
   "source": [
    "extra_metadata = results\n",
    "print(len(extra_metadata))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "7e490170",
   "metadata": {},
   "outputs": [],
   "source": [
    "# save this extra metadata as a npz file\n",
    "np.savez(os.path.join(base_dir, \"extra_metadata.npz\"), **extra_metadata)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "ca3f06ea",
   "metadata": {},
   "outputs": [],
   "source": [
    "extra_metadata_load = np.load(os.path.join(base_dir, \"extra_metadata.npz\"))\n",
    "print(len(extra_metadata_load))\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "eb9936d4",
   "metadata": {},
   "outputs": [],
   "source": [
    "# total deltas\n",
    "total_deltas = []\n",
    "pos_total_deltas = []\n",
    "neg_total_deltas = []\n",
    "pos_mse = []\n",
    "neg_mse = []\n",
    "for clip_id, vals in results.items():\n",
    "    total_deltas.append(vals[0][\"total_delta\"])\n",
    "    total_deltas.append(vals[1][\"total_delta\"])\n",
    "    pos_total_deltas.append(vals[1][\"total_delta\"])\n",
    "    neg_total_deltas.append(vals[0][\"total_delta\"])\n",
    "    pos_mse.append(vals[1][\"mse\"])\n",
    "    neg_mse.append(vals[0][\"mse\"])\n",
    "\n",
    "print(len(total_deltas))\n",
    "print(len(pos_total_deltas))\n",
    "print(len(neg_total_deltas))\n",
    "print(len(pos_mse))\n",
    "print(len(neg_mse))\n",
    "# save the total deltas\n",
    "\n",
    "# save the total deltas\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a362af5e",
   "metadata": {},
   "outputs": [],
   "source": [
    "import matplotlib.pyplot as plt\n",
    "\n",
    "#plt.hist(total_deltas, bins=100)\n",
    "#plt.show()\n",
    "\n",
    "plt.hist(pos_total_deltas, bins=100, alpha=0.5, label=\"pos\")    \n",
    "plt.hist(neg_total_deltas, bins=100, alpha=0.5, label=\"neg\")\n",
    "# compute 95% and 5% of the pos tital deltas\n",
    "pos_95 = np.percentile(pos_total_deltas, 95)\n",
    "pos_5 = np.percentile(pos_total_deltas, 5)\n",
    "neg_95 = np.percentile(neg_total_deltas, 95)\n",
    "neg_5 = np.percentile(neg_total_deltas, 5)\n",
    "\n",
    "plt.axvline(pos_95, color=\"red\", linestyle=\"--\")\n",
    "plt.axvline(pos_5, color=\"red\", linestyle=\"--\")\n",
    "#plt.axvline(neg_95, color=\"blue\", linestyle=\"--\")\n",
    "#plt.axvline(neg_5, color=\"blue\", linestyle=\"--\")\n",
    "plt.legend()\n",
    "plt.show()\n",
    "\n",
    "# save the total deltas\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "90396690",
   "metadata": {},
   "outputs": [],
   "source": [
    "# histogram of pos_mse\n",
    "plt.hist(pos_mse, bins=100, alpha=0.5)\n",
    "plt.hist(neg_mse, bins=100, alpha=0.5)\n",
    "# compute 95% and 5% of the pos tital deltas\n",
    "pos_95 = np.percentile(pos_mse, 95)\n",
    "pos_5 = np.percentile(pos_mse, 5)\n",
    "neg_95 = np.percentile(neg_mse, 95)\n",
    "neg_5 = np.percentile(neg_mse, 5)\n",
    "\n",
    "plt.axvline(pos_95, color=\"red\", linestyle=\"--\")\n",
    "plt.axvline(pos_5, color=\"red\", linestyle=\"--\")\n",
    "plt.legend([\"pos\", \"neg\"])\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "9df1f36a",
   "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",
    "\n",
    "# load the val metas\n",
    "import numpy as np\n",
    "from suno_utils.utils.text import read_jsonl\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "e6d2d456",
   "metadata": {},
   "outputs": [],
   "source": [
    "val_metas = read_jsonl(os.path.join(base_dir, f\"metas_tr_{version}.jsonl\"))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3845ba50",
   "metadata": {},
   "outputs": [],
   "source": [
    "# pick a random meta, then decode the positive and negative vae_latents\n",
    "idx = np.random.randint(0, len(val_metas))\n",
    "meta = val_metas[idx]\n",
    "# decode the positive and negative vae_latents\n",
    "pos_vae_latents = np.load(meta[\"pos_vae_latents_filepath\"])[\"vae_latents\"]\n",
    "neg_vae_latents = np.load(meta[\"neg_vae_latents_filepath\"])[\"vae_latents\"]\n",
    "\n",
    "# decode the positive and negative vae_latents\n",
    "pos_audio = codec_decode(pos_vae_latents)\n",
    "neg_audio = codec_decode(neg_vae_latents)\n",
    "\n",
    "print(idx)\n",
    "#print(\"agreement\", meta[\"agreement\"])\n",
    "#print(\"num_ratings\", meta[\"num_ratings\"]) \n",
    "#print(\"chosen_file_index\", meta[\"chosen_file_index\"])\n",
    "#print(\"unchosen_file_index\", meta[\"unchosen_file_index\"])\n",
    "\n",
    "# play the positive and negative audios\n",
    "print(\"pos_audio\")\n",
    "pos_audio.play()\n",
    "print(\"neg_audio\")\n",
    "neg_audio.play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "ee3e26e3",
   "metadata": {},
   "outputs": [],
   "source": [
    "def process_dir(base_dir, dirname):\n",
    "    results = {}\n",
    "    semantic_codes_filepath = os.path.join(base_dir, dirname, f\"{dirname}_semantic.npz\")\n",
    "    neg_vae_filepath = os.path.join(base_dir, dirname, f\"{dirname}_{model_name}_0_upsampled_vae.npz\")\n",
    "    pos_vae_filepath = os.path.join(base_dir, dirname, f\"{dirname}_{model_name}_1_upsampled_vae.npz\")\n",
    "\n",
    "    # decode the vae latents\n",
    "    pos_vae_latents = np.load(pos_vae_filepath)[\"vae_latents\"]\n",
    "    neg_vae_latents = np.load(neg_vae_filepath)[\"vae_latents\"]\n",
    " \n",
    "    # codec decode the vae latents\n",
    "    print(\"pos audio\")\n",
    "    pos_audio = codec_decode(pos_vae_latents).play()\n",
    "    print(\"neg audio\")\n",
    "    neg_audio = codec_decode(neg_vae_latents).play()\n",
    "\n",
    "\n",
    "    "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2eeacac7",
   "metadata": {},
   "outputs": [],
   "source": [
    "# sort the results by pos_mse\n",
    "sorted_results = sorted(results.items(), key=lambda x: x[1][1][\"total_delta\"], reverse=True)\n",
    "\n",
    "for i in range(3):\n",
    "    clip_id = sorted_results[i][0]\n",
    "    print(clip_id)\n",
    "    pos_total_delta = sorted_results[i][1][1][\"total_delta\"]\n",
    "    neg_total_delta = sorted_results[i][1][0][\"total_delta\"]\n",
    "    print(\"pos total delta\", pos_total_delta)\n",
    "    print(\"neg total delta\", neg_total_delta)\n",
    "    process_dir(base_dir, clip_id)\n",
    "\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "b47dae20",
   "metadata": {},
   "outputs": [],
   "source": [
    "import matplotlib.pyplot as plt\n",
    "\n",
    "# --- Plot deltas (as before) ---\n",
    "plot_data_deltas = []\n",
    "feature_titles = {\n",
    "    \"shimmer_score\": \"Shimmer score\",\n",
    "    \"lufs_db\": \"Loudness\",\n",
    "    \"ear_score\": \"Ear score\",\n",
    "    \"ear_v3_score\": \"Ear v3 score\",\n",
    "    \"hoot_cer\": \"Hoot cer\",\n",
    "    \"stereo_width\": \"Stereo width\",\n",
    "    \"total_delta\": \"Total frequency delta\",\n",
    "    \"mse\" : \"Spec. Reference MSE\"\n",
    "}\n",
    "feature_xlabels_deltas = {\n",
    "    \"shimmer_score\": \"Shimmer Score Delta\",\n",
    "    \"lufs_db\": \"Loudness Delta\",\n",
    "    \"ear_score\": \"Ear Score Delta\",\n",
    "    \"ear_v3_score\": \"Ear v3 Score Delta\",\n",
    "    \"hoot_cer\": \"Hoot CER Delta\",\n",
    "    \"stereo_width\": \"Stereo Width Delta\",\n",
    "    \"total_delta\": \"Total frequency delta\",\n",
    "    \"mse\" : \"Spec. Reference MSE\"\n",
    "}\n",
    "feature_xlabels_pos = {\n",
    "    \"shimmer_score\": \"Shimmer Score (pos)\",\n",
    "    \"lufs_db\": \"Loudness (pos)\",\n",
    "    \"ear_score\": \"Ear Score (pos)\",\n",
    "    \"ear_v3_score\": \"Ear v3 Score (pos)\",\n",
    "    \"hoot_cer\": \"Hoot CER (pos)\",\n",
    "    \"stereo_width\": \"Stereo Width (pos)\",\n",
    "    \"total_delta\": \"Total frequency delta (pos)\",\n",
    "    \"mse\" : \"Spec. Reference MSE (pos)\"\n",
    "}\n",
    "\n",
    "for f in features:\n",
    "    plot_data_deltas.append({\n",
    "        \"data\": all_deltas[f],\n",
    "        \"title\": f\"{feature_titles.get(f, f)} deltas\",\n",
    "        \"xlabel\": feature_xlabels_deltas.get(f, f\"{f} Delta\")\n",
    "    })\n",
    "\n",
    "n_plots = len(plot_data_deltas)\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_deltas):\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",
    "    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",
    "    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()\n",
    "\n",
    "# --- Plot pos values in a separate set of subplots ---\n",
    "plot_data_pos = []\n",
    "for f in features:\n",
    "    plot_data_pos.append({\n",
    "        \"data\": all_pos_features[f],\n",
    "        \"title\": f\"{feature_titles.get(f, f)} (pos)\",\n",
    "        \"xlabel\": feature_xlabels_pos.get(f, f\"{f} (pos)\")\n",
    "    })\n",
    "\n",
    "n_plots_pos = len(plot_data_pos)\n",
    "n_cols_pos = 2\n",
    "n_rows_pos = (n_plots_pos + 1) // n_cols_pos\n",
    "\n",
    "fig_pos, axs_pos = plt.subplots(n_rows_pos, n_cols_pos, figsize=(8, 2.5 * n_rows_pos))\n",
    "axs_pos = axs_pos.flatten()\n",
    "\n",
    "for i, plot_info in enumerate(plot_data_pos):\n",
    "    if len(plot_info[\"data\"]) == 0:\n",
    "        continue\n",
    "    data = np.array(plot_info[\"data\"])\n",
    "    ax = axs_pos[i]\n",
    "    ax.hist(data, bins=50, color='lightgreen', 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",
    "    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",
    "    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_pos)):\n",
    "    fig_pos.delaxes(axs_pos[j])\n",
    "\n",
    "plt.tight_layout()\n",
    "plt.show()"
   ]
  }
 ],
 "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
}
