{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 1,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T13:35:05.844764Z",
     "iopub.status.busy": "2025-11-20T13:35:05.844615Z",
     "iopub.status.idle": "2025-11-20T13:35:05.857756Z",
     "shell.execute_reply": "2025-11-20T13:35:05.857345Z",
     "shell.execute_reply.started": "2025-11-20T13:35:05.844748Z"
    }
   },
   "outputs": [],
   "source": [
    "# setup autoload\n",
    "%load_ext autoreload\n",
    "%autoreload 2"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 2,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:58:21.040680Z",
     "start_time": "2024-05-16T13:58:19.777010Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T13:35:05.859369Z",
     "iopub.status.busy": "2025-11-20T13:35:05.859250Z",
     "iopub.status.idle": "2025-11-20T13:35:08.374967Z",
     "shell.execute_reply": "2025-11-20T13:35:08.374422Z",
     "shell.execute_reply.started": "2025-11-20T13:35:05.859356Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "The autoreload extension is already loaded. To reload it, use:\n",
      "  %reload_ext autoreload\n"
     ]
    }
   ],
   "source": [
    "import ast\n",
    "import os\n",
    "import shutil\n",
    "import sys\n",
    "from collections import defaultdict\n",
    "\n",
    "import json\n",
    "import numpy as np\n",
    "import pandas as pd\n",
    "from preference_data_preparation_auk import *\n",
    "from preference_helper import *\n",
    "from sklearn.model_selection import train_test_split\n",
    "from suno_utils.utils.s3 import download_s3_files\n",
    "from suno_utils.utils.text import read_json, read_jsonl, write_json, write_jsonl\n",
    "from tqdm import tqdm\n",
    "\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",
    "# setup autoload\n",
    "%load_ext autoreload\n",
    "%autoreload 2\n",
    "\n",
    "\n",
    "def custom_parse(x):\n",
    "    try:\n",
    "        return json.loads(x)\n",
    "    except:\n",
    "        return {}"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 3,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:58:21.082172Z",
     "start_time": "2024-05-16T13:58:21.041926Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T13:35:08.376956Z",
     "iopub.status.busy": "2025-11-20T13:35:08.376827Z",
     "iopub.status.idle": "2025-11-20T13:35:08.565948Z",
     "shell.execute_reply": "2025-11-20T13:35:08.565476Z",
     "shell.execute_reply.started": "2025-11-20T13:35:08.376942Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "N_TOKENS_AUDIO 12000\n"
     ]
    }
   ],
   "source": [
    "OUT_DATA_DIR = \"/app2/suno/data/dpo/auk_t1_v101\"\n",
    "os.makedirs(OUT_DATA_DIR, exist_ok=True)\n",
    "shutil.copyfile(\n",
    "    \"/app/suno/data/dpo/7v_v20_full/tokenizer_60k.json\",\n",
    "    os.path.join(OUT_DATA_DIR, \"tokenizer_60k.json\"),\n",
    ")\n",
    "NPZ_DIR = \"/app2/suno/data/dpo/auk_t1_npz\"\n",
    "N_TOKENS_AUDIO = 25 * 8 * 60\n",
    "print(\"N_TOKENS_AUDIO\", N_TOKENS_AUDIO)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 4,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T13:35:08.567699Z",
     "iopub.status.busy": "2025-11-20T13:35:08.567574Z",
     "iopub.status.idle": "2025-11-20T13:41:53.370285Z",
     "shell.execute_reply": "2025-11-20T13:41:53.369709Z",
     "shell.execute_reply.started": "2025-11-20T13:35:08.567685Z"
    }
   },
   "outputs": [],
   "source": [
    "df = pd.read_pickle(\n",
    "    \"/home/tony/Data/Preference/auk_t1/fully_merged_auk_t1_20251118.pkl\"\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 5,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T13:41:53.372659Z",
     "iopub.status.busy": "2025-11-20T13:41:53.372523Z",
     "iopub.status.idle": "2025-11-20T13:42:11.317428Z",
     "shell.execute_reply": "2025-11-20T13:42:11.316850Z",
     "shell.execute_reply.started": "2025-11-20T13:41:53.372644Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "after dropna (12119078, 87)\n"
     ]
    }
   ],
   "source": [
    "df = df.dropna(axis=1, how=\"all\")\n",
    "print(\"after dropna\", df.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 6,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:58:56.199480Z",
     "start_time": "2024-05-16T13:58:53.963687Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T13:42:11.318177Z",
     "iopub.status.busy": "2025-11-20T13:42:11.318023Z",
     "iopub.status.idle": "2025-11-20T13:47:35.849778Z",
     "shell.execute_reply": "2025-11-20T13:47:35.849007Z",
     "shell.execute_reply.started": "2025-11-20T13:42:11.318161Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "12184904\n",
      "12184904\n",
      "pre-downloaded df (12119078, 87)\n",
      "downloaded df (12119072, 87)\n"
     ]
    }
   ],
   "source": [
    "converted_paths = os.listdir(NPZ_DIR)\n",
    "print(len(converted_paths))\n",
    "\n",
    "converted_paths = set([f.replace(\".npz\", \"\") for f in converted_paths])\n",
    "print(len(converted_paths))\n",
    "\n",
    "print(\"pre-downloaded df\", df.shape)\n",
    "df[df[\"s3_id\"].isin(converted_paths)].shape\n",
    "df = df[df[\"s3_id\"].isin(converted_paths)].copy()\n",
    "print(\"downloaded df\", df.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 7,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:58:56.467253Z",
     "start_time": "2024-05-16T13:58:56.207647Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T13:47:35.850840Z",
     "iopub.status.busy": "2025-11-20T13:47:35.850535Z",
     "iopub.status.idle": "2025-11-20T13:47:37.304132Z",
     "shell.execute_reply": "2025-11-20T13:47:37.303579Z",
     "shell.execute_reply.started": "2025-11-20T13:47:35.850823Z"
    }
   },
   "outputs": [
    {
     "data": {
      "text/plain": [
       "task\n",
       "                       8062037\n",
       "cover                  2152600\n",
       "artist_consistency     1001459\n",
       "artist_cover            382784\n",
       "extend                  249530\n",
       "upload_extend           223826\n",
       "artist_extend            46806\n",
       "underpainting               16\n",
       "cover_extend                 6\n",
       "overpainting                 6\n",
       "artist_cover_extend          2\n",
       "Name: count, dtype: int64"
      ]
     },
     "execution_count": 7,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "df[\"task\"].value_counts()"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# LET's do the data prep"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 8,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:58:56.592883Z",
     "start_time": "2024-05-16T13:58:56.470781Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T13:47:37.305120Z",
     "iopub.status.busy": "2025-11-20T13:47:37.304797Z",
     "iopub.status.idle": "2025-11-20T13:47:44.707463Z",
     "shell.execute_reply": "2025-11-20T13:47:44.706735Z",
     "shell.execute_reply.started": "2025-11-20T13:47:37.305104Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "preference  model_name  \n",
      "False       chirp-auk-t1    6059792\n",
      "True        chirp-auk-t1    6059280\n",
      "Name: count, dtype: int64\n",
      "before filter on model name (12119072, 87)\n",
      "after filter on model name (12119072, 87)\n",
      "is_public\n",
      "False    11522477\n",
      "True       596595\n",
      "Name: count, dtype: int64\n"
     ]
    }
   ],
   "source": [
    "## for 13b this is easy for now\n",
    "print(df.groupby([\"preference\"])[\"model_name\"].value_counts())\n",
    "print(\"before filter on model name\", df.shape)\n",
    "df = df[df[\"model_name\"].isin([\"chirp-auk-t1\"])]\n",
    "# df = df[df[\"model_name\"].isin([\"chirp-v3p5-engine-t-6\"])]\n",
    "print(\"after filter on model name\", df.shape)\n",
    "print(df[\"is_public\"].value_counts())"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 9,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:58:56.909539Z",
     "start_time": "2024-05-16T13:58:56.595736Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T13:47:44.708494Z",
     "iopub.status.busy": "2025-11-20T13:47:44.708213Z",
     "iopub.status.idle": "2025-11-20T13:48:17.356782Z",
     "shell.execute_reply": "2025-11-20T13:48:17.356023Z",
     "shell.execute_reply.started": "2025-11-20T13:47:44.708477Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "before filter on request id pairs (12119072, 87)\n",
      "after filter on request id pairs (12119066, 87)\n",
      "preference  model_name  \n",
      "False       chirp-auk-t1    6059792\n",
      "True        chirp-auk-t1    6059274\n",
      "Name: count, dtype: int64\n"
     ]
    }
   ],
   "source": [
    "print(\"before filter on request id pairs\", df.shape)\n",
    "df = df[\n",
    "    df[\"request_id\"].isin(\n",
    "        df[\"request_id\"].value_counts().index[df[\"request_id\"].value_counts() == 2]\n",
    "    )\n",
    "]\n",
    "print(\"after filter on request id pairs\", df.shape)\n",
    "print(df.groupby([\"preference\"])[\"model_name\"].value_counts())\n",
    "assert df.shape[0] == df[\"request_id\"].nunique() * 2"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 10,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T13:48:17.357898Z",
     "iopub.status.busy": "2025-11-20T13:48:17.357588Z",
     "iopub.status.idle": "2025-11-20T14:31:25.511594Z",
     "shell.execute_reply": "2025-11-20T14:31:25.504092Z",
     "shell.execute_reply.started": "2025-11-20T13:48:17.357879Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "unique_requests 6059533\n",
      "before removing duplicates (12119066, 202)\n",
      "after removing duplicates (12119066, 195)\n"
     ]
    }
   ],
   "source": [
    "# Let's use the old selection for now -- for quality assurance\n",
    "# expand the metadata columns -- this takes forever...~ 6 mins\n",
    "# test_slice = df[\"metadata\"].apply(lambda x: ast.literal_eval(str(x)))\n",
    "# test_slice = df[\"metadata\"].apply(lambda x: custom_parse(x))\n",
    "test_slice = df[\"metadata\"]  # .apply(lambda x: custom_parse(x))\n",
    "test_slice_series = test_slice.apply(pd.Series)\n",
    "df = pd.concat([df, test_slice_series], axis=1, join=\"inner\")\n",
    "print(\"unique_requests\", df[\"request_id\"].nunique())\n",
    "# remove the duplicates\n",
    "print(\"before removing duplicates\", df.shape)\n",
    "df = df.loc[:, ~df.columns.duplicated()].copy()\n",
    "print(\"after removing duplicates\", df.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 11,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T14:31:25.555680Z",
     "iopub.status.busy": "2025-11-20T14:31:25.555413Z",
     "iopub.status.idle": "2025-11-20T14:32:24.488168Z",
     "shell.execute_reply": "2025-11-20T14:32:24.487655Z",
     "shell.execute_reply.started": "2025-11-20T14:31:25.555661Z"
    }
   },
   "outputs": [
    {
     "data": {
      "text/plain": [
       "pos_diff_preference\n",
       " 1.0    4132060\n",
       " 2.0    1927111\n",
       " 0.0         88\n",
       "-1.0         15\n",
       "Name: count, dtype: int64"
      ]
     },
     "execution_count": 11,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "df = df.sort_values(by=[\"request_id\", \"preference\", \"diff_preference\"])\n",
    "df[\"pos_diff_preference\"] = df[\"diff_preference\"].diff()\n",
    "# df[\"cer_diff_preference\"] = df[\"cer\"].diff()\n",
    "df[df[\"preference\"]][\"pos_diff_preference\"].value_counts()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 12,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T14:32:24.488928Z",
     "iopub.status.busy": "2025-11-20T14:32:24.488773Z",
     "iopub.status.idle": "2025-11-20T14:32:38.803899Z",
     "shell.execute_reply": "2025-11-20T14:32:38.803344Z",
     "shell.execute_reply.started": "2025-11-20T14:32:24.488913Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "positive param_experiment\n",
      "n_tag_2                127936\n",
      "cfg_steps_240          127319\n",
      "n_tag_1                126656\n",
      "tag_cfg_1              125938\n",
      "temp_s_95              125790\n",
      "cfg_steps_60           125756\n",
      "temp_s_85              125409\n",
      "temp_s_80              123850\n",
      "cfg_steps_10           121981\n",
      "tag_cfg_3              120745\n",
      "mask_control_slider     61189\n",
      "Name: count, dtype: int64\n"
     ]
    }
   ],
   "source": [
    "try:\n",
    "    print(\"positive\", df[df[\"preference\"]][\"param_experiment\"].value_counts())\n",
    "except:\n",
    "    pass"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 13,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T14:32:38.804645Z",
     "iopub.status.busy": "2025-11-20T14:32:38.804497Z",
     "iopub.status.idle": "2025-11-20T14:37:49.886197Z",
     "shell.execute_reply": "2025-11-20T14:37:49.885644Z",
     "shell.execute_reply.started": "2025-11-20T14:32:38.804630Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Found 356757 duplicated prompts 178383 unique requests\n",
      "Found 85608 request_ids with duplicate prompts but not highest play counts in their group\n",
      "['de68121d-03f0-4e92-96a4-dbfc192103dc', '151ca9d7-1353-4a2c-bb14-a3d61cc79c34', '341aad29-3d56-4963-be32-bdf3757e6390', '17e29d69-1e92-4304-a9f5-d3a60c62b006', '3d19c8b1-edc1-4b0c-90e5-39dc4679cff2', '1e32db78-9c98-4995-8492-a5e98aaf65ac', '730936a6-f6bb-4f43-a54d-d698b88b8cb6', '6deaeb48-8c9b-4964-af81-fd5ef0aaf20a', 'fd2c071e-3d91-4149-aa31-28256a78ce9b', '6e3c55be-939c-41b0-9e57-f997925a5d5e']\n",
      "Before dedup user gen requests 12119066\n",
      "After dedup user gen requests 12119066\n"
     ]
    }
   ],
   "source": [
    "# Find duplicated prompts with count > 2\n",
    "duplicate_entries = df.groupby(\n",
    "    [\"user_id\", \"prompt_text\", \"tags\", \"task\", \"edited_clip_id\"]\n",
    ").filter(lambda x: len(x) > 2)\n",
    "print(\n",
    "    \"Found\",\n",
    "    len(duplicate_entries),\n",
    "    \"duplicated prompts\",\n",
    "    len(duplicate_entries[\"request_id\"].unique()),\n",
    "    \"unique requests\",\n",
    ")\n",
    "\n",
    "# Group by user_id, prompt_text, and tags to find duplicate prompt groups\n",
    "prompt_groups = duplicate_entries.groupby(\n",
    "    [\"user_id\", \"prompt_text\", \"tags\", \"task\", \"edited_clip_id\"]\n",
    ")\n",
    "\n",
    "# For each prompt group, find the request_id with the highest total reaction_play_count\n",
    "low_play_count_request_ids = []\n",
    "for prompt_key, prompt_group in prompt_groups:\n",
    "    # Get the sum of reaction_play_count for each request_id in this group\n",
    "    request_play_counts = prompt_group.groupby(\"request_id\")[\n",
    "        \"reaction_play_count\"\n",
    "    ].sum()\n",
    "\n",
    "    # Find the max play count in this group\n",
    "    max_play_count = request_play_counts.max()\n",
    "\n",
    "    # Add request_ids that don't have the max play count to our filter list\n",
    "    lower_play_count_request_ids = request_play_counts[\n",
    "        request_play_counts < max_play_count\n",
    "    ].index.tolist()\n",
    "    low_play_count_request_ids.extend(lower_play_count_request_ids)\n",
    "\n",
    "# Display the filtered request IDs\n",
    "print(\n",
    "    f\"Found {len(low_play_count_request_ids)} request_ids with duplicate prompts but not highest play counts in their group\"\n",
    ")\n",
    "print(\n",
    "    low_play_count_request_ids[:10]\n",
    "    if len(low_play_count_request_ids) > 10\n",
    "    else low_play_count_request_ids\n",
    ")\n",
    "print(\"Before dedup user gen requests\", df.shape[0])\n",
    "# df = df[~df[\"request_id\"].isin(low_play_count_request_ids)]\n",
    "print(\"After dedup user gen requests\", df.shape[0])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 14,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:40.799375Z",
     "start_time": "2024-05-16T13:59:36.394236Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T14:37:49.886899Z",
     "iopub.status.busy": "2025-11-20T14:37:49.886753Z",
     "iopub.status.idle": "2025-11-20T14:40:20.216907Z",
     "shell.execute_reply": "2025-11-20T14:40:20.216339Z",
     "shell.execute_reply.started": "2025-11-20T14:37:49.886885Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "192713\n",
      "good_continue_at\n",
      "True     12091443\n",
      "False       27623\n",
      "Name: count, dtype: int64\n",
      "\n",
      " Check some basics... \n",
      " preference\n",
      "False    6059792\n",
      "True     6059274\n",
      "Name: count, dtype: int64 model_name\n",
      "chirp-auk-t1    12119066\n",
      "Name: count, dtype: int64 preference  model_name  \n",
      "False       chirp-auk-t1    6059792\n",
      "True        chirp-auk-t1    6059274\n",
      "Name: count, dtype: int64\n",
      "task\n",
      "                       8062032\n",
      "cover                  2152600\n",
      "artist_consistency     1001458\n",
      "artist_cover            382784\n",
      "extend                  249530\n",
      "upload_extend           223826\n",
      "artist_extend            46806\n",
      "underpainting               16\n",
      "cover_extend                 6\n",
      "overpainting                 6\n",
      "artist_cover_extend          2\n",
      "Name: count, dtype: int64\n"
     ]
    }
   ],
   "source": [
    "df[\"id\"] = df[\"str_id\"]\n",
    "# get the original duration of the clips, if they are concacted\n",
    "df[\"original_duration_s\"] = df[\"total_start_s\"] + df[\"duration\"]\n",
    "# classify the continue at behavoirs by the duration choice\n",
    "audio_prompt_id_to_continue_at = {}\n",
    "\n",
    "for _, row in df[~df[\"continued_parent\"].isna()].iterrows():\n",
    "    audio_prompt_id = row[\"continued_parent\"]\n",
    "    if audio_prompt_id not in audio_prompt_id_to_continue_at:\n",
    "        audio_prompt_id_to_continue_at[audio_prompt_id] = row[\"continue_at\"]\n",
    "    else:\n",
    "        # pick the max\n",
    "        audio_prompt_id = max(\n",
    "            audio_prompt_id_to_continue_at[audio_prompt_id], row[\"continue_at\"]\n",
    "        )\n",
    "print(len(audio_prompt_id_to_continue_at))\n",
    "df[\"has_continue_and_start_continue_at\"] = df[\"id\"].apply(\n",
    "    lambda x: audio_prompt_id_to_continue_at.get(x)\n",
    ")\n",
    "# we want continue at to be at most of the clip...\n",
    "df[\"good_continue_at\"] = (\n",
    "    (df[\"has_continue_and_start_continue_at\"] / df[\"duration\"]) > 0.9\n",
    ") | df[\"has_continue_and_start_continue_at\"].isna()\n",
    "print(df[\"good_continue_at\"].value_counts())\n",
    "\n",
    "\n",
    "print(\n",
    "    \"\\n Check some basics... \\n\",\n",
    "    df[\"preference\"].value_counts(),\n",
    "    df[\"model_name\"].value_counts(),\n",
    "    df.groupby([\"preference\"])[\"model_name\"].value_counts(),\n",
    ")\n",
    "\n",
    "df = df.sort_values(by=[\"request_id\", \"preference\"])\n",
    "df[\"duration_rel_diff\"] = df[\"duration\"].diff()\n",
    "df[\"play_rel_diff\"] = df[\"reaction_play_count\"].diff()\n",
    "print(df[\"task\"].value_counts())\n",
    "\n",
    "# df[\"post_infill_duration\"] = (\n",
    "#     df[\"duration\"]\n",
    "#     + df[\"infill_context_end_s\"]\n",
    "#     - df[\"infill_context_start_s\"]\n",
    "#     - df[\"include_future_s\"]\n",
    "#     - df[\"include_history_s\"]\n",
    "#     - df[\"infill_dur_s\"]\n",
    "# )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 21,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T15:45:38.054779Z",
     "iopub.status.busy": "2025-11-20T15:45:38.054430Z",
     "iopub.status.idle": "2025-11-20T15:46:07.061300Z",
     "shell.execute_reply": "2025-11-20T15:46:07.060728Z",
     "shell.execute_reply.started": "2025-11-20T15:45:38.054761Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "(12119066, 202)\n",
      "(12118312, 202)\n"
     ]
    }
   ],
   "source": [
    "# sth maybe happening with clip_id duplication?\n",
    "print(df.shape)\n",
    "df[\"request_sum\"] = df.groupby(\"request_id\")[\"preference\"].transform(\"sum\")\n",
    "# creation of interesting_clips\n",
    "df = df[(df[\"request_sum\"] == 1)]\n",
    "print(df.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 68,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T16:25:44.949990Z",
     "iopub.status.busy": "2025-11-20T16:25:44.949638Z",
     "iopub.status.idle": "2025-11-20T16:26:10.248676Z",
     "shell.execute_reply": "2025-11-20T16:26:10.248115Z",
     "shell.execute_reply.started": "2025-11-20T16:25:44.949972Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "positive preference_score\n",
      "1    3148675\n",
      "2    1771225\n",
      "0     777011\n",
      "3     338192\n",
      "4      22654\n",
      "5       1399\n",
      "Name: count, dtype: int64\n",
      "negative preference_score\n",
      "0    5943538\n",
      "1     105598\n",
      "2       9384\n",
      "3        618\n",
      "4         18\n",
      "Name: count, dtype: int64\n"
     ]
    }
   ],
   "source": [
    "df[\"preference_score\"] = (\n",
    "    df[\"upvoted\"].astype(int)\n",
    "    + df[\"has_action\"].astype(int)\n",
    "    + df[\"part_of_concat\"].astype(int)\n",
    "    + df[\"is_in_playlist\"].astype(int)\n",
    "    + (df[\"n_edits\"] >= 10).astype(int)\n",
    ")\n",
    "print(\"positive\", df[df[\"preference\"]][\"preference_score\"].value_counts())\n",
    "print(\"negative\", df[~df[\"preference\"]][\"preference_score\"].value_counts())"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 73,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:41.035167Z",
     "start_time": "2024-05-16T13:59:40.801098Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T16:28:53.783585Z",
     "iopub.status.busy": "2025-11-20T16:28:53.783033Z",
     "iopub.status.idle": "2025-11-20T16:29:48.512613Z",
     "shell.execute_reply": "2025-11-20T16:29:48.512014Z",
     "shell.execute_reply.started": "2025-11-20T16:28:53.783567Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "after duration 0.9995443259754329\n",
      "after infill duration 1.0\n",
      "neg_filter_reaction_play_count 1.0\n",
      "neg_filter_upvote_count 0.9892\n",
      "neg_filter_norm_play_frac 1.0\n",
      "neg_filter_continues 0.9999\n",
      "neg_filter_preference_score 0.9809\n",
      "----------------\n",
      "pos_filter_continues 0.9955\n",
      "pos_filter_reaction_play_count 1.0\n",
      "pos_filter_relative_play_count 0.9762\n",
      "pos_filter_cer_diff_preference 1.0\n",
      "pos_filter_bad_flags 0.9998\n",
      "after filter on play counts 0.292\n",
      "after filter on higher quality 0.2584\n",
      "pos_filter_preference_score 0.3521\n",
      "----------------\n",
      "negative 5936195 positive 504212\n",
      "----------------\n",
      "total pair requests 6059156  --> selected pair requests 491594 frac 0.081  --> total intitial users 435182\n"
     ]
    }
   ],
   "source": [
    "normal_pos_play_count = 10\n",
    "# this is lower, cause a concat is probably already ensuring that it is good\n",
    "concat_pos_play_count = 2\n",
    "# this is a filter on the concated clip\n",
    "concat_total_play_count = 10\n",
    "\n",
    "all_fitlers = (df[\"duration\"] >= 10) & (df[\"duration\"] <= 480)\n",
    "print(\"after duration\", all_fitlers.sum() / df.shape[0])\n",
    "infill_duration_filter = ~df[\"task\"].isin(\n",
    "    [\n",
    "        \"infill\",\n",
    "        \"infill_intro\",\n",
    "        \"infill_outro\",\n",
    "    ]\n",
    ")  # | (df[\"post_infill_duration\"] <= 239)\n",
    "print(\"after infill duration\", infill_duration_filter.sum() / df.shape[0])\n",
    "# negative fitlers\n",
    "total_negative = df[~df[\"preference\"]].shape[0]\n",
    "neg_filter_reaction_play_count = (~df[\"preference\"]) & (df[\"reaction_play_count\"] >= 1)\n",
    "print(\n",
    "    \"neg_filter_reaction_play_count\",\n",
    "    round(neg_filter_reaction_play_count.sum() / total_negative, 4),\n",
    ")\n",
    "neg_filter_upvote_count = (~df[\"preference\"]) & (df[\"upvote_count\"] == 0)\n",
    "print(\n",
    "    \"neg_filter_upvote_count\",\n",
    "    round(neg_filter_upvote_count.sum() / total_negative, 4),\n",
    ")\n",
    "neg_filter_norm_play_frac = (~df[\"preference\"]) & (df[\"norm_play_frac\"] <= 3.1)\n",
    "print(\n",
    "    \"neg_filter_norm_play_frac\",\n",
    "    round(neg_filter_norm_play_frac.sum() / total_negative, 4),\n",
    ")\n",
    "neg_filter_continues = (~df[\"preference\"]) & (\n",
    "    df[\"has_continue_and_start_continue_at\"].isna()\n",
    ")\n",
    "print(\n",
    "    \"neg_filter_continues\",\n",
    "    round(neg_filter_continues.sum() / total_negative, 4),\n",
    ")\n",
    "neg_filter_preference_score = (~df[\"preference\"]) & (\n",
    "    df[\"preference_score\"] == 0\n",
    ")\n",
    "print(\n",
    "    \"neg_filter_preference_score\",\n",
    "    round(neg_filter_preference_score.sum() / total_negative, 4),\n",
    ")\n",
    "\n",
    "neg_filter_selection_mask = (\n",
    "    all_fitlers\n",
    "    & infill_duration_filter\n",
    "    & neg_filter_reaction_play_count\n",
    "    & neg_filter_upvote_count\n",
    "    & neg_filter_norm_play_frac\n",
    "    & neg_filter_continues\n",
    "    & neg_filter_preference_score\n",
    ")\n",
    "\n",
    "print(\"----------------\")\n",
    "total_positive = df[df[\"preference\"]].shape[0]\n",
    "assert total_positive == total_negative\n",
    "pos_filter_continues = (df[\"preference\"]) & (df[\"good_continue_at\"])\n",
    "print(\"pos_filter_continues\", round(pos_filter_continues.sum() / total_positive, 4))\n",
    "pos_filter_reaction_play_count = (df[\"preference\"]) & (df[\"reaction_play_count\"] >= 1)\n",
    "print(\n",
    "    \"pos_filter_reaction_play_count\",\n",
    "    round(pos_filter_reaction_play_count.sum() / total_positive, 4),\n",
    ")\n",
    "pos_filter_relative_play_count = (df[\"preference\"]) & (df[\"play_rel_diff\"] >= 0)\n",
    "print(\n",
    "    \"pos_filter_relative_play_count\",\n",
    "    round(pos_filter_relative_play_count.sum() / total_positive, 4),\n",
    ")\n",
    "pos_filter_cer_diff_preference = (\n",
    "    df[\n",
    "        \"preference\"\n",
    "    ]  # & (df[\"pos_diff_preference\"] == 2) # & (df[\"cer_diff_preference\"] < 0.5) & (df[\"cer\"] < 0.99)\n",
    ")\n",
    "print(\n",
    "    \"pos_filter_cer_diff_preference\",\n",
    "    round(pos_filter_cer_diff_preference.sum() / total_positive, 4),\n",
    ")\n",
    "pos_filter_bad_flags = (\n",
    "    (df[\"preference\"]) & (df[\"flag_count\"] == 0) & (df[\"dislike_count\"] == 0)\n",
    ")\n",
    "print(\n",
    "    \"pos_filter_bad_flags\",\n",
    "    round(pos_filter_bad_flags.sum() / total_positive, 4),\n",
    ")\n",
    "pos_filter_play_counts = (df[\"preference\"]) & (\n",
    "    (\n",
    "        (df[\"part_of_concat\"])\n",
    "        & (df[\"reaction_play_count\"] >= concat_pos_play_count)\n",
    "        & (df[\"concat_play_counts\"] >= concat_total_play_count)\n",
    "    )\n",
    "    | (\n",
    "        (~df[\"part_of_concat\"]) & (df[\"reaction_play_count\"] >= normal_pos_play_count)\n",
    "        # & (df[\"norm_play_frac\"] >= 2.1)  # this is a bit of a luxury cut...\n",
    "    )\n",
    "    | (df[\"task\"].isin([\"infill\", \"infill_intro\", \"infill_outro\"]))\n",
    ")\n",
    "print(\n",
    "    \"after filter on play counts\",\n",
    "    round(pos_filter_play_counts.sum() / total_positive, 4),\n",
    ")\n",
    "high_quality_tasks_filter = (\n",
    "    (\n",
    "        df[\"task\"].isin(\n",
    "            [\n",
    "                \"cover\",\n",
    "                \"upload_extend\",\n",
    "                \"cover_extend\",\n",
    "                \"artist_cover\",\n",
    "                \"extend\",\n",
    "                \"artist_consistency\",\n",
    "                \"artist_extend\",\n",
    "                \"\",\n",
    "            ]\n",
    "        )\n",
    "    )\n",
    "    & (\n",
    "        (df[\"upvote_count\"] >= 1)  # (df[\"upvote_count\"] >= 1)\n",
    "        | (df[\"reaction_play_count\"] >= 10)\n",
    "        | (df[\"concat_play_counts\"] >= 10)\n",
    "    )\n",
    "    & (\n",
    "        (df[\"part_of_concat\"])\n",
    "        | (\n",
    "            (~df[\"part_of_concat\"])\n",
    "            & (df[\"norm_play_frac\"] >= 5.1)  # this is a bit of a luxury cut...\n",
    "            & (\n",
    "                df[\"norm_play_frac\"] >= df[\"reaction_play_count\"] / 3\n",
    "            )  # play duration is not low on average\n",
    "        )\n",
    "    )\n",
    ")\n",
    "medium_quality_tasks_filter = (\n",
    "    df[\"task\"].isin(\n",
    "        [\n",
    "            \"infill\",\n",
    "            \"infill_intro\",\n",
    "            \"infill_outro\",\n",
    "        ]\n",
    "    )\n",
    ") & (  # let more infill through only in this case...\n",
    "    (\n",
    "        df[\"upvote_count\"] >= 0\n",
    "    )  # (df[\"upvote_count\"] >= 1)  (df[\"pos_diff_preference\"] == 2)\n",
    "    | (df[\"reaction_play_count\"] >= 1)\n",
    "    | (df[\"concat_play_counts\"] >= 1)\n",
    ")\n",
    "pos_filter_higher_quality = (df[\"preference\"]) & (\n",
    "    high_quality_tasks_filter | medium_quality_tasks_filter\n",
    ")\n",
    "print(\n",
    "    \"after filter on higher quality\",\n",
    "    round(pos_filter_higher_quality.sum() / total_positive, 4),\n",
    ")\n",
    "\n",
    "user_gen_filter = (\n",
    "    df[\"user_n_clips\"] >= 40\n",
    ")  # user needs to have genereated at least 100 over the time period\n",
    "\n",
    "pos_filter_preference_score = (df[\"preference\"]) & (\n",
    "    df[\"preference_score\"] >= 2\n",
    ")\n",
    "print(\n",
    "    \"pos_filter_preference_score\",\n",
    "    round(pos_filter_preference_score.sum() / total_positive, 4),\n",
    ")\n",
    "\n",
    "print(\"----------------\")\n",
    "pos_filter_selectin_mask = (\n",
    "    (df[\"preference\"])  # get basics aligned\n",
    "    & all_fitlers\n",
    "    & infill_duration_filter\n",
    "    & pos_filter_continues\n",
    "    & pos_filter_reaction_play_count\n",
    "    & pos_filter_relative_play_count\n",
    "    & pos_filter_cer_diff_preference\n",
    "    & pos_filter_bad_flags\n",
    "    & pos_filter_play_counts\n",
    "    & pos_filter_higher_quality\n",
    "    & user_gen_filter\n",
    "    & pos_filter_preference_score\n",
    ")\n",
    "print(\n",
    "    \"negative\",\n",
    "    sum(neg_filter_selection_mask),\n",
    "    \"positive\",\n",
    "    sum(pos_filter_selectin_mask),\n",
    ")\n",
    "\n",
    "neg_filter_requests = df[neg_filter_selection_mask][\"request_id\"].unique()\n",
    "pos_filter_requests = df[pos_filter_selectin_mask][\"request_id\"].unique()\n",
    "# looking for very strong signal here:\n",
    "# listen to the positive/negative more than once\n",
    "# disliked one of the clips\n",
    "unique_requests = set(pos_filter_requests).intersection(neg_filter_requests)\n",
    "print(\"----------------\")\n",
    "print(\n",
    "    \"total pair requests\",\n",
    "    df[\"request_id\"].nunique(),\n",
    "    \" --> selected pair requests\",\n",
    "    len(unique_requests),\n",
    "    f\"frac {len(unique_requests) / df['request_id'].nunique():.3f}\",\n",
    "    \" --> total intitial users\",\n",
    "    df[\"user_id\"].nunique(),\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 78,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:41.250737Z",
     "start_time": "2024-05-16T13:59:41.036434Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T16:32:25.804129Z",
     "iopub.status.busy": "2025-11-20T16:32:25.803791Z",
     "iopub.status.idle": "2025-11-20T16:32:35.243839Z",
     "shell.execute_reply": "2025-11-20T16:32:35.243246Z",
     "shell.execute_reply.started": "2025-11-20T16:32:25.804112Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "auk_t1_v101 requests 491594 clips 983188 total khrs 51.733; N gpus for 1000 iters 61.449; 16 gpus for x iters 960.145; n unique users 116703 n pro users 116179\n"
     ]
    }
   ],
   "source": [
    "df_slice = df[df[\"request_id\"].isin(set(unique_requests))].copy()\n",
    "print(\n",
    "    f\"{os.path.basename(OUT_DATA_DIR)} requests\",\n",
    "    df_slice[\"request_id\"].nunique(),\n",
    "    \"clips\",\n",
    "    df_slice.shape[0],\n",
    "    f\"total khrs {sum(df_slice['duration'] / 3600 / 1000):.3f};\",\n",
    "    f\"N gpus for 1000 iters {df_slice.shape[0] / 8 / 2 / 1000:.3f};\",\n",
    "    # f\"4 gpus for x iters {df_slice.shape[0] / 8 / 2 / 4:.3f};\",\n",
    "    f\"16 gpus for x iters {df_slice.shape[0] / 8 / 8 / 16:.3f};\",\n",
    "    f\"n unique users {df_slice['user_id'].nunique()}\",\n",
    "    f\"n pro users {df_slice[df_slice['is_pro_user']]['user_id'].nunique()}\",\n",
    ")\n",
    "# auk_mix_t1_v2 requests 102002 clips 204004 total khrs 9.191; N gpus for 1000 iters 12.750; 4 gpus for x iters 3187.562; n unique users 36408 n pro users 34038\n",
    "# auk_t1_v1 requests 9179 clips 18358 total khrs 0.854; N gpus for 1000 iters 1.147; 4 gpus for x iters 286.844; n unique users 6288 n pro users 6275\n",
    "# auk_t1_v2 requests 40903 clips 81806 total khrs 3.864; N gpus for 1000 iters 5.113; 4 gpus for x iters 1278.219; n unique users 21079 n pro users 20966\n",
    "# auk_t1_v3 requests 102015 clips 204030 total khrs 9.703; N gpus for 1000 iters 12.752; 4 gpus for x iters 3187.969; n unique users 42837 n pro users 42462\n",
    "# auk_t1_v4 requests 211452 clips 422904 total khrs 20.216; N gpus for 1000 iters 26.431; 4 gpus for x iters 6607.875; n unique users 71182 n pro users 70059\n",
    "# auk_t1_v9 requests 547743 clips 1095486 total khrs 56.567; N gpus for 1000 iters 68.468; 4 gpus for x iters 17116.969; n unique users 129354 n pro users 125126\n",
    "# auk_t1_v13 requests 627646 clips 1255292 total khrs 64.530; N gpus for 1000 iters 78.456; 4 gpus for x iters 19613.938; n unique users 139873 n pro users 134738\n",
    "# auk_t1_v17 requests 494902 clips 989804 total khrs 51.445; N gpus for 1000 iters 61.863; 4 gpus for x iters 15465.688; n unique users 114240 n pro users 109708\n",
    "# auk_t1_v19 requests 740881 clips 1481762 total khrs 76.133; N gpus for 1000 iters 92.610; 4 gpus for x iters 23152.531; n unique users 154570 n pro users 147261\n",
    "# auk_t1_v24 requests 717745 clips 1435490 total khrs 74.515; N gpus for 1000 iters 89.718; 4 gpus for x iters 22429.531; n unique users 139344 n pro users 130740\n",
    "# auk_t1_v29 requests 1097586 clips 2195172 total khrs 112.614; N gpus for 1000 iters 137.198; 4 gpus for x iters 34299.562; n unique users 193138 n pro users 179429\n",
    "# auk_t1_v30 requests 1224422 clips 2448844 total khrs 125.330; N gpus for 1000 iters 153.053; 4 gpus for x iters 38263.188; n unique users 203138 n pro users 186758\n",
    "# auk_t1_v31 requests 389265 clips 778530 total khrs 40.761; N gpus for 1000 iters 48.658; 4 gpus for x iters 12164.531; n unique users 85198 n pro users 78445\n",
    "# auk_t1_v33 requests 1285261 clips 2570522 total khrs 131.585; N gpus for 1000 iters 160.658; 4 gpus for x iters 40164.406; n unique users 208932 n pro users 191380\n",
    "# auk_t1_v33 requests 1086639 clips 2173278 total khrs 110.641; N gpus for 1000 iters 135.830; 4 gpus for x iters 33957.469; n unique users 197293 n pro users 180968 -- play dur from /3 to /2\n",
    "# auk_t1_v33 requests 806877 clips 1613754 total khrs 81.964; N gpus for 1000 iters 100.860; 4 gpus for x iters 25214.906; n unique users 156679 n pro users 144253 -- filter to web\n",
    "# auk_t1_v37 requests 1339132 clips 2678264 total khrs 137.055; N gpus for 1000 iters 167.392; 4 gpus for x iters 41847.875; n unique users 214313 n pro users 196815\n",
    "# auk_t1_v38 requests 979451 clips 1958902 total khrs 102.038; N gpus for 1000 iters 122.431; 4 gpus for x iters 30607.844; n unique users 170773 n pro users 156737\n",
    "# auk_t1_v43 requests 184775 clips 369550 total khrs 19.089; N gpus for 1000 iters 23.097; 4 gpus for x iters 5774.219; n unique users 60652 n pro users 59126\n",
    "# auk_t1_v45 requests 1176216 clips 2352432 total khrs 122.102; N gpus for 1000 iters 147.027; 4 gpus for x iters 36756.750; n unique users 176751 n pro users 163097\n",
    "# auk_t1_v48 requests 294407 clips 588814 total khrs 29.992; N gpus for 1000 iters 36.801; 4 gpus for x iters 9200.219; n unique users 80991 n pro users 78825\n",
    "# auk_t1_v101 requests 491594 clips 983188 total khrs 51.733; N gpus for 1000 iters 61.449; 16 gpus for x iters 960.145; n unique users 116703 n pro users 116179"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 79,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T16:32:35.244840Z",
     "iopub.status.busy": "2025-11-20T16:32:35.244675Z",
     "iopub.status.idle": "2025-11-20T16:32:35.457369Z",
     "shell.execute_reply": "2025-11-20T16:32:35.456906Z",
     "shell.execute_reply.started": "2025-11-20T16:32:35.244824Z"
    }
   },
   "outputs": [
    {
     "data": {
      "text/plain": [
       "source\n",
       "web        775414\n",
       "ios        134002\n",
       "android     73772\n",
       "Name: count, dtype: int64"
      ]
     },
     "execution_count": 79,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "df_slice[\"source\"].value_counts()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 80,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:41.277006Z",
     "start_time": "2024-05-16T13:59:41.252105Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T16:32:35.458046Z",
     "iopub.status.busy": "2025-11-20T16:32:35.457904Z",
     "iopub.status.idle": "2025-11-20T16:32:36.150838Z",
     "shell.execute_reply": "2025-11-20T16:32:36.150286Z",
     "shell.execute_reply.started": "2025-11-20T16:32:35.458032Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "positive in playlist (271101, 203)\n",
      "task\n",
      "                      672680\n",
      "cover                 147404\n",
      "artist_consistency    101424\n",
      "artist_cover           36132\n",
      "extend                 12306\n",
      "upload_extend          10622\n",
      "artist_extend           2620\n",
      "Name: count, dtype: int64\n"
     ]
    }
   ],
   "source": [
    "test_mask = (df_slice[\"preference\"]) & (\n",
    "    (df_slice[\"is_in_playlist\"]) | (df_slice[\"concat_in_playlist\"])\n",
    ")\n",
    "print(\"positive in playlist\", df_slice[test_mask].shape)\n",
    "print(df_slice[\"task\"].value_counts())"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 81,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T16:32:36.152051Z",
     "iopub.status.busy": "2025-11-20T16:32:36.151887Z",
     "iopub.status.idle": "2025-11-20T16:32:36.171393Z",
     "shell.execute_reply": "2025-11-20T16:32:36.170939Z",
     "shell.execute_reply.started": "2025-11-20T16:32:36.152035Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "is_public\n",
      "False    845788\n",
      "True     137400\n",
      "Name: count, dtype: int64\n"
     ]
    }
   ],
   "source": [
    "print(df_slice[\"is_public\"].value_counts())"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 82,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T16:32:36.171992Z",
     "iopub.status.busy": "2025-11-20T16:32:36.171840Z",
     "iopub.status.idle": "2025-11-20T16:32:36.571966Z",
     "shell.execute_reply": "2025-11-20T16:32:36.571415Z",
     "shell.execute_reply.started": "2025-11-20T16:32:36.171978Z"
    }
   },
   "outputs": [],
   "source": [
    "df_slice[\"npz_path\"] = df_slice[\"s3_id\"].map(lambda x: f\"{NPZ_DIR}/{x}.npz\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 83,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T16:32:36.572696Z",
     "iopub.status.busy": "2025-11-20T16:32:36.572547Z",
     "iopub.status.idle": "2025-11-20T16:32:36.587701Z",
     "shell.execute_reply": "2025-11-20T16:32:36.587249Z",
     "shell.execute_reply.started": "2025-11-20T16:32:36.572681Z"
    }
   },
   "outputs": [],
   "source": [
    "# tr_metas_t1_v7 = read_jsonl(\n",
    "#     os.path.join(\"/app2/suno/data/dpo/auk_t1_v33\", f\"meta_tr.jsonl\")\n",
    "# )\n",
    "# tr_metas_t1_v6 = read_jsonl(\n",
    "#     os.path.join(\"/app2/suno/data/dpo/auk_t1_v34\", f\"meta_tr.jsonl\")\n",
    "# )\n",
    "# tr_metas_t1_v19 = read_jsonl(\n",
    "#     os.path.join(\"/app2/suno/data/dpo/auk_t1_v35\", f\"meta_tr.jsonl\")\n",
    "# )\n",
    "# known_train_ids = set()\n",
    "# for prev_tr_meta in tr_metas_t1_v7:\n",
    "#     known_train_ids.add(prev_tr_meta[\"id\"])\n",
    "# for prev_tr_meta in tr_metas_t1_v6:\n",
    "#     known_train_ids.add(prev_tr_meta[\"id\"])\n",
    "# for prev_tr_meta in tr_metas_t1_v19:\n",
    "#     known_train_ids.add(prev_tr_meta[\"id\"])\n",
    "# print(len(known_train_ids))\n",
    "# print(df_slice.shape)\n",
    "# df_slice = df_slice[~df_slice[\"id\"].isin(known_train_ids)].copy()\n",
    "# print(df_slice.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 84,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T16:32:36.588324Z",
     "iopub.status.busy": "2025-11-20T16:32:36.588181Z",
     "iopub.status.idle": "2025-11-20T16:32:41.987360Z",
     "shell.execute_reply": "2025-11-20T16:32:41.986782Z",
     "shell.execute_reply.started": "2025-11-20T16:32:36.588310Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "(983188, 204)\n"
     ]
    }
   ],
   "source": [
    "df_total = df_slice.copy()\n",
    "print(df_total.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 121,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T19:53:16.928116Z",
     "iopub.status.busy": "2025-11-20T19:53:16.927691Z",
     "iopub.status.idle": "2025-11-20T19:53:55.129449Z",
     "shell.execute_reply": "2025-11-20T19:53:55.128789Z",
     "shell.execute_reply.started": "2025-11-20T19:53:16.928094Z"
    }
   },
   "outputs": [],
   "source": [
    "# df_total.to_pickle(\"/home/tony/Data/Preference/auk_t1/fully_merged_auk_t1_20251118_super_v101.pkl\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 86,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T14:00:20.866354Z",
     "start_time": "2024-05-16T14:00:12.443344Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T16:32:42.004151Z",
     "iopub.status.busy": "2025-11-20T16:32:42.004011Z",
     "iopub.status.idle": "2025-11-20T16:32:42.112324Z",
     "shell.execute_reply": "2025-11-20T16:32:42.111745Z",
     "shell.execute_reply.started": "2025-11-20T16:32:42.004136Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "(983188, 204)\n",
      "task\n",
      "                      672680\n",
      "cover                 147404\n",
      "artist_consistency    101424\n",
      "artist_cover           36132\n",
      "extend                 12306\n",
      "upload_extend          10622\n",
      "artist_extend           2620\n",
      "Name: count, dtype: int64\n"
     ]
    },
    {
     "ename": "NameError",
     "evalue": "name 'BREAK' is not defined",
     "output_type": "error",
     "traceback": [
      "\u001b[0;31m---------------------------------------------------------------------------\u001b[0m",
      "\u001b[0;31mNameError\u001b[0m                                 Traceback (most recent call last)",
      "Cell \u001b[0;32mIn[86], line 6\u001b[0m\n\u001b[1;32m      4\u001b[0m \u001b[38;5;28mprint\u001b[39m(df_slice\u001b[38;5;241m.\u001b[39mshape)\n\u001b[1;32m      5\u001b[0m \u001b[38;5;28mprint\u001b[39m(df_slice[\u001b[38;5;124m\"\u001b[39m\u001b[38;5;124mtask\u001b[39m\u001b[38;5;124m\"\u001b[39m]\u001b[38;5;241m.\u001b[39mvalue_counts())\n\u001b[0;32m----> 6\u001b[0m \u001b[43mBREAK\u001b[49m\n",
      "\u001b[0;31mNameError\u001b[0m: name 'BREAK' is not defined"
     ]
    }
   ],
   "source": [
    "# df_slice.to_pickle(\n",
    "#     \"/home/tony/Data/Preference/30b_v6/interesting_clips_v4_h_t_6_20250426_full_long_slice.pkl\"\n",
    "# )\n",
    "print(df_slice.shape)\n",
    "print(df_slice[\"task\"].value_counts())\n",
    "BREAK"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Need to kick out the ones has gpt prompt -- these are pairs with different text inputs"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 87,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T16:32:55.148907Z",
     "iopub.status.busy": "2025-11-20T16:32:55.148557Z",
     "iopub.status.idle": "2025-11-20T16:33:00.576758Z",
     "shell.execute_reply": "2025-11-20T16:33:00.576157Z",
     "shell.execute_reply.started": "2025-11-20T16:32:55.148888Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "(983188, 204)\n",
      "(983188, 204)\n",
      "(983188, 204)\n"
     ]
    }
   ],
   "source": [
    "print(df_slice.shape)\n",
    "df_slice = df_slice[df_slice[\"request_id\"].apply(lambda x: len(x) > 3)]\n",
    "print(df_slice.shape)\n",
    "# df_slice = df_slice[df_slice[\"is_pro_user\"]].copy()\n",
    "print(df_slice.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 88,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:41.932296Z",
     "start_time": "2024-05-16T13:59:41.932287Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T16:33:00.577767Z",
     "iopub.status.busy": "2025-11-20T16:33:00.577599Z",
     "iopub.status.idle": "2025-11-20T16:33:00.641171Z",
     "shell.execute_reply": "2025-11-20T16:33:00.640691Z",
     "shell.execute_reply.started": "2025-11-20T16:33:00.577751Z"
    }
   },
   "outputs": [],
   "source": [
    "# don't have continue at\n",
    "df_slice[\"request_id\"] = df_slice[\"request_id\"].astype(str)\n",
    "# df_slice[\"npz_path\"] = df_slice[\"npz_path\"].apply(lambda x: str(x).replace(\"_npz\", \"_npz/\"))\n",
    "# df_slice[df_slice[\"continue_at\"].isna()][\"request_id\"].nunique(), df_slice[\"request_id\"].nunique()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 89,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:41.932966Z",
     "start_time": "2024-05-16T13:59:41.932957Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T16:33:00.642031Z",
     "iopub.status.busy": "2025-11-20T16:33:00.641698Z",
     "iopub.status.idle": "2025-11-20T16:33:00.892114Z",
     "shell.execute_reply": "2025-11-20T16:33:00.891579Z",
     "shell.execute_reply.started": "2025-11-20T16:33:00.642016Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "491594\n"
     ]
    }
   ],
   "source": [
    "final_filtered_requests = df_slice[\"request_id\"].astype(str).unique()\n",
    "# final_filtered_requests = df_slice[df_slice[\"is_pro_user\"]][\"request_id\"].astype(str).unique()\n",
    "print(len(final_filtered_requests))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 90,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:41.934277Z",
     "start_time": "2024-05-16T13:59:41.934268Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T16:33:00.893288Z",
     "iopub.status.busy": "2025-11-20T16:33:00.893128Z",
     "iopub.status.idle": "2025-11-20T16:33:13.384780Z",
     "shell.execute_reply": "2025-11-20T16:33:13.384216Z",
     "shell.execute_reply.started": "2025-11-20T16:33:00.893272Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "486678 4916\n",
      "(973356, 204) (9832, 204)\n"
     ]
    }
   ],
   "source": [
    "train_requests, val_requests = train_test_split(\n",
    "    sorted(list(final_filtered_requests)), test_size=0.01, random_state=42\n",
    ")\n",
    "print(len(train_requests), len(val_requests))\n",
    "\n",
    "train_df = df_slice[df_slice[\"request_id\"].astype(str).isin(set(train_requests))].copy()\n",
    "val_df = df_slice[df_slice[\"request_id\"].astype(str).isin(set(val_requests))].copy()\n",
    "train_df = train_df.sort_values(by=[\"request_id\", \"preference\"])\n",
    "train_df = train_df  # .reset_index()\n",
    "val_df = val_df.sort_values(by=[\"request_id\", \"preference\"])\n",
    "val_df = val_df  # .reset_index()\n",
    "train_df = train_df.reset_index(drop=True)\n",
    "val_df = val_df.reset_index(drop=True)\n",
    "print(train_df.shape, val_df.shape)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Actually make"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 91,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:41.935620Z",
     "start_time": "2024-05-16T13:59:41.935613Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T16:33:13.385512Z",
     "iopub.status.busy": "2025-11-20T16:33:13.385359Z",
     "iopub.status.idle": "2025-11-20T16:33:53.288415Z",
     "shell.execute_reply": "2025-11-20T16:33:53.287886Z",
     "shell.execute_reply.started": "2025-11-20T16:33:13.385496Z"
    }
   },
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "100%|██████████████████████████████████████████████████████████████████████████████████████████████████████| 973356/973356 [00:39<00:00, 24941.12it/s]"
     ]
    },
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "51,213 hours of 973356 clips, 60.83475 nodes, 950.54296875 iters\n"
     ]
    },
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "\n"
     ]
    }
   ],
   "source": [
    "total_duration = 0\n",
    "for i, row in tqdm(train_df.iterrows(), total=len(train_df)):\n",
    "    # we need to alternate between preference: neg, pos\n",
    "    # print(i, row)\n",
    "    try:\n",
    "        assert row[\"preference\"] == (i % 2 == 1)\n",
    "        total_duration += row[\"duration\"]\n",
    "    except Exception as E:\n",
    "        print(i, row)\n",
    "        print(E)\n",
    "        raise ValueError()\n",
    "\n",
    "print(\n",
    "    f\"{round(total_duration / 60 / 60):,} hours of {train_df.shape[0]} clips, {train_df.shape[0] / 8 / 2 / 1000} nodes, {train_df.shape[0] / 8 / 8 / 16} iters\"\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 92,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:41.936268Z",
     "start_time": "2024-05-16T13:59:41.936260Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T16:33:53.289131Z",
     "iopub.status.busy": "2025-11-20T16:33:53.288980Z",
     "iopub.status.idle": "2025-11-20T16:35:55.597078Z",
     "shell.execute_reply": "2025-11-20T16:35:55.596559Z",
     "shell.execute_reply.started": "2025-11-20T16:33:53.289115Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "t_data_memmap is set to: 12000\n"
     ]
    },
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████| 9832/9832 [02:02<00:00, 80.42it/s]\n"
     ]
    },
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Total 9832 clips, 1 different prompts, 0 different tags, 0 different negative tags\n",
      "334 hours of False\n",
      "335 hours of True\n",
      "gen: 360.9 hours\n",
      "cover: 144.6 hours\n",
      "extend: 13.8 hours\n",
      "artist_consistency: 103.9 hours\n",
      "artist_cover: 43.3 hours\n",
      "artist_extend: 2.4 hours\n",
      "\n",
      "--- Gender Distribution ---\n",
      "  unspecified: 9,832 (100.0%)\n",
      "\n",
      "--- Negative Tags Usage ---\n",
      "  has_neg_tags: 750 (7.6%)\n",
      "  no_neg_tags: 9,082 (92.4%)\n",
      "\n",
      "--- Control Slider Usage ---\n",
      "  has_control_slider: 1,476 (15.0% of clips)\n",
      "  no_control_slider: 8,356 (85.0% of clips)\n",
      "Done\n"
     ]
    }
   ],
   "source": [
    "make_dataset(\n",
    "    val_df, OUT_DATA_DIR, is_val=True, npz_dir=NPZ_DIR, t_data_memmap=N_TOKENS_AUDIO\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 93,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T16:35:55.597787Z",
     "iopub.status.busy": "2025-11-20T16:35:55.597634Z",
     "iopub.status.idle": "2025-11-20T16:35:56.105039Z",
     "shell.execute_reply": "2025-11-20T16:35:56.104574Z",
     "shell.execute_reply.started": "2025-11-20T16:35:55.597770Z"
    }
   },
   "outputs": [],
   "source": [
    "# test_npz = np.load(\"/app/suno/data/dpo/30b_npz/26d19085-18da-4701-af43-122684543891.npz\")\n",
    "# for k in test_npz.keys():\n",
    "#     print(k)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 94,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:41.936964Z",
     "start_time": "2024-05-16T13:59:41.936957Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T16:35:56.105701Z",
     "iopub.status.busy": "2025-11-20T16:35:56.105551Z",
     "iopub.status.idle": "2025-11-20T19:50:18.624083Z",
     "shell.execute_reply": "2025-11-20T19:50:18.623485Z",
     "shell.execute_reply.started": "2025-11-20T16:35:56.105686Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "t_data_memmap is set to: 12000\n"
     ]
    },
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      " 24%|█████████████████████████                                                                              | 236936/973356 [46:19<2:20:54, 87.10it/s]"
     ]
    },
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "236922, 'history_arr is not a file in the archive', upload_extend, /app2/suno/data/dpo/auk_t1_npz/99804665-d14c-4495-98b8-8fc86d883d4d.npz.\n",
      "WTF --> 236923, skip, preference: True, be7c277f-51d5-4b10-b263-37c924ad23af, task: upload_extend.\n"
     ]
    },
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "100%|███████████████████████████████████████████████████████████████████████████████████████████████████████| 973356/973356 [3:14:19<00:00, 83.48it/s]\n"
     ]
    },
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Total 973354 clips, 168 different prompts, 7 different tags, 0 different negative tags\n",
      "32,871 hours of False\n",
      "32,999 hours of True\n",
      "gen: 35653.3 hours\n",
      "artist_consistency: 10456.5 hours\n",
      "cover: 14048.9 hours\n",
      "extend: 1352.3 hours\n",
      "artist_cover: 4091.6 hours\n",
      "artist_extend: 267.9 hours\n",
      "🚨 Error upload_extend: 1\n",
      "\n",
      "--- Gender Distribution ---\n",
      "  female: 4 (0.0%)\n",
      "  male: 6 (0.0%)\n",
      "  unspecified: 973,345 (100.0%)\n",
      "\n",
      "--- Negative Tags Usage ---\n",
      "  has_neg_tags: 77,878 (8.0%)\n",
      "  no_neg_tags: 895,477 (92.0%)\n",
      "\n",
      "--- Control Slider Usage ---\n",
      "  has_control_slider: 138,462 (14.2% of clips)\n",
      "  no_control_slider: 834,893 (85.8% of clips)\n",
      "Done\n"
     ]
    }
   ],
   "source": [
    "make_dataset(\n",
    "    train_df, OUT_DATA_DIR, is_val=False, npz_dir=NPZ_DIR, t_data_memmap=N_TOKENS_AUDIO\n",
    ")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-01-29T19:46:47.549860Z",
     "start_time": "2024-01-29T19:46:47.548015Z"
    }
   },
   "source": [
    "# Validation"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 95,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:41.937879Z",
     "start_time": "2024-05-16T13:59:41.937870Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T19:50:18.624876Z",
     "iopub.status.busy": "2025-11-20T19:50:18.624717Z",
     "iopub.status.idle": "2025-11-20T19:50:19.790549Z",
     "shell.execute_reply": "2025-11-20T19:50:19.790000Z",
     "shell.execute_reply.started": "2025-11-20T19:50:18.624860Z"
    }
   },
   "outputs": [],
   "source": [
    "# verify\n",
    "mm = np.memmap(os.path.join(OUT_DATA_DIR, f\"data_val.bin\"), dtype=np.uint16, mode=\"r\")\n",
    "test_metas = read_jsonl(os.path.join(OUT_DATA_DIR, f\"meta_val.jsonl\"))\n",
    "test_info = read_json(os.path.join(OUT_DATA_DIR, f\"info_val.json\"))\n",
    "mm = mm.reshape(-1, N_TOKENS_AUDIO, 1)\n",
    "assert len(mm) == len(test_metas)\n",
    "assert mm[:100, :, 0].min() >= 0\n",
    "assert mm[:100, :, 0].max() <= 4000"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 96,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T19:50:19.792480Z",
     "iopub.status.busy": "2025-11-20T19:50:19.792296Z",
     "iopub.status.idle": "2025-11-20T19:50:19.814692Z",
     "shell.execute_reply": "2025-11-20T19:50:19.814213Z",
     "shell.execute_reply.started": "2025-11-20T19:50:19.792464Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Counter({None: 6722, 'cover': 1490, 'artist_consistency': 990, 'artist_cover': 376, 'extend': 230, 'artist_extend': 24})\n"
     ]
    }
   ],
   "source": [
    "task_counts = Counter()\n",
    "for test_meta in test_metas:\n",
    "    task_counts[test_meta.get(\"task\")] += 1\n",
    "print(task_counts)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 97,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:41.938629Z",
     "start_time": "2024-05-16T13:59:41.938621Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T19:50:19.817965Z",
     "iopub.status.busy": "2025-11-20T19:50:19.817827Z",
     "iopub.status.idle": "2025-11-20T19:50:19.831390Z",
     "shell.execute_reply": "2025-11-20T19:50:19.830967Z",
     "shell.execute_reply.started": "2025-11-20T19:50:19.817952Z"
    }
   },
   "outputs": [],
   "source": [
    "# # randomly listen to some stuff\n",
    "# from suno_utils.tasks.dac_2c_12cb import preload_models as preload_codec_models\n",
    "# from suno_utils.tasks.dac_2c_12cb import (\n",
    "#     encode as codec_encode,\n",
    "#     decode_stream_to_full_audio as codec_decode,\n",
    "#     EMBEDDING_RATE as CODEC_EMBEDDING_RATE,\n",
    "#     decode as decode\n",
    "# )\n",
    "# os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"0\"\n",
    "# _ = preload_codec_models(\"/app/suno/data/dpo/models/dac_2c_25x12.pt\", device=\"cuda\")\n",
    "# assert len(test_metas) == len(mm)\n",
    "# idx_list = list(range(len(test_metas)))\n",
    "# # random.shuffle(idx_list)\n",
    "# # idx_list = [idx for idx in idx_list if \"text\" in test_metas[idx]]\n",
    "# print(len(mm))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 98,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:41.939205Z",
     "start_time": "2024-05-16T13:59:41.939198Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T19:50:19.832216Z",
     "iopub.status.busy": "2025-11-20T19:50:19.831863Z",
     "iopub.status.idle": "2025-11-20T19:50:19.843777Z",
     "shell.execute_reply": "2025-11-20T19:50:19.843355Z",
     "shell.execute_reply.started": "2025-11-20T19:50:19.832202Z"
    }
   },
   "outputs": [],
   "source": [
    "# import random\n",
    "# idx = random.choice(test_info[\"perference_0\"][\"idx_list\"])\n",
    "# assert \"original_duration_s\" in test_metas[idx]\n",
    "# # positive index should be shifted by 1\n",
    "# pos_idx = idx + 1\n",
    "# print(\n",
    "#     \"tags:\",\n",
    "#     test_metas[idx].get(\"tags\") == test_metas[pos_idx].get(\"tags\"),\n",
    "#     test_metas[idx].get(\"tags\"),\n",
    "# )\n",
    "# arr = mm[idx, 1:].copy().astype(np.int16)[:, 1:]\n",
    "# pos_arr = mm[pos_idx, 1:].copy().astype(np.int16)[:, 1:]\n",
    "# pad_idx_arr = np.where(arr == COARSE_PAD_TOKEN)[0]\n",
    "# if len(pad_idx_arr) > 0:\n",
    "#     arr = arr[: pad_idx_arr[0], :]\n",
    "# pos_pad_idx_arr = np.where(pos_arr == COARSE_PAD_TOKEN)[0]\n",
    "# if len(pos_pad_idx_arr) > 0:\n",
    "#     pos_arr = pos_arr[: pos_pad_idx_arr[0], :]\n",
    "# a = decode(arr)\n",
    "# print(\"\\n negative example \\n\", test_metas[idx])\n",
    "# a.play(compress=False)\n",
    "# pos_a = decode(pos_arr)\n",
    "# print(\"\\n positive example \\n\", test_metas[pos_idx])\n",
    "# pos_a.play(compress=False)\n",
    "# print(\n",
    "#     \"text:\",\n",
    "#     test_metas[idx].get(\"text\") == test_metas[pos_idx].get(\"text\"),\n",
    "#     test_metas[idx].get(\"text\"),\n",
    "# )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 99,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:41.939977Z",
     "start_time": "2024-05-16T13:59:41.939969Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T19:50:19.844391Z",
     "iopub.status.busy": "2025-11-20T19:50:19.844250Z",
     "iopub.status.idle": "2025-11-20T19:50:19.855883Z",
     "shell.execute_reply": "2025-11-20T19:50:19.855447Z",
     "shell.execute_reply.started": "2025-11-20T19:50:19.844378Z"
    }
   },
   "outputs": [],
   "source": [
    "# val_df[val_df[\"tags\"] == 'a vibrant blend of experimental jazz fusion, drum-and-bass and swagger fuzzed-out guitars']"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 100,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:41.940610Z",
     "start_time": "2024-05-16T13:59:41.940603Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T19:50:19.856581Z",
     "iopub.status.busy": "2025-11-20T19:50:19.856374Z",
     "iopub.status.idle": "2025-11-20T19:50:19.867965Z",
     "shell.execute_reply": "2025-11-20T19:50:19.867536Z",
     "shell.execute_reply.started": "2025-11-20T19:50:19.856499Z"
    }
   },
   "outputs": [],
   "source": [
    "# from collections import Counter\n",
    "# c = Counter()\n",
    "# for _, row in df_slice.iterrows():\n",
    "#     # print(row[\"metadata\"])\n",
    "#     for k in ast.literal_eval(row[\"metadata\"]).keys():\n",
    "#         c[k] += 1"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 101,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:41.941167Z",
     "start_time": "2024-05-16T13:59:41.941159Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T19:50:19.868582Z",
     "iopub.status.busy": "2025-11-20T19:50:19.868445Z",
     "iopub.status.idle": "2025-11-20T19:50:19.880135Z",
     "shell.execute_reply": "2025-11-20T19:50:19.879701Z",
     "shell.execute_reply.started": "2025-11-20T19:50:19.868569Z"
    }
   },
   "outputs": [],
   "source": [
    "# original_npz_path = f\"/app/suno/data/dpo/7b_npz/{test_metas[idx]['id']}.npz\"\n",
    "# original_npz_path = \"/app/suno/data/dpo/7b_npz/729c3011-f672-4ccd-8d82-1cbf2b52ff69.npz\"\n",
    "# original_arr = np.load(original_npz_path)[\"v2_raw\"]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 102,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:41.941801Z",
     "start_time": "2024-05-16T13:59:41.941793Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T19:50:19.880864Z",
     "iopub.status.busy": "2025-11-20T19:50:19.880722Z",
     "iopub.status.idle": "2025-11-20T19:50:19.903956Z",
     "shell.execute_reply": "2025-11-20T19:50:19.903495Z",
     "shell.execute_reply.started": "2025-11-20T19:50:19.880850Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "4916 0\n"
     ]
    }
   ],
   "source": [
    "def validation_on_metas(input_metas):\n",
    "    total_bad = 0\n",
    "    total_good = 0\n",
    "    for idx in range(len(input_metas)):\n",
    "        if idx % 2 == 0:\n",
    "            pos_idx = idx + 1\n",
    "            if input_metas[idx].get(\"tags\") != input_metas[pos_idx].get(\"tags\"):\n",
    "                print(input_metas[idx].get(\"id\"), input_metas[idx].get(\"tags\"),input_metas[pos_idx].get(\"id\"), input_metas[pos_idx].get(\"tags\"), )\n",
    "                # print(test_metas[idx].get(\"text\") == test_metas[pos_idx].get(\"text\"), test_metas[idx].get(\"tags\"), test_metas[pos_idx].get(\"tags\"))\n",
    "                total_bad += 1\n",
    "            elif input_metas[idx].get(\"text\") != input_metas[pos_idx].get(\"text\"):\n",
    "                # print(test_metas[idx].get(\"text\") == test_metas[pos_idx].get(\"text\"), test_metas[idx].get(\"tags\"), test_metas[pos_idx].get(\"tags\"))\n",
    "                total_bad += 1\n",
    "            else:\n",
    "                total_good += 1\n",
    "    print(total_good, total_bad)\n",
    "    return\n",
    "\n",
    "\n",
    "validation_on_metas(test_metas)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 103,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:41.942520Z",
     "start_time": "2024-05-16T13:59:41.942511Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T19:50:19.904577Z",
     "iopub.status.busy": "2025-11-20T19:50:19.904440Z",
     "iopub.status.idle": "2025-11-20T19:50:19.983648Z",
     "shell.execute_reply": "2025-11-20T19:50:19.983177Z",
     "shell.execute_reply.started": "2025-11-20T19:50:19.904564Z"
    }
   },
   "outputs": [],
   "source": [
    "train_info = read_json(os.path.join(OUT_DATA_DIR, f\"info_tr.json\"))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 104,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:41.943072Z",
     "start_time": "2024-05-16T13:59:41.943065Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T19:50:19.984302Z",
     "iopub.status.busy": "2025-11-20T19:50:19.984162Z",
     "iopub.status.idle": "2025-11-20T19:50:20.032036Z",
     "shell.execute_reply": "2025-11-20T19:50:20.031572Z",
     "shell.execute_reply.started": "2025-11-20T19:50:19.984288Z"
    }
   },
   "outputs": [],
   "source": [
    "n_neg_tr = train_info[\"perference_0\"][\"idx_list\"]\n",
    "n_pos_tr = train_info[\"perference_1\"][\"idx_list\"]\n",
    "assert len(n_pos_tr) == len(n_neg_tr)\n",
    "# make sure they are offset by 1 and exactly 1\n",
    "for i, j in zip(n_neg_tr, n_pos_tr):\n",
    "    assert i == j - 1"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 105,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:41.944246Z",
     "start_time": "2024-05-16T13:59:41.944237Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T19:50:20.032835Z",
     "iopub.status.busy": "2025-11-20T19:50:20.032561Z",
     "iopub.status.idle": "2025-11-20T19:50:20.045316Z",
     "shell.execute_reply": "2025-11-20T19:50:20.044869Z",
     "shell.execute_reply.started": "2025-11-20T19:50:20.032819Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "total samples 973354 (973356, 204)\n"
     ]
    }
   ],
   "source": [
    "total_iters = len(n_neg_tr) + len(n_pos_tr)\n",
    "print(\"total samples\", total_iters, train_df.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 106,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T19:50:20.045914Z",
     "iopub.status.busy": "2025-11-20T19:50:20.045781Z",
     "iopub.status.idle": "2025-11-20T19:51:16.373784Z",
     "shell.execute_reply": "2025-11-20T19:51:16.373004Z",
     "shell.execute_reply.started": "2025-11-20T19:50:20.045900Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "486677 0\n"
     ]
    }
   ],
   "source": [
    "metas_tr = read_jsonl(os.path.join(OUT_DATA_DIR, \"meta_tr.jsonl\"))\n",
    "validation_on_metas(metas_tr)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 107,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T19:51:16.374717Z",
     "iopub.status.busy": "2025-11-20T19:51:16.374540Z",
     "iopub.status.idle": "2025-11-20T19:51:16.510204Z",
     "shell.execute_reply": "2025-11-20T19:51:16.509576Z",
     "shell.execute_reply.started": "2025-11-20T19:51:16.374700Z"
    }
   },
   "outputs": [],
   "source": [
    "# new_metas_tr = []\n",
    "# for index, l in enumerate(metas_tr):\n",
    "#     if index % 2 == 1:\n",
    "#         last_l = new_metas_tr[-1]\n",
    "#         if l[\"tags\"] != last_l[\"tags\"]:\n",
    "#             print(l[\"tags\"], last_l[\"tags\"])\n",
    "#             l[\"tags\"] = last_l[\"tags\"]\n",
    "#     new_metas_tr.append(l)\n",
    "# validation_on_metas(new_metas_tr)\n",
    "# write_jsonl(new_metas_tr, os.path.join(OUT_DATA_DIR, \"meta_tr.jsonl\"))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 108,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:41.945249Z",
     "start_time": "2024-05-16T13:59:41.945241Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T19:51:16.511107Z",
     "iopub.status.busy": "2025-11-20T19:51:16.510856Z",
     "iopub.status.idle": "2025-11-20T19:51:16.528000Z",
     "shell.execute_reply": "2025-11-20T19:51:16.527430Z",
     "shell.execute_reply.started": "2025-11-20T19:51:16.511090Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "1 epoch per batch 4, total 950.541015625\n"
     ]
    }
   ],
   "source": [
    "print(\"1 epoch per batch 4, total\", total_iters / 16 / 8 / 8)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 109,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T19:51:16.528749Z",
     "iopub.status.busy": "2025-11-20T19:51:16.528593Z",
     "iopub.status.idle": "2025-11-20T19:51:16.542818Z",
     "shell.execute_reply": "2025-11-20T19:51:16.542307Z",
     "shell.execute_reply.started": "2025-11-20T19:51:16.528733Z"
    }
   },
   "outputs": [],
   "source": [
    "# import time\n",
    "# time.sleep(60 * 60 * 1)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 110,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:41.945972Z",
     "start_time": "2024-05-16T13:59:41.945964Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T19:51:16.543702Z",
     "iopub.status.busy": "2025-11-20T19:51:16.543411Z",
     "iopub.status.idle": "2025-11-20T19:51:16.555185Z",
     "shell.execute_reply": "2025-11-20T19:51:16.554664Z",
     "shell.execute_reply.started": "2025-11-20T19:51:16.543685Z"
    }
   },
   "outputs": [],
   "source": [
    "# !cd /home/tony/Work/tony/slurm/bluejay && sbatch sbatch_ipo_bluejay"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 111,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T19:51:16.555897Z",
     "iopub.status.busy": "2025-11-20T19:51:16.555753Z",
     "iopub.status.idle": "2025-11-20T19:51:16.578715Z",
     "shell.execute_reply": "2025-11-20T19:51:16.578163Z",
     "shell.execute_reply.started": "2025-11-20T19:51:16.555883Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Cache kept!\n"
     ]
    }
   ],
   "source": [
    "import shutil\n",
    "\n",
    "# Basic file copy\n",
    "shutil.copy(\n",
    "    \"/home/tony/Work/tony/Preference/make_dataset_auk_t1_super.ipynb\",\n",
    "    os.path.join(OUT_DATA_DIR, \"make_dataset.ipynb\"),\n",
    ")\n",
    "print(\"Cache kept!\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# some gymathtics loading prev data"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 112,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-05-16T13:59:41.946562Z",
     "start_time": "2024-05-16T13:59:41.946555Z"
    },
    "execution": {
     "iopub.execute_input": "2025-11-20T19:51:16.581476Z",
     "iopub.status.busy": "2025-11-20T19:51:16.581143Z",
     "iopub.status.idle": "2025-11-20T19:51:16.593036Z",
     "shell.execute_reply": "2025-11-20T19:51:16.592511Z",
     "shell.execute_reply.started": "2025-11-20T19:51:16.581459Z"
    }
   },
   "outputs": [],
   "source": [
    "# prev_v3_data = \"/app/suno/data/dpo/7v_v20_full/\"\n",
    "\n",
    "# test_val_metas = read_jsonl(os.path.join(prev_v3_data, f\"meta_val.jsonl\"))\n",
    "# test_tr_metas = read_jsonl(os.path.join(prev_v3_data, f\"meta_tr.jsonl\"))\n",
    "\n",
    "# all_ids = set()\n",
    "# for meta in test_val_metas:\n",
    "#     all_ids.add(meta[\"id\"])\n",
    "# for meta in test_tr_metas:\n",
    "#     all_ids.add(meta[\"id\"])\n",
    "# print(len(all_ids), len(test_val_metas) + len(test_tr_metas))\n",
    "\n",
    "# all_ids = list(all_ids)\n",
    "# with open(\"/home/tony/Data/Preference/7b_v2/7v_v20_full_recut_id.json\", \"w\") as fp:\n",
    "#     json.dump(all_ids, fp)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 113,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T19:51:16.593927Z",
     "iopub.status.busy": "2025-11-20T19:51:16.593631Z",
     "iopub.status.idle": "2025-11-20T19:51:16.605789Z",
     "shell.execute_reply": "2025-11-20T19:51:16.605278Z",
     "shell.execute_reply.started": "2025-11-20T19:51:16.593912Z"
    }
   },
   "outputs": [],
   "source": [
    "# x_data = train_df[train_df[\"preference\"]][\"similarity\"]\n",
    "# y_data = train_df[~train_df[\"preference\"]][\"similarity\"]\n",
    "# from matplotlib.colors import LogNorm\n",
    "\n",
    "# fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(24, 10))\n",
    "\n",
    "# # 2D Histogram\n",
    "# h = ax1.hist2d(\n",
    "#     x_data,\n",
    "#     y_data,\n",
    "#     bins=(50, 50),\n",
    "#     cmap=\"coolwarm\",\n",
    "#     range=[[0, 1], [0, 1]],\n",
    "#     norm=LogNorm(),\n",
    "# )\n",
    "\n",
    "# ax1.set_xlabel(\"Semantic Distance (Preferred)\")\n",
    "# ax1.set_ylabel(\"Semantic Distance (Non-Preferred)\")\n",
    "# ax1.set_title(\n",
    "#     \"2D Histogram of Semantic Distances: Preferred vs Non-Preferred (Log Scale)\"\n",
    "# )\n",
    "\n",
    "# cbar1 = plt.colorbar(h[3], ax=ax1)\n",
    "# cbar1.set_label(\"Number of Request IDs (Log Scale)\")\n",
    "\n",
    "# # Scatter plot\n",
    "# ax2.scatter(x_data, y_data, alpha=0.1, s=1)\n",
    "# ax2.set_xlabel(\"Semantic Distance (Preferred)\")\n",
    "# ax2.set_ylabel(\"Semantic Distance (Non-Preferred)\")\n",
    "# ax2.set_title(\"Scatter Plot of Semantic Distances: Preferred vs Non-Preferred\")\n",
    "# ax2.set_xlim(0, 1)\n",
    "# ax2.set_ylim(0, 1)\n",
    "\n",
    "# plt.tight_layout()\n",
    "# plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 114,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T19:51:16.606661Z",
     "iopub.status.busy": "2025-11-20T19:51:16.606362Z",
     "iopub.status.idle": "2025-11-20T19:51:16.617897Z",
     "shell.execute_reply": "2025-11-20T19:51:16.617374Z",
     "shell.execute_reply.started": "2025-11-20T19:51:16.606645Z"
    }
   },
   "outputs": [],
   "source": [
    "# train_metas = read_jsonl(os.path.join(OUT_DATA_DIR, f\"meta_tr.jsonl\"))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 115,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T19:51:16.618565Z",
     "iopub.status.busy": "2025-11-20T19:51:16.618419Z",
     "iopub.status.idle": "2025-11-20T19:51:16.631133Z",
     "shell.execute_reply": "2025-11-20T19:51:16.630633Z",
     "shell.execute_reply.started": "2025-11-20T19:51:16.618551Z"
    }
   },
   "outputs": [
    {
     "data": {
      "text/plain": [
       "dict_keys(['perference_0', 'perference_1'])"
      ]
     },
     "execution_count": 115,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "train_info.keys()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 116,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T19:51:16.631860Z",
     "iopub.status.busy": "2025-11-20T19:51:16.631707Z",
     "iopub.status.idle": "2025-11-20T19:51:19.180054Z",
     "shell.execute_reply": "2025-11-20T19:51:19.179418Z",
     "shell.execute_reply.started": "2025-11-20T19:51:16.631845Z"
    }
   },
   "outputs": [],
   "source": [
    "import torch\n",
    "a = torch.tensor([6.2500e-04, 3.9062e-05, 2.3462e-03])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 117,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T19:51:19.180884Z",
     "iopub.status.busy": "2025-11-20T19:51:19.180717Z",
     "iopub.status.idle": "2025-11-20T19:51:19.994370Z",
     "shell.execute_reply": "2025-11-20T19:51:19.993827Z",
     "shell.execute_reply.started": "2025-11-20T19:51:19.180866Z"
    }
   },
   "outputs": [
    {
     "data": {
      "text/plain": [
       "tensor(0.0010)"
      ]
     },
     "execution_count": 117,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "a.mean()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 118,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T19:51:19.995135Z",
     "iopub.status.busy": "2025-11-20T19:51:19.994978Z",
     "iopub.status.idle": "2025-11-20T19:51:20.014098Z",
     "shell.execute_reply": "2025-11-20T19:51:20.013630Z",
     "shell.execute_reply.started": "2025-11-20T19:51:19.995119Z"
    }
   },
   "outputs": [],
   "source": [
    "# import time\n",
    "# time.sleep(3600 * 3)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 119,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-11-20T19:51:20.014922Z",
     "iopub.status.busy": "2025-11-20T19:51:20.014764Z",
     "iopub.status.idle": "2025-11-20T19:51:20.031091Z",
     "shell.execute_reply": "2025-11-20T19:51:20.030658Z",
     "shell.execute_reply.started": "2025-11-20T19:51:20.014906Z"
    }
   },
   "outputs": [],
   "source": [
    "# !cd /home/tony/Work/tony/slurm/diffusion && sbatch run_diffusion_infill.sh"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python 3 (ipykernel)",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "codemirror_mode": {
    "name": "ipython",
    "version": 3
   },
   "file_extension": ".py",
   "mimetype": "text/x-python",
   "name": "python",
   "nbconvert_exporter": "python",
   "pygments_lexer": "ipython3",
   "version": "3.10.15"
  },
  "toc": {
   "base_numbering": 1,
   "nav_menu": {},
   "number_sections": true,
   "sideBar": true,
   "skip_h1_title": false,
   "title_cell": "Table of Contents",
   "title_sidebar": "Contents",
   "toc_cell": false,
   "toc_position": {},
   "toc_section_display": true,
   "toc_window_display": false
  }
 },
 "nbformat": 4,
 "nbformat_minor": 4
}
