{
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Select Preference Data"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "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": {},
   "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",
    "        \"private_key_file\": \"/home/tony/.aws/rsa_key.p8\",\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": {},
   "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-11-19 16:00:00\"  # v4 out\n",
    "# cutoff_date = \"2024-12-16 20:40:00\"  # v4 s32 out\n",
    "# cutoff_date = \"2025-01-28 15:30:00\"  # diff v4 out\n",
    "# cutoff_date = \"2025-01-30 01:45:00\"  # diff v4 out with cfg...\n",
    "# cutoff_date = \"2025-02-21 22:15:00\"  # diff v5 out\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-24 23:15:00\"  #  diff v2 data collection out\n",
    "# cutoff_date = \"2025-05-01 00:00:00\"  # auk out\n",
    "cutoff_date = \"2025-06-03 18:30:00\"  # stems out, slider out\n",
    "# cutoff_date = \"2025-06-06 00:00:00\"  # diff v2 d4 data collection out\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-v3p5-engine-t-6\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "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": {},
   "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": {},
   "outputs": [],
   "source": [
    "# gathered_data = gather_data(engine, cutoff_date)\n",
    "gathered_data = gather_data_with_snowflake(\n",
    "    snow_session,\n",
    "    cutoff_date,\n",
    "    filter_play_count=1,  # filter_user_n_clips=40\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "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": {},
   "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": {},
   "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": {},
   "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": {},
   "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": {},
   "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": {},
   "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": {},
   "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()), 2\n",
    "    ),\n",
    ")\n",
    "# Get value counts\n",
    "print_out_value_counts_nicely(clip_df, \"source\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "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\")\n",
    "    | (clip_df[\"clip_type\"] == \"concat_infilling\")\n",
    "    | (clip_df[\"clip_type\"] == \"stem_mix\")\n",
    "    | (clip_df[\"clip_type\"] == \"edit_v3_export\")\n",
    "    | (clip_df[\"clip_type\"] == \"edit_speed\")\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": {},
   "outputs": [],
   "source": [
    "# check the model conts\n",
    "print_out_value_counts_nicely(clip_df, \"model_name\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "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",
    "    & (clip_df[\"model_name\"] != \"chirp-auk\")\n",
    "    & (clip_df[\"model_name\"] != \"chirp-ahi\")\n",
    "    & (clip_df[\"model_name\"] != \"chirp-v4-h-t-6-cfg-null\")\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": {},
   "outputs": [],
   "source": [
    "print_out_value_counts_nicely(clip_df, \"task\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "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": {},
   "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": {},
   "outputs": [],
   "source": [
    "# set user number of clips generated\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": {},
   "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": {},
   "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": {},
   "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": {},
   "outputs": [],
   "source": [
    "print(\n",
    "    \"has upvoted in exp\",\n",
    "    round(clip_df[\"upvoted\"].value_counts(normalize=True)[True], 5),\n",
    ")\n",
    "print(\n",
    "    \"has downvoted out of exp\",\n",
    "    round(clip_df[\"downvoted\"].value_counts(normalize=True)[True], 5),\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "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": {},
   "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\"]\n",
    "    # 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",
    "print(len(has_action_ids))\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": {},
   "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": {},
   "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": {},
   "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",
    ")\n",
    "# print_out_value_counts_nicely(clip_df, \"n_edits\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "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": {},
   "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": {},
   "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": {},
   "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": {},
   "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": {},
   "outputs": [],
   "source": [
    "# filter out 8 stems for now?\n",
    "clip_df[\"request_count\"] = clip_df.groupby(\"request_id\")[\"request_id\"].transform(\n",
    "    \"count\"\n",
    ")\n",
    "# creation of interesting_clips\n",
    "interesting_clips = clip_df[\n",
    "    (clip_df[\"request_id\"].isin(requests)) & (clip_df[\"request_count\"] == 2)\n",
    "].copy()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "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": {},
   "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": {},
   "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": {},
   "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": {},
   "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": {},
   "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": {},
   "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": {},
   "outputs": [],
   "source": [
    "print_out_value_counts_nicely(interesting_clips, \"model_name\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "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": {},
   "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": {},
   "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": {},
   "outputs": [],
   "source": [
    "get_preference_counts(interesting_clips)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "plot_clip_basic_distributions(interesting_clips)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "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": {},
   "outputs": [],
   "source": [
    "# subselect interesting clips\n",
    "interesting_clips_masks = (\n",
    "    interesting_clips[\"model_name\"].str.contains(\"v3p5|v4|v5|auk|ahi\")\n",
    ") & (interesting_clips[\"reaction_play_count\"] > 0)\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": {},
   "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": {},
   "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": {},
   "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": {},
   "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": {},
   "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",
    "        or model_name.startswith(\"chirp-auk\")\n",
    "        or model_name.startswith(\"chirp-ahi\")\n",
    "    ):\n",
    "        if \"param_experiment\" in metadata:\n",
    "            exp = metadata.get(\"param_experiment\", \"\")\n",
    "            if exp:\n",
    "                if exp == \"mask_control_slider\" and not metadata.get(\n",
    "                    \"control_sliders\", None\n",
    "                ):\n",
    "                    return model_name\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": {},
   "outputs": [],
   "source": [
    "get_preference_counts(user_intersting_clips_3p5)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "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": {},
   "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": {},
   "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": {},
   "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": {},
   "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": {},
   "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": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(\"is upload\")\n",
    "get_preference_counts(\n",
    "    user_intersting_clips_3p5[(user_intersting_clips_3p5[\"task\"] == \"upload_extend\")],\n",
    "    \"is upload extend\",\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(\"fixed_infill\")\n",
    "upload_extend_slice_df = user_intersting_clips_3p5[\n",
    "    (user_intersting_clips_3p5[\"task\"] == \"fixed_infill\")\n",
    "].copy()\n",
    "if upload_extend_slice_df.shape[0] > 0:\n",
    "    get_preference_counts(\n",
    "        upload_extend_slice_df,\n",
    "        title_name=\"fixed_infill\",\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": {},
   "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",
    "# TODO: concat play count / play duraiton needs to be somehow counted as well\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": {},
   "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": {},
   "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\"] >= 40)\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": {},
   "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": {},
   "outputs": [],
   "source": [
    "validate_preference_data(final_interesting_clips)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "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": {},
   "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 = 3\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=3\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": {},
   "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=4,\n",
    ")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Alpha testing user selection"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "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": {},
   "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": {},
   "outputs": [],
   "source": [
    "top_users = clip_df[clip_df[\"user_n_clips\"] >= 20][\"user_id\"].unique()\n",
    "print(len(top_users))"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Snow flake access"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "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": {},
   "outputs": [],
   "source": [
    "print_out_value_counts_nicely(final_interesting_clips, \"source\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "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(\"match model\", 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",
    "            | (\n",
    "                subset_v4_clips_df_full[\"upvote_count\"] >= 1\n",
    "            )  # 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 '3 HOUR')\n",
    "            and p_hour = hour(SYSDATE() - INTERVAL '3 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",
    "            (\n",
    "                subset_v4_clips_df_test[\"reaction_play_count\"] >= min_play_cut\n",
    "            )  # used to be 3 -- increase to 5\n",
    "            | (\n",
    "                subset_v4_clips_df_test[\"concat_play_counts\"] >= min_play_cut\n",
    "            )  # used to be 3 -- increase to 5\n",
    "            | (\n",
    "                subset_v4_clips_df_test[\"upvote_count\"] >= 1\n",
    "            )  # 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": {},
   "outputs": [],
   "source": [
    "# for x in final_interesting_clips[final_interesting_clips[\"model_name\"] == \"chirp-v5-stem-v0\"][[\"id\"]].values:\n",
    "#     print(x[0])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print_out_value_counts_nicely(final_interesting_clips, \"model_name\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# for this we eat the quality cost and get a bit more data\n",
    "final_subset_auk_og_clips_df = analyze_clip_data_with_snowflake(\n",
    "    final_interesting_clips, \"chirp-auk-t0\", top_users, snow_session, min_play_cut=3\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# final_subset_upsample_diff_v1_df = analyze_clip_data_with_snowflake(\n",
    "#     final_interesting_clips, \"chirp-v4-up-u-7\", top_users, snow_session, min_play_cut=5\n",
    "# )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# final_subset_upsample_diff_v1_df.to_pickle(\n",
    "#     f\"/home/tony/Data/Preference/up_v2_d3/interesting_clips_diff_v1_20250528.pkl\",\n",
    "# )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# final_subset_auk_infill_30b_clips_df = analyze_clip_data_with_snowflake(\n",
    "#     final_interesting_clips, \"chirp-auk-infill\", top_users, snow_session, min_play_cut=3\n",
    "# )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# final_subset_stem_clips_df = analyze_clip_data_with_snowflake(\n",
    "#     final_interesting_clips, \"chirp-v5-stem-v0\", top_users, snow_session, min_play_cut=1\n",
    "# )\n",
    "final_subset_upsample_ahi_clips_df = analyze_clip_data_with_snowflake(\n",
    "    final_interesting_clips, \"chirp-ahi-up-2\", top_users, snow_session, min_play_cut=5\n",
    ")\n",
    "final_subset_upsample_ahi_clips_df_2 = analyze_clip_data_with_snowflake(\n",
    "    final_interesting_clips,\n",
    "    \"chirp-v4-up-u-d-2-4\",\n",
    "    top_users,\n",
    "    snow_session,\n",
    "    min_play_cut=5,\n",
    ")\n",
    "final_subset_auk_clips_df = analyze_clip_data_with_snowflake(\n",
    "    final_interesting_clips, \"chirp-auk-t1\", top_users, snow_session, min_play_cut=5\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# final_subset_upsample_clips_df = analyze_clip_data_with_snowflake(\n",
    "#     final_interesting_clips, \"chirp-v4-up-u-7\", top_users, snow_session, min_play_cut=5\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, min_play_cut=5\n",
    "# )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(\"auk\", final_subset_auk_clips_df.shape)\n",
    "print(\"ahi\", final_subset_upsample_ahi_clips_df.shape)\n",
    "# print(\"ahi sneaked\", final_subset_upsample_ahi_clips_df_2.shape)\n",
    "# print(\"diff stem\", final_subset_stem_clips_df.shape)\n",
    "# cut at 5\n",
    "# s-32 (1221770, 95)\n",
    "# t-6 (409626, 95)\n",
    "# upsample (211348, 95)\n",
    "# vs cut at 10\n",
    "# s-32 (712578, 96)\n",
    "# t-6 (258662, 96)\n",
    "# upsample (80826, 96)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print_out_value_counts_nicely(final_subset_auk_clips_df, \"source\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "total_ahi_df = pd.concat(\n",
    "    [final_subset_upsample_ahi_clips_df, final_subset_upsample_ahi_clips_df_2]\n",
    ")\n",
    "# total_ahi_df = final_subset_upsample_ahi_clips_df.copy()\n",
    "print(\"total ahi\", total_ahi_df.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# final_subset_auk_clips_df.to_pickle(\n",
    "#     \"/home/tony/Data/Preference/auk_t1/interesting_clips_auk_t1_20250502.pkl\",\n",
    "# )\n",
    "# print(\"auk_t1\", final_subset_auk_clips_df.shape)\n",
    "# total_ahi_df.to_pickle(\n",
    "#     \"/home/tony/Data/Preference/up_v2_d3/interesting_clips_ahi_d3_20250502.pkl\",\n",
    "# )\n",
    "# print(\"ahi_d3\", total_ahi_df.shape)\n",
    "# print(\"Saving done!\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Task usage stats"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "clip_df[\"task\"].value_counts()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# plot_clip_distribution(clip_df[(clip_df[\"task\"] == \"cover\") & (clip_df[\"model_name\"] == \"chirp-v4-h-s-32\")& (clip_df[\"created_at\"] > \"2025-04-01 00:45:00\")])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# bad_clip_df = clip_df[(clip_df[\"task\"] == \"cover\") & (clip_df[\"model_name\"] == \"chirp-v4-h-s-32\")& (clip_df[\"created_at\"] > \"2025-04-01 00:45:00\")]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "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": {},
   "outputs": [],
   "source": [
    "task_mask_any = (\n",
    "    (clip_df[\"task\"] == \"infill\")\n",
    "    | (clip_df[\"task\"] == \"infill_intro\")\n",
    "    | (clip_df[\"task\"] == \"infill_outro\")\n",
    "    | (clip_df[\"task\"] == \"cover\")\n",
    "    | (clip_df[\"task\"] == \"artist_consistency\")\n",
    "    | (clip_df[\"task\"] == \"extend\")\n",
    "    | (clip_df[\"task\"] == \"upload_extend\")\n",
    "    | (clip_df[\"task\"] == \"upsample\")\n",
    ") & (clip_df[\"is_pro_user\"])\n",
    "print_task_usage_stats(clip_df, task_mask_any, \"any task\", n_pro_created)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "special_users = set(clip_df[clip_df[\"is_pro_user\"]][\"user_id\"].unique()).difference(\n",
    "    clip_df[task_mask_any][\"user_id\"].unique()\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(len(special_users))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "potential_special_bot_mask = clip_df[\"user_id\"].isin(special_users)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "clip_df[potential_special_bot_mask][\"user_id\"].value_counts().describe()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "clip_df[potential_special_bot_mask][\"user_id\"].value_counts()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "clip_df[task_mask_image][\"user_id\"].value_counts().head(n=5)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "clip_df[task_mask_cover][\"user_id\"].value_counts().head(n=5)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "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": {},
   "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=6,\n",
    "# )"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Infill test"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "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": {},
   "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": {},
   "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": {},
   "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": {},
   "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": {},
   "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": {},
   "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": {},
   "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": {},
   "outputs": [],
   "source": [
    "playlist_df[\"user_id\"].nunique() / playlist_df.shape[0]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "playlist_id_to_user_id = playlist_df.set_index(\"id\")[\"user_id\"].to_dict()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "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": {},
   "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": {},
   "outputs": [],
   "source": [
    "unique_clips_in_playlist = playlist_clip_df[\"clip_id\"].unique()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "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": {},
   "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": {},
   "outputs": [],
   "source": [
    "# v4_users = clip_df[clip_df[\"model_name\"].str.contains(\"v4\")][\"user_id\"].unique()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "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": {},
   "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": {},
   "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": {},
   "outputs": [],
   "source": [
    "# selected_indices = total_clip_df[\"prompt_text\"].apply(lambda x: bool(re.search(r'Wir ziehen durch die Straßen und die Clubs dieser', str(x), re.IGNORECASE)))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "final_interesting_clips[\"model_name\"] = final_interesting_clips.apply(\n",
    "    lambda row: modify_model_name(row[\"model_name\"], row[\"metadata\"]), axis=1\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# requests_with_vol = final_interesting_clips[final_interesting_clips[\"model_name\"].str.contains(\"chirp-v4-h-s-32-u-4-6\")][\"request_id\"].unique()\n",
    "# print(len(requests_with_vol))\n",
    "# vol_final_interesting_clips = final_interesting_clips[final_interesting_clips[\"request_id\"].isin(requests_with_vol)].copy()\n",
    "# get_preference_counts(\n",
    "#     vol_final_interesting_clips,\n",
    "#     title_name=\"subset test\",\n",
    "# )\n",
    "# # vol_final_interesting_clips.to_pickle(\n",
    "# #     \"/home/tony/Data/Preference/13b_v32/interesting_clips_exp_20250219_full.pkl\",\n",
    "# # )\n",
    "# print(\"vol exps\", vol_final_interesting_clips.shape)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## find a song with matching lyrics"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# # takes ~ 3 mins\n",
    "# lyrics_session_query = snow_session.sql(\n",
    "#     f\"\"\"select *\n",
    "#     from CLIP\n",
    "#     where REGEXP_LIKE(prompt_text, 'kwaśna.*')\n",
    "#     \"\"\"\n",
    "# )\n",
    "# lyrics_matched_df = pd.DataFrame(lyrics_session_query.collect())\n",
    "# print(lyrics_matched_df.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# lyrics_matched_df"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Find a specific user's creations"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# user_session_query = snow_session.sql(\n",
    "#     f\"\"\"select *\n",
    "#     from CLIP\n",
    "#     where user_id=62804651\n",
    "#     \"\"\"\n",
    "# )\n",
    "# user_matched_df = pd.DataFrame(user_session_query.collect())\n",
    "# print(user_matched_df.shape)\n",
    "# # user_matched_df.to_csv(\"/home/tony/Data/for_minz_20250202.csv\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "final_interesting_clips = user_intersting_clips[\n",
    "    user_intersting_clips[\"request_id\"].isin(final_good_enough_requests)\n",
    "].copy()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "get_preference_counts(\n",
    "    final_interesting_clips,\n",
    "    title_name=\"subset test\",\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "clip_df[\"model_name\"].value_counts()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Ensure we have datetime columns\n",
    "if (\n",
    "    \"created_at\" not in clip_df.columns\n",
    "    and \"p_date\" in clip_df.columns\n",
    "    and \"p_hour\" in clip_df.columns\n",
    "):\n",
    "    # Create a datetime from p_date and p_hour if created_at is not available\n",
    "    clip_df[\"datetime\"] = pd.to_datetime(clip_df[\"p_date\"]) + pd.to_timedelta(\n",
    "        clip_df[\"p_hour\"], unit=\"h\"\n",
    "    )\n",
    "else:\n",
    "    clip_df[\"datetime\"] = pd.to_datetime(clip_df[\"created_at\"])\n",
    "\n",
    "# Extract date and hour for grouping\n",
    "clip_df[\"date\"] = clip_df[\"datetime\"].dt.date\n",
    "clip_df[\"hour\"] = clip_df[\"datetime\"].dt.hour\n",
    "\n",
    "model_mask = clip_df[\"model_name\"].str.contains(\"chirp-v4-h-s-32\")\n",
    "# Group by date and hour and count clips\n",
    "date_hour_counts = clip_df[model_mask].groupby([\"date\", \"hour\"]).size()\n",
    "\n",
    "# Calculate like rate per date and hour if we have like information\n",
    "if \"upvote_count\" in clip_df.columns:\n",
    "    date_hour_likes = (\n",
    "        clip_df[model_mask].groupby([\"date\", \"hour\"])[\"upvote_count\"].mean()\n",
    "    )\n",
    "else:\n",
    "    # If we don't have direct like information, create a placeholder\n",
    "    date_hour_likes = pd.Series(0, index=date_hour_counts.index)\n",
    "    print(\"No like information found in the dataframe\")\n",
    "\n",
    "# Create a figure with a single plot that has two y-axes\n",
    "fig, ax1 = plt.subplots(figsize=(14, 10))\n",
    "\n",
    "# Create x-axis positions for the bars\n",
    "date_hour_index = date_hour_counts.index.tolist()\n",
    "x = np.arange(len(date_hour_index))\n",
    "\n",
    "# Plot number of clips per date and hour (bars)\n",
    "bars = ax1.bar(x, date_hour_counts.values, color=\"skyblue\", alpha=0.7)\n",
    "ax1.set_xlabel(\"Date and Hour\")\n",
    "ax1.set_ylabel(\"Number of Clips\", color=\"blue\")\n",
    "ax1.tick_params(axis=\"y\", labelcolor=\"blue\")\n",
    "\n",
    "# Set x-axis ticks and labels\n",
    "# Only show a subset of ticks for readability\n",
    "tick_step = max(1, len(date_hour_index) // 20)  # Show at most 20 ticks\n",
    "tick_positions = [i for i in range(0, len(date_hour_index), tick_step)]\n",
    "tick_labels = [\n",
    "    f\"{date_hour_index[i][0].strftime('%Y-%m-%d')} {date_hour_index[i][1]:02d}:00\"\n",
    "    for i in tick_positions\n",
    "]\n",
    "ax1.set_xticks(tick_positions)\n",
    "ax1.set_xticklabels(tick_labels, rotation=45, ha=\"right\")\n",
    "ax1.grid(axis=\"y\", linestyle=\"--\", alpha=0.3)\n",
    "\n",
    "# Add vertical lines to separate dates\n",
    "unique_dates = sorted(list(set([idx[0] for idx in date_hour_index])))\n",
    "date_boundaries = [0]\n",
    "for i in range(1, len(unique_dates)):\n",
    "    # Find the first index where the date changes\n",
    "    for j, (date, _) in enumerate(date_hour_index):\n",
    "        if date == unique_dates[i]:\n",
    "            date_boundaries.append(j - 0.5)\n",
    "            break\n",
    "\n",
    "# Draw vertical lines at date boundaries\n",
    "for boundary in date_boundaries[1:]:  # Skip the first boundary (0)\n",
    "    ax1.axvline(x=boundary, color=\"gray\", linestyle=\"-\", alpha=0.3)\n",
    "\n",
    "# Create a second y-axis for like rate\n",
    "ax2 = ax1.twinx()\n",
    "line = ax2.plot(\n",
    "    x,\n",
    "    date_hour_likes.values,\n",
    "    color=\"red\",\n",
    "    marker=\"o\",\n",
    "    linestyle=\"-\",\n",
    "    linewidth=2,\n",
    "    markersize=4,\n",
    ")\n",
    "ax2.set_ylabel(\"Like Rate\", color=\"red\")\n",
    "ax2.tick_params(axis=\"y\", labelcolor=\"red\")\n",
    "max_like_rate = date_hour_likes.max() if not date_hour_likes.empty else 0\n",
    "ax2.set_ylim(0, max_like_rate * 1.1 if max_like_rate > 0 else 1)\n",
    "\n",
    "# Add title and adjust layout\n",
    "plt.title(\n",
    "    \"Clip Generation Count and Like Rate by Date and Hour, model: chirp-v4-h-s-32\",\n",
    "    fontsize=14,\n",
    ")\n",
    "fig.tight_layout()\n",
    "plt.show()\n",
    "\n",
    "# Print some statistics\n",
    "print(f\"Total clips: {len(clip_df)}\")\n",
    "if not date_hour_counts.empty:\n",
    "    max_count_idx = date_hour_counts.idxmax()\n",
    "    max_date, max_hour = max_count_idx\n",
    "    print(\n",
    "        f\"Date and hour with most clips: {max_date.strftime('%Y-%m-%d')} {max_hour:02d}:00 ({date_hour_counts.max()} clips)\"\n",
    "    )\n",
    "if not date_hour_likes.empty and date_hour_likes.max() > 0:\n",
    "    max_like_idx = date_hour_likes.idxmax()\n",
    "    max_like_date, max_like_hour = max_like_idx\n",
    "    print(\n",
    "        f\"Date and hour with highest like rate: {max_like_date.strftime('%Y-%m-%d')} {max_like_hour:02d}:00 ({date_hour_likes.max():.2%})\"\n",
    "    )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Ensure we have datetime columns\n",
    "if (\n",
    "    \"created_at\" not in clip_df.columns\n",
    "    and \"p_date\" in clip_df.columns\n",
    "    and \"p_hour\" in clip_df.columns\n",
    "):\n",
    "    # Create a datetime from p_date and p_hour if created_at is not available\n",
    "    clip_df[\"datetime\"] = pd.to_datetime(clip_df[\"p_date\"]) + pd.to_timedelta(\n",
    "        clip_df[\"p_hour\"], unit=\"h\"\n",
    "    )\n",
    "else:\n",
    "    clip_df[\"datetime\"] = pd.to_datetime(clip_df[\"created_at\"])\n",
    "\n",
    "# Extract date and hour for grouping\n",
    "clip_df[\"date\"] = clip_df[\"datetime\"].dt.date\n",
    "clip_df[\"hour\"] = clip_df[\"datetime\"].dt.hour\n",
    "\n",
    "model_mask = clip_df[\"model_name\"].str.contains(\"chirp-v3p5-engine-s-8\")\n",
    "# Group by date and hour and count clips\n",
    "date_hour_counts = clip_df[model_mask].groupby([\"date\", \"hour\"]).size()\n",
    "\n",
    "# Calculate like rate per date and hour if we have like information\n",
    "if \"has_action\" in clip_df.columns:\n",
    "    date_hour_likes = clip_df[model_mask].groupby([\"date\", \"hour\"])[\"has_action\"].mean()\n",
    "else:\n",
    "    # If we don't have direct like information, create a placeholder\n",
    "    date_hour_likes = pd.Series(0, index=date_hour_counts.index)\n",
    "    print(\"No like information found in the dataframe\")\n",
    "\n",
    "# Create a figure with a single plot that has two y-axes\n",
    "fig, ax1 = plt.subplots(figsize=(14, 10))\n",
    "\n",
    "# Create x-axis positions for the bars\n",
    "date_hour_index = date_hour_counts.index.tolist()\n",
    "x = np.arange(len(date_hour_index))\n",
    "\n",
    "# Plot number of clips per date and hour (bars)\n",
    "bars = ax1.bar(x, date_hour_counts.values, color=\"skyblue\", alpha=0.7)\n",
    "ax1.set_xlabel(\"Date and Hour\")\n",
    "ax1.set_ylabel(\"Number of Clips\", color=\"blue\")\n",
    "ax1.tick_params(axis=\"y\", labelcolor=\"blue\")\n",
    "\n",
    "# Set x-axis ticks and labels\n",
    "# Only show a subset of ticks for readability\n",
    "tick_step = max(1, len(date_hour_index) // 20)  # Show at most 20 ticks\n",
    "tick_positions = [i for i in range(0, len(date_hour_index), tick_step)]\n",
    "tick_labels = [\n",
    "    f\"{date_hour_index[i][0].strftime('%Y-%m-%d')} {date_hour_index[i][1]:02d}:00\"\n",
    "    for i in tick_positions\n",
    "]\n",
    "ax1.set_xticks(tick_positions)\n",
    "ax1.set_xticklabels(tick_labels, rotation=45, ha=\"right\")\n",
    "ax1.grid(axis=\"y\", linestyle=\"--\", alpha=0.3)\n",
    "\n",
    "# Add vertical lines to separate dates\n",
    "unique_dates = sorted(list(set([idx[0] for idx in date_hour_index])))\n",
    "date_boundaries = [0]\n",
    "for i in range(1, len(unique_dates)):\n",
    "    # Find the first index where the date changes\n",
    "    for j, (date, _) in enumerate(date_hour_index):\n",
    "        if date == unique_dates[i]:\n",
    "            date_boundaries.append(j - 0.5)\n",
    "            break\n",
    "\n",
    "# Draw vertical lines at date boundaries\n",
    "for boundary in date_boundaries[1:]:  # Skip the first boundary (0)\n",
    "    ax1.axvline(x=boundary, color=\"gray\", linestyle=\"-\", alpha=0.3)\n",
    "\n",
    "# Create a second y-axis for like rate\n",
    "ax2 = ax1.twinx()\n",
    "line = ax2.plot(\n",
    "    x,\n",
    "    date_hour_likes.values,\n",
    "    color=\"red\",\n",
    "    marker=\"o\",\n",
    "    linestyle=\"-\",\n",
    "    linewidth=2,\n",
    "    markersize=4,\n",
    ")\n",
    "ax2.set_ylabel(\"Like Rate\", color=\"red\")\n",
    "ax2.tick_params(axis=\"y\", labelcolor=\"red\")\n",
    "max_like_rate = date_hour_likes.max() if not date_hour_likes.empty else 0\n",
    "ax2.set_ylim(0, max_like_rate * 1.1 if max_like_rate > 0 else 1)\n",
    "\n",
    "# Add title and adjust layout\n",
    "plt.title(\n",
    "    \"Clip Generation Count and Like Rate by Date and Hour, model: chirp-v3p5-engine-s-8\",\n",
    "    fontsize=14,\n",
    ")\n",
    "fig.tight_layout()\n",
    "plt.show()\n",
    "\n",
    "# Print some statistics\n",
    "print(f\"Total clips: {len(clip_df)}\")\n",
    "if not date_hour_counts.empty:\n",
    "    max_count_idx = date_hour_counts.idxmax()\n",
    "    max_date, max_hour = max_count_idx\n",
    "    print(\n",
    "        f\"Date and hour with most clips: {max_date.strftime('%Y-%m-%d')} {max_hour:02d}:00 ({date_hour_counts.max()} clips)\"\n",
    "    )\n",
    "if not date_hour_likes.empty and date_hour_likes.max() > 0:\n",
    "    max_like_idx = date_hour_likes.idxmax()\n",
    "    max_like_date, max_like_hour = max_like_idx\n",
    "    print(\n",
    "        f\"Date and hour with highest like rate: {max_like_date.strftime('%Y-%m-%d')} {max_like_hour:02d}:00 ({date_hour_likes.max():.2%})\"\n",
    "    )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Calculate like rate per user\n",
    "user_stats = clip_df.groupby(\"user_id\").agg({\"is_deleted\": [\"count\", \"sum\"]})\n",
    "user_stats.columns = [\"total_clips\", \"liked_clips\"]\n",
    "user_stats[\"like_rate\"] = user_stats[\"liked_clips\"] / user_stats[\"total_clips\"]\n",
    "\n",
    "# Filter out users with very few clips for more meaningful analysis\n",
    "min_clips = 5\n",
    "filtered_users = user_stats[user_stats[\"total_clips\"] >= min_clips]\n",
    "\n",
    "# Plot the distribution of like rates\n",
    "plt.figure(figsize=(12, 8))\n",
    "plt.hist(\n",
    "    filtered_users[\"like_rate\"], bins=50, alpha=0.7, color=\"skyblue\", edgecolor=\"black\"\n",
    ")\n",
    "plt.xlabel(\"Like Rate\", fontsize=12)\n",
    "plt.ylabel(\"Number of Users\", fontsize=12)\n",
    "plt.title(\n",
    "    f\"Distribution of Like Rates per User (Users with {min_clips}+ clips)\", fontsize=14\n",
    ")\n",
    "plt.grid(axis=\"y\", linestyle=\"--\", alpha=0.3)\n",
    "\n",
    "# Add statistics as text\n",
    "mean_like_rate = filtered_users[\"like_rate\"].mean()\n",
    "median_like_rate = filtered_users[\"like_rate\"].median()\n",
    "stats_text = (\n",
    "    f\"Mean: {mean_like_rate:.2%}\\n\"\n",
    "    f\"Median: {median_like_rate:.2%}\\n\"\n",
    "    f\"Users analyzed: {len(filtered_users)} (with {min_clips}+ clips)\"\n",
    ")\n",
    "plt.annotate(\n",
    "    stats_text,\n",
    "    xy=(0.95, 0.95),\n",
    "    xycoords=\"axes fraction\",\n",
    "    ha=\"right\",\n",
    "    va=\"top\",\n",
    "    bbox=dict(boxstyle=\"round\", fc=\"white\", alpha=0.7),\n",
    ")\n",
    "\n",
    "plt.tight_layout()\n",
    "plt.show()\n",
    "\n",
    "# Print additional statistics\n",
    "print(f\"Total users with at least {min_clips} clips: {len(filtered_users)}\")\n",
    "print(f\"Mean like rate: {mean_like_rate:.2%}\")\n",
    "print(f\"Median like rate: {median_like_rate:.2%}\")\n",
    "\n",
    "# Show top users by like rate (with minimum clip count)\n",
    "temp_top_users = filtered_users.sort_values(\"like_rate\", ascending=False).head(10)\n",
    "print(\"\\nTop 10 users by like rate (minimum 5 clips):\")\n",
    "print(temp_top_users[[\"total_clips\", \"liked_clips\", \"like_rate\"]].reset_index())\n",
    "\n",
    "# Show users with most clips\n",
    "most_active = user_stats.sort_values(\"total_clips\", ascending=False).head(10)\n",
    "print(\"\\nTop 10 most active users:\")\n",
    "print(most_active[[\"total_clips\", \"liked_clips\", \"like_rate\"]].reset_index())"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Find users with 100% like rate (like_rate = 1.0)\n",
    "perfect_likers = filtered_users[filtered_users[\"like_rate\"] == 1.0]\n",
    "print(f\"Number of users with 100% like rate: {len(perfect_likers)}\")\n",
    "\n",
    "# Get the user IDs of these perfect likers\n",
    "perfect_liker_ids = perfect_likers.index.tolist()\n",
    "\n",
    "# Filter the clip dataframe to only include clips from these users\n",
    "perfect_clips = clip_df[clip_df[\"user_id\"].isin(perfect_liker_ids)]\n",
    "\n",
    "# Check how many clips we have from these users\n",
    "print(f\"Total clips from users with 100% like rate: {len(perfect_clips)}\")\n",
    "\n",
    "# Extract datetime information\n",
    "perfect_clips[\"created_datetime\"] = pd.to_datetime(perfect_clips[\"created_at\"])\n",
    "perfect_clips[\"date_hour\"] = perfect_clips[\"created_datetime\"].dt.strftime(\n",
    "    \"%Y-%m-%d %H:00\"\n",
    ")\n",
    "\n",
    "# Count clips by date-hour\n",
    "date_hour_counts = perfect_clips[\"date_hour\"].value_counts().sort_index()\n",
    "\n",
    "# Create a single plot for date-hour distribution\n",
    "plt.figure(figsize=(16, 8))\n",
    "date_hour_counts.plot(kind=\"bar\", color=\"skyblue\", edgecolor=\"black\", alpha=0.7)\n",
    "plt.title(\n",
    "    \"Generation Time Distribution by Date-Hour\\nfor Users with 100% Like Rate\",\n",
    "    fontsize=14,\n",
    ")\n",
    "plt.xlabel(\"Date-Hour\", fontsize=12)\n",
    "plt.ylabel(\"Number of Clips\", fontsize=12)\n",
    "\n",
    "# Reduce the number of x-labels to avoid overcrowding\n",
    "total_labels = len(date_hour_counts)\n",
    "if total_labels > 0:\n",
    "    # Show only about 10 labels evenly distributed\n",
    "    step_size = max(1, total_labels // 10)\n",
    "    plt.xticks(\n",
    "        range(0, total_labels, step_size),\n",
    "        [date_hour_counts.index[i] for i in range(0, total_labels, step_size)],\n",
    "        rotation=45,\n",
    "        ha=\"right\",\n",
    "    )\n",
    "\n",
    "plt.grid(axis=\"y\", linestyle=\"--\", alpha=0.3)\n",
    "\n",
    "# Add statistics as text annotation\n",
    "avg_clips_per_user = len(perfect_clips) / len(perfect_liker_ids)\n",
    "stats_text = (\n",
    "    f\"Users with 100% like rate: {len(perfect_liker_ids)}\\n\"\n",
    "    f\"Total clips: {len(perfect_clips)}\\n\"\n",
    "    f\"Avg clips per user: {avg_clips_per_user:.1f}\"\n",
    ")\n",
    "\n",
    "# Add text box to the figure\n",
    "plt.annotate(\n",
    "    stats_text,\n",
    "    xy=(0.95, 0.95),\n",
    "    xycoords=\"axes fraction\",\n",
    "    ha=\"right\",\n",
    "    va=\"top\",\n",
    "    bbox=dict(boxstyle=\"round\", fc=\"white\", alpha=0.7),\n",
    ")\n",
    "\n",
    "plt.tight_layout()\n",
    "plt.show()\n",
    "\n",
    "# Print the top 5 date-hours with the most activity from perfect likers\n",
    "print(\"\\nTop 5 date-hours with most activity from users with 100% action rate:\")\n",
    "print(date_hour_counts.head(5))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Get play count statistics\n",
    "print(perfect_clips[\"is_deleted\"].value_counts())"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "perfect_clips[\"is_deleted\"].value_counts()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# infill_clip_df = final_interesting_clips[(final_interesting_clips[\"task\"] == \"infill\") & (final_interesting_clips[\"model_name\"] == \"chirp-v4-6b-t-03\")]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# infill_clip_df[infill_clip_df[\"part_of_concat\"]].head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# all_pairs = clip_df[clip_df[\"request_id\"].isin(infill_clip_df[\"request_id\"].unique())].copy()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# all_pairs[\"flagged\"].value_counts()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "selected_model_name = \"chirp-auk-t1\"\n",
    "test_clip_auk_df = final_interesting_clips[\n",
    "    (final_interesting_clips[\"model_name\"] == selected_model_name)\n",
    "]\n",
    "print(\"selected_model_name\", test_clip_auk_df.shape)\n",
    "test_clip_df = final_interesting_clips[\n",
    "    final_interesting_clips[\"request_id\"].isin(test_clip_auk_df[\"request_id\"].unique())\n",
    "].copy()\n",
    "print(\"selected_model_name\", test_clip_df.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# test_clip_df.to_pickle(\n",
    "#     \"/home/tony/Data/Preference/auk/interesting_clips_exp_20250422_auk_t1.pkl\",\n",
    "# )\n",
    "# print(\"auk exps\", test_clip_df.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# subset_request_ids = user_intersting_clips_3p5[user_intersting_clips_3p5[\"model_name\"].isin([\"chirp-v4-6b-t-21_a_c_c_1\", \"chirp-v4-6b-t-21_a_c_c_2\"])][\"request_id\"].unique()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# test_subset_df = user_intersting_clips_3p5[user_intersting_clips_3p5[\"request_id\"].isin(subset_request_ids)]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# test_subset_df.to_pickle(\n",
    "#     \"/home/tony/Data/Preference/auk/interesting_clips_exp_20250422_auk_t1_sara_cfg.pkl\",\n",
    "# )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# from collections import defaultdict\n",
    "# from suno_utils.audio import Audio\n",
    "# audio_loundesses = defaultdict(list)\n",
    "# loundess_models = [\"chirp-v4-up-u-d-2-3\", \"chirp-ahi-up-1\", \"chirp-v4-up-u-7\"]\n",
    "# for test_model in loundess_models:\n",
    "#     subset_clip_df = final_interesting_clips[(final_interesting_clips[\"model_name\"] == test_model)][\"s3_id\"]\n",
    "#     print(subset_clip_df.shape)\n",
    "#     for index, s3_id in tqdm.tqdm(enumerate(subset_clip_df.unique())):\n",
    "#         if index > 50:\n",
    "#             break\n",
    "#         try:\n",
    "#             audio = Audio.from_s3(f\"s3://suno-data-uploads/studio/uploads/{s3_id}.mp3\", n_channels=2)\n",
    "#             loudness = audio.loudness\n",
    "#             audio_loundesses[test_model].append(loudness)\n",
    "#         except:\n",
    "#             pass\n",
    "# plt.clf()\n",
    "# for test_model in loundess_models:\n",
    "#     plt.hist(audio_loundesses[test_model], label=f\"{test_model}, mean {round(np.mean(audio_loundesses[test_model]), 2)}\", alpha=0.5, bins=np.linspace(-20, -10, 50))\n",
    "# plt.legend()\n",
    "# plt.title(\"loudness war\")\n",
    "# plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# final_interesting_clips[\"created_datetime\"] = pd.to_datetime(final_interesting_clips[\"created_at\"])\n",
    "# final_interesting_clips[\"hour\"] = final_interesting_clips[\"created_datetime\"].dt.strftime(\"%H\")\n",
    "# subset_request_ids = final_interesting_clips[final_interesting_clips[\"model_name\"].str.contains(\"tech\")][\"request_id\"].unique()\n",
    "# subset_final_interesting_clips = final_interesting_clips[final_interesting_clips[\"request_id\"].isin(subset_request_ids)].copy()\n",
    "# for fixed_hour in sorted(final_interesting_clips[\"hour\"].unique()):\n",
    "#     print(\"Fixed hour\", fixed_hour)\n",
    "#     get_preference_counts(\n",
    "#         subset_final_interesting_clips[subset_final_interesting_clips[\"hour\"] == fixed_hour],\n",
    "#         title_name=\"subset test\",\n",
    "#     )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# subset_request_ids = user_intersting_clips_3p5[user_intersting_clips_3p5[\"model_name\"].isin([\"chirp-auk-t1-d6\"])][\"request_id\"].unique()\n",
    "\n",
    "# subset_dur_user_intersting_clips_3p5 = user_intersting_clips_3p5[user_intersting_clips_3p5[\"request_id\"].isin(subset_request_ids)].copy()\n",
    "\n",
    "# subset_dur_user_intersting_clips_3p5[\"model_name\"].value_counts()\n",
    "\n",
    "# subset_dur_user_intersting_clips_3p5[subset_dur_user_intersting_clips_3p5[\"model_name\"] == \"chirp-auk-t1-d6\"][\"duration\"].describe()\n",
    "\n",
    "# subset_dur_user_intersting_clips_3p5[subset_dur_user_intersting_clips_3p5[\"model_name\"] == \"chirp-auk-t1\"][\"duration\"].describe()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(\"auk\", final_subset_auk_clips_df.shape)\n",
    "print(\"ahi\", final_subset_upsample_ahi_clips_df.shape)\n",
    "print(\"ahi sneaked\", final_subset_upsample_ahi_clips_df_2.shape)\n",
    "todays_save_date = \"20250608\"\n",
    "# total_ahi_df.to_pickle(\n",
    "#     f\"/home/tony/Data/Preference/up_v2_d4/interesting_clips_ahi_d4_{todays_save_date}.pkl\",\n",
    "# )\n",
    "# print(\"ahi_d3\", total_ahi_df.shape)\n",
    "##\n",
    "# final_subset_auk_og_clips_df.to_pickle(\n",
    "#     f\"/home/tony/Data/Preference/auk_t0/interesting_clips_auk_t0_{todays_save_date}.pkl\",\n",
    "# )\n",
    "# print(\"auk og\", final_subset_auk_og_clips_df.shape)\n",
    "\n",
    "# final_subset_auk_clips_df.to_pickle(\n",
    "#     f\"/home/tony/Data/Preference/auk_t1/interesting_clips_auk_t1_{todays_save_date}.pkl\",\n",
    "# )\n",
    "# final_subset_auk_infill_30b_clips_df.to_pickle(\n",
    "#     f\"/home/tony/Data/Preference/30b_t7/interesting_clips_30_infill_t1_{todays_save_date}.pkl\",\n",
    "# )\n",
    "print(\"auk_t1\", final_subset_auk_clips_df.shape)\n",
    "print(f\"Saving done! to {todays_save_date}\")\n",
    "\n",
    "# auk (4743256, 90)\n",
    "# ahi (367122, 90)\n",
    "# ahi sneaked (22902, 90)\n",
    "# auk og (116910, 90)\n",
    "# ahi_d3 (390024, 90)\n",
    "# Saving done!"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# final_interesting_clips[final_interesting_clips[\"model_name\"] == \"chirp-ahi-up-2\"][\"task\"].value_counts()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "suno_env_dev",
   "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
}
