{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import pandas as pd\n",
    "\n",
    "filepath = \"/home/tony/Data/Preference/30b_v3/interesting_clips_v4_t_4_20240923_full_with_sem_distance_and_similarity.pkl\"\n",
    "train_df = pd.read_pickle(filepath)\n",
    "train_df = train_df.sort_values(by=[\"request_id\", \"preference\"])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "train_df.tail(n=2)[[\"s3_id\", \"prompt_text\", \"similarity\", \"preference\"]]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "from suno_utils.audio import Audio\n",
    "from suno_utils.tasks.mert_25 import preload_models, encode\n",
    "_ = preload_models(    \n",
    "    checkpoint_filepath=\"s3://suno-data/georg/models/semantic/mert_25.pt\",\n",
    "    centroids_filepath=\"s3://suno-data/georg/models/semantic/mert_25_2x4k.npy\",\n",
    "    device=\"cuda\",\n",
    ")\n",
    "\n",
    "\n",
    "def measure_similarity(\n",
    "        input_s3_id: str, \n",
    "        cover_a_s3_id: str, \n",
    "        cover_b_s3_id: str, \n",
    "        s3_base_path: str = \"s3://suno-data-uploads/studio/uploads/\", \n",
    "        feature: str = \"semantic\"\n",
    "    ):\n",
    "    input_audio = Audio.from_s3(os.path.join(s3_base_path, input_s3_id + \".mp3\"))\n",
    "    cover_a_audio = Audio.from_s3(os.path.join(s3_base_path, cover_a_s3_id + \".mp3\"))\n",
    "    cover_b_audio = Audio.from_s3(os.path.join(s3_base_path, cover_b_s3_id + \".mp3\"))\n",
    "\n",
    "    if feature == \"semantic\":\n",
    "        input_sem = encode(input_audio, do_clustering=False)\n",
    "        cover_a_sem = encode(cover_a_audio, do_clustering=False)\n",
    "        cover_b_sem = encode(cover_b_audio, do_clustering=False)\n",
    "\n",
    "        input_mean = input_sem.mean(axis=0)\n",
    "        cover_a_mean = cover_a_sem.mean(axis=0)\n",
    "        cover_b_mean = cover_b_sem.mean(axis=0)\n",
    "    elif feature == \"crema\":\n",
    "        input_crema = analyze(y=input_audio.float_array, sr=input_audio.sample_rate)\n",
    "        cover_a_crema = analyze(y=cover_a_audio.float_array, sr=cover_a_audio.sample_rate)\n",
    "        cover_b_crema = analyze(y=cover_b_audio.float_array, sr=cover_b_audio.sample_rate)\n",
    "        print(input_crema.shape)\n",
    "\n",
    "        input_mean = input_crema.mean(axis=0)\n",
    "        cover_a_mean = cover_a_crema.mean(axis=0)\n",
    "        cover_b_mean = cover_b_crema.mean(axis=0)\n",
    "    else:\n",
    "        raise ValueError(f\"Feature {feature} not supported\")\n",
    "\n",
    "    # mse distance between input and cover_a\n",
    "    input_to_cover_a_distance = ((input_mean - cover_a_mean) ** 2).sum()\n",
    "    input_to_cover_b_distance = ((input_mean - cover_b_mean) ** 2).sum()\n",
    "\n",
    "    return input_to_cover_a_distance, input_to_cover_b_distance\n",
    "\n",
    "\n",
    "def analyze_row(index: int):\n",
    "    # get the cover id \n",
    "    row1_metadata = dict(train_df.iloc[index][\"metadata\"])\n",
    "    input_s3_id = row1_metadata[\"cover_clip_id\"]\n",
    "    cover_s3_id = train_df.iloc[index][\"s3_id\"]\n",
    "    s3_base_path: str = \"s3://suno-data-uploads/studio/uploads/\"\n",
    "\n",
    "    input_audio = Audio.from_s3(os.path.join(s3_base_path, input_s3_id + \".mp3\"))\n",
    "    cover_audio = Audio.from_s3(os.path.join(s3_base_path, cover_s3_id + \".mp3\"))\n",
    "\n",
    "    input_sem = encode(input_audio, do_clustering=False)\n",
    "    cover_sem = encode(cover_audio, do_clustering=False)\n",
    "\n",
    "    input_mean = input_sem.mean(axis=0)\n",
    "    cover_mean = cover_sem.mean(axis=0)\n",
    "\n",
    "    # mse distance between input and cover\n",
    "    input_to_cover_distance = ((input_mean - cover_mean) ** 2).sum()\n",
    "\n",
    "    return input_to_cover_distance"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# print all of the key names \n",
    "print(train_df.iloc[0].keys())\n",
    "\n",
    "# put the metadata object into a dictionary\n",
    "cover_id = train_df.iloc[0][\"s3_id\"]\n",
    "metadata = train_df.iloc[0][\"metadata\"]\n",
    "metadata_dict = dict(metadata)\n",
    "print(cover_id)\n",
    "print(metadata_dict[\"cover_clip_id\"])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# iterate over all rows and measure similarity between cover a and cover b\n",
    "# one example is comprised of two adjacent rows in the dataframe\n",
    "# so iterate over the dataframe in steps of 2 and get two rows at a time\n",
    "for index in range(0, len(train_df), 2):\n",
    "    print(f\"Processing index {index}\")\n",
    "    row1 = train_df.iloc[index]\n",
    "    row2 = train_df.iloc[index + 1]\n",
    "\n",
    "    a_tony_sim = row1[\"similarity\"]\n",
    "    b_tony_sim = row2[\"similarity\"]\n",
    "    a_preference = row1[\"preference\"]\n",
    "    b_preference = row2[\"preference\"]\n",
    "\n",
    "    # get the cover id \n",
    "    row1_metadata = dict(row1[\"metadata\"])\n",
    "    input_s3_id = row1_metadata[\"cover_clip_id\"]\n",
    "\n",
    "    cover_a_s3_id = row1[\"s3_id\"]\n",
    "    cover_b_s3_id = row2[\"s3_id\"]\n",
    "\n",
    "    input_to_cover_a_distance, input_to_cover_b_distance = measure_similarity(input_s3_id, cover_a_s3_id, cover_b_s3_id)\n",
    "    print(f\"a dist: {input_to_cover_a_distance}, b dist: {input_to_cover_b_distance}\")\n",
    "    print(f\"a tony sim: {a_tony_sim}, b tony sim: {b_tony_sim}\")\n",
    "    print(f\"a preference: {a_preference}, b preference: {b_preference}\")\n",
    "    print(\"\")\n",
    "\n",
    "    # add similarity back to the dataframe\n",
    "    train_df.at[index, \"distance\"] = input_to_cover_a_distance\n",
    "    train_df.at[index + 1, \"distance\"] = input_to_cover_b_distance\n",
    "\n",
    "    if index > 10:\n",
    "        break"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "suno_env",
   "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.9"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
