{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# setup autoload\n",
    "%load_ext autoreload\n",
    "%autoreload 2"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import ast\n",
    "import os\n",
    "import shutil\n",
    "import sys\n",
    "from collections import defaultdict\n",
    "import json\n",
    "\n",
    "import numpy as np\n",
    "import pandas as pd\n",
    "from preference_data_preparation_diff import *\n",
    "from sklearn.model_selection import train_test_split\n",
    "from suno_utils.utils.s3 import download_s3_files\n",
    "from suno_utils.utils.text import read_json, read_jsonl, write_json, write_jsonl\n",
    "from tqdm import tqdm\n",
    "import matplotlib.pyplot as plt\n",
    "from suno_utils.audio import Audio\n",
    "from suno_analytics.preference_helper import get_preference_counts\n",
    "\n",
    "pd.set_option(\"display.max_rows\", 500)\n",
    "pd.set_option(\"display.max_columns\", 500)\n",
    "pd.set_option(\"display.width\", 1000)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "OUT_DATA_DIR = \"/app/suno/data/dpo/diff_v6_t3_comb/\"\n",
    "os.makedirs(OUT_DATA_DIR, exist_ok=True)\n",
    "NPZ_DIR = \"/app/suno/data/dpo/diff_v6\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "df = pd.read_pickle(\n",
    "    \"/home/tony/Data/Preference/up_v6/interesting_clips_up_u_6_20250320_full.pkl\"\n",
    ")  # , engine='python')\n",
    "print(\"Preference data shape\", df.shape)\n",
    "print(\"unique users\", df[\"user_id\"].nunique())"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(df[\"is_public\"].value_counts())\n",
    "# # remove public for now cause fucking users\n",
    "# df = df[~df[\"is_public\"]]\n",
    "# print(df[\"is_public\"].value_counts())"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "df[\"upsample_clip_id\"] = df[\"metadata\"].apply(lambda x: x.get(\"upsample_clip_id\", \"\"))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "all_converted_paths = os.listdir(NPZ_DIR)\n",
    "print(len(all_converted_paths))\n",
    "\n",
    "converted_paths = set(\n",
    "    [f.replace(\".npz\", \"\") for f in all_converted_paths if \"vae\" not in f]\n",
    ")\n",
    "print(len(converted_paths))\n",
    "vae_converted_paths = set(\n",
    "    [f.replace(\"_vae.npz\", \"\") for f in all_converted_paths if \"vae\" in f]\n",
    ")\n",
    "print(len(vae_converted_paths))\n",
    "\n",
    "print(\"pre-downloaded df\", df.shape)\n",
    "df[df[\"upsample_clip_id\"].isin(converted_paths)].shape\n",
    "df = df[df[\"upsample_clip_id\"].isin(converted_paths)].copy()\n",
    "print(\"downloaded df\", df.shape)\n",
    "df[df[\"s3_id\"].isin(vae_converted_paths)].shape\n",
    "df = df[df[\"s3_id\"].isin(vae_converted_paths)].copy()\n",
    "print(\"vae downloaded df\", df.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "df[\"is_up\"] = df[\"model_name\"].str.contains(\"up\")\n",
    "df[\"is_up\"].value_counts()"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# LET's do the data prep"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "## for 13b this is easy for now\n",
    "print(df.groupby([\"preference\"])[\"model_name\"].value_counts())\n",
    "print(df.shape)\n",
    "df = df[df[\"model_name\"].isin([\"chirp-v4-up-u-6\"])]\n",
    "print(df.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(df.shape)\n",
    "df = df[\n",
    "    df[\"request_id\"].isin(\n",
    "        df[\"request_id\"].value_counts().index[df[\"request_id\"].value_counts() == 2]\n",
    "    )\n",
    "]\n",
    "print(df.shape)\n",
    "print(df.groupby([\"preference\"])[\"model_name\"].value_counts())\n",
    "assert df.shape[0] == df[\"request_id\"].nunique() * 2"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "with open(\"/home/tony/Data/Preference/up_v6/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",
    "unpacked_pair_quality = {}\n",
    "for request_id, pairs_of_qualities in full_pair_quality.items():\n",
    "    for clip_id, pair_quality in pairs_of_qualities.items():\n",
    "        unpacked_pair_quality[clip_id] = pair_quality\n",
    "print(\"Total unpacked pair quality scores:\", len(unpacked_pair_quality))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "clip_diffs = []\n",
    "clip_ratios = []\n",
    "for request_id, pairs_of_qualities in full_pair_quality.items():\n",
    "    mean_neg_scores = []\n",
    "    mean_pos_scores = []\n",
    "    for i, (clip_id, pair_quality) in enumerate(pairs_of_qualities.items()):\n",
    "        if i % 2 == 0:\n",
    "            mean_neg_scores.append(np.mean(pair_quality[\"ear_v2_quality_scores\"] if pair_quality else 0))\n",
    "        if i % 2 == 1:\n",
    "            mean_pos_scores.append(np.mean(pair_quality[\"ear_v2_quality_scores\"] if pair_quality else 0))\n",
    "    ratios = [(pos - neg) / (pos + 0.0001) for pos, neg in zip(mean_pos_scores, mean_neg_scores)]\n",
    "    pos_diffs = [(pos - prev_pos) / (prev_pos + 0.0001) for prev_pos, pos in zip(mean_pos_scores, mean_pos_scores[1:])]\n",
    "    neg_diffs = [(neg - prev_neg) / (prev_neg + 0.0001) for prev_neg, neg in zip(mean_neg_scores, mean_neg_scores[1:])]\n",
    "    \n",
    "    for i in range(1, len(ratios)):\n",
    "        clip_ratios.append(ratios[i])\n",
    "        clip_diffs.append(pos_diffs[i - 1] - neg_diffs[i - 1])\n",
    "x = clip_ratios\n",
    "y = clip_diffs"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Create the 2D histogram (heatmap)\n",
    "plt.figure(figsize=(10, 8))\n",
    "\n",
    "# Create a 2D histogram\n",
    "bin_edges = np.linspace(-0.5, 0.5, 101)  # 30 bins from -1 to 1\n",
    "hist, x_edges, y_edges = np.histogram2d(\n",
    "    x, y, \n",
    "    bins=[bin_edges, bin_edges],  # Same bins for both x and y\n",
    "    range=[[-0.5, 0.5], [-0.5, 0.5]]      # Ensure range is from -1 to 1 for both axes\n",
    ")\n",
    "\n",
    "# Create a heatmap using pcolormesh for better control\n",
    "X, Y = np.meshgrid(x_edges[:-1], y_edges[:-1])\n",
    "plt.pcolormesh(X, Y, hist.T, cmap='viridis', shading='auto')\n",
    "\n",
    "# Add a color bar\n",
    "cbar = plt.colorbar()\n",
    "cbar.set_label('Counts', rotation=270, labelpad=20, fontsize=12)\n",
    "\n",
    "# Add labels and title\n",
    "plt.xlabel('Clip quality diff ratios', fontsize=12)\n",
    "plt.ylabel('Clip quality diff ratio difference with prev', fontsize=12)\n",
    "plt.title('2D Histogram (Heatmap) of Correlated Data', fontsize=14)\n",
    "\n",
    "# Show the plot\n",
    "plt.tight_layout()\n",
    "plt.savefig('2d_histogram.png', dpi=300)  # Save to file (optional)\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "_ = plt.hist(clip_diffs, bins=np.linspace(-1, 1, 100))\n",
    "for percentage in [0.1, 0.2, 0.5, 0.8, 0.9]:\n",
    "    print(percentage, \"--->\", np.quantile(sorted(clip_diffs), percentage))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "_ = plt.hist(clip_ratios, bins=np.linspace(-1, 1, 100))\n",
    "for percentage in [0.1, 0.2, 0.5, 0.8, 0.9]:\n",
    "    print(percentage, \"--->\", np.quantile(sorted(clip_ratios), percentage))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def get_audio_quality_measures(s3_id):\n",
    "    audio_quality = unpacked_pair_quality.get(s3_id, [])\n",
    "    if not audio_quality:\n",
    "        return [None for _ in range(11)]\n",
    "    return [\n",
    "        np.mean(audio_quality[\"ear_v2_quality_scores\"]), # float(audio_quality[\"ear_v2_quality_scores\"]),\n",
    "        float(audio_quality[\"shimmer_score\"]),\n",
    "        float(audio_quality[\"loudness_factor\"]),\n",
    "        audio_quality[\"spectral_character\"],\n",
    "        float(audio_quality[\"spectral_centroid\"]),\n",
    "        float(audio_quality[\"bass_ratio\"]),\n",
    "        float(audio_quality[\"mid_ratio\"]),\n",
    "        float(audio_quality[\"high_ratio\"]),\n",
    "        float(audio_quality[\"stereo_width\"]),\n",
    "        int(audio_quality[\"total_clips\"]),\n",
    "        float(audio_quality[\"clips_per_second\"]),\n",
    "    ]\n",
    "\n",
    "\n",
    "df[\n",
    "    [\n",
    "        \"pair_quality\",\n",
    "        \"total_shimmer_score\",\n",
    "        \"loudness_factor\",\n",
    "        \"spectral_character\",\n",
    "        \"spectral_centroid\",\n",
    "        \"bass_ratio\",\n",
    "        \"mid_ratio\",\n",
    "        \"high_ratio\",\n",
    "        \"stereo_width\",\n",
    "        \"total_clips\",\n",
    "        \"clips_per_second\",\n",
    "    ]\n",
    "] = pd.DataFrame(df[\"s3_id\"].apply(get_audio_quality_measures).tolist(), index=df.index)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(df.shape)\n",
    "df = df.dropna(subset=[\n",
    "        \"pair_quality\",\n",
    "        \"total_shimmer_score\",\n",
    "        \"loudness_factor\",\n",
    "        \"spectral_character\",\n",
    "        \"spectral_centroid\",\n",
    "        \"bass_ratio\",\n",
    "        \"mid_ratio\",\n",
    "        \"high_ratio\",\n",
    "        \"stereo_width\",\n",
    "        \"total_clips\",\n",
    "        \"clips_per_second\",\n",
    "    ])\n",
    "print(df.shape)\n",
    "df = df[\n",
    "    df[\"request_id\"].isin(\n",
    "        df[\"request_id\"].value_counts().index[df[\"request_id\"].value_counts() == 2]\n",
    "    )\n",
    "]\n",
    "print(df.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Let's use the old selection for now -- for quality assurance\n",
    "# expand the metadata columns -- this takes forever...~ 6 mins\n",
    "# test_slice = df[\"metadata\"].apply(lambda x: ast.literal_eval(str(x)))\n",
    "# test_slice = df[\"metadata\"]  # .apply(lambda x: custom_parse(x))\n",
    "# test_slice_series = test_slice.apply(pd.Series)\n",
    "# df = pd.concat([df, test_slice_series], axis=1, join=\"inner\")\n",
    "print(\"unique_requests\", df[\"request_id\"].nunique())"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "df = df.loc[:, ~df.columns.duplicated()].copy()\n",
    "# get the original duration of the clips, if they are concacted\n",
    "df[\"original_duration_s\"] = df[\"total_start_s\"] + df[\"duration\"]\n",
    "# classify the continue at behavoirs by the duration choice\n",
    "audio_prompt_id_to_continue_at = {}\n",
    "for _, row in df[~df[\"continued_parent\"].isna()].iterrows():\n",
    "    audio_prompt_id = row[\"continued_parent\"]\n",
    "    if audio_prompt_id not in audio_prompt_id_to_continue_at:\n",
    "        audio_prompt_id_to_continue_at[audio_prompt_id] = row[\"continue_at\"]\n",
    "    else:\n",
    "        # pick the max\n",
    "        audio_prompt_id = max(\n",
    "            audio_prompt_id_to_continue_at[audio_prompt_id], row[\"continue_at\"]\n",
    "        )\n",
    "print(len(audio_prompt_id_to_continue_at))\n",
    "df[\"has_continue_and_start_continue_at\"] = df[\"s3_id\"].apply(\n",
    "    lambda x: audio_prompt_id_to_continue_at.get(x)\n",
    ")\n",
    "# we want continue at to be at most of the clip...\n",
    "df[\"good_continue_at\"] = (\n",
    "    (df[\"has_continue_and_start_continue_at\"] / df[\"duration\"]) > 0.9\n",
    ") | df[\"has_continue_and_start_continue_at\"].isna()\n",
    "print(df[\"good_continue_at\"].value_counts())\n",
    "\n",
    "\n",
    "print(\n",
    "    \"\\n Check some basics... \\n\",\n",
    "    df[\"preference\"].value_counts(),\n",
    "    df[\"model_name\"].value_counts(),\n",
    "    df.groupby([\"preference\"])[\"model_name\"].value_counts(),\n",
    ")\n",
    "\n",
    "df = df.sort_values(by=[\"request_id\", \"preference\"])\n",
    "df[\"duration_rel_diff\"] = df[\"duration\"].diff()\n",
    "df[\"play_rel_diff\"] = df[\"reaction_play_count\"].diff()\n",
    "print(df[\"task\"].value_counts())"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "df = df.sort_values(by=[\"request_id\", \"preference\", \"diff_preference\"])\n",
    "df[\"pos_diff_preference\"] = df[\"diff_preference\"].diff()\n",
    "df[df[\"preference\"]][\"pos_diff_preference\"].value_counts()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "lookup_percentiles = [5, 10, 20, 50, 80, 90, 95]\n",
    "df[\"shimmer_score_diff\"] = df[\"total_shimmer_score\"].diff()\n",
    "percentiles = np.percentile(\n",
    "    df[df[\"preference\"]][\"shimmer_score_diff\"].dropna(), lookup_percentiles\n",
    ")\n",
    "plt.hist(\n",
    "    df[df[\"preference\"]][\"shimmer_score_diff\"],\n",
    "    label=f\"pos, mean: {np.mean(df[df['preference']]['shimmer_score_diff']):.2f}\",\n",
    "    bins=np.linspace(-1, 1, 100),\n",
    "    alpha=0.5,\n",
    ")\n",
    "textstr = \"\\n\".join(\n",
    "    [\n",
    "        f\"{lookup_percentiles[i]}th: {percentile:.2f}\"\n",
    "        for i, percentile in enumerate(percentiles)\n",
    "    ]\n",
    ")\n",
    "plt.gcf().text(\n",
    "    0.15,\n",
    "    0.98,\n",
    "    textstr,\n",
    "    fontsize=10,\n",
    "    verticalalignment=\"top\",\n",
    "    horizontalalignment=\"left\",\n",
    "    bbox=dict(facecolor=\"white\", alpha=0.5),\n",
    ")\n",
    "for percentile in percentiles:\n",
    "    plt.axvline(x=percentile, color=\"r\", linestyle=\"dashed\", linewidth=1)\n",
    "plt.title(\n",
    "    f\"Shimmer score difference --> {lookup_percentiles[-1]}th, {percentiles[-1]:.2f}\"\n",
    ")\n",
    "# plt.yscale(\"log\")\n",
    "plt.legend()\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "lookup_percentiles = [5, 10, 20, 50, 80, 90, 95]\n",
    "percentiles = np.percentile(\n",
    "    df[df[\"preference\"]][\"pair_quality\"].dropna(), lookup_percentiles\n",
    ")\n",
    "plt.hist(\n",
    "    df[df[\"preference\"]][\"pair_quality\"],\n",
    "    label=f\"pos, mean: {np.mean(df[df['preference']]['pair_quality']):.2f}\",\n",
    "    bins=np.linspace(10, 30, 100),\n",
    "    alpha=0.5,\n",
    ")\n",
    "plt.hist(\n",
    "    df[~df[\"preference\"]][\"pair_quality\"],\n",
    "    label=f\"neg, mean: {np.mean(df[df['preference']]['pair_quality']):.2f}\",\n",
    "    bins=np.linspace(10, 30, 100),\n",
    "    alpha=0.5,\n",
    ")\n",
    "textstr = \"\\n\".join(\n",
    "    [\n",
    "        f\"{lookup_percentiles[i]}th: {percentile:.2f}\"\n",
    "        for i, percentile in enumerate(percentiles)\n",
    "    ]\n",
    ")\n",
    "plt.gcf().text(\n",
    "    0.15,\n",
    "    0.98,\n",
    "    textstr,\n",
    "    fontsize=10,\n",
    "    verticalalignment=\"top\",\n",
    "    horizontalalignment=\"left\",\n",
    "    bbox=dict(facecolor=\"white\", alpha=0.5),\n",
    ")\n",
    "for percentile in percentiles:\n",
    "    plt.axvline(x=percentile, color=\"r\", linestyle=\"dashed\", linewidth=1)\n",
    "plt.title(f\"Pair quality --> {lookup_percentiles[0]}th, {percentiles[0]:.2f}\")\n",
    "# plt.yscale(\"log\")\n",
    "plt.legend()\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "df[\"pair_quality_diff\"] = df[\"pair_quality\"].diff() / df[\"pair_quality\"]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "lookup_percentiles = [5, 10, 20, 50, 80, 90, 95]\n",
    "percentiles = np.percentile(\n",
    "    df[df[\"preference\"]][\"pair_quality_diff\"].dropna(), lookup_percentiles\n",
    ")\n",
    "plt.hist(\n",
    "    df[df[\"preference\"]][\"pair_quality_diff\"],\n",
    "    label=f\"pos, mean: {np.mean(df[df['preference']]['pair_quality_diff']):.2f}\",\n",
    "    bins=np.linspace(-0.5, 0.5, 100),\n",
    "    alpha=0.5,\n",
    ")\n",
    "textstr = \"\\n\".join(\n",
    "    [\n",
    "        f\"{lookup_percentiles[i]}th: {percentile:.2f}\"\n",
    "        for i, percentile in enumerate(percentiles)\n",
    "    ]\n",
    ")\n",
    "plt.gcf().text(\n",
    "    0.15,\n",
    "    0.98,\n",
    "    textstr,\n",
    "    fontsize=10,\n",
    "    verticalalignment=\"top\",\n",
    "    horizontalalignment=\"left\",\n",
    "    bbox=dict(facecolor=\"white\", alpha=0.5),\n",
    ")\n",
    "for percentile in percentiles:\n",
    "    plt.axvline(x=percentile, color=\"r\", linestyle=\"dashed\", linewidth=1)\n",
    "plt.title(f\"Pair quality diff --> {lookup_percentiles[0]}th, {percentiles[0]:.2f}\")\n",
    "# plt.yscale(\"log\")\n",
    "plt.legend()\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "subset_requests = df[(df[\"pair_quality_diff\"] > 0.2) & (df[\"preference\"])][\"request_id\"].unique()\n",
    "df[df[\"request_id\"].isin(subset_requests)][[\"id\", \"request_id\", \"pair_quality\", \"preference\", \"total_shimmer_score\"]]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# CS said -1 means stero to mono; 1 means mono to stereo\n",
    "# only cut off the left side\n",
    "df[\"stereo_width_diff\"] = df[\"stereo_width\"].diff()\n",
    "lookup_percentiles = [5, 10, 20, 50, 80, 90, 95]\n",
    "percentiles = np.percentile(\n",
    "    df[df[\"preference\"]][\"stereo_width_diff\"].dropna(), lookup_percentiles\n",
    ")\n",
    "plt.hist(\n",
    "    df[df[\"preference\"]][\"stereo_width_diff\"],\n",
    "    label=f\"pos, mean: {np.mean(df[df['preference']]['stereo_width_diff']):.2f}\",\n",
    "    bins=np.linspace(-1, 1, 400),\n",
    "    alpha=0.5,\n",
    ")\n",
    "textstr = \"\\n\".join(\n",
    "    [\n",
    "        f\"{lookup_percentiles[i]}th: {percentile:.2f}\"\n",
    "        for i, percentile in enumerate(percentiles)\n",
    "    ]\n",
    ")\n",
    "plt.gcf().text(\n",
    "    0.15,\n",
    "    0.98,\n",
    "    textstr,\n",
    "    fontsize=10,\n",
    "    verticalalignment=\"top\",\n",
    "    horizontalalignment=\"left\",\n",
    "    bbox=dict(facecolor=\"white\", alpha=0.5),\n",
    ")\n",
    "for percentile in percentiles:\n",
    "    plt.axvline(x=percentile, color=\"r\", linestyle=\"dashed\", linewidth=1)\n",
    "plt.title(\n",
    "    f\"Stereo Width Difference --> {lookup_percentiles[0]}th, {percentiles[0]:.2f}\"\n",
    ")\n",
    "# plt.yscale(\"log\")\n",
    "plt.legend()\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# CS said no tails is good\n",
    "lookup_percentiles = [5, 10, 20, 50, 80, 90, 95]\n",
    "percentiles = np.percentile(\n",
    "    df[df[\"preference\"]][\"spectral_centroid\"].dropna(), lookup_percentiles\n",
    ")\n",
    "percentiles = np.percentile(\n",
    "    df[~df[\"preference\"]][\"spectral_centroid\"].dropna(), lookup_percentiles\n",
    ")\n",
    "plt.hist(\n",
    "    df[df[\"preference\"]][\"spectral_centroid\"],\n",
    "    label=f\"pos, mean: {np.mean(df[df['preference']]['spectral_centroid']):.2f}\",\n",
    "    bins=np.linspace(0, 8000, 400),\n",
    "    alpha=0.5,\n",
    ")\n",
    "plt.hist(\n",
    "    df[~df[\"preference\"]][\"spectral_centroid\"],\n",
    "    label=f\"neg, mean: {np.mean(df[~df['preference']]['spectral_centroid']):.2f}\",\n",
    "    bins=np.linspace(0, 8000, 400),\n",
    "    alpha=0.5,\n",
    ")\n",
    "textstr = \"\\n\".join(\n",
    "    [\n",
    "        f\"{lookup_percentiles[i]}th: {percentile:.2f}\"\n",
    "        for i, percentile in enumerate(percentiles)\n",
    "    ]\n",
    ")\n",
    "plt.gcf().text(\n",
    "    0.15,\n",
    "    0.98,\n",
    "    textstr,\n",
    "    fontsize=10,\n",
    "    verticalalignment=\"top\",\n",
    "    horizontalalignment=\"left\",\n",
    "    bbox=dict(facecolor=\"white\", alpha=0.5),\n",
    ")\n",
    "for percentile in percentiles:\n",
    "    plt.axvline(x=percentile, color=\"r\", linestyle=\"dashed\", linewidth=1)\n",
    "plt.title(\"Spectral Centroid\")\n",
    "# plt.yscale(\"log\")\n",
    "plt.legend()\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# CS said no tails is good\n",
    "# take the relative centroid diff\n",
    "df[\"spectral_centroid_diff\"] = df[\"spectral_centroid\"].diff() / df[\"spectral_centroid\"]\n",
    "lookup_percentiles = [5, 10, 20, 50, 80, 90, 95]\n",
    "percentiles = np.percentile(\n",
    "    df[df[\"preference\"]][\"spectral_centroid_diff\"].dropna(), lookup_percentiles\n",
    ")\n",
    "plt.hist(\n",
    "    df[df[\"preference\"]][\"spectral_centroid_diff\"],\n",
    "    label=f\"pos, mean: {np.mean(df[df['preference']]['spectral_centroid_diff']):.2f}\",\n",
    "    bins=np.linspace(-1, 1, 400),\n",
    "    alpha=0.5,\n",
    ")\n",
    "textstr = \"\\n\".join(\n",
    "    [\n",
    "        f\"{lookup_percentiles[i]}th: {percentile:.2f}\"\n",
    "        for i, percentile in enumerate(percentiles)\n",
    "    ]\n",
    ")\n",
    "plt.gcf().text(\n",
    "    0.15,\n",
    "    0.98,\n",
    "    textstr,\n",
    "    fontsize=10,\n",
    "    verticalalignment=\"top\",\n",
    "    horizontalalignment=\"left\",\n",
    "    bbox=dict(facecolor=\"white\", alpha=0.5),\n",
    ")\n",
    "for percentile in percentiles:\n",
    "    plt.axvline(x=percentile, color=\"r\", linestyle=\"dashed\", linewidth=1)\n",
    "plt.title(\n",
    "    f\"Spectral Centroid Difference ratio --> {lookup_percentiles[-1]}th, {percentiles[-1]:.2f}\"\n",
    ")\n",
    "# plt.yscale(\"log\")\n",
    "plt.legend()\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "df = df.loc[:, ~df.columns.duplicated()].copy()\n",
    "# get the original duration of the clips, if they are concacted\n",
    "df[\"original_duration_s\"] = df[\"total_start_s\"] + df[\"duration\"]\n",
    "# classify the continue at behavoirs by the duration choice\n",
    "audio_prompt_id_to_continue_at = {}\n",
    "for _, row in df[~df[\"continued_parent\"].isna()].iterrows():\n",
    "    audio_prompt_id = row[\"continued_parent\"]\n",
    "    if audio_prompt_id not in audio_prompt_id_to_continue_at:\n",
    "        audio_prompt_id_to_continue_at[audio_prompt_id] = row[\"continue_at\"]\n",
    "    else:\n",
    "        # pick the max\n",
    "        audio_prompt_id = max(\n",
    "            audio_prompt_id_to_continue_at[audio_prompt_id], row[\"continue_at\"]\n",
    "        )\n",
    "print(len(audio_prompt_id_to_continue_at))\n",
    "df[\"has_continue_and_start_continue_at\"] = df[\"s3_id\"].apply(\n",
    "    lambda x: audio_prompt_id_to_continue_at.get(x)\n",
    ")\n",
    "# we want continue at to be at most of the clip...\n",
    "df[\"good_continue_at\"] = (\n",
    "    (df[\"has_continue_and_start_continue_at\"] / df[\"duration\"]) > 0.9\n",
    ") | df[\"has_continue_and_start_continue_at\"].isna()\n",
    "print(df[\"good_continue_at\"].value_counts())\n",
    "\n",
    "\n",
    "print(\n",
    "    \"\\n Check some basics... \\n\",\n",
    "    df[\"preference\"].value_counts(),\n",
    "    df[\"model_name\"].value_counts(),\n",
    "    df.groupby([\"preference\"])[\"model_name\"].value_counts(),\n",
    ")\n",
    "\n",
    "df = df.sort_values(by=[\"request_id\", \"preference\"])\n",
    "df[\"duration_rel_diff\"] = df[\"duration\"].diff()\n",
    "df[\"play_rel_diff\"] = df[\"reaction_play_count\"].diff()\n",
    "print(df[\"task\"].value_counts())"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "plt.hist(\n",
    "    df[df[\"preference\"]][\"total_shimmer_score\"],\n",
    "    label=f\"pos, mean: {np.mean(df[df['preference']]['total_shimmer_score']):.2f}\",\n",
    "    bins=np.linspace(0, 10, 100),\n",
    "    alpha=0.5,\n",
    ")\n",
    "plt.hist(\n",
    "    df[~df[\"preference\"]][\"total_shimmer_score\"],\n",
    "    label=f\"neg, mean: {np.mean(df[~df['preference']]['total_shimmer_score']):.2f}\",\n",
    "    bins=np.linspace(0, 10, 100),\n",
    "    alpha=0.5,\n",
    ")\n",
    "percentiles = np.percentile(df[df[\"preference\"]][\"total_shimmer_score\"].dropna(), [50, 75, 90])\n",
    "for percentile in percentiles:\n",
    "    # print(percentile)\n",
    "    plt.axvline(x=percentile, color=\"r\", linestyle=\"dashed\", linewidth=1)\n",
    "# plt.yscale(\"log\")\n",
    "plt.title(f\"Shimmer score\")\n",
    "plt.legend()\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "normal_pos_play_count = 3\n",
    "# this is lower, cause a concat is probably already ensuring that it is good\n",
    "concat_pos_play_count = 2\n",
    "# this is a filter on the concated clip\n",
    "concat_total_play_count = 3\n",
    "\n",
    "neg_filter_selection_mask = (\n",
    "    (~df[\"preference\"])  # get basics aligned\n",
    "    & (df[\"reaction_play_count\"] >= 1)  # has to be played once\n",
    "    # & (df[\"play_count\"] <= 3)  # if it is actually bad, shouldn't be listened often\n",
    "    & (df[\"duration\"] >= 30)  # can't be too short, otherwise it is obvious\n",
    "    # & (df[\"duration\"] <= 60)  # can't be badly long\n",
    "    & (df[\"has_continue_and_start_continue_at\"].isna())  # won't have any continues\n",
    "    & (df[\"norm_play_frac\"] <= 2.1)\n",
    "    # & (df[\"sum_total_play_duration_5\"] >= 31)\n",
    "    # & (df[\"dislike_count\"] >= 1) # this is kinda strict\n",
    "    #     & (\n",
    "    #         (df_slice[\"is_in_playlist\"] == False)\n",
    "    #         & (df_slice[\"concat_in_playlist\"] == False)\n",
    "    #     )  # can't be part of a playlist -- otherwise there are some like signal in it?\n",
    ")\n",
    "pos_filter_selectin_mask = (\n",
    "    (df[\"preference\"])  # get basics aligned\n",
    "    & (\n",
    "        df[\"good_continue_at\"]\n",
    "    )  # if continue, needs to continue off a certain percentage\n",
    "    & (df[\"reaction_play_count\"] >= 1)\n",
    "    & (df[\"play_rel_diff\"] >= 0)  # this is more like quality assurance\n",
    "    & (df[\"duration\"] >= 30)  # can't be too short, otherwise it is obvious\n",
    "    & (df[\"dislike_count\"] == 0)  # can't have dislikes\n",
    "    & (df[\"flag_count\"] == 0)  # can't have issues\n",
    "    & (\n",
    "        (\n",
    "            (df[\"part_of_concat\"])\n",
    "            & (df[\"reaction_play_count\"] >= concat_pos_play_count)\n",
    "            & (df[\"concat_play_counts\"] >= concat_total_play_count)\n",
    "        )\n",
    "        | (\n",
    "            (~df[\"part_of_concat\"])\n",
    "            & (df[\"reaction_play_count\"] >= normal_pos_play_count)\n",
    "            & (df[\"norm_play_frac\"] >= 2.1)  # this is a bit of a luxury cut...\n",
    "            & (df[\"sum_total_play_duration_5\"] >= 31)\n",
    "        )\n",
    "    )\n",
    "    # & (df[\"norm_play_frac\"] >= 1.9)\n",
    "    # & (df[\"user_n_clips\"] >= 100)  # user needs to have genereated at least 20\n",
    "    # & (df[\"duration_rel_diff\"] < 10) # positive isn't just longer\n",
    "    # & ((df[\"task\"] == \"\") | (df[\"task\"] == \"extend\"))\n",
    "    & (\n",
    "        (df[\"upvote_count\"] >= 1)\n",
    "        | (df[\"reaction_play_count\"] >= 5)\n",
    "        | (df[\"concat_play_counts\"] >= 5)\n",
    "    )\n",
    "    # & (df[\"pos_diff_preference\"] == 2)\n",
    "    # & ((0 < df[\"similarity\"]) &  (df[\"similarity\"] <= 0.99))\n",
    "    # & (\n",
    "    #     (df[\"cer_diff_preference\"] < 0.25) & (df[\"cer\"] < 0.8)\n",
    "    # )  # cut on hoot cer difference and abs cer\n",
    "    # & (df[\"pair_quality\"] > 0.31)  # bottom 5%\n",
    "    # & ((df[\"total_shimmer_score\"] < 1) | (df[\"shimmer_score_diff\"] < 0.4))\n",
    "    # & (df[\"stereo_width_diff\"] > -0.2)  # cut off bottom 5%\n",
    "    # & (df[\"spectral_centroid_diff\"] < 0.25)  # crop off the top 5%\n",
    ")\n",
    "print(\n",
    "    \"negative\",\n",
    "    sum(neg_filter_selection_mask),\n",
    "    \"positive\",\n",
    "    sum(pos_filter_selectin_mask),\n",
    ")\n",
    "\n",
    "neg_filter_requests = df[neg_filter_selection_mask][\"request_id\"].unique()\n",
    "pos_filter_requests = df[pos_filter_selectin_mask][\"request_id\"].unique()\n",
    "# looking for very strong signal here:\n",
    "# listen to the positive/negative more than once\n",
    "# disliked one of the clips\n",
    "unique_requests = set(pos_filter_requests).intersection(neg_filter_requests)\n",
    "print(\n",
    "    \"total pair requests\",\n",
    "    df[\"request_id\"].nunique(),\n",
    "    \"selected pair requests\",\n",
    "    len(unique_requests),\n",
    "    f\"frac {len(unique_requests) / df['request_id'].nunique():.3f}\",\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "df_slice = df[df[\"request_id\"].isin(set(unique_requests))].copy()\n",
    "print(\n",
    "    f\"{os.path.basename(OUT_DATA_DIR)} requests\",\n",
    "    df_slice[\"request_id\"].nunique(),\n",
    "    \"clips\",\n",
    "    df_slice.shape[0],\n",
    "    f\"total khrs {sum(df_slice['duration'] / 3600 / 1000):.3f};\",\n",
    "    f\"N gpus for 1000 iters {df_slice.shape[0] / 8 / 2 / 1000:.3f};\",\n",
    "    f\"4 gpus for x iters {df_slice.shape[0] / 8 / 2 / 4:.3f};\",\n",
    "    f\"n unique users {df_slice['user_id'].nunique()}\",\n",
    "    f\"n pro users {df_slice[df_slice['is_pro_user']]['user_id'].nunique()}\",\n",
    ")\n",
    "# up t7 requests 37735 clips 75470 total khrs 3.964; N gpus for 1000 iters 4.717; 4 gpus for x iters 1179.219; n unique users 19093 n pro users 16217\n",
    "# up t17 requests 49262 clips 98524 total khrs 5.187; N gpus for 1000 iters 6.158; 4 gpus for x iters 1539.438; n unique users 24093 n pro users 20326\n",
    "# up v2 t1 requests 6008 clips 12016 total khrs 0.642; N gpus for 1000 iters 0.751; 4 gpus for x iters 187.750; n unique users 3846 n pro users 3645\n",
    "# up v2 t2 requests 10772 clips 21544 total khrs 1.154; N gpus for 1000 iters 1.347; 4 gpus for x iters 336.625; n unique users 6527 n pro users 6027\n",
    "# up v3 t10 requests 27102 clips 54204 total khrs 2.938; N gpus for 1000 iters 3.388; 4 gpus for x iters 846.938; n unique users 13932 n pro users 12187\n",
    "# up v4 t1  requests 3201 clips 6402 total khrs 0.343; N gpus for 1000 iters 0.400; 4 gpus for x iters 100.031; n unique users 2363 n pro users 2321\n",
    "# up v4 t2  requests 12818 clips 25636 total khrs 1.371; N gpus for 1000 iters 1.602; 4 gpus for x iters 400.562; n unique users 7753 n pro users 7524\n",
    "# up v4 t3  requests 15332 clips 30664 total khrs 1.637; N gpus for 1000 iters 1.917; 4 gpus for x iters 479.125; n unique users 8960 n pro users 8656\n",
    "# up v4 t4  requests 18368 clips 36736 total khrs 1.962; N gpus for 1000 iters 2.296; 4 gpus for x iters 574.000; n unique users 10427 n pro users 10049\n",
    "# up v4 t5  requests 22037 clips 44074 total khrs 2.347; N gpus for 1000 iters 2.755; 4 gpus for x iters 688.656; n unique users 12121 n pro users 11551\n",
    "# up v4 t6  requests 27374 clips 54748 total khrs 2.915; N gpus for 1000 iters 3.422; 4 gpus for x iters 855.438; n unique users 14463 n pro users 13678\n",
    "# up v4 t7  requests 31797 clips 63594 total khrs 3.384; N gpus for 1000 iters 3.975; 4 gpus for x iters 993.656; n unique users 16398 n pro users 15325\n",
    "# up v5 t1  requests 17035 clips 34070 total khrs 1.788; N gpus for 1000 iters 2.129; 4 gpus for x iters 532.344; n unique users 10858 n pro users 8065\n",
    "# up v5 t2  requests 20030 clips 40060 total khrs 2.108; N gpus for 1000 iters 2.504; 4 gpus for x iters 625.938; n unique users 12395 n pro users 9156\n",
    "# up v6 t1  requests 17941 clips 35882 total khrs 1.889; N gpus for 1000 iters 2.243; 4 gpus for x iters 560.656; n unique users 11471 n pro users 8374\n",
    "# up v6 t2  requests 23563 clips 47126 total khrs 2.473; N gpus for 1000 iters 2.945; 4 gpus for x iters 736.344; n unique users 14550 n pro users 10353\n",
    "# up v6 t3  requests 34761 clips 69522 total khrs 3.647; N gpus for 1000 iters 4.345; 4 gpus for x iters 1086.281; n unique users 20264 n pro users 14000"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "test_mask = (df_slice[\"preference\"]) & (\n",
    "    (df_slice[\"is_in_playlist\"]) | (df_slice[\"concat_in_playlist\"])\n",
    ")\n",
    "print(\"positive in playlist\", df_slice[test_mask].shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# save positive ids\n",
    "# positive_preference_ids = df_slice[df_slice[\"preference\"] == False][\"s3_id\"].to_json(orient='values')\n",
    "# with open('/home/tony/Data/Preference/7b_v2/7v_v20_full_recut_id_negative.json', 'w') as file:\n",
    "#     file.write(positive_preference_ids)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# def modify_model_name(model_name, metadata):\n",
    "#     if (\n",
    "#         model_name.startswith(\"chirp-v3p5-engine-t\")\n",
    "#         or model_name.startswith(\"chirp-v3p5-engine-s\")\n",
    "#         or model_name.startswith(\"chirp-v4\")\n",
    "#         or model_name.startswith(\"chirp-v3p5-h-s-31\")\n",
    "#     ):\n",
    "#         if \"param_experiment\" in metadata:\n",
    "#             exp = metadata.get(\"param_experiment\", \"\")\n",
    "#             if exp:\n",
    "#                 return f\"{model_name}_{exp}\"\n",
    "#     return model_name\n",
    "\n",
    "# metrics_check_df_slice = df_slice.copy()\n",
    "# metrics_check_df_slice[\"model_name\"] = metrics_check_df_slice.apply(\n",
    "#     lambda row: modify_model_name(row[\"model_name\"], row[\"metadata\"]), axis=1\n",
    "# )\n",
    "# get_preference_counts(metrics_check_df_slice)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(df_slice[\"source\"].value_counts())"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# df_slice.to_csv(\"/home/tony/Data/Preference/13b_v0/interesting_clips_v3p5_s_8_20240828_slice.csv\")\n",
    "BREAK"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Need to kick out the ones has gpt prompt -- these are pairs with different text inputs"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "final_filtered_requests = df_slice[\"request_id\"].unique()\n",
    "print(len(final_filtered_requests))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "train_requests, val_requests = train_test_split(\n",
    "    sorted(list(final_filtered_requests)), test_size=0.01, random_state=42\n",
    ")\n",
    "print(len(train_requests), len(val_requests))\n",
    "\n",
    "train_df = df_slice[df_slice[\"request_id\"].isin(set(train_requests))].copy()\n",
    "val_df = df_slice[df_slice[\"request_id\"].isin(set(val_requests))].copy()\n",
    "train_df = train_df.sort_values(by=[\"request_id\", \"preference\"])\n",
    "train_df = train_df  # .reset_index()\n",
    "val_df = val_df.sort_values(by=[\"request_id\", \"preference\"])\n",
    "val_df = val_df  # .reset_index()\n",
    "\n",
    "print(train_df.shape, val_df.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# BREAK"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Actually make"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# val_df[[\"request_id\", \"metadata\", \"updated_at\", \"user_id\", \"preference\"]].head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "total_duration = 0\n",
    "for i, row in tqdm(train_df.iterrows(), total=len(train_df)):\n",
    "    # we need to alternate between preference: neg, pos\n",
    "    # print(i, row)\n",
    "    try:\n",
    "        assert row[\"preference\"] == (i % 2 == 1)\n",
    "    except:\n",
    "        print(i, row)\n",
    "    total_duration += row[\"duration\"]\n",
    "print(\n",
    "    f\"{round(total_duration / 60 / 60):,} hours of {train_df.shape[0]} clips, {train_df.shape[0] / 8 / 2 / 1000} nodes, {train_df.shape[0] / 8 / 2 / 4} steps\"\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "make_dataset(\n",
    "    val_df,\n",
    "    OUT_DATA_DIR,\n",
    "    is_val=True,\n",
    "    npz_dir=NPZ_DIR,\n",
    "    do_extend_chunks=True,\n",
    "    clip_id_to_quality_scores=unpacked_pair_quality,\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "make_dataset(\n",
    "    train_df,\n",
    "    OUT_DATA_DIR,\n",
    "    is_val=False,\n",
    "    npz_dir=NPZ_DIR,\n",
    "    do_extend_chunks=True,\n",
    "    clip_id_to_quality_scores=unpacked_pair_quality,\n",
    ")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-01-29T19:46:47.549860Z",
     "start_time": "2024-01-29T19:46:47.548015Z"
    }
   },
   "source": [
    "# Validation"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# verify\n",
    "metas_val = read_jsonl(os.path.join(OUT_DATA_DIR, \"metas_val.jsonl\"))\n",
    "print(len(metas_val) / 2)\n",
    "mm_semantic_val = np.memmap(\n",
    "    os.path.join(OUT_DATA_DIR, \"data_semantic_val.bin\"), dtype=np.uint16, mode=\"r\"\n",
    ")\n",
    "mm_vae_val = np.memmap(\n",
    "    os.path.join(OUT_DATA_DIR, \"data_vae_val.bin\"), dtype=np.float16, mode=\"r\"\n",
    ")\n",
    "\n",
    "\n",
    "mm_vae_val = mm_vae_val.reshape(-1, VAE_MEMMAP_SIZE, VAE_DIM)\n",
    "print(mm_vae_val.shape)\n",
    "\n",
    "mm_semantic_val = mm_semantic_val.reshape(-1, SEMANTIC_MEMMAP_SIZE)\n",
    "print(mm_semantic_val.shape)\n",
    "\n",
    "assert len(metas_val) == mm_vae_val.shape[0] == mm_semantic_val.shape[0]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# # load codec for decoding\n",
    "# from suno_utils.tasks.dac_vae_100hz_peaq import (  # NOTE: works for 25hz as well\n",
    "#     preload_models as preload_codec_models,\n",
    "#     decode as codec_decode,\n",
    "#     encode as codec_encode,\n",
    "#     get_embedding_rate,\n",
    "#     load_model as load_codec_model,\n",
    "# )\n",
    "\n",
    "# CODEC_FILEPATH = \"s3://suno-data/christian/25hz_vae_peaq_kl_0.005.pth\"\n",
    "# preload_codec_models(CODEC_FILEPATH)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# # decode some audio\n",
    "idx = 108\n",
    "# # ensure even index\n",
    "assert idx % 2 == 0\n",
    "# print(metas_val[idx])\n",
    "# print(\"negative\")\n",
    "# audio = codec_decode(mm_vae_val[idx])\n",
    "# audio.normalize_volume().play()\n",
    "\n",
    "# print(metas_val[idx + 1])\n",
    "# print(\"positive\")\n",
    "# audio = codec_decode(mm_vae_val[idx + 1])\n",
    "# audio.normalize_volume().play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "torch.equal(torch.tensor(mm_semantic_val[idx]), torch.tensor(mm_semantic_val[idx + 1]))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# original_npz_path = f\"/app/suno/data/dpo/7b_npz/{test_metas[idx]['id']}.npz\"\n",
    "# original_npz_path = \"/app/suno/data/dpo/7b_npz/729c3011-f672-4ccd-8d82-1cbf2b52ff69.npz\"\n",
    "# original_arr = np.load(original_npz_path)[\"v2_raw\"]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def validation_on_metas(input_metas):\n",
    "    total_bad = 0\n",
    "    total_good = 0\n",
    "    for idx in range(len(input_metas)):\n",
    "        if idx % 2 == 0:\n",
    "            pos_idx = idx + 1\n",
    "            if input_metas[idx].get(\"tags\") != input_metas[pos_idx].get(\"tags\"):\n",
    "                # print(test_metas[idx].get(\"text\") == test_metas[pos_idx].get(\"text\"), test_metas[idx].get(\"tags\"), test_metas[pos_idx].get(\"tags\"))\n",
    "                total_bad += 1\n",
    "            else:\n",
    "                total_good += 1\n",
    "    print(total_good, total_bad)\n",
    "    return"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "metas_tr = read_jsonl(os.path.join(OUT_DATA_DIR, \"metas_tr.jsonl\"))\n",
    "validation_on_metas(metas_tr)\n",
    "print(len(metas_tr))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "sum(len(meta[\"tags\"][0]) == 0 for meta in metas_tr)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "!cd /home/tony/Work/tony/slurm/diffusion && sbatch run_diffusion.sh"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import shutil\n",
    "\n",
    "# Basic file copy\n",
    "shutil.copy(\n",
    "    \"/home/tony/Work/tony/Preference/make_dataset_diff_upsample_v1_r5_comb.ipynb\",\n",
    "    os.path.join(OUT_DATA_DIR, \"make_dataset.ipynb\"),\n",
    ")\n",
    "print(\"Cache kept!\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Inspections "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# df[df[\"preference\"] & (df[\"shimmer_score_diff\"] > 3)][\n",
    "#     [\n",
    "#         \"index\",\n",
    "#         \"s3_id\",\n",
    "#         \"total_shimmer_score\",\n",
    "#         \"shimmer_score_diff\",\n",
    "#         \"request_id\",\n",
    "#         \"preference\",\n",
    "#     ]\n",
    "# ].tail()\n",
    "\n",
    "# df[df[\"preference\"] & (df[\"pair_quality\"] < 0.1)][\n",
    "#     [\n",
    "#         \"index\",\n",
    "#         \"s3_id\",\n",
    "#         \"total_shimmer_score\",\n",
    "#         \"shimmer_score_diff\",\n",
    "#         \"pair_quality\",\n",
    "#         \"request_id\",\n",
    "#         \"preference\",\n",
    "#     ]\n",
    "# ].tail()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# test_pair_df = df[df[\"request_id\"] == \"621c8b02-a905-48f1-a2d5-a4e8423d1505\"]\n",
    "# print(\n",
    "#     test_pair_df[\n",
    "#         [\n",
    "#             \"s3_id\",\n",
    "#             \"total_shimmer_score\",\n",
    "#             \"pair_quality\",\n",
    "#             \"request_id\",\n",
    "#             \"preference\",\n",
    "#             \"prompt_text\",\n",
    "#         ]\n",
    "#     ]\n",
    "# )\n",
    "# negative_audio = Audio.from_s3(\n",
    "#     f\"s3://suno-data-uploads/studio/uploads/{test_pair_df['s3_id'].values[0]}.mp3\",\n",
    "#     n_channels=2,\n",
    "# )\n",
    "# print(\"negative\")\n",
    "# negative_audio.get_segment(0, 30).play()\n",
    "# positive_audio = Audio.from_s3(\n",
    "#     f\"s3://suno-data-uploads/studio/uploads/{test_pair_df['s3_id'].values[1]}.mp3\",\n",
    "#     n_channels=2,\n",
    "# )\n",
    "# print(\"positive\")\n",
    "# positive_audio.get_segment(0, 30).play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# total_dict = {}\n",
    "# total_dict.update(pair_quality_dict)\n",
    "# total_dict.update(pair_quality_1_dict)\n",
    "# total_dict.update(pair_quality_2_dict)\n",
    "# total_dict.update(pair_quality_3_dict)\n",
    "# len(total_dict)\n",
    "# with open(\n",
    "#     os.path.join(\"/home/tony/Data/Preference/up_v1\", \"pair_quality.json\"), \"w\"\n",
    "# ) as fp:\n",
    "#     json.dump(total_dict, fp, indent=4)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# import numpy as np"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# test_arr = np.load(\"/home/tony/Data/test_npz/diffusion_input_tensor([ 18, 182]).npy\")\n",
    "# test_arr.shape"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# mm_vae_val = np.memmap(\n",
    "#     os.path.join(OUT_DATA_DIR, \"data_vae_val.bin\"), dtype=np.float16, mode=\"r\"\n",
    "# )\n",
    "\n",
    "# mm_vae_val = mm_vae_val.reshape(-1, VAE_MEMMAP_SIZE, VAE_DIM)\n",
    "# print(mm_vae_val.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# print(\"negative\")\n",
    "# audio = codec_decode(test_arr[0].T / 2.5)\n",
    "# audio.normalize_volume().play()\n",
    "\n",
    "# print(\"positive\")\n",
    "# audio = codec_decode(test_arr[1].T / 2.5)\n",
    "# audio.normalize_volume().play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# print(\"negative\")\n",
    "# audio = codec_decode(mm_vae_val[18])\n",
    "# audio.normalize_volume().play()\n",
    "\n",
    "# print(\"positive\")\n",
    "# audio = codec_decode(mm_vae_val[19])\n",
    "# audio.normalize_volume().play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch\n",
    "\n",
    "rng = torch.quasirandom.SobolEngine(1, scramble=True, seed=0)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "t = rng.draw(4)[:, 0].to(torch.bfloat16)\n",
    "print(t)\n",
    "# Replace 1% of t with ones to ensure training on terminal SNR\n",
    "t = torch.where(torch.rand_like(t) < 0.5, torch.ones_like(t), t)\n",
    "print(t)\n",
    "t = torch.repeat_interleave(t, repeats=2, dim=0)\n",
    "print(t)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "(t * 32).to(int) / 32"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "t[0] = 0.99"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "torch.round(t * 32) / 32"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# import json\n",
    "# # with open(f\"/home/tony/Data/Preference/up_v6/full_pair_quality.json\", \"r\") as f:\n",
    "# #    result = json.load(f)\n",
    "# result = {}\n",
    "# print(len(result))\n",
    "# for job_idx in range(8):\n",
    "#     with open(f\"/home/tony/Data/Preference/up_v6/full_pair_quality_{job_idx}.json\", \"r\") as fp:\n",
    "#         current_result = json.load(fp)\n",
    "#         result.update(current_result)\n",
    "# print(len(result))\n",
    "# with open(f\"/home/tony/Data/Preference/up_v6/full_pair_quality.json\", \"w\") as f:\n",
    "#     json.dump(result, f, indent=4)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# import torch\n",
    "# semantic_codes_chunk = torch.ones((1, 100))\n",
    "# semantic_skip_phase = 0\n",
    "# semantic_skip_factor = 4\n",
    "# mask = torch.ones_like(semantic_codes_chunk, dtype=torch.bool)\n",
    "# indices = (\n",
    "#     torch.arange(semantic_codes_chunk.size(1)) + semantic_skip_phase\n",
    "# ) % semantic_skip_factor == 0\n",
    "# mask[:, indices] = False\n",
    "# mask"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "df_slice.columns"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# df_slice[(df_slice[\"preference\"]) & (df_slice[\"total_shimmer_score\"] < 0.5)][[\"s3_id\", \"total_shimmer_score\"]].head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# df_slice[(df_slice[\"preference\"]) & (df_slice[\"total_shimmer_score\"] > 2)][[\"s3_id\", \"total_shimmer_score\"]].head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 4
}
