{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"3\"\n",
    "\n",
    "import torch\n",
    "import torchaudio\n",
    "\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"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "input_filepath = \"/home/christian/code/christian/outputs/covers/rock-to-my-roll-short.mp3\"\n",
    "cover_a_filepath = \"/home/christian/code/christian/outputs/covers/004-A.mp3\"\n",
    "cover_b_filepath = \"/home/christian/code/christian/outputs/covers/004-B.mp3\"\n",
    "\n",
    "input_audio = Audio.from_file(input_filepath)\n",
    "cover_a_audio = Audio.from_file(cover_a_filepath)\n",
    "cover_b_audio = Audio.from_file(cover_b_filepath)\n",
    "\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",
    "print(input_sem.shape, cover_a_sem.shape, cover_b_sem.shape)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "s3_base_path = \"s3://suno-data-uploads/studio/uploads/\"\n",
    "input_s3_id = \"e8895c5c-4b6d-4dcd-a239-0406221ab337.mp3\"\n",
    "cover_a_s3_id = \"30bd8134-2bb3-4876-8ffb-6be944b6dc06.mp3\"\n",
    "cover_b_s3_id = \"a044752a-f9ec-4c8a-b310-2553d0ca464d.mp3\"\n",
    "\n",
    "\n",
    "input_audio = Audio.from_s3(os.path.join(s3_base_path, input_s3_id))\n",
    "cover_a_audio = Audio.from_s3(os.path.join(s3_base_path, cover_a_s3_id))\n",
    "cover_b_audio = Audio.from_s3(os.path.join(s3_base_path, cover_b_s3_id))\n",
    "\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",
    "print(input_sem.shape, cover_a_sem.shape, cover_b_sem.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "input_mean = input_sem.mean(axis=0)\n",
    "input_std = input_sem.std(axis=0)\n",
    "\n",
    "cover_a_mean = cover_a_sem.mean(axis=0)\n",
    "cover_a_std = cover_a_sem.std(axis=0)\n",
    "\n",
    "cover_b_mean = cover_b_sem.mean(axis=0)\n",
    "cover_b_std = cover_b_sem.std(axis=0)\n",
    "\n",
    "print(input_mean[:10])\n",
    "print(cover_a_mean[:10])\n",
    "print(cover_b_mean[:10])\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",
    "print(input_to_cover_a_distance, input_to_cover_b_distance)\n",
    "\n",
    "# distance threshold above 1.0 seems to indicate a valid cover\n",
    "# so we are looking for cases where one has distance below and the other above \n",
    "# we also need to check for silence in the covers"
   ]
  },
  {
   "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
}
