{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "0",
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.gpt import chirp_v2_5 as chirp_v3\n",
    "from suno_utils.tasks.ditto_v2 import preload_models\n",
    "from suno_utils.tasks.ditto_v2 import encode_overlap as encode, SAMPLE_RATE, load_model\n",
    "from suno_utils.utils.clip import SunoClip\n",
    "import numpy as np\n",
    "import pandas as pd\n",
    "from tqdm import tqdm\n",
    "import boto3\n",
    "import matplotlib.pyplot as plt\n",
    "import json\n",
    "import pickle\n",
    "\n",
    "tqdm.pandas()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "1",
   "metadata": {},
   "outputs": [],
   "source": [
    "PREF_PATH = \"/home/tony/Data/Preference/auk_t1/interesting_clips_auk_t1_20250604.pkl\"\n",
    "PAIR_PATH = \"/app2/suno/data/sara/cover_pref_filtering/pariwise_preferences.pkl\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2",
   "metadata": {},
   "outputs": [],
   "source": [
    "pickle_path = \"/app2/suno/data/sara/cover_pref_filtering/batch_encode_results.pkl\"\n",
    "\n",
    "with open(pickle_path, \"rb\") as f:\n",
    "    result_dict = pickle.load(f)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3",
   "metadata": {},
   "outputs": [],
   "source": [
    "df = pd.read_pickle(PREF_PATH)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "4",
   "metadata": {},
   "outputs": [],
   "source": [
    "df.head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "5",
   "metadata": {},
   "outputs": [],
   "source": [
    "df[\"task\"] = df[\"metadata\"].apply(lambda x: x.get(\"task\", None))\n",
    "df_covers = df[df[\"task\"] == \"cover\"].copy()\n",
    "df_covers[\"cover_clip_id\"] = df_covers[\"metadata\"].apply(\n",
    "    lambda x: x.get(\"cover_clip_id\", None)\n",
    ")\n",
    "df_cover_prefs = df_covers[\n",
    "    [\n",
    "        \"pos_preference\",\n",
    "        \"neg_preference\",\n",
    "        \"diff_preference\",\n",
    "        \"preference\",\n",
    "        \"s3_id\",\n",
    "        \"cover_clip_id\",\n",
    "    ]\n",
    "]\n",
    "print(f\"Found {df_cover_prefs.shape[0]} preference covers from {df.shape[0]} rows\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "6",
   "metadata": {},
   "outputs": [],
   "source": [
    "df_cover_prefs.head(6)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "7",
   "metadata": {},
   "outputs": [],
   "source": [
    "def combine_pairwise_preferences(df):\n",
    "    \"\"\"\n",
    "    Combine pairwise preference pairs into single rows.\n",
    "    Each pair shares the same cover_clip_id and represents a comparison.\n",
    "\n",
    "    Returns a clean dataframe with:\n",
    "    - cover_clip_id: shared identifier\n",
    "    - s3_id_1: first item's s3_id\n",
    "    - s3_id_2: second item's s3_id\n",
    "    - diff_preference_1: first item's diff_preference\n",
    "    - diff_preference_2: second item's diff_preference\n",
    "    \"\"\"\n",
    "\n",
    "    combined_rows = []\n",
    "    print(df.shape[0])\n",
    "\n",
    "    # Process consecutive pairs based on original ordering\n",
    "    for i in range(0, len(df), 2):\n",
    "        if i + 1 < len(df):  # Ensure we have a pair\n",
    "            row1 = df.iloc[i]\n",
    "            row2 = df.iloc[i + 1]\n",
    "\n",
    "            preference = None\n",
    "            if row1[\"diff_preference\"] == 1 and row2[\"diff_preference\"] == -1:\n",
    "                preference = \"A\"\n",
    "            elif row2[\"diff_preference\"] == 1 and row1[\"diff_preference\"] == -1:\n",
    "                preference = \"B\"\n",
    "\n",
    "            # Create combined row with clean column names\n",
    "            combined_row = {\n",
    "                \"cover_clip_id\": row1[\n",
    "                    \"cover_clip_id\"\n",
    "                ],  # Assuming pairs share the same cover_clip_id\n",
    "                \"s3_id_1\": row1[\"s3_id\"],\n",
    "                \"s3_id_2\": row2[\"s3_id\"],\n",
    "                \"preference\": preference,\n",
    "                \"diff_preference_1\": row1[\"diff_preference\"],\n",
    "                \"diff_preference_2\": row2[\"diff_preference\"],\n",
    "            }\n",
    "\n",
    "            combined_rows.append(combined_row)\n",
    "        else:\n",
    "            print(f\"Warning: Odd number of rows, last row (index {i}) will be skipped\")\n",
    "\n",
    "    return pd.DataFrame(combined_rows)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "8",
   "metadata": {},
   "outputs": [],
   "source": [
    "df_pairwise_rows = combine_pairwise_preferences(df_cover_prefs)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "9",
   "metadata": {},
   "outputs": [],
   "source": [
    "df_pairwise_rows = df_pairwise_rows[df_pairwise_rows[\"preference\"].notna()]\n",
    "print(df_pairwise_rows.shape[0])\n",
    "df_pairwise_rows.head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "10",
   "metadata": {},
   "outputs": [],
   "source": [
    "# df_pairwise_rows.to_pickle(PAIR_PATH)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "11",
   "metadata": {},
   "outputs": [],
   "source": [
    "def cosine_similarity(a, b):\n",
    "    dot_product = np.dot(a, b)\n",
    "\n",
    "    magnitude_a = np.sqrt(np.dot(a, a))\n",
    "    magnitude_b = np.sqrt(np.dot(b, b))\n",
    "\n",
    "    return dot_product / (magnitude_a * magnitude_b)\n",
    "\n",
    "\n",
    "def score_row(row):\n",
    "    source_embed = result_dict.get(row[\"cover_clip_id\"], None)\n",
    "    cover_embed = result_dict.get(row[\"s3_id\"], None)\n",
    "    if source_embed is not None and cover_embed is not None:\n",
    "        cos_sim = cosine_similarity(source_embed, cover_embed)\n",
    "        return cos_sim\n",
    "    return None\n",
    "\n",
    "\n",
    "def score_combined_row(row):\n",
    "    source_embed = result_dict.get(row[\"cover_clip_id\"], None)\n",
    "    cover_1_embed = result_dict.get(row[\"s3_id_1\"], None)\n",
    "    cover_2_embed = result_dict.get(row[\"s3_id_2\"], None)\n",
    "\n",
    "    if source_embed is None or cover_1_embed is None or cover_2_embed is None:\n",
    "        return pd.Series(\n",
    "            [None, None, None], index=[\"score_pos\", \"score_neg\", \"sim_score_diff\"]\n",
    "        )\n",
    "\n",
    "    score_1 = cosine_similarity(source_embed, cover_1_embed)\n",
    "    score_2 = cosine_similarity(source_embed, cover_2_embed)\n",
    "\n",
    "    final_score = 0\n",
    "    score_pos = 0\n",
    "    score_neg = 0\n",
    "    if row[\"preference\"] == \"A\":\n",
    "        final_score = score_1 - score_2\n",
    "        score_pos = score_1\n",
    "        score_neg = score_2\n",
    "    elif row[\"preference\"] == \"B\":\n",
    "        final_score = score_2 - score_1\n",
    "        score_pos = score_2\n",
    "        score_neg = score_1\n",
    "\n",
    "    return pd.Series(\n",
    "        [score_pos, score_neg, final_score],\n",
    "        index=[\"score_pos\", \"score_neg\", \"sim_score_diff\"],\n",
    "    )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "12",
   "metadata": {},
   "outputs": [],
   "source": [
    "df_cover_prefs = df_cover_prefs[df_cover_prefs[\"diff_preference\"] != 0]\n",
    "print(df_cover_prefs.shape[0])\n",
    "subset = df_cover_prefs"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "13",
   "metadata": {},
   "outputs": [],
   "source": [
    "subset[\"sim_score\"] = subset.progress_apply(score_row, axis=1)\n",
    "subset = subset[subset[\"sim_score\"].notna()]\n",
    "print(subset.shape[0])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "14",
   "metadata": {},
   "outputs": [],
   "source": [
    "subset.head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "15",
   "metadata": {},
   "outputs": [],
   "source": [
    "def plot_distribution_comparison(\n",
    "    df,\n",
    "    value_col,\n",
    "    group_col,\n",
    "    group_values,\n",
    "    bins=20,\n",
    "    figsize=(10, 8),\n",
    "    colors=None,\n",
    "    show_stats=True,\n",
    "    alpha=0.7,\n",
    "):\n",
    "    if colors is None:\n",
    "        colors = [\"skyblue\", \"salmon\"]\n",
    "\n",
    "    # Filter data for each group\n",
    "    group1_data = df[df[group_col] == group_values[0]][value_col]\n",
    "    group2_data = df[df[group_col] == group_values[1]][value_col]\n",
    "\n",
    "    # Create figure with subplots stacked vertically\n",
    "    fig, (ax1, ax2) = plt.subplots(2, 1, figsize=figsize, sharex=True)\n",
    "\n",
    "    # Define common bins for consistent comparison\n",
    "    all_data = pd.concat([group1_data, group2_data])\n",
    "    bin_range = (all_data.min(), all_data.max())\n",
    "\n",
    "    # Plot histogram for first group\n",
    "    ax1.hist(\n",
    "        group1_data,\n",
    "        bins=bins,\n",
    "        range=bin_range,\n",
    "        alpha=alpha,\n",
    "        color=colors[0],\n",
    "        edgecolor=\"black\",\n",
    "    )\n",
    "    ax1.set_title(f\"{value_col} Distribution: {group_col} = {group_values[0]}\")\n",
    "    ax1.set_ylabel(\"Frequency\")\n",
    "    ax1.grid(True, alpha=0.3)\n",
    "\n",
    "    # Plot histogram for second group\n",
    "    ax2.hist(\n",
    "        group2_data,\n",
    "        bins=bins,\n",
    "        range=bin_range,\n",
    "        alpha=alpha,\n",
    "        color=colors[1],\n",
    "        edgecolor=\"black\",\n",
    "    )\n",
    "    ax2.set_title(f\"{value_col} Distribution: {group_col} = {group_values[1]}\")\n",
    "    ax2.set_xlabel(f\"{value_col} Value\")\n",
    "    ax2.set_ylabel(\"Frequency\")\n",
    "    ax2.grid(True, alpha=0.3)\n",
    "\n",
    "    # Ensure same y-scale for better comparison\n",
    "    if len(group1_data) > 0 and len(group2_data) > 0:\n",
    "        max_freq = max(ax1.get_ylim()[1], ax2.get_ylim()[1])\n",
    "    elif len(group1_data) > 0:\n",
    "        max_freq = ax1.get_ylim()[1]\n",
    "    elif len(group2_data) > 0:\n",
    "        max_freq = ax2.get_ylim()[1]\n",
    "    else:\n",
    "        max_freq = 1  # Default value if both groups are empty\n",
    "    ax1.set_ylim(0, max_freq)\n",
    "    ax2.set_ylim(0, max_freq)\n",
    "\n",
    "    # Add sample size information\n",
    "    ax1.text(\n",
    "        0.02,\n",
    "        0.95,\n",
    "        f\"n = {len(group1_data)}\",\n",
    "        transform=ax1.transAxes,\n",
    "        verticalalignment=\"top\",\n",
    "        bbox=dict(boxstyle=\"round\", facecolor=\"white\", alpha=0.8),\n",
    "    )\n",
    "    ax2.text(\n",
    "        0.02,\n",
    "        0.95,\n",
    "        f\"n = {len(group2_data)}\",\n",
    "        transform=ax2.transAxes,\n",
    "        verticalalignment=\"top\",\n",
    "        bbox=dict(boxstyle=\"round\", facecolor=\"white\", alpha=0.8),\n",
    "    )\n",
    "\n",
    "    plt.tight_layout()\n",
    "\n",
    "    # Print summary statistics\n",
    "    if show_stats:\n",
    "        print(\"Summary Statistics:\")\n",
    "        print(f\"\\n{group_col} = {group_values[0]}:\")\n",
    "        print(f\"  Count: {len(group1_data)}\")\n",
    "        if len(group1_data) > 0:\n",
    "            print(f\"  Mean: {group1_data.mean():.4f}\")\n",
    "            print(f\"  Std: {group1_data.std():.4f}\")\n",
    "            print(f\"  Min: {group1_data.min():.4f}\")\n",
    "            print(f\"  Max: {group1_data.max():.4f}\")\n",
    "\n",
    "        print(f\"\\n{group_col} = {group_values[1]}:\")\n",
    "        print(f\"  Count: {len(group2_data)}\")\n",
    "        if len(group2_data) > 0:\n",
    "            print(f\"  Mean: {group2_data.mean():.4f}\")\n",
    "            print(f\"  Std: {group2_data.std():.4f}\")\n",
    "            print(f\"  Min: {group2_data.min():.4f}\")\n",
    "            print(f\"  Max: {group2_data.max():.4f}\")\n",
    "\n",
    "    plt.show()\n",
    "    return fig, (ax1, ax2)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "16",
   "metadata": {},
   "outputs": [],
   "source": [
    "# Example usage with your data:\n",
    "fig, axes = plot_distribution_comparison(\n",
    "    subset, \"sim_score\", \"diff_preference\", [1, -1]\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "17",
   "metadata": {},
   "outputs": [],
   "source": [
    "df_pairwise_rows[[\"score_pos\", \"score_neg\", \"sim_score_diff\"]] = (\n",
    "    df_pairwise_rows.progress_apply(score_combined_row, axis=1)\n",
    ")\n",
    "df_pairwise_rows = df_pairwise_rows[df_pairwise_rows[\"sim_score_diff\"].notna()]\n",
    "print(df_pairwise_rows.shape[0])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "18",
   "metadata": {},
   "outputs": [],
   "source": [
    "df_pair_high_pref = df_pairwise_rows[df_pairwise_rows[\"score_pos\"] > 0.9]\n",
    "df_pair_high_pref = df_pair_high_pref[abs(df_pair_high_pref[\"sim_score_diff\"]) > 0.1]\n",
    "df_pair_high_pref = df_pair_high_pref[abs(df_pair_high_pref[\"score_neg\"]) > 0.5]\n",
    "print(df_pair_high_pref.shape[0] / df_pairwise_rows.shape[0])\n",
    "df_pair_high_pref.head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "19",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "20",
   "metadata": {},
   "outputs": [],
   "source": [
    "plt.figure(figsize=(10, 6))\n",
    "plt.hist(\n",
    "    df_pairwise_rows[\"sim_score_diff\"],\n",
    "    bins=30,\n",
    "    alpha=0.7,\n",
    "    color=\"skyblue\",\n",
    "    edgecolor=\"black\",\n",
    ")\n",
    "plt.xlabel(\"Ditto self_sim delta\")\n",
    "plt.ylabel(\"Frequency\")\n",
    "plt.title(\"Preference Pair sel_sim Delta\")\n",
    "plt.grid(True, alpha=0.3)\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "21",
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "suno_clean",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "codemirror_mode": {
    "name": "ipython",
    "version": 3
   },
   "file_extension": ".py",
   "mimetype": "text/x-python",
   "name": "python",
   "nbconvert_exporter": "python",
   "pygments_lexer": "ipython3",
   "version": "3.10.15"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
