{
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Select Preference Data"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:11:04.757824Z",
     "start_time": "2024-05-26T00:11:04.555293Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.807Z"
    }
   },
   "outputs": [],
   "source": [
    "# setup tailscale if you haven't\n",
    "# https://tailscale.com/kb/1031/install-linux\n",
    "!sudo tailscale up --accept-routes=true\n",
    "\n",
    "# setup autoload\n",
    "%load_ext autoreload\n",
    "%autoreload 2"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:11:08.392310Z",
     "start_time": "2024-05-26T00:11:04.759383Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.807Z"
    }
   },
   "outputs": [],
   "source": [
    "# make sure sqlalchemy is >=2\n",
    "# pip install psycopg2-binary\n",
    "# pip install \"sqlalchemy>=2\"\n",
    "import os\n",
    "import datetime\n",
    "from collections import defaultdict, Counter\n",
    "import json\n",
    "from urllib.parse import quote\n",
    "import time\n",
    "\n",
    "import boto3\n",
    "import matplotlib.pyplot as plt\n",
    "import numpy as np\n",
    "import pandas as pd\n",
    "import sqlalchemy\n",
    "import tqdm\n",
    "from botocore.exceptions import ClientError\n",
    "from suno_analytics.preference_helper import get_preference_counts\n",
    "from suno_analytics.preference_data_selection import (\n",
    "    gather_data,\n",
    "    gather_data_with_snowflake,\n",
    "    plot_clip_distribution,\n",
    "    parse_metadata_for_basics,\n",
    "    get_concat_clip_ids,\n",
    "    validate_preference_data,\n",
    "    run_bot_detection,\n",
    "    print_out_value_counts_nicely,\n",
    "    merge_concat_clips_with_reactions,\n",
    "    plot_clip_basic_distributions,\n",
    ")\n",
    "\n",
    "\n",
    "# setup some pandas display stuff\n",
    "pd.set_option(\"display.max_rows\", 500)\n",
    "pd.set_option(\"display.max_columns\", 500)\n",
    "pd.set_option(\"display.width\", 1000)\n",
    "\n",
    "\n",
    "def get_secret():\n",
    "    secret_name = \"app-user-main-db-secret\"\n",
    "    region_name = \"us-east-2\"\n",
    "    # Create a Secrets Manager client\n",
    "    session = boto3.session.Session()\n",
    "    client = session.client(service_name=\"secretsmanager\", region_name=region_name)\n",
    "    try:\n",
    "        get_secret_value_response = client.get_secret_value(SecretId=secret_name)\n",
    "    except ClientError as e:\n",
    "        raise e\n",
    "    secret = get_secret_value_response[\"SecretString\"]\n",
    "    return json.loads(secret)\n",
    "\n",
    "\n",
    "my_secrets = get_secret()\n",
    "\n",
    "# alternative...\n",
    "engine = sqlalchemy.create_engine(\n",
    "    \"postgresql://suno:%s@suno-main-postgres-prod-analytics.cnfvffydbwvc.us-east-2.rds.amazonaws.com/suno_main\"\n",
    "    % quote(my_secrets[\"password\"]),\n",
    ")\n",
    "\n",
    "\n",
    "home_dir = os.path.expanduser(\"~\")\n",
    "snow_password_path = os.path.join(home_dir, \".aws\", \"snow_pw.txt\")\n",
    "if os.path.exists(snow_password_path):\n",
    "    # !pip install snowflake\n",
    "    from snowflake.core import Root\n",
    "    from snowflake.snowpark import Session\n",
    "\n",
    "    with open(snow_password_path, \"r\") as fp:\n",
    "        fp_lines = fp.readlines()\n",
    "        snow_password = fp_lines[0].strip()\n",
    "        snow_username = fp_lines[1].strip()\n",
    "\n",
    "    CONNECTION_PARAMETERS = {\n",
    "        \"account\": \"fu90569.us-east-2.aws\",\n",
    "        \"user\": snow_username,\n",
    "        \"password\": snow_password,\n",
    "        \"role\": \"ACCOUNTADMIN\",\n",
    "        \"database\": \"SUNO_PROD\",\n",
    "        \"warehouse\": \"SUNO_PROD_LARGE\",\n",
    "        \"schema\": \"PROD\",\n",
    "    }"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Validate some info"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:11:08.550447Z",
     "start_time": "2024-05-26T00:11:08.397196Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.807Z"
    }
   },
   "outputs": [],
   "source": [
    "# there are 4 hr time difference between eastern time and utc\n",
    "# cutoff_date = \"2024-08-26 21:00:00\"  # v4-t3 out\n",
    "# cutoff_date = \"2024-09-12 21:00:00\"  # covers beta out\n",
    "# cutoff_date = \"2024-09-22 00:00:00\"  # pre fe exp out\n",
    "# cutoff_date = \"2024-09-26 12:00:00\"  # s29 out\n",
    "# cutoff_date = \"2024-10-09 15:20:00\"  # 30b t5 out\n",
    "# cutoff_date = \"2024-10-31 16:00:00\"  # 30b t6 out\n",
    "# cutoff_date = \"2024-11-12 13:00:00\"  # 30b t6-2 out\n",
    "# cutoff_date = \"2024-10-19 16:00:00\"  # v4 out\n",
    "# cutoff_date = \"2024-11-19 16:00:00\"  # v4 out\n",
    "# cutoff_date = \"2024-12-16 20:40:00\"  # v4 s32 out\n",
    "# cutoff_date = \"2025-02-21 22:15:00\"  # diff v5 out -- until 0411\n",
    "# cutoff_date = \"2025-03-06 19:00:00\"  # diff v6 out\n",
    "# cutoff_date = \"2025-03-24 00:00:00\"  #  diff v7 out\n",
    "# cutoff_date = \"2025-03-25 00:40:00\"  # test\n",
    "# cutoff_date = (\n",
    "#     (datetime.datetime.now() - datetime.timedelta(hours=4))\n",
    "#     .astimezone(datetime.timezone.utc)\n",
    "#     .strftime(\"%Y-%m-%d %H:%M:%S\")\n",
    "# )\n",
    "print(datetime.datetime.now(), time.time(), cutoff_date)\n",
    "\n",
    "# target_model_name = \"chirp-v4-h-t-6\"\n",
    "# target_model_name = \"chirp-v4-up-u-1\"\n",
    "target_model_name = \"chirp-v4-h-s-32\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:11:09.349857Z",
     "start_time": "2024-05-26T00:11:08.551408Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.807Z"
    }
   },
   "outputs": [],
   "source": [
    "df_all_tables = pd.read_sql_query(\n",
    "    \"SELECT table_name FROM information_schema.tables WHERE table_schema = 'public'\",\n",
    "    engine,\n",
    ")\n",
    "# should have all the basic table names here\n",
    "assert df_all_tables[\"table_name\"].nunique() >= 61"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Query the DB"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.807Z"
    }
   },
   "outputs": [],
   "source": [
    "if not os.path.exists(snow_password_path):\n",
    "    raise Exception(\"you are not authorized to access snowflake -- please setup\")\n",
    "\n",
    "snow_session = Session.builder.configs(CONNECTION_PARAMETERS).create()\n",
    "\n",
    "snow_root = Root(snow_session)\n",
    "snow_schema = snow_root.databases[\"SUNO_PROD\"].schemas[\"PROD\"]\n",
    "print(snow_schema.name)\n",
    "\n",
    "# from snowflake.snowpark.functions import col\n",
    "# !pip install \"snowflake-connector-python[pandas]\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.808Z"
    }
   },
   "outputs": [],
   "source": [
    "# gathered_data = gather_data(engine, cutoff_date)\n",
    "gathered_data = gather_data_with_snowflake(\n",
    "    snow_session, \n",
    "    cutoff_date,\n",
    "    # filter_model_name=target_model_name, \n",
    "    filter_play_count=1, \n",
    "    filter_user_n_clips=100\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.808Z"
    }
   },
   "outputs": [],
   "source": [
    "bots_action_df = gathered_data[\"bots_action_df\"]\n",
    "reaction_df = gathered_data[\"reaction_df\"]\n",
    "total_clip_df = gathered_data[\"total_clip_df\"]\n",
    "playlist_clip_df = gathered_data[\"playlist_clip_df\"]\n",
    "discord_info_df = gathered_data[\"discord_info_df\"]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.808Z"
    }
   },
   "outputs": [],
   "source": [
    "# parse out the necessary metadata early\n",
    "total_clip_df[\n",
    "    [\"continued_parent\", \"duration\", \"source\", \"clip_type\", \"task\", \"edited_clip_id\"]\n",
    "] = pd.DataFrame(\n",
    "    total_clip_df[\"metadata\"].map(parse_metadata_for_basics).tolist(),\n",
    "    index=total_clip_df.index,\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.808Z"
    }
   },
   "outputs": [],
   "source": [
    "# filter on versions\n",
    "clip_df = total_clip_df.copy()\n",
    "total_clip_counts = clip_df.shape[0]\n",
    "print(f\"total clips: {total_clip_counts}\")\n",
    "print_out_value_counts_nicely(clip_df, \"clip_type\")\n",
    "# check the number of audio uploads\n",
    "upload_clip_df = total_clip_df[total_clip_df[\"clip_type\"] == \"upload\"].copy()\n",
    "stem_clip_df = total_clip_df[total_clip_df[\"clip_type\"] == \"stem\"].copy()\n",
    "print(\"total without model:\", (total_clip_df[\"model_name\"] == \"\").sum())\n",
    "\n",
    "# Call the function\n",
    "plot_clip_distribution(total_clip_df)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Proceed with feature engineering and cleaning up"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.808Z"
    }
   },
   "outputs": [],
   "source": [
    "upvoted_df = reaction_df[reaction_df[\"reaction_type\"] == \"L\"].copy()\n",
    "print(f\"number of upvoates: {upvoted_df.shape[0]:,} rows\")\n",
    "upvoted_ids = upvoted_df[\"clip_id\"]\n",
    "\n",
    "flagged_df = reaction_df[reaction_df[\"flagged\"]].copy()\n",
    "print(f\"number of flagged reports: {flagged_df.shape[0]:,} rows\")\n",
    "flagged_ids = flagged_df[\"clip_id\"]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.808Z"
    }
   },
   "outputs": [],
   "source": [
    "# this is probably the right way to figure out the pro user group\n",
    "pro_users = set(discord_info_df[\"user_id\"].unique())\n",
    "reaction_df[\"is_pro_user\"] = reaction_df[\"user_id\"].isin(pro_users)\n",
    "clip_df[\"is_pro_user\"] = clip_df[\"user_id\"].isin(pro_users)\n",
    "\n",
    "# this is very interesting....\n",
    "# reaction check\n",
    "print(\"Reactions fraction by pro user:\")\n",
    "print_out_value_counts_nicely(reaction_df, \"is_pro_user\")\n",
    "print(\"------------\")\n",
    "# clip check\n",
    "print(\"Clip generated fraction by pro user:\")\n",
    "print_out_value_counts_nicely(clip_df, \"is_pro_user\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.808Z"
    }
   },
   "outputs": [],
   "source": [
    "# find out the stem parent ids\n",
    "stem_parent_ids = set(\n",
    "    stem_clip_df[\"metadata\"].apply(lambda x: x.get(\"stem_from_id\", \"xxx\"))\n",
    ")\n",
    "print(\"stem parent ids:\", len(stem_parent_ids))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:23:12.019778Z",
     "start_time": "2024-05-26T00:22:57.637371Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.808Z"
    }
   },
   "outputs": [],
   "source": [
    "# add clip is in playlist feature\n",
    "clip_df[\"is_in_playlist\"] = clip_df[\"id\"].isin(playlist_clip_df[\"clip_id\"].unique())\n",
    "print(\"Clips in a splaylist:\")\n",
    "print_out_value_counts_nicely(clip_df, \"is_in_playlist\")\n",
    "clip_df[\"has_stems\"] = clip_df[\"id\"].astype(str).isin(stem_parent_ids)\n",
    "print(\"------------\")\n",
    "print(\"Clips has stem children:\")\n",
    "print_out_value_counts_nicely(clip_df, \"has_stems\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:23:17.912428Z",
     "start_time": "2024-05-26T00:23:12.021726Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.808Z"
    }
   },
   "outputs": [],
   "source": [
    "# parse the metadata for histories and types\n",
    "clip_history_df = clip_df[~clip_df[\"continued_parent\"].isna()].copy()\n",
    "# these are the direct parent's ids -- not grandparents\n",
    "has_continued_children_ids = clip_history_df[\"continued_parent\"]\n",
    "print(\n",
    "    \"clips that have children:\",\n",
    "    len(has_continued_children_ids),\n",
    "    \"\\nclips that are parents:\",\n",
    "    has_continued_children_ids.nunique(),\n",
    "    \"\\n\",\n",
    "    \"Average continues from clip = \",\n",
    "    round(\n",
    "        len(has_continued_children_ids) / (len(has_continued_children_ids.unique()) + 1), 2\n",
    "    ),\n",
    ")\n",
    "# Get value counts\n",
    "print_out_value_counts_nicely(clip_df, \"source\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:23:22.092835Z",
     "start_time": "2024-05-26T00:23:17.914456Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.808Z"
    }
   },
   "outputs": [],
   "source": [
    "print(\"total uploads:\", (clip_df[\"model_name\"] == \"\").sum())\n",
    "print(\"clips without request id:\", (clip_df[\"request_id\"].isna()).sum())\n",
    "# the nans are concats, we want to drop them for now\n",
    "concated_clips = clip_df[\n",
    "    (clip_df[\"clip_type\"] == \"concat\") | (clip_df[\"clip_type\"] == \"concat_infilling\")\n",
    "].copy()\n",
    "non_request_clips = clip_df[clip_df[\"request_id\"].isna()].copy()\n",
    "print(\n",
    "    \"clips without request id:\",\n",
    "    non_request_clips.shape[0],\n",
    "    non_request_clips[\"clip_type\"].value_counts(),\n",
    ")\n",
    "# need to kick them out...\n",
    "clip_df = clip_df[~clip_df[\"request_id\"].isna()]\n",
    "print(\n",
    "    f\"Clips without request id (concat, uploads...) frac = {concated_clips.shape[0] / total_clip_counts:.5f}\"\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:23:22.833148Z",
     "start_time": "2024-05-26T00:23:22.094796Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.808Z"
    }
   },
   "outputs": [],
   "source": [
    "# check the model conts\n",
    "print_out_value_counts_nicely(clip_df, \"model_name\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:23:26.317699Z",
     "start_time": "2024-05-26T00:23:22.835074Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.808Z"
    }
   },
   "outputs": [],
   "source": [
    "print(\"pre-filter model type clip_df shape:\", clip_df.shape)\n",
    "clip_df = clip_df[\n",
    "    (clip_df[\"model_name\"] != \"chirp-v3-5\")\n",
    "    & (clip_df[\"model_name\"] != \"chirp-v3-0\")\n",
    "    & (clip_df[\"model_name\"] != \"chirp-v3-5-tau\")\n",
    "    & (clip_df[\"model_name\"] != \"chirp-v3-5-upload\")\n",
    "    & (clip_df[\"model_name\"] != \"chirp-v3-5-short\")\n",
    "    & (clip_df[\"model_name\"] != \"chirp-v4\")\n",
    "    & (clip_df[\"model_name\"] != \"chirp-v4-tau\")\n",
    "    & (clip_df[\"model_name\"] != \"chirp-up\")\n",
    "]\n",
    "print(\"post-filter model type clip_df shape:\", clip_df.shape)\n",
    "print_out_value_counts_nicely(clip_df, \"model_name\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.808Z"
    }
   },
   "outputs": [],
   "source": [
    "concated_clips = merge_concat_clips_with_reactions(concated_clips, reaction_df)\n",
    "# TODO: why so many clips are concats without plays??? -- oh probably they concat multiple times?\n",
    "print(\"All concats\", concated_clips.shape[0])\n",
    "concated_clips = concated_clips[concated_clips[\"reaction_play_count\"] > 0]\n",
    "print(\"total concats with plays\", concated_clips.shape[0])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:23:38.974479Z",
     "start_time": "2024-05-26T00:23:34.385507Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.808Z"
    }
   },
   "outputs": [],
   "source": [
    "concat_clips_ids = get_concat_clip_ids(concated_clips, clip_df, upload_clip_df)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Features"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:23:40.236690Z",
     "start_time": "2024-05-26T00:23:39.995713Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.808Z"
    }
   },
   "outputs": [],
   "source": [
    "# set user number of clips generated\n",
    "if \"user_n_clips\" not in clip_df.columns:\n",
    "    clip_df[\"user_n_clips\"] = clip_df[\"user_id\"].map(clip_df[\"user_id\"].value_counts())\n",
    "    print(clip_df[\"user_n_clips\"].describe())"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:23:44.826860Z",
     "start_time": "2024-05-26T00:23:40.238330Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.808Z"
    }
   },
   "outputs": [],
   "source": [
    "# add upvoted column\n",
    "clip_df[\"upvoted\"] = clip_df[\"id\"].isin(upvoted_ids)\n",
    "print(\n",
    "    \"has upvoted\",\n",
    "    clip_df[\"upvoted\"].value_counts(),\n",
    "    clip_df[\"upvoted\"].value_counts(normalize=True),\n",
    "    (clip_df[\"upvote_count\"] >= 1).value_counts(normalize=True),\n",
    "    (clip_df[\"upvote_count\"] > 1).value_counts(normalize=True),\n",
    ")\n",
    "# clip_df = clip_df.drop(columns=['upvote_count'])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:23:48.404433Z",
     "start_time": "2024-05-26T00:23:44.857788Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.808Z"
    }
   },
   "outputs": [],
   "source": [
    "disliked_ids = reaction_df[reaction_df[\"reaction_type\"] == \"D\"][\"clip_id\"].unique()\n",
    "\n",
    "clip_df[\"downvoted\"] = clip_df[\"id\"].isin(disliked_ids)\n",
    "print(\"downvoted fraction by category:\")\n",
    "print_out_value_counts_nicely(clip_df, \"downvoted\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:23:56.590022Z",
     "start_time": "2024-05-26T00:23:48.405668Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.808Z"
    }
   },
   "outputs": [],
   "source": [
    "# add continued column -- uuid and str are not compatible X.x\n",
    "clip_df[\"has_continued\"] = (\n",
    "    clip_df[\"id\"].astype(str).isin(set(list(has_continued_children_ids)))\n",
    ")\n",
    "print(\"has_continued fraction by category:\")\n",
    "print_out_value_counts_nicely(clip_df, \"has_continued\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.808Z"
    }
   },
   "outputs": [],
   "source": [
    "print(\n",
    "    \"has upvoted in exp\",\n",
    "    # clip_df[in_exp_mask][\"upvoted\"].value_counts(),\n",
    "    round(clip_df[\"upvoted\"].value_counts(normalize=True)[True], 5),\n",
    "    # (clip_df[clip_df[\"in_fe_exp\"]][\"upvote_count\"] >= 1).value_counts(normalize=True),\n",
    ")\n",
    "print(\n",
    "    \"has downvoted out of exp\",\n",
    "    # clip_df[out_exp_mask][\"downvoted\"].value_counts(),\n",
    "    round(clip_df[\"downvoted\"].value_counts(normalize=True)[True], 5),\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:24:04.704306Z",
     "start_time": "2024-05-26T00:23:56.591353Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.808Z"
    }
   },
   "outputs": [],
   "source": [
    "# add concat column\n",
    "clip_df[\"part_of_concat\"] = clip_df[\"id\"].astype(str).isin(concat_clips_ids)\n",
    "print(\"part_of_concat fraction by category:\")\n",
    "print_out_value_counts_nicely(clip_df, \"part_of_concat\")\n",
    "\n",
    "print(\"------------\")\n",
    "print(\"Model distribution for part_of_concat clips:\")\n",
    "for model, fraction in (\n",
    "    clip_df[clip_df[\"part_of_concat\"]][\"model_name\"]\n",
    "    .value_counts(normalize=True)\n",
    "    .items()\n",
    "):\n",
    "    print(f\"{model}: {fraction:.2%}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:24:09.954233Z",
     "start_time": "2024-05-26T00:24:04.705554Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.808Z"
    }
   },
   "outputs": [],
   "source": [
    "# verify bots action are all non-empty\n",
    "bots_action_df.fillna(0, inplace=True)\n",
    "action_mask = (\n",
    "    bots_action_df[\"download_audio_count\"]\n",
    "    + bots_action_df[\"download_video_count\"]\n",
    "    + bots_action_df[\"download_audio_wav_count\"]\n",
    "    + bots_action_df[\"share_count\"] # will remove share cause it can be negative, just can be...\n",
    ") >= 1\n",
    "has_action_ids = set(i for i in bots_action_df[action_mask][\"clip_id\"].unique())\n",
    "clip_df[\"has_action\"] = clip_df[\"id\"].isin(has_action_ids)\n",
    "print(\"has_action fraction by category:\")\n",
    "print_out_value_counts_nicely(clip_df, \"has_action\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:24:14.579841Z",
     "start_time": "2024-05-26T00:24:09.955568Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.808Z"
    }
   },
   "outputs": [],
   "source": [
    "# add downvoted column\n",
    "clip_df[\"flagged\"] = clip_df[\"id\"].isin(flagged_ids)\n",
    "print(\"flagged fraction by category:\")\n",
    "print_out_value_counts_nicely(clip_df, \"flagged\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.809Z"
    }
   },
   "outputs": [],
   "source": [
    "clip_df[\"deleted\"] = clip_df[\"is_deleted\"]\n",
    "print(\"deleted fraction by category:\")\n",
    "print_out_value_counts_nicely(clip_df, \"deleted\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.809Z"
    }
   },
   "outputs": [],
   "source": [
    "edit_id_counts = clip_df[\"edited_clip_id\"].value_counts()\n",
    "clip_df[\"n_edits\"] = clip_df[\"id\"].map(edit_id_counts)\n",
    "print(\n",
    "    \"number of edits per clip:\",\n",
    "    clip_df[\"n_edits\"].mean(),\n",
    "    \"median\",\n",
    "    clip_df[\"n_edits\"].median(),\n",
    "    \"std\",\n",
    "    clip_df[\"n_edits\"].std(),\n",
    ")\n",
    "# print_out_value_counts_nicely(clip_df, \"n_edits\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:24:15.382315Z",
     "start_time": "2024-05-26T00:24:14.581073Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.809Z"
    }
   },
   "outputs": [],
   "source": [
    "# This is probably the most important cell of this notebook -- what are good labels, and not having good label makes it a bad label\n",
    "must_be_positive_mask = (\n",
    "    (clip_df[\"upvoted\"])\n",
    "    | (clip_df[\"has_action\"])\n",
    "    | (clip_df[\"part_of_concat\"])\n",
    "    | (clip_df[\"is_in_playlist\"])\n",
    "    | (\n",
    "        clip_df[\"n_edits\"] >= 10\n",
    "    )  # has more edit operations (upsample, cover, extend, etc)\n",
    ")\n",
    "must_be_not_negative_mask = (\n",
    "    (~clip_df[\"downvoted\"]) & (~clip_df[\"deleted\"]) & (~clip_df[\"flagged\"])\n",
    ")\n",
    "must_be_negative_mask = (\n",
    "    (clip_df[\"downvoted\"]) | (clip_df[\"flagged\"]) | (clip_df[\"deleted\"])\n",
    ")\n",
    "total_clips_count = clip_df.shape[0]\n",
    "must_be_positive_count = sum(must_be_positive_mask)\n",
    "definitely_not_negative_count = sum(must_be_not_negative_mask)\n",
    "must_be_negative_count = sum(must_be_negative_mask)\n",
    "\n",
    "print(\n",
    "    f\"Total clips: {total_clips_count:,}\\n\"\n",
    "    f\"Must be positive: {must_be_positive_count:,} ({must_be_positive_count/total_clips_count:.2%})\\n\"\n",
    "    f\"Definitely not negative: {definitely_not_negative_count:,} ({definitely_not_negative_count/total_clips_count:.2%})\\n\"\n",
    "    f\"Must be negative: {must_be_negative_count:,} ({must_be_negative_count/total_clips_count:.2%})\"\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:24:31.095990Z",
     "start_time": "2024-05-26T00:24:15.383572Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.809Z"
    }
   },
   "outputs": [],
   "source": [
    "mask = must_be_positive_mask & must_be_not_negative_mask\n",
    "total_unique_requests = clip_df[\"request_id\"].nunique()\n",
    "liked_requests = clip_df[mask][\"request_id\"].unique()  # requests with at least 1 like\n",
    "unliked_requests = clip_df[~mask][\"request_id\"].unique()  # requests without like\n",
    "has_liked_requests = set(liked_requests).intersection(\n",
    "    set(unliked_requests)\n",
    ")  # the request must have 1 like and one without like\n",
    "print(f\"Liked requests: {len(liked_requests):,}\")\n",
    "print(f\"Not liked requests: {len(unliked_requests):,}\")\n",
    "print(f\"Requests with preference paired generations: {len(has_liked_requests):,}\")\n",
    "print(\n",
    "    f\"Percentage of total unique requests: {len(has_liked_requests) / total_unique_requests:.2%}\"\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.809Z"
    }
   },
   "outputs": [],
   "source": [
    "# introduce a negative preference count\n",
    "has_disliked_half_requests = clip_df[must_be_negative_mask][\n",
    "    \"request_id\"\n",
    "].unique()  # requests with at least 1 dislike\n",
    "not_have_disliked_requests = clip_df[~must_be_negative_mask][\n",
    "    \"request_id\"\n",
    "].unique()  # request without dislike\n",
    "has_disliked_requests = set(has_disliked_half_requests).intersection(\n",
    "    set(not_have_disliked_requests)\n",
    ")  # the request must have 1 dislike and one without dislike\n",
    "print(f\"Disliked requests: {len(has_disliked_half_requests):,}\")\n",
    "print(f\"Not disliked requests: {len(not_have_disliked_requests):,}\")\n",
    "print(f\"Requests with preference paired generations: {len(has_disliked_requests):,}\")\n",
    "print(\n",
    "    f\"Percentage of total unique requests: {len(has_disliked_requests) / total_unique_requests:.2%}\"\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:24:31.099372Z",
     "start_time": "2024-05-26T00:24:31.097244Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.809Z"
    }
   },
   "outputs": [],
   "source": [
    "requests = has_liked_requests.union(has_disliked_requests)\n",
    "print(f\"Total selected pairs of requests: {len(requests):,}\")\n",
    "print(\n",
    "    f\"Percentage of total unique requests: {len(requests) / total_unique_requests:.2%}\"\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:24:31.239254Z",
     "start_time": "2024-05-26T00:24:31.100389Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.809Z"
    }
   },
   "outputs": [],
   "source": [
    "# this used to be a terrible bug...X.x\n",
    "assert mask.shape[0] == clip_df.shape[0]\n",
    "clip_df[\"pos_preference\"] = mask\n",
    "clip_df[\"neg_preference\"] = must_be_negative_mask\n",
    "# note that this is along the same row, so a positive clip can't be negative\n",
    "clip_df[\"diff_preference\"] = clip_df[\"pos_preference\"].astype(int) - clip_df[\n",
    "    \"neg_preference\"\n",
    "].astype(int)\n",
    "print(\"Difference in preference counts:\")\n",
    "value_counts = clip_df[\"diff_preference\"].value_counts()\n",
    "total = value_counts.sum()\n",
    "for value, count in value_counts.items():\n",
    "    fraction = count / total\n",
    "    print(f\"{value}: {count:,} ({fraction:.2%})\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:24:37.322031Z",
     "start_time": "2024-05-26T00:24:31.240829Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.809Z"
    }
   },
   "outputs": [],
   "source": [
    "# creation of interesting_clips\n",
    "interesting_clips = clip_df[clip_df[\"request_id\"].isin(requests)].copy()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.809Z"
    }
   },
   "outputs": [],
   "source": [
    "interesting_clips = interesting_clips.sort_values(\n",
    "    by=[\"request_id\", \"diff_preference\"]\n",
    ").reset_index()\n",
    "interesting_clips[\n",
    "    [\"request_id\", \"pos_preference\", \"neg_preference\", \"diff_preference\"]\n",
    "].head(n=6)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.809Z"
    }
   },
   "outputs": [],
   "source": [
    "# this is a mix now\n",
    "value_counts = interesting_clips[\"diff_preference\"].value_counts()\n",
    "total = value_counts.sum()\n",
    "for value, count in value_counts.items():\n",
    "    fraction = count / total\n",
    "    print(f\"Difference {value}: {count:,} ({fraction:.2%})\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.809Z"
    }
   },
   "outputs": [],
   "source": [
    "diff_series = interesting_clips[\"diff_preference\"].diff()\n",
    "value_counts = diff_series[1::2].value_counts()\n",
    "total = value_counts.sum()\n",
    "for value, count in value_counts.items():\n",
    "    fraction = count / total\n",
    "    print(f\"Value {value}: {count:,} ({fraction:.2%})\")\n",
    "# 1 is pos, not neg pair or nothing, neg; 2 is pos / neg (hence the larger difference)\n",
    "# there are only two values for this positive pair\n",
    "assert diff_series[1::2].nunique() == 2"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:24:38.960130Z",
     "start_time": "2024-05-26T00:24:37.323369Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.809Z"
    }
   },
   "outputs": [],
   "source": [
    "# assign the labels now\n",
    "interesting_clips[\"preference\"] = interesting_clips.index % 2 == 1\n",
    "# get df of requests -- let's move on!\n",
    "print(f\"Number of unique request_ids: {interesting_clips['request_id'].nunique():,}\")\n",
    "print(f\"Number of unique ids: {interesting_clips['id'].nunique():,}\")\n",
    "validate_preference_data(interesting_clips)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:24:43.332222Z",
     "start_time": "2024-05-26T00:24:43.166461Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.809Z"
    }
   },
   "outputs": [],
   "source": [
    "# # listen to some pairs\n",
    "# test_requests = interesting_clips[\"request_id\"].sample(10)\n",
    "\n",
    "# for i in range(1):\n",
    "#     rows = interesting_clips[interesting_clips[\"request_id\"] == test_requests.iloc[i]]\n",
    "#     assert rows.shape[0] == 2\n",
    "#     # Audio.from_s3(f\"s3://suno-data-uploads/studio/uploads/{row['s3_id']}.mp3\").play()\n",
    "#     # sort by likes\n",
    "#     rows = rows.sort_values(\"upvoted\", ascending=True)\n",
    "#     print(rows.iloc[0][\"prompt_text\"])\n",
    "#     print(rows.iloc[0][\"metadata\"])\n",
    "#     for _, row in rows.iterrows():\n",
    "#         print(row[\"id\"], row[\"preference\"], row[\"upvoted\"])\n",
    "#         Audio.from_s3(\n",
    "#             f\"s3://suno-data-uploads/studio/uploads/{row['s3_id']}.mp3\"\n",
    "#         ).play()\n",
    "#         with open_from_s3(\n",
    "#             f\"s3://suno-data-uploads/studio/uploads/{row['s3_id']}.npz\", as_binary=True\n",
    "#         ) as f:\n",
    "#             # read numpy array\n",
    "#             npz_a = np.load(f)\n",
    "#             if \"v1_raw\" in npz_a:\n",
    "#                 a = np.load(f)[\"v1_raw\"]\n",
    "#             elif \"v3.0_raw\" in npz_a:\n",
    "#                 a = np.load(f)[\"v3.0_raw\"]\n",
    "#             else:\n",
    "#                 print(\"npz_a\", npz_a)\n",
    "#                 raise ValueError\n",
    "#             print(a.shape)\n",
    "#     print()"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Further cuts and selections"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.809Z"
    }
   },
   "outputs": [],
   "source": [
    "# need the reaction play counts\n",
    "# Filter reaction_df for relevant clip_ids\n",
    "partial_reaction_df = reaction_df[\n",
    "    reaction_df[\"clip_id\"].isin(set(interesting_clips[\"id\"]))\n",
    "].copy()\n",
    "\n",
    "# Calculate total play counts\n",
    "total_play_counts = (\n",
    "    partial_reaction_df.groupby(\"clip_id\")[\"play_count\"].sum().reset_index()\n",
    ")\n",
    "total_play_counts = total_play_counts.rename(\n",
    "    columns={\"clip_id\": \"id\", \"play_count\": \"reaction_play_count\"}\n",
    ")\n",
    "\n",
    "# Calculate pro user play counts\n",
    "pro_play_counts = (\n",
    "    partial_reaction_df[partial_reaction_df[\"is_pro_user\"]]\n",
    "    .groupby(\"clip_id\")[\"play_count\"]\n",
    "    .sum()\n",
    "    .reset_index()\n",
    ")\n",
    "pro_play_counts = pro_play_counts.rename(\n",
    "    columns={\"clip_id\": \"id\", \"play_count\": \"reaction_pro_play_count\"}\n",
    ")\n",
    "\n",
    "# Merge with user_intersting_clips\n",
    "interesting_clips = interesting_clips.merge(total_play_counts, on=\"id\", how=\"left\")\n",
    "interesting_clips = interesting_clips.merge(pro_play_counts, on=\"id\", how=\"left\")\n",
    "\n",
    "print(f\"Number of interesting clips: {len(interesting_clips):,}\")\n",
    "# Get unique counts for request_id and id\n",
    "unique_request_ids = interesting_clips[\"request_id\"].nunique()\n",
    "unique_clip_ids = interesting_clips[\"id\"].nunique()\n",
    "\n",
    "# Print the results in a formatted manner\n",
    "print(\"Unique request and clip counts in interesting_clips:\")\n",
    "print(f\"{'Request IDs:':<15} {unique_request_ids:,}\")\n",
    "print(f\"{'Clip IDs:':<15} {unique_clip_ids:,}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:24:43.737067Z",
     "start_time": "2024-05-26T00:24:43.563216Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.809Z"
    }
   },
   "outputs": [],
   "source": [
    "preference_counts = interesting_clips.groupby(\"batch_index\")[\n",
    "    \"preference\"\n",
    "].value_counts()\n",
    "total_counts = preference_counts.groupby(level=0).sum()\n",
    "\n",
    "print(\"Preference counts and fractions by batch index:\")\n",
    "print(\"-\" * 50)\n",
    "for batch_index in [0, 1]:\n",
    "    print(f\"Batch Index: {batch_index}\")\n",
    "    for preference in [False, True]:\n",
    "        count = preference_counts[batch_index, preference]\n",
    "        fraction = count / total_counts[batch_index]\n",
    "        print(f\"  Preference {preference}: Count: {count:,} Fraction: {fraction:.2%}\")\n",
    "    print()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:24:44.304573Z",
     "start_time": "2024-05-26T00:24:43.973218Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.809Z"
    }
   },
   "outputs": [],
   "source": [
    "print_out_value_counts_nicely(interesting_clips, \"model_name\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:24:44.533983Z",
     "start_time": "2024-05-26T00:24:44.305722Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.809Z"
    }
   },
   "outputs": [],
   "source": [
    "print(\"Time Validation:\")\n",
    "print(\"-\" * 20)\n",
    "print(\"Interesting Clips:\")\n",
    "print(f\"  Earliest: {interesting_clips['created_at'].min()}\")\n",
    "print(f\"  Latest:   {interesting_clips['created_at'].max()}\")\n",
    "print(\"\\nAll Clips:\")\n",
    "print(f\"  Earliest: {clip_df['created_at'].min()}\")\n",
    "print(f\"  Latest:   {clip_df['created_at'].max()}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:24:50.674838Z",
     "start_time": "2024-05-26T00:24:46.366677Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.809Z"
    }
   },
   "outputs": [],
   "source": [
    "# make sure we sort here before proceed\n",
    "interesting_clips = interesting_clips.sort_values(by=[\"request_id\", \"preference\"])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:24:51.600621Z",
     "start_time": "2024-05-26T00:24:50.676186Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.809Z"
    }
   },
   "outputs": [],
   "source": [
    "# Calculate the ratio of preferred clips to total clips for each model\n",
    "clip_df_model_counts = clip_df[\"model_name\"].value_counts()\n",
    "preference_ratio = (\n",
    "    interesting_clips[interesting_clips[\"preference\"]][\"model_name\"].value_counts()\n",
    "    / clip_df_model_counts\n",
    ")\n",
    "\n",
    "# Print the results in a formatted manner\n",
    "print(\"Ratio of preferred clips to total clips for each model:\")\n",
    "print(\"-\" * 60)\n",
    "for model, ratio in preference_ratio.items():\n",
    "    n = clip_df_model_counts[model]\n",
    "    uncertainty = (ratio * (1 - ratio) / n) ** 0.5\n",
    "    print(f\"{model:<30} {ratio:.2%} ± {uncertainty:.2%}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:24:56.302599Z",
     "start_time": "2024-05-26T00:24:51.601863Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.812Z"
    }
   },
   "outputs": [],
   "source": [
    "get_preference_counts(interesting_clips)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:24:57.443236Z",
     "start_time": "2024-05-26T00:24:56.306473Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.812Z"
    }
   },
   "outputs": [],
   "source": [
    "plot_clip_basic_distributions(interesting_clips)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:24:58.153310Z",
     "start_time": "2024-05-26T00:24:57.858363Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.812Z"
    }
   },
   "outputs": [],
   "source": [
    "# FUCK THIS FOR NOW\n",
    "# MAX_PREFERENCE_PER_USER = 400\n",
    "# grouped_interesting_clips = interesting_clips.groupby([\"user_id\"])\n",
    "# user_top_df = (\n",
    "#     interesting_clips.sort_values(\n",
    "#         [\"preference\", \"upvote_count\", \"part_of_concat\", \"is_in_playlist\"], ascending=False\n",
    "#     )\n",
    "#     .groupby(\"user_id\")\n",
    "#     .head(MAX_PREFERENCE_PER_USER)\n",
    "# )\n",
    "# print(user_top_df.shape, interesting_clips.shape)\n",
    "\n",
    "# user_top_requests = user_top_df[\"request_id\"].unique()\n",
    "# user_intersting_clips = interesting_clips[\n",
    "#     interesting_clips[\"request_id\"].isin(user_top_requests)\n",
    "# ].copy()\n",
    "# print(user_intersting_clips.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:24:58.890520Z",
     "start_time": "2024-05-26T00:24:58.154350Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "# subselect interesting clips\n",
    "interesting_clips_masks = (interesting_clips[\"model_name\"].str.contains(\"v3p5|v4\")) & (\n",
    "    interesting_clips[\"reaction_play_count\"] > 0\n",
    ")\n",
    "# make sure we have pairs\n",
    "extra_compare_mask = interesting_clips[interesting_clips_masks][\"request_id\"].isin(\n",
    "    interesting_clips[interesting_clips_masks][\"request_id\"]\n",
    "    .value_counts()\n",
    "    .index[interesting_clips[interesting_clips_masks][\"request_id\"].value_counts() == 2]\n",
    ")\n",
    "user_intersting_clips = interesting_clips[\n",
    "    interesting_clips_masks & extra_compare_mask\n",
    "].copy()\n",
    "\n",
    "print(\"Number of clips in interesting_clips:\")\n",
    "print(f\"{interesting_clips.shape[0]:,}\")\n",
    "print(\"Number of clips in user_interesting_clips:\")\n",
    "print(f\"{user_intersting_clips.shape[0]:,}\")\n",
    "# Calculate the ratio of preferred clips to total clips for each model\n",
    "preference_ratio = (\n",
    "    user_intersting_clips[user_intersting_clips[\"preference\"]][\n",
    "        \"model_name\"\n",
    "    ].value_counts()\n",
    "    / clip_df_model_counts\n",
    ")\n",
    "\n",
    "# Print the results in a formatted manner\n",
    "print(\"Ratio of preferred clips to total clips for each model:\")\n",
    "print(\"-\" * 60)\n",
    "for model, ratio in preference_ratio.items():\n",
    "    n = clip_df_model_counts[model]\n",
    "    uncertainty = (ratio * (1 - ratio) / n) ** 0.5\n",
    "    print(f\"{model:<30} {ratio:.2%} ± {uncertainty:.2%}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:24:59.204547Z",
     "start_time": "2024-05-26T00:24:58.891851Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "# Calculate the number of preferences per user\n",
    "preferences_per_user = user_intersting_clips[\"user_id\"].value_counts()\n",
    "\n",
    "# Determine the maximum number of preferences\n",
    "max_preferences = preferences_per_user.max()\n",
    "\n",
    "# Choose bins using Sturges' rule, but ensure a minimum of 15 bins and a maximum of 30\n",
    "n_bins = max(30, min(100, int(np.ceil(np.log2(len(preferences_per_user)) + 1))))\n",
    "\n",
    "# Calculate bin edges using a linear scale\n",
    "bin_edges = np.linspace(preferences_per_user.min(), max_preferences, n_bins)\n",
    "\n",
    "plt.figure(figsize=(10, 6))\n",
    "plt.hist(preferences_per_user, bins=bin_edges, edgecolor=\"black\")\n",
    "plt.yscale(\"log\")\n",
    "plt.xlabel(\"Number of preferences per user\")\n",
    "plt.ylabel(\"Number of users (log scale)\")\n",
    "plt.title(\"Distribution of User Preferences\")\n",
    "plt.grid(axis=\"both\", linestyle=\"--\", alpha=0.7)\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:24:59.207707Z",
     "start_time": "2024-05-26T00:24:59.205659Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "print(\"Summary of user_interesting_clips:\")\n",
    "print(f\"Total requests: {user_intersting_clips.shape[0]:,}\")\n",
    "print(f\"Unique clips: {user_intersting_clips.shape[0] // 2:,}\")\n",
    "print(\n",
    "    f\"Fraction of total clips: {user_intersting_clips.shape[0] / total_clip_counts:.2%}\"\n",
    ")\n",
    "print(\"Time Validation:\")\n",
    "print(f\"Earliest timestamp: {user_intersting_clips['created_at'].min()}\")\n",
    "print(f\"Latest timestamp:   {user_intersting_clips['created_at'].max()}\")\n",
    "model_to_test = target_model_name\n",
    "print(\n",
    "    f\"Earliest timestamp: {user_intersting_clips[user_intersting_clips['model_name'] == model_to_test]['created_at'].min()}\"\n",
    ")\n",
    "print(\n",
    "    f\"Latest timestamp:   {user_intersting_clips[user_intersting_clips['model_name'] == model_to_test]['created_at'].max()}\"\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "def parse_for_instrumental(x):\n",
    "    if \"make_instrumental\" not in x:\n",
    "        return False\n",
    "    out = x.get(\"make_instrumental\", False)\n",
    "    return out\n",
    "\n",
    "\n",
    "# from suno_analytics.preference_data_selection import parse_for_tag, parse_for_one_box\n",
    "# user_intersting_clips[\"tags\"] = user_intersting_clips[\"metadata\"].apply(parse_for_tag)\n",
    "# user_intersting_clips[\"is_onebox\"] = user_intersting_clips[\"metadata\"].apply(parse_for_one_box)\n",
    "# user_intersting_clips[\"is_instrumental\"] = user_intersting_clips[\"metadata\"].apply(parse_for_instrumental)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:25:00.893917Z",
     "start_time": "2024-05-26T00:25:00.485775Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "user_compare_mask = (\n",
    "    user_intersting_clips[\"created_at\"] >= cutoff_date\n",
    "    # & (\n",
    "    #     (user_intersting_clips[\"model_name\"].str.startswith(\"chirp-v3p5-engine-t\"))\n",
    "    #     | (user_intersting_clips[\"model_name\"].str.startswith(\"chirp-v3p5-engine-s\"))\n",
    "    # )\n",
    "    # & (~user_intersting_clips[\"is_pro_user\"])\n",
    "    # & (~user_intersting_clips[\"is_onebox\"])\n",
    "    # & user_intersting_clips[\"is_instrumental\"]\n",
    ")\n",
    "# # this is fucked up sometimes one box doesn't give prompt to one generation\n",
    "extra_compare_mask = user_intersting_clips[user_compare_mask][\"request_id\"].isin(\n",
    "    user_intersting_clips[user_compare_mask][\"request_id\"]\n",
    "    .value_counts()\n",
    "    .index[user_intersting_clips[user_compare_mask][\"request_id\"].value_counts() == 2]\n",
    ")\n",
    "\n",
    "user_compare_mask = user_compare_mask & extra_compare_mask"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:25:01.140472Z",
     "start_time": "2024-05-26T00:25:00.895575Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "user_intersting_clips_3p5 = (\n",
    "    user_intersting_clips[user_compare_mask].reset_index().copy()\n",
    ")\n",
    "\n",
    "\n",
    "def modify_model_name(model_name, metadata):\n",
    "    if (\n",
    "        model_name.startswith(\"chirp-v3p5-engine-t\")\n",
    "        or model_name.startswith(\"chirp-v3p5-engine-s\")\n",
    "        or model_name.startswith(\"chirp-v4\")\n",
    "        or model_name.startswith(\"chirp-v3p5-h-s-31\")\n",
    "    ):\n",
    "        if \"param_experiment\" in metadata:\n",
    "            exp = metadata.get(\"param_experiment\", \"\")\n",
    "            if exp:\n",
    "                return f\"{model_name}_{exp}\"\n",
    "    return model_name\n",
    "\n",
    "\n",
    "user_intersting_clips_3p5[\"model_name\"] = user_intersting_clips_3p5.apply(\n",
    "    lambda row: modify_model_name(row[\"model_name\"], row[\"metadata\"]), axis=1\n",
    ")\n",
    "user_intersting_clips_3p5 = user_intersting_clips_3p5.sort_values(\n",
    "    by=[\"request_id\", \"preference\"]\n",
    ")\n",
    "print(user_intersting_clips_3p5.shape)\n",
    "model_counts = user_intersting_clips_3p5[\"model_name\"].value_counts()\n",
    "model_fracs = model_counts / model_counts.sum()\n",
    "\n",
    "print(\"Model Name Value Counts and Fractions:\")\n",
    "print_out_value_counts_nicely(user_intersting_clips_3p5, \"model_name\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:25:02.009450Z",
     "start_time": "2024-05-26T00:25:01.523515Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "get_preference_counts(user_intersting_clips_3p5)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:25:02.332947Z",
     "start_time": "2024-05-26T00:25:02.010694Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "print(\"first gen\")\n",
    "first_gen_slice_df = user_intersting_clips_3p5[\n",
    "    (user_intersting_clips_3p5[\"continued_parent\"].isna())\n",
    "    & (user_intersting_clips_3p5[\"task\"] == \"\")\n",
    "].copy()\n",
    "if first_gen_slice_df.shape[0] > 0:\n",
    "    get_preference_counts(\n",
    "        first_gen_slice_df,\n",
    "        title_name=\"first generation\",\n",
    "    )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:25:02.574659Z",
     "start_time": "2024-05-26T00:25:02.334214Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "print(\"is continue\")\n",
    "get_preference_counts(\n",
    "    user_intersting_clips_3p5[(user_intersting_clips_3p5[\"task\"] == \"extend\")],\n",
    "    \"is extend\",\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "# print(\"is cover\")\n",
    "# get_preference_counts(\n",
    "#     user_intersting_clips_3p5[(user_intersting_clips_3p5[\"task\"] == \"cover\")],\n",
    "#     \"is cover\",\n",
    "# )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "# print(\"is infill\")\n",
    "# get_preference_counts(\n",
    "#     user_intersting_clips_3p5[(user_intersting_clips_3p5[\"task\"] == \"infill\")],\n",
    "#     \"is infill\",\n",
    "# )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "# print(\"is artist\")\n",
    "# get_preference_counts(\n",
    "#     user_intersting_clips_3p5[\n",
    "#         (user_intersting_clips_3p5[\"task\"] == \"artist_consistency\")\n",
    "#     ],\n",
    "#     \"is artist\",\n",
    "# )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "# print(\"upsample\")\n",
    "# upsample_slice_df = user_intersting_clips_3p5[\n",
    "#     (user_intersting_clips_3p5[\"task\"] == \"upsample\")\n",
    "# ].copy()\n",
    "# if upsample_slice_df.shape[0] > 0:\n",
    "#     get_preference_counts(\n",
    "#         upsample_slice_df,\n",
    "#         title_name=\"upsample\",\n",
    "#     )"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Clean up SHIT\n",
    "\n",
    "to get the right play conts, we need the right df..."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "def unpack_dict(x):\n",
    "    if v := concat_clips_ids.get(str(x)):\n",
    "        return v\n",
    "    else:\n",
    "        return {\n",
    "            \"total_start_s\": None,\n",
    "            \"total_clip_s\": None,\n",
    "            \"concat_play_counts\": None,\n",
    "            \"concat_in_playlist\": None,\n",
    "            \"concat_likes\": None,\n",
    "            \"concat_dislikes\": None,\n",
    "        }\n",
    "\n",
    "\n",
    "extra_cols = user_intersting_clips[\"id\"].apply(unpack_dict)\n",
    "extra_cols_df = pd.DataFrame.from_records(extra_cols.values, index=extra_cols.index)\n",
    "user_intersting_clips[\n",
    "    [\n",
    "        \"total_start_s\",\n",
    "        \"total_clip_s\",\n",
    "        \"concat_play_counts\",\n",
    "        \"concat_in_playlist\",\n",
    "        \"concat_likes\",\n",
    "        \"concat_dislikes\",\n",
    "    ]\n",
    "] = extra_cols_df"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "user_intersting_clips = user_intersting_clips.sort_values(by=[\"request_id\", \"preference\", \"diff_preference\"])\n",
    "user_intersting_clips[\"pos_diff_preference\"] = user_intersting_clips[\"diff_preference\"].diff()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:25:08.361849Z",
     "start_time": "2024-05-26T00:25:08.199583Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "pos_too_much_data_mask = (\n",
    "    (user_intersting_clips[\"preference\"])\n",
    "    & (\n",
    "        (\n",
    "            user_intersting_clips[\"reaction_play_count\"] >= 3\n",
    "        )  # single play is super catchy\n",
    "        | (\n",
    "            user_intersting_clips[\"concat_play_counts\"] >= 3\n",
    "        )  # or the concat play is super catchy\n",
    "    )\n",
    "    & (user_intersting_clips[\"user_n_clips\"] >= 100)\n",
    "    # & (user_intersting_clips[\"pos_diff_preference\"] == 2)\n",
    "    # & (user_intersting_clips[\"continued_parent\"].isna())\n",
    ")\n",
    "neg_too_much_data_mask = (\n",
    "    (~user_intersting_clips[\"preference\"])\n",
    "    & (user_intersting_clips[\"reaction_play_count\"] >= 1)  # single play is super catchy\n",
    "    # & (user_intersting_clips[\"user_n_clips\"] >= 40)\n",
    "    # & (user_intersting_clips[\"continued_parent\"].isna())\n",
    ")\n",
    "# Calculate and print the proportion of data that meets our criteria\n",
    "pos_proportion = pos_too_much_data_mask.sum() / user_intersting_clips.shape[0] * 2\n",
    "print(f\"Positive proportion of data meeting criteria: {pos_proportion:.2%}\")\n",
    "neg_proportion = neg_too_much_data_mask.sum() / user_intersting_clips.shape[0] * 2\n",
    "print(f\"Negative proportion of data meeting criteria: {neg_proportion:.2%}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "user_intersting_clips[\"diff_preference\"].value_counts()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:25:08.523643Z",
     "start_time": "2024-05-26T00:25:08.367348Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "final_good_enough_requests = set(\n",
    "    user_intersting_clips[pos_too_much_data_mask][\"request_id\"].unique()\n",
    ").intersection(\n",
    "    set(user_intersting_clips[neg_too_much_data_mask][\"request_id\"].unique())\n",
    ")\n",
    "final_interesting_clips = user_intersting_clips[\n",
    "    user_intersting_clips[\"request_id\"].isin(final_good_enough_requests)\n",
    "].copy()\n",
    "# Get the value counts of model_name for preferred clips\n",
    "model_counts = final_interesting_clips[final_interesting_clips[\"preference\"]][\n",
    "    \"model_name\"\n",
    "].value_counts()\n",
    "\n",
    "# Print the results in a nicely formatted way\n",
    "total_count = model_counts.sum()\n",
    "print(\"Model Name Value Counts for Preferred Clips:\")\n",
    "print(\"-\" * 70)\n",
    "print(f\"{'Model':<30} {'Count':>10} {'Fraction':>15}\")\n",
    "print(\"-\" * 70)\n",
    "for model, count in model_counts.items():\n",
    "    fraction = count / total_count\n",
    "    print(f\"{model:<30} {count:>10,d} {fraction:>15.2%}\")\n",
    "print(\"-\" * 70)\n",
    "print(f\"{'Total':<30} {total_count:>10,d} {1:>15.2%}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:25:09.198504Z",
     "start_time": "2024-05-26T00:25:08.885507Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "validate_preference_data(final_interesting_clips)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "# Get the number of rows for final_interesting_clips with the specific model\n",
    "row_count = final_interesting_clips[\n",
    "    final_interesting_clips[\"model_name\"] == target_model_name\n",
    "].shape[0]\n",
    "\n",
    "# Print the row count in a nicely formatted way\n",
    "print(f\"Number of rows in final_interesting_clips for model {target_model_name}:\")\n",
    "print(f\"{row_count:,}\")\n",
    "print(\"done\", final_interesting_clips.shape)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# For faster processing once"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-26T00:25:10.265241Z",
     "start_time": "2024-05-26T00:25:09.934743Z"
    },
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "# Calculate the number of unique users\n",
    "total_unique_users = clip_df[\"user_id\"].nunique()\n",
    "\n",
    "# Print the result in a nicely formatted way\n",
    "print(\"Total Unique Users:\")\n",
    "print(\"-\" * 20)\n",
    "print(f\"{total_unique_users:,}\")\n",
    "print(\"-\" * 20)\n",
    "\n",
    "# This can take a while cause we have a lot of users...\n",
    "# query = \"\"\"\n",
    "# SELECT *\n",
    "# FROM auth_user\n",
    "# \"\"\"\n",
    "# user_df = pd.read_sql_query(query, engine)\n",
    "# user_df.head()\n",
    "\n",
    "test_user_id = 4688272\n",
    "print(\n",
    "    clip_df[clip_df[\"user_id\"] == test_user_id][\"created_at\"]\n",
    "    .apply(lambda x: str(x)[:10])\n",
    "    .value_counts()\n",
    ")\n",
    "print(clip_df[clip_df[\"user_id\"] == test_user_id].shape)\n",
    "query = \"\"\"\n",
    "SELECT *\n",
    "FROM auth_user\n",
    "WHERE id=25400222\n",
    "\"\"\"\n",
    "# 3 keenan\n",
    "# 6 martin\n",
    "# 8 tony -- that's me!\n",
    "# 186417 georg\n",
    "# test_user_df = pd.read_sql_query(query, engine)\n",
    "# test_user_df"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Find some weird generations"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "# not_known_bot_gens_mask = clip_df[\"model_name\"] != \"chirp-v3p5-engine-b\"\n",
    "# # run_bot_detection(clip_df[not_known_bot_gens_mask], reaction_df, write_to_file=True, cut_off_freq=0.95)\n",
    "# run_bot_detection(\n",
    "#     clip_df,\n",
    "#     reaction_df,\n",
    "#     write_to_file=True,\n",
    "#     cut_off_freq=0.95,\n",
    "#     min_generations_for_no_reaction=10,\n",
    "# )"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Alpha testing user selection"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "# Alpha testing user selection\n",
    "# we focus on the folks who are good good\n",
    "\n",
    "# # 0526 is v2 -- prod\n",
    "# # 0529 is v4 -- still good IMO, more data\n",
    "# early_v3p5_data = pd.read_csv(\"/home/tony/Data/Preference/13b_v0/interesting_clips_20240529.csv\")\n",
    "\n",
    "# print(\"uqniue users for vp5\", early_v3p5_data[\"user_id\"].nunique())\n",
    "\n",
    "# early_v3_data = pd.read_csv(\"/home/tony/Data/Preference/7b_v0_interesting_clips.csv\")\n",
    "\n",
    "# print(\"uqniue users for v3\", early_v3_data[\"user_id\"].nunique())\n",
    "\n",
    "# early_v2_data = pd.read_csv(\"/home/tony/Data/Preference/3b_v0_interesting_clips.csv\")\n",
    "\n",
    "# print(\"uqniue users for v2\", early_v2_data[\"user_id\"].nunique())\n",
    "\n",
    "# intersection_user_ids_super = set(early_v3p5_data[\"user_id\"].unique()).intersection(set(early_v3_data[\"user_id\"].unique())).intersection(set(early_v2_data[\"user_id\"].unique()))\n",
    "\n",
    "# intersection_user_ids_v3_on = set(early_v3p5_data[\"user_id\"].unique()).intersection(set(early_v3_data[\"user_id\"].unique())).difference(intersection_user_ids_super)\n",
    "\n",
    "# print(len(intersection_user_ids_super), len(intersection_user_ids_v3_on))\n",
    "\n",
    "# super_user_df = user_df[user_df[\"id\"].isin(intersection_user_ids_super)].copy()\n",
    "# print(super_user_df.shape)\n",
    "# v3_onward_user_df = user_df[user_df[\"id\"].isin(intersection_user_ids_v3_on)].copy()\n",
    "# print(v3_onward_user_df.shape)\n",
    "# super_user_df.to_csv(\"/home/tony/Data/Preference/alpha_users/super_user.csv\", index=False)\n",
    "# v3_onward_user_df.to_csv(\"/home/tony/Data/Preference/alpha_users/v3_onward_user.csv\", index=False)\n",
    "# print(\"Done!!\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# User generated clips lifetime filter"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.813Z"
    }
   },
   "outputs": [],
   "source": [
    "# query = \"\"\"\n",
    "# SELECT *\n",
    "# FROM bots_userstats\n",
    "# WHERE total_clips>=100\n",
    "# \"\"\"\n",
    "# user_stats_df = pd.read_sql_query(query, engine)\n",
    "# print(user_stats_df.shape)\n",
    "# user_stats_df[\"total_clips\"].describe()\n",
    "# top_users = user_stats_df[user_stats_df[\"total_clips\"] >= 100][\"user_id\"].unique()\n",
    "# print(len(top_users))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "top_users = clip_df[clip_df[\"user_n_clips\"] >= 50][\"user_id\"].unique()\n",
    "print(len(top_users))"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2024-06-21T19:35:17.755108Z",
     "iopub.status.busy": "2024-06-21T19:35:17.754937Z",
     "iopub.status.idle": "2024-06-21T19:35:17.774581Z",
     "shell.execute_reply": "2024-06-21T19:35:17.774106Z",
     "shell.execute_reply.started": "2024-06-21T19:35:17.755091Z"
    }
   },
   "source": [
    "# Snow flake access"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "if not os.path.exists(snow_password_path):\n",
    "    raise Exception(\"you are not authorized to access snowflake -- please setup\")\n",
    "\n",
    "snow_session = Session.builder.configs(CONNECTION_PARAMETERS).create()\n",
    "\n",
    "snow_root = Root(snow_session)\n",
    "snow_schema = snow_root.databases[\"SUNO_PROD\"].schemas[\"PROD\"]\n",
    "print(snow_schema.name)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "print_out_value_counts_nicely(final_interesting_clips, \"source\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "def analyze_clip_data_with_snowflake(\n",
    "    final_interesting_clips, target_model_name, top_users, snow_session, min_play_cut=5\n",
    "):\n",
    "    # select the df we want to squery for play counts\n",
    "    subset_v4_clips_df_full = final_interesting_clips[\n",
    "        final_interesting_clips[\"model_name\"] == target_model_name\n",
    "    ].copy()\n",
    "    print(subset_v4_clips_df_full.shape)\n",
    "\n",
    "    pre_play_duration_mask = (\n",
    "        subset_v4_clips_df_full[\"preference\"]\n",
    "        & (subset_v4_clips_df_full[\"user_id\"].isin(top_users))\n",
    "        & (\n",
    "            (subset_v4_clips_df_full[\"reaction_play_count\"] >= min_play_cut)\n",
    "            | (subset_v4_clips_df_full[\"concat_play_counts\"] >= min_play_cut)\n",
    "            | (subset_v4_clips_df_full[\"upvote_count\"] >= 1)  # positive signal leakage (strongest)\n",
    "        )\n",
    "    ) | (\n",
    "        (~subset_v4_clips_df_full[\"preference\"])\n",
    "        & (subset_v4_clips_df_full[\"user_id\"].isin(top_users))\n",
    "    )\n",
    "    subset_v4_clips_df_all = subset_v4_clips_df_full[pre_play_duration_mask].copy()\n",
    "    print(subset_v4_clips_df_all.shape)\n",
    "\n",
    "    # Filter for pairs\n",
    "    pair_request_mask = subset_v4_clips_df_all[\"request_id\"].isin(\n",
    "        subset_v4_clips_df_all[\"request_id\"]\n",
    "        .value_counts()\n",
    "        .index[subset_v4_clips_df_all[\"request_id\"].value_counts() == 2]\n",
    "    )\n",
    "    subset_v4_clips_df = subset_v4_clips_df_all[pair_request_mask].copy()\n",
    "    print(subset_v4_clips_df.shape)\n",
    "\n",
    "    # Get clip IDs and query Snowflake in batches\n",
    "    v4_clip_ids = list(str(s) for s in subset_v4_clips_df[\"id\"].unique())\n",
    "    snow_batch_size = 100_000\n",
    "    snow_results = []\n",
    "\n",
    "    for clip_ids_chunk in tqdm.tqdm(\n",
    "        [\n",
    "            v4_clip_ids[i : i + snow_batch_size]\n",
    "            for i in range(0, len(v4_clip_ids), snow_batch_size)\n",
    "        ]\n",
    "    ):\n",
    "        id_query_str = \",\".join(\"'\" + x + \"'\" for x in clip_ids_chunk)\n",
    "        print(f\"Number of clip IDs in this chunk: {len(clip_ids_chunk)}\")\n",
    "        print(f\"Length of the ID query string: {len(id_query_str)}\")\n",
    "\n",
    "        session_query = snow_session.sql(\n",
    "            f\"\"\"select *\n",
    "            from ML_SONG_SUMMARY_INFO\n",
    "            where p_date = DATE(SYSDATE() - INTERVAL '2 HOUR')\n",
    "            and p_hour = hour(SYSDATE() - INTERVAL '2 HOUR')\n",
    "            and song_id in ({id_query_str})\n",
    "            order by p_hour desc;\"\"\"\n",
    "        )\n",
    "        temp_df_snow_test = pd.DataFrame(session_query.collect())\n",
    "        snow_results.append(temp_df_snow_test)\n",
    "    print(len(snow_results))\n",
    "\n",
    "    # Process Snowflake results\n",
    "    df_snow_test = pd.concat(snow_results)\n",
    "    df_snow_test = df_snow_test.rename(columns=lambda x: x.lower())\n",
    "    df_snow_test = df_snow_test.rename(columns={\"song_id\": \"str_id\"})\n",
    "    print(\"Shape of df_snow_test:\")\n",
    "    print(f\"Rows: {df_snow_test.shape[0]}\")\n",
    "    print(f\"Columns: {df_snow_test.shape[1]}\")\n",
    "\n",
    "    # Merge data and calculate normalized play fractions\n",
    "    subset_v4_clips_df[\"str_id\"] = subset_v4_clips_df[\"id\"].astype(str)\n",
    "    subset_v4_clips_df_test = subset_v4_clips_df.merge(\n",
    "        df_snow_test, on=\"str_id\", how=\"left\"\n",
    "    )\n",
    "    subset_v4_clips_df_test[\"norm_play_frac\"] = (\n",
    "        subset_v4_clips_df_test[\"sum_total_play_duration_5\"].fillna(0)\n",
    "        / subset_v4_clips_df_test[\"duration\"]\n",
    "    )\n",
    "\n",
    "    # Create visualization\n",
    "    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(16, 6))\n",
    "\n",
    "    # First subplot: Total play duration\n",
    "    pos_play_time = subset_v4_clips_df_test[subset_v4_clips_df_test[\"preference\"]][\n",
    "        \"sum_total_play_duration_5\"\n",
    "    ]\n",
    "    neg_play_time = subset_v4_clips_df_test[~subset_v4_clips_df_test[\"preference\"]][\n",
    "        \"sum_total_play_duration_5\"\n",
    "    ]\n",
    "\n",
    "    pos_play_time.hist(\n",
    "        bins=np.linspace(0, 400, 100),\n",
    "        alpha=0.5,\n",
    "        label=f\"pos (mean={pos_play_time.mean():.2f}, median={pos_play_time.median():.2f})\",\n",
    "        ax=ax1,\n",
    "    )\n",
    "    neg_play_time.hist(\n",
    "        bins=np.linspace(0, 400, 100),\n",
    "        alpha=0.5,\n",
    "        label=f\"neg (mean={neg_play_time.mean():.2f}, median={neg_play_time.median():.2f})\",\n",
    "        ax=ax1,\n",
    "    )\n",
    "    ax1.legend()\n",
    "    ax1.set_xlabel(\"Total play duration in seconds\")\n",
    "    ax1.set_ylabel(\"counts\")\n",
    "    ax1.set_title(\"Play duration comparison\")\n",
    "\n",
    "    # Second subplot: Normalized play fraction\n",
    "    pos_norm_play_frac = subset_v4_clips_df_test[subset_v4_clips_df_test[\"preference\"]][\n",
    "        \"norm_play_frac\"\n",
    "    ]\n",
    "    neg_norm_play_frac = subset_v4_clips_df_test[\n",
    "        ~subset_v4_clips_df_test[\"preference\"]\n",
    "    ][\"norm_play_frac\"]\n",
    "\n",
    "    pos_norm_play_frac.hist(\n",
    "        bins=np.linspace(0, 10, 100),\n",
    "        alpha=0.5,\n",
    "        label=f\"pos (mean={pos_norm_play_frac.mean():.2f}, median={pos_norm_play_frac.median():.2f})\",\n",
    "        ax=ax2,\n",
    "    )\n",
    "    neg_norm_play_frac.hist(\n",
    "        bins=np.linspace(0, 10, 100),\n",
    "        alpha=0.5,\n",
    "        label=f\"neg (mean={neg_norm_play_frac.mean():.2f}, median={neg_norm_play_frac.median():.2f})\",\n",
    "        ax=ax2,\n",
    "    )\n",
    "    ax2.legend()\n",
    "    ax2.set_xlabel(\"Normalized play counts (play duration/duration)\")\n",
    "    ax2.set_ylabel(\"Log counts\")\n",
    "    ax2.set_yscale(\"log\")\n",
    "    ax2.set_title(\"Normalized play duration comparison (Log scale)\")\n",
    "\n",
    "    plt.tight_layout()\n",
    "    plt.show()\n",
    "\n",
    "    # Apply filters and analyze results\n",
    "    play_duration_mask = (\n",
    "        subset_v4_clips_df_test[\"preference\"]\n",
    "        & (subset_v4_clips_df_test[\"norm_play_frac\"] >= 0.95)\n",
    "        & (subset_v4_clips_df_test[\"sum_total_play_duration_5\"] >= 10)\n",
    "        & (subset_v4_clips_df_test[\"user_id\"].isin(top_users))\n",
    "        & (\n",
    "            (subset_v4_clips_df_test[\"reaction_play_count\"] >= min_play_cut)  # used to be 3 -- increase to 5\n",
    "            | (subset_v4_clips_df_test[\"concat_play_counts\"] >= min_play_cut) # used to be 3 -- increase to 5\n",
    "            | (subset_v4_clips_df_test[\"upvote_count\"] >= 1)  # positive signal leakage (strongest)\n",
    "        )\n",
    "    ) | (\n",
    "        (~subset_v4_clips_df_test[\"preference\"])\n",
    "        & (subset_v4_clips_df_test[\"norm_play_frac\"] <= 3.1)\n",
    "        & (subset_v4_clips_df_test[\"sum_total_play_duration_5\"] >= 10)\n",
    "        & (subset_v4_clips_df_test[\"user_id\"].isin(top_users))\n",
    "    )\n",
    "\n",
    "    # Calculate and print statistics\n",
    "    frac_pass_play_duration = (\n",
    "        play_duration_mask.sum() / subset_v4_clips_df_test.shape[0]\n",
    "    )\n",
    "    print(\n",
    "        f\"Fraction of clips that pass the play duration cut: {frac_pass_play_duration:.4f}\"\n",
    "    )\n",
    "\n",
    "    unique_requests_pass_play_durations = subset_v4_clips_df_test[play_duration_mask][\n",
    "        \"request_id\"\n",
    "    ].unique()\n",
    "    print(\n",
    "        f\"Number of unique requests passing play duration criteria: {len(unique_requests_pass_play_durations)}\"\n",
    "    )\n",
    "\n",
    "    fraction_requests_pass = (\n",
    "        len(unique_requests_pass_play_durations)\n",
    "        / subset_v4_clips_df_test[\"request_id\"].nunique()\n",
    "    )\n",
    "    print(\n",
    "        f\"Fraction of unique requests that pass play duration criteria: {fraction_requests_pass:.4f}\"\n",
    "    )\n",
    "\n",
    "    # Final filtering and analysis\n",
    "    subset_v4_clips_df_pass_duration = subset_v4_clips_df_test[\n",
    "        play_duration_mask\n",
    "    ].copy()\n",
    "    play_duration_mask_request_mask = subset_v4_clips_df_pass_duration[\n",
    "        \"request_id\"\n",
    "    ].isin(\n",
    "        subset_v4_clips_df_pass_duration[\"request_id\"]\n",
    "        .value_counts()\n",
    "        .index[subset_v4_clips_df_pass_duration[\"request_id\"].value_counts() == 2]\n",
    "    )\n",
    "    final_subset_v4_clips_df = subset_v4_clips_df_pass_duration[\n",
    "        play_duration_mask_request_mask\n",
    "    ].copy()\n",
    "\n",
    "    unique_request_count = final_subset_v4_clips_df[\"request_id\"].nunique()\n",
    "    print(\n",
    "        f\"Number of unique request IDs: {unique_request_count:,}, Total {subset_v4_clips_df_test['request_id'].nunique()}\"\n",
    "    )\n",
    "\n",
    "    print(\"Start --------------------------\")\n",
    "    print_out_value_counts_nicely(subset_v4_clips_df_test, \"task\")\n",
    "    print(\"End --------------------------\")\n",
    "    print_out_value_counts_nicely(final_subset_v4_clips_df, \"task\")\n",
    "\n",
    "    return final_subset_v4_clips_df"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "# final_subset_upsample_clips_df = analyze_clip_data_with_snowflake(\n",
    "#     final_interesting_clips, \"chirp-v4-up-u-1\", top_users, snow_session, min_play_cut=10\n",
    "# )\n",
    "final_subset_s32_clips_df = analyze_clip_data_with_snowflake(\n",
    "    final_interesting_clips, \"chirp-v4-h-s-32\", top_users, snow_session, min_play_cut=5\n",
    ")\n",
    "final_subset_t6_clips_df = analyze_clip_data_with_snowflake(\n",
    "    final_interesting_clips, \"chirp-v4-h-t-6\", top_users, snow_session\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "print(\"s-32\", final_subset_s32_clips_df.shape)\n",
    "print(\"t-6\", final_subset_t6_clips_df.shape)\n",
    "# print(\"upsample\", final_subset_upsample_clips_df.shape)\n",
    "# s-32 (1341572, 96)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "print_out_value_counts_nicely(final_subset_s32_clips_df, \"source\")\n",
    "# print_out_value_counts_nicely(final_subset_t6_clips_df, \"source\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "# final_subset_upsample_clips_df.to_pickle(\n",
    "#     \"/home/tony/Data/Preference/up_v1/interesting_clips_up_u_1_20250125_full.pkl\",\n",
    "# )\n",
    "# print(\"up_v1\", final_subset_upsample_clips_df.shape)\n",
    "# final_subset_upsample_clips_df.to_pickle(\n",
    "#     \"/home/tony/Data/Preference/up_v3/interesting_clips_up_u_3_20250106_full.pkl\",\n",
    "# )\n",
    "# print(\"up_v3\", final_subset_upsample_clips_df.shape)\n",
    "# final_subset_t6_clips_df.to_pickle(\n",
    "#      \"/home/tony/Data/Preference/30b_v6/interesting_clips_v4_h_t_6_20250501_full_long.pkl\",\n",
    "# )\n",
    "# print(\"t-6\", final_subset_t6_clips_df.shape)\n",
    "# final_subset_s32_clips_df.to_pickle(\n",
    "#     \"/home/tony/Data/Preference/13b_v32/interesting_clips_v4_h_s_32_20250501_full_long.pkl\",\n",
    "# )\n",
    "# print(\"s-32\", final_subset_s32_clips_df.shape)\n",
    "# print(\"Saving done!\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Task usage stats"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "clip_df[\"task\"].value_counts()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "n_pro_created = clip_df[clip_df[\"is_pro_user\"]][\"user_id\"].nunique()\n",
    "task_mask_cover = clip_df[\"task\"] == \"cover\"\n",
    "task_mask_artist = clip_df[\"task\"] == \"artist_consistency\"\n",
    "task_mask_infill = (\n",
    "    (clip_df[\"task\"] == \"infill\")\n",
    "    | (clip_df[\"task\"] == \"infill_intro\")\n",
    "    | (clip_df[\"task\"] == \"infill_outro\")\n",
    ")\n",
    "task_mask_image = (clip_df[\"task\"] == \"image_to_song\") | (\n",
    "    clip_df[\"task\"] == \"video_to_song\"\n",
    ")\n",
    "\n",
    "\n",
    "def print_task_usage_stats(clip_df, task_mask, task_name, n_pro_created):\n",
    "    n_created = clip_df[task_mask][\"user_id\"].nunique()\n",
    "    print(\n",
    "        f\"{task_name} usage: {n_created} out of {n_pro_created} ({round(n_created / n_pro_created, 4)})\",\n",
    "        \"\\n\",\n",
    "        \"-------------->\",\n",
    "    )\n",
    "    print_out_value_counts_nicely(clip_df[task_mask], \"model_name\")\n",
    "    print(\"\\n\", \"--------------------------\")\n",
    "\n",
    "\n",
    "print_task_usage_stats(clip_df, task_mask_cover, \"cover\", n_pro_created)\n",
    "print_task_usage_stats(clip_df, task_mask_infill, \"infill\", n_pro_created)\n",
    "print_task_usage_stats(clip_df, task_mask_artist, \"artist\", n_pro_created)\n",
    "print_task_usage_stats(clip_df, task_mask_image, \"image/video\", n_pro_created)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "clip_df[task_mask_image][\"user_id\"].value_counts().head(n=5)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "clip_df[task_mask_cover][\"user_id\"].value_counts().head(n=5)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "# v4_clip_ids = list(str(s) for s in clip_df[\"s3_id\"].unique())\n",
    "# snow_batch_size = 100_000\n",
    "# snow_results = []\n",
    "\n",
    "# for clip_ids_chunk in tqdm.tqdm(\n",
    "#     [\n",
    "#         v4_clip_ids[i : i + snow_batch_size]\n",
    "#         for i in range(0, len(v4_clip_ids), snow_batch_size)\n",
    "#     ]\n",
    "# ):\n",
    "#     id_query_str = \",\".join(\"'\" + x + \"'\" for x in clip_ids_chunk)\n",
    "#     print(f\"Number of clip IDs in this chunk: {len(clip_ids_chunk)}\")\n",
    "#     print(f\"Length of the ID query string: {len(id_query_str)}\")\n",
    "\n",
    "#     session_query = snow_session.sql(\n",
    "#         f\"\"\" select *\n",
    "#         from ML_SONG_SUMMARY_INFO\n",
    "#         where p_date = DATE(SYSDATE() - INTERVAL '2 HOUR')\n",
    "#         and p_hour = hour(SYSDATE() - INTERVAL '2 HOUR')\n",
    "#         and song_id in ({id_query_str})\n",
    "#         order by p_hour desc;\"\"\"\n",
    "#     )\n",
    "#     temp_df_snow_test = pd.DataFrame(session_query.collect())\n",
    "#     snow_results.append(temp_df_snow_test)\n",
    "# print(len(snow_results))\n",
    "\n",
    "# # Process Snowflake results\n",
    "# df_snow_test = pd.concat(snow_results)\n",
    "# df_snow_test = df_snow_test.rename(columns=lambda x: x.lower())\n",
    "# df_snow_test = df_snow_test.rename(columns={\"song_id\": \"str_id\"})\n",
    "# print(\"Shape of df_snow_test:\")\n",
    "# print(f\"Rows: {df_snow_test.shape[0]}\")\n",
    "# print(f\"Columns: {df_snow_test.shape[1]}\")\n",
    "# df_snow_test[\"clip_id\"] = df_snow_test[\"str_id\"]\n",
    "# run_bot_detection(\n",
    "#     clip_df,\n",
    "#     df_snow_test[df_snow_test[\"total_play_time\"] >= 5],\n",
    "#     write_to_file=True,\n",
    "#     cut_off_freq=0.95,\n",
    "#     min_generations_for_no_reaction=10,\n",
    "# )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "total_clip_df[\"is_pro_user\"] = total_clip_df[\"user_id\"].isin(pro_users)\n",
    "run_bot_detection(\n",
    "    total_clip_df,\n",
    "    reaction_df,\n",
    "    write_to_file=True,\n",
    "    cut_off_freq=0.95,\n",
    "    min_generations_for_no_reaction=10,\n",
    ")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Infill test"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "# def get_infill_type(x):\n",
    "#     if max(x[\"infll_start_context\"], x[\"infll_end_context\"]) <= 30:\n",
    "#         return \"short\"\n",
    "#     elif max(x[\"infll_start_context\"], x[\"infll_end_context\"]) <= 60:\n",
    "#         return \"mid\"\n",
    "#     else:\n",
    "#         return \"long\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "# clip_df_infill_task_mask_infill = (\n",
    "#     (clip_df[\"task\"] == \"infill\")\n",
    "#     | (clip_df[\"task\"] == \"infill_intro\")\n",
    "#     | (clip_df[\"task\"] == \"infill_outro\")\n",
    "# ) & (clip_df[\"created_at\"] >= \"2024-11-06 02:00:00\")\n",
    "# clip_infill_df = clip_df[clip_df_infill_task_mask_infill].copy()\n",
    "# ##\n",
    "# user_intersting_clips_3p5_task_mask_infill = (\n",
    "#     (user_intersting_clips_3p5[\"task\"] == \"infill\")\n",
    "#     | (user_intersting_clips_3p5[\"task\"] == \"infill_intro\")\n",
    "#     | (user_intersting_clips_3p5[\"task\"] == \"infill_outro\")\n",
    "# ) & (user_intersting_clips_3p5[\"created_at\"] >= \"2024-11-06 02:00:00\")\n",
    "# user_intersting_clips_3p5_infill = user_intersting_clips_3p5[\n",
    "#     user_intersting_clips_3p5_task_mask_infill\n",
    "# ].copy()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "# test_slice_series = clip_infill_df[\"metadata\"].apply(pd.Series)\n",
    "# df = pd.concat([clip_infill_df, test_slice_series], axis=1, join=\"inner\")\n",
    "# print(df.shape)\n",
    "# df = df.loc[:, ~df.columns.duplicated()].copy()\n",
    "# df[\"infll_start_context\"] = df[\"infill_start_s\"] - df[\"infill_context_start_s\"]\n",
    "# df[\"infll_end_context\"] = df[\"infill_context_end_s\"] - df[\"infill_end_s\"]\n",
    "# df[\"infill_type\"] = df[[\"infll_start_context\", \"infll_end_context\"]].apply(\n",
    "#     lambda x: get_infill_type(x), axis=1\n",
    "# )\n",
    "# clip_df_model_counts = df[\"infill_type\"].value_counts()\n",
    "# print(clip_df_model_counts)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "# plt.hist(df[\"infill_context_start_s\"], bins=np.linspace(-10, 300, 100))\n",
    "# plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "# test_slice_series = user_intersting_clips_3p5_infill[\"metadata\"].apply(pd.Series)\n",
    "# df = pd.concat([user_intersting_clips_3p5_infill, test_slice_series], axis=1, join=\"inner\")\n",
    "# print(df.shape)\n",
    "# df = df.loc[:, ~df.columns.duplicated()].copy()\n",
    "# df[\"infll_start_context\"] = df[\"infill_start_s\"]  - df[\"infill_context_start_s\"]\n",
    "# df[\"infll_end_context\"] = df[\"infill_context_end_s\"] -  df[\"infill_end_s\"]\n",
    "# df[\"infill_type\"] = df[[\"infll_start_context\", \"infll_end_context\"]].apply(lambda x:  get_infill_type(x), axis=1)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "# plt.hist(df[\"infill_context_start_s\"], bins=np.linspace(-10, 300, 100))\n",
    "# plt.show()\n",
    "# plt.hist(df[\"infill_context_end_s\"], bins=np.linspace(-10, 300, 100))\n",
    "# plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "# model_counts = df[df[\"part_of_concat\"]][\"infill_type\"].value_counts()\n",
    "\n",
    "# # Print the results in a nicely formatted way\n",
    "# total_count = model_counts.sum()\n",
    "# print(\"Model Name Value Counts for Preferred Clips:\")\n",
    "# print(\"-\" * 70)\n",
    "# print(f\"{'Model':<30} {'Count':>10} {'Fraction':>15}\")\n",
    "# print(\"-\" * 70)\n",
    "# for model, count in model_counts.items():\n",
    "#     fraction = count / total_count\n",
    "#     print(f\"{model:<30} {count:>10,d} {fraction:>15.2%}\")\n",
    "# print(\"-\" * 70)\n",
    "# print(f\"{'Total':<30} {total_count:>10,d} {1:>15.2%}\")\n",
    "\n",
    "# # Calculate the ratio of preferred clips to total clips for each model\n",
    "# preference_ratio = (\n",
    "#     df[df[\"part_of_concat\"]][\"infill_type\"].value_counts() / clip_df_model_counts\n",
    "# )\n",
    "\n",
    "\n",
    "# # Print the results in a formatted manner\n",
    "# print(\"\\n Ratio of preferred clips to total clips for each model:\")\n",
    "# print(\"-\" * 60)\n",
    "# for model, ratio in preference_ratio.items():\n",
    "#     n = clip_df_model_counts[model]\n",
    "#     uncertainty = (ratio * (1 - ratio) / n) ** 0.5\n",
    "#     print(f\"{model:<30} {ratio:.2%} ± {uncertainty:.2%}\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Playlists"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "from snowflake.snowpark.functions import col as snow_col\n",
    "\n",
    "playlist_df = (\n",
    "    snow_session.table(\"rds_playlist\")\n",
    "    .select(\"*\")\n",
    "    .filter((snow_col(\"updated_at\") >= cutoff_date))\n",
    "    .collect_nowait()\n",
    "    .result(result_type=\"pandas\")\n",
    "    .rename(columns=lambda x: x.lower())\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.814Z"
    }
   },
   "outputs": [],
   "source": [
    "playlist_df[\"user_id\"].nunique() / playlist_df.shape[0]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.815Z"
    }
   },
   "outputs": [],
   "source": [
    "playlist_id_to_user_id = playlist_df.set_index(\"id\")[\"user_id\"].to_dict()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.815Z"
    }
   },
   "outputs": [],
   "source": [
    "playlist_clip_df[\"user_id\"] = playlist_clip_df[\"playlist_id\"].apply(\n",
    "    lambda x: playlist_id_to_user_id.get(x)\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.815Z"
    }
   },
   "outputs": [],
   "source": [
    "playlist_clip_df[playlist_clip_df[\"user_id\"].isna()]"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Other ppl's clip in playlists"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.815Z"
    }
   },
   "outputs": [],
   "source": [
    "unique_clips_in_playlist = playlist_clip_df[\"clip_id\"].unique()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.815Z"
    }
   },
   "outputs": [],
   "source": [
    "total_clip_id_to_user_id = (\n",
    "    total_clip_df[total_clip_df[\"id\"].isin(unique_clips_in_playlist)]\n",
    "    .set_index(\"s3_id\")[\"user_id\"]\n",
    "    .to_dict()\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.815Z"
    }
   },
   "outputs": [],
   "source": [
    "playlist_clip_df[\"clip_user_id\"] = playlist_clip_df[\"clip_id\"].apply(\n",
    "    lambda x: total_clip_id_to_user_id.get(x)\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.815Z"
    }
   },
   "outputs": [],
   "source": [
    "playlist_clip_df"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.815Z"
    }
   },
   "outputs": [],
   "source": [
    "# v4_users = clip_df[clip_df[\"model_name\"].str.contains(\"v4\")][\"user_id\"].unique()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.815Z"
    }
   },
   "outputs": [],
   "source": [
    "# discord_info_df[discord_info_df[\"user_id\"].isin(v4_users)][[\"user_id\", \"subscription_status\", \"extra_credits_balance\", \"display_name\", \"handle\"]]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.815Z"
    }
   },
   "outputs": [],
   "source": [
    "# bad_ids_dict = run_bot_detection(\n",
    "#     total_clip_df,\n",
    "#     reaction_df,\n",
    "#     write_to_file=False,\n",
    "#     cut_off_freq=0.5,\n",
    "#     min_generations_for_no_reaction=10,\n",
    "#     return_bad_user_ids=True\n",
    "# )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.815Z"
    }
   },
   "outputs": [],
   "source": [
    "# print(len(bad_ids_dict[\"bad_pro_user_ids\"]))\n",
    "# print(len(pro_users))\n",
    "# good_pro_users = set(pro_users).difference(bad_ids_dict[\"bad_pro_user_ids\"])\n",
    "# print(len(good_pro_users))\n",
    "# with open(\"/home/tony/Work/good_pro_user_2024_11_22.json\", \"w\") as fp:\n",
    "#     json.dump([int(x) for x in sorted(good_pro_users)], fp)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.815Z"
    }
   },
   "outputs": [],
   "source": [
    "# run_bot_detection(\n",
    "#     total_clip_df,\n",
    "#     reaction_df,\n",
    "#     write_to_file=False,\n",
    "#     cut_off_freq=0.95,\n",
    "#     min_generations_for_no_reaction=20,\n",
    "# )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "execution_failed": "2025-05-06T15:22:38.815Z"
    }
   },
   "outputs": [],
   "source": [
    "clip_df[clip_df[\"is_pro_user\"]][\"user_id\"].nunique()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python 3 (ipykernel)",
   "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"
  },
  "toc": {
   "base_numbering": 1,
   "nav_menu": {},
   "number_sections": true,
   "sideBar": true,
   "skip_h1_title": false,
   "title_cell": "Table of Contents",
   "title_sidebar": "Contents",
   "toc_cell": false,
   "toc_position": {},
   "toc_section_display": true,
   "toc_window_display": false
  }
 },
 "nbformat": 4,
 "nbformat_minor": 4
}
