{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "04d5371e",
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"3\"\n",
    "\n",
    "import torch\n",
    "import numpy as np\n",
    "import json\n",
    "import torchaudio"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "14746dea",
   "metadata": {},
   "outputs": [],
   "source": [
    "base_dir = \"/app2/suno/data/christian/outputs/v3-base-data-ctx-t3\"\n",
    "labels_path = \"/home/christian/code/christian/metadata/reward_model/evals/v3-base-data-ctx-t3-cjs.json\"\n",
    "model_name = \"v3_flow_sft_t8_rd1_pair_t3_1E6_beta100_n16_bt2_noise_1k_last\"\n",
    "\n",
    "# load the labels\n",
    "with open(labels_path, \"r\") as f:\n",
    "    labels = json.load(f)[\"results\"]\n",
    "\n",
    "print(f\"Loaded {len(labels)} labels\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "66773685",
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.tasks.ear_v3 import load_checkpoint\n",
    "\n",
    "#ear_model_filepath = \"/app2/suno/checkpoints/2025-09-05_16-13-40_s691/last_ckpt.pt\" # 47%\n",
    "#ear_model_filepath = \"/app2/suno/checkpoints/2025-09-05_20-07-23_s6954/last_ckpt.pt\" # 59%\n",
    "#ear_model_filepath = \"/app2/suno/checkpoints/2025-09-05_20-07-23_s6954/best_ckpt.pt\" # 60%\n",
    "#ear_model_filepath = \"/app2/suno/checkpoints/2025-09-08_14-23-56_s8558/last_ckpt.pt\" # 61%\n",
    "#ear_model_filepath = \"/app2/suno/checkpoints/2025-09-08_14-23-56_s8558/best_ckpt.pt\" # 61%\n",
    "#ear_model_filepath = \"/app2/suno/checkpoints/2025-09-08_17-40-27_s5826/last_ckpt.pt\" # 67% (68%)\n",
    "#ear_model_filepath = \"/app2/suno/checkpoints/2025-09-08_22-03-25_s9042/last_ckpt.pt\" # 62% (58%)\n",
    "#ear_model_filepath = \"/app2/suno/checkpoints/2025-09-09_14-48-28_s45/best_ckpt.pt\" # 62% (69%)\n",
    "#ear_model_filepath = \"/app2/suno/checkpoints/2025-09-18_10-32-07_s4862/best_ckpt.pt\"\n",
    "#ear_model_filepath = \"/app2/suno/checkpoints/2025-09-18_10-33-55_s9767/best_ckpt.pt\"\n",
    "#ear_model_filepath = \"/app2/suno/checkpoints/2025-09-22_18-13-22_s635/last_ckpt.pt\"\n",
    "#ear_model_filepath = \"/app2/suno/checkpoints/2025-09-30_17-25-36_s3326/last_ckpt.pt\"\n",
    "\n",
    "#ear_model_filepath = \"/app2/suno/checkpoints/2025-10-02_23-34-30_s400/last_ckpt.pt\"\n",
    "#ear_model_filepath = \"/app2/suno/checkpoints/2025-10-03_00-13-10_s7805/last_ckpt.pt\"\n",
    "#ear_model_filepath = \"/app2/suno/checkpoints/2025-10-03_00-59-14_s9683/best_ckpt.pt\"\n",
    "#ear_model_filepath = \"/app2/suno/checkpoints/2025-10-03_10-36-34_s5844/last_ckpt.pt\"\n",
    "#ear_model_filepath = \"/app2/suno/checkpoints/2025-10-03_10-52-33_s7400/last_ckpt.pt\"\n",
    "#ear_model_filepath = \"/app2/suno/checkpoints/2025-10-03_13-39-33_s2150/last_ckpt.pt\" # only degrade\n",
    "ear_model_filepath = \"/app2/suno/checkpoints/2025-10-07_21-32-12_s4698/best_ckpt.pt\"\n",
    "ear_model_filepath = \"/app2/suno/checkpoints/2025-10-08_11-52-17_s8263/best_ckpt.pt\"\n",
    "model_v3 = load_checkpoint(ear_model_filepath, \"cuda:0\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "4f8bd3f1",
   "metadata": {},
   "outputs": [],
   "source": [
    "# lets load some data of full clips "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "ba6ec38e",
   "metadata": {},
   "outputs": [],
   "source": [
    "scores = []\n",
    "from tqdm import tqdm\n",
    "\n",
    "use_v2 = False\n",
    "skip_history = False\n",
    "vae_scale_factor = 0.4\n",
    "\n",
    "pbar = tqdm(labels.items())\n",
    "for clip_id, label in pbar:\n",
    "    if label == \"skip\":\n",
    "        continue\n",
    "\n",
    "    # skip ones where history is not None\n",
    "    history_vae_filepath = os.path.join(base_dir, clip_id, f\"{clip_id}_history_vae.npz\")\n",
    "    if skip_history and os.path.exists(history_vae_filepath):\n",
    "        continue\n",
    "\n",
    "    # for each label, there is a positive and negative example\n",
    "    a_mp3_filepath = os.path.join(base_dir, clip_id, f\"{clip_id}_{model_name}_0.mp3\")\n",
    "    b_mp3_filepath = os.path.join(base_dir, clip_id, f\"{clip_id}_{model_name}_1.mp3\")\n",
    "\n",
    "    a_vae_filepath = os.path.join(base_dir, clip_id, f\"{clip_id}_{model_name}_0_upsampled_vae.npz\")\n",
    "    b_vae_filepath = os.path.join(base_dir, clip_id, f\"{clip_id}_{model_name}_1_upsampled_vae.npz\")\n",
    "\n",
    "    # load the vae latents\n",
    "    a_vae_latents = torch.from_numpy(np.load(a_vae_filepath)[\"vae_latents\"]).float()\n",
    "    b_vae_latents = torch.from_numpy(np.load(b_vae_filepath)[\"vae_latents\"]).float()\n",
    "\n",
    "    # apply scale factor to vae\n",
    "    a_vae_latents = a_vae_latents * vae_scale_factor\n",
    "    b_vae_latents = b_vae_latents * vae_scale_factor\n",
    "\n",
    "    # assert \n",
    "    assert a_vae_latents.shape[0] == 750\n",
    "    assert b_vae_latents.shape[0] == 750\n",
    "\n",
    "    if use_v2:\n",
    "        a_score = model_v2.get_score(a_mp3_filepath)\n",
    "        b_score = model_v2.get_score(b_mp3_filepath)\n",
    "    else:   \n",
    "        a_score = model_v3(a_vae_latents.unsqueeze(0).cuda()).mean().item()\n",
    "        b_score = model_v3(b_vae_latents.unsqueeze(0).cuda()).mean().item()\n",
    "\n",
    "    # now i want to check if the a_result is higher than the b_result\n",
    "    # if so, then the model predicts A as the label, otherwise predicts B\n",
    "    if a_score > b_score:\n",
    "        pred_label = 0\n",
    "    else:\n",
    "        pred_label = 1\n",
    "    \n",
    "    # check if the pred_label is correct\n",
    "    if pred_label == label:\n",
    "        scores.append(1)\n",
    "    else:\n",
    "        scores.append(0)\n",
    "\n",
    "    pbar.set_description(f\"Score: {np.mean(scores)}\")\n",
    "\n",
    "print(np.mean(scores))\n",
    "print(len(scores))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "1104e30d",
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "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
}
