{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import json\n",
    "import numpy as np\n",
    "import pandas as pd\n",
    "from datetime import datetime\n",
    "from dataclasses import dataclass, asdict\n",
    "from langdetect import detect"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "####### UPDATE THESE ##########\n",
    "\n",
    "model_A = \"/app2/suno/data/ditto_evals/2025_07_29-02_41_53/labelbox_auk_genre_mappings_2025_07_29-02_41_53.json\"\n",
    "model_B = \"/app2/suno/data/ditto_evals/2025_07_29-02_41_53/labelbox_bluejay_genre_mappings_2025_07_29-02_41_53.json\"\n",
    "\n",
    "model_A_name = \"auk_t1\"\n",
    "model_B_name = \"bluejay_t2\"\n",
    "\n",
    "model_bad = \"/app2/suno/data/ditto_evals/2025_07_29-02_41_53/labelbox_v35_genre_mappings_2025_07_29-02_41_53.json\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def load_gen_json(filepath):\n",
    "    with open(filepath, \"r\", encoding=\"utf-8\") as file:\n",
    "        data = json.load(file)\n",
    "    by_item = []\n",
    "    for genre, outputs in data.items():\n",
    "        for val in outputs:\n",
    "            val[\"genre\"] = genre\n",
    "            by_item.append(val)\n",
    "            val[\"instrumental\"] = len(val[\"lyrics\"]) < 15\n",
    "            if val[\"instrumental\"]:\n",
    "                val[\"lyrics\"] = \"[Instrumental]\"\n",
    "    df = pd.DataFrame(by_item)\n",
    "\n",
    "    return df"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "A_df = load_gen_json(model_A)\n",
    "B_df = load_gen_json(model_B)\n",
    "bad_df = load_gen_json(model_bad)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(len(A_df), len(B_df))\n",
    "B_df.head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "bad_df.head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "@dataclass\n",
    "class ClipDetails:\n",
    "    s3_id: str\n",
    "    lyrics: str\n",
    "    tags: str\n",
    "    genre: str\n",
    "    instrumental: bool\n",
    "\n",
    "\n",
    "def is_english(df):\n",
    "    filtered = df[df[\"instrumental\"] == False]\n",
    "    for text in filtered.lyrics.tolist():\n",
    "        assert detect(text) == \"en\"\n",
    "\n",
    "\n",
    "def make_clip_metadata(clip, model_name):\n",
    "    if isinstance(clip, ClipDetails):\n",
    "        clip_dict = asdict(clip)\n",
    "    else:\n",
    "        clip_dict = clip.to_dict()\n",
    "    clip_details = ClipDetails(**clip_dict)\n",
    "    clip_metadata = {\n",
    "        \"key\": f\"clip_{model_name}\",\n",
    "        \"name\": None,\n",
    "        \"url\": f\"https://cdn1.suno.ai/{clip_details.s3_id}.mp3\",\n",
    "        \"metadata\": asdict(clip_details),\n",
    "    }\n",
    "    clip_metadata[\"metadata\"][\"source\"] = model_name\n",
    "    return clip_metadata\n",
    "\n",
    "\n",
    "is_english(A_df)\n",
    "is_english(B_df)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "timestamp = datetime.now().strftime(\"%Y%m%d\")\n",
    "test_label = f\"{model_A_name}-vs-{model_B_name}-{timestamp}\"\n",
    "\n",
    "\n",
    "def make_pair(clip_0, clip_1, model_0_name, model_1_name):\n",
    "    clip_0_metadata = make_clip_metadata(clip_0, model_0_name)\n",
    "    clip_1_metadata = make_clip_metadata(clip_1, model_1_name)\n",
    "\n",
    "    example = []\n",
    "    if np.random.rand() > 0.5:  # display model A first\n",
    "        clip_0_metadata[\"name\"] = \"Clip A\"\n",
    "        example.append(clip_0_metadata)\n",
    "        clip_1_metadata[\"name\"] = \"Clip B\"\n",
    "        example.append(clip_1_metadata)\n",
    "    else:  # display model B first\n",
    "        clip_1_metadata[\"name\"] = \"Clip A\"\n",
    "        example.append(clip_1_metadata)\n",
    "        clip_0_metadata[\"name\"] = \"Clip B\"\n",
    "        example.append(clip_0_metadata)\n",
    "\n",
    "    return example\n",
    "\n",
    "\n",
    "metadata = []\n",
    "for i in range(0, len(A_df)):\n",
    "    clip_0 = A_df.iloc[i]\n",
    "    clip_1 = B_df.iloc[i]\n",
    "\n",
    "    example = make_pair(clip_0, clip_1, model_A_name, model_B_name)\n",
    "\n",
    "    metadata.append(example)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def generate_unique_pairings(model_B, model_bad, num_pairings=20):\n",
    "    \"\"\"Generate unique pairings ensuring no s3_id appears twice.\"\"\"\n",
    "    match_columns = [\"tags\", \"lyrics\", \"genre\", \"instrumental\"]\n",
    "\n",
    "    # Create match keys for both dataframes\n",
    "    for df in [model_B, model_bad]:\n",
    "        df[\"match_key\"] = df[match_columns].apply(lambda x: \"|||\".join(x.astype(str)), axis=1)\n",
    "\n",
    "    used_ids = {\"B\": set(), \"bad\": set()}\n",
    "    matches = []\n",
    "\n",
    "    # Shuffle for randomness and iterate through model_B\n",
    "    for _, row_B in model_B.sample(frac=1, random_state=42).iterrows():\n",
    "        if row_B[\"s3_id\"] in used_ids[\"B\"] or len(matches) >= num_pairings:\n",
    "            continue\n",
    "\n",
    "        # Find available matches in model_bad\n",
    "        available_matches = model_bad[\n",
    "            (model_bad[\"match_key\"] == row_B[\"match_key\"])\n",
    "            & (model_bad[\"s3_id\"] != row_B[\"s3_id\"])\n",
    "            & (~model_bad[\"s3_id\"].isin(used_ids[\"bad\"]))\n",
    "        ]\n",
    "\n",
    "        if not available_matches.empty:\n",
    "            selected = available_matches.sample(n=1, random_state=np.random.randint(0, 10000)).iloc[0]\n",
    "\n",
    "            matches.append(\n",
    "                {\n",
    "                    \"model_B_s3_id\": row_B[\"s3_id\"],\n",
    "                    \"model_bad_s3_id\": selected[\"s3_id\"],\n",
    "                    **{col: row_B[col] for col in match_columns},\n",
    "                }\n",
    "            )\n",
    "\n",
    "            used_ids[\"B\"].add(row_B[\"s3_id\"])\n",
    "            used_ids[\"bad\"].add(selected[\"s3_id\"])\n",
    "\n",
    "    if not matches:\n",
    "        print(\"No matching records found!\")\n",
    "    elif len(matches) < num_pairings:\n",
    "        print(f\"Only {len(matches)} unique pairs available.\")\n",
    "\n",
    "    return pd.DataFrame(matches)\n",
    "\n",
    "\n",
    "def create_honeypots(pairings_df, model_b_genre=\"model_b\", model_bad_genre=\"model_bad\"):\n",
    "    \"\"\"Create honeypots list with ClipDetails objects.\"\"\"\n",
    "\n",
    "    honeypots = []\n",
    "    for _, row in pairings_df.iterrows():\n",
    "        pair = [\n",
    "            ClipDetails(\n",
    "                s3_id=row[\"model_B_s3_id\"],\n",
    "                tags=row[\"tags\"],\n",
    "                lyrics=row[\"lyrics\"],\n",
    "                genre=model_b_genre,\n",
    "                instrumental=row[\"instrumental\"],\n",
    "            ),\n",
    "            ClipDetails(\n",
    "                s3_id=row[\"model_bad_s3_id\"],\n",
    "                tags=row[\"tags\"],\n",
    "                lyrics=row[\"lyrics\"],\n",
    "                genre=model_bad_genre,\n",
    "                instrumental=row[\"instrumental\"],\n",
    "            ),\n",
    "        ]\n",
    "        honeypots.append(pair)\n",
    "\n",
    "    return honeypots\n",
    "\n",
    "\n",
    "unique_pairs = generate_unique_pairings(B_df, bad_df, 20)\n",
    "honeypots = create_honeypots(unique_pairs, \"honeypot_good\", \"honeypot_bad\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "for honeypot in honeypots:\n",
    "    good = honeypot[0]\n",
    "    bad = honeypot[1]\n",
    "    hp_example = make_pair(good, bad, \"honeypot_good\", \"honeypot_bad\")\n",
    "    metadata.append(hp_example)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "metadata[-1]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# save json metadata\n",
    "metadata_filepath = f\"./outputs_internal/metadata-{test_label}.json\"\n",
    "with open(metadata_filepath, \"w\") as fp:\n",
    "    json.dump(metadata, fp, indent=2)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "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": 2
}
