{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 1,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T20:54:50.533571Z",
     "start_time": "2024-04-12T20:54:48.912399Z"
    }
   },
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "/tmp/ipykernel_3199211/4201459797.py:1: DeprecationWarning: \n",
      "Pyarrow will become a required dependency of pandas in the next major release of pandas (pandas 3.0),\n",
      "(to allow more performant data types, such as the Arrow string type, and better interoperability with other libraries)\n",
      "but was not found to be installed on your system.\n",
      "If this would cause problems for you,\n",
      "please provide us feedback at https://github.com/pandas-dev/pandas/issues/54466\n",
      "        \n",
      "  import pandas as pd\n"
     ]
    }
   ],
   "source": [
    "import pandas as pd\n",
    "import numpy as np\n",
    "import os\n",
    "from tqdm import tqdm\n",
    "from sklearn.model_selection import train_test_split\n",
    "from suno_utils.utils.s3 import download_s3_files\n",
    "import sys\n",
    "from collections import defaultdict\n",
    "from suno_utils.utils.text import (\n",
    "    write_jsonl,\n",
    "    read_jsonl,\n",
    "    write_json,\n",
    "    read_json,\n",
    ")\n",
    "import shutil\n",
    "import ast\n",
    "from preference_helper import *\n",
    "\n",
    "sys.path.insert(0, \"/home/tony/Work/glockenspiel/sunoGPT/scripts/\")\n",
    "\n",
    "from data_preparation_7b import *\n",
    "import numpy as np\n",
    "\n",
    "pd.set_option('display.max_rows', 500)\n",
    "pd.set_option('display.max_columns', 500)\n",
    "pd.set_option('display.width', 1000)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 2,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T20:54:50.569535Z",
     "start_time": "2024-04-12T20:54:50.534934Z"
    }
   },
   "outputs": [],
   "source": [
    "OUT_DATA_DIR = \"/app/suno/data/dpo/7v_r2_reprov3_before/\"\n",
    "# OUT_DATA_DIR = \"/home/tony/Data/test/7v_v6_full/\"\n",
    "os.makedirs(OUT_DATA_DIR, exist_ok=True)\n",
    "shutil.copyfile(\"/app/suno/data/dpo/7v_v1_full/tokenizer_60k.json\", os.path.join(OUT_DATA_DIR, \"tokenizer_60k.json\"))\n",
    "NPZ_DIR = \"/app/suno/data/dpo/v3_npz\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 3,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T20:54:50.596358Z",
     "start_time": "2024-04-12T20:54:50.571584Z"
    }
   },
   "outputs": [],
   "source": [
    "# df_origin = pd.read_csv(\"/home/tony/Data/Preference/7b_v1/interesting_clips.csv\")\n",
    "# df_origin[df_origin[\"id\"] == \"75f69d46-1d81-4327-9253-a8a4886777d1\"]\n",
    "# df_origin[df_origin[\"request_id\"] == \"9414f356-88d6-40ba-b95e-9f4c5004a4aa\"]\n",
    "# df_origin[df_origin[\"request_id\"] == \"b2a150c1-f058-4bff-a013-1961e574631d\"]\n",
    "# df_origin[df_origin[\"id\"] == \"574ba431-18a1-4e28-ac7c-d3e401dba6fe\"]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 4,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T20:57:48.676217Z",
     "start_time": "2024-04-12T20:54:50.597712Z"
    }
   },
   "outputs": [
    {
     "data": {
      "text/plain": [
       "(3320972, 74)"
      ]
     },
     "execution_count": 4,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "df = pd.read_csv(\"/home/tony/Data/Preference/7b_v2/pre_model_20240412.csv\", engine='python')\n",
    "df.shape\n",
    "# v0: (68746, 31)\n",
    "# v5: (938008, 37)\n",
    "# v6: (1097024, 38)\n",
    "# v21: (1822704, 41)\n",
    "# r1 \n",
    "# v2: 2620268 (pre fixes...lots of imperfections...)\n",
    "# v3: 1552512"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 5,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T20:57:49.391609Z",
     "start_time": "2024-04-12T20:57:48.677565Z"
    }
   },
   "outputs": [
    {
     "data": {
      "text/plain": [
       "is_7b\n",
       "True    3320970\n",
       "Name: count, dtype: int64"
      ]
     },
     "execution_count": 5,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "df[\"is_7b\"] = df[\"model_name\"].str.contains(\"v3\")\n",
    "df[\"is_7b\"].value_counts()"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# LET's do the data prep"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 6,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T20:57:49.395147Z",
     "start_time": "2024-04-12T20:57:49.393430Z"
    }
   },
   "outputs": [],
   "source": [
    "date_cut = '2024-03-22 04:30:00'"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 7,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T20:57:51.834496Z",
     "start_time": "2024-04-12T20:57:49.396771Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "preference  model_name         \n",
      "False       chirp-v3-engine-i      1171846\n",
      "            chirp-v3-engine-v0      249904\n",
      "            chirp-v3-engine-d       211292\n",
      "            chirp-v3-engine-s        23985\n",
      "            chirp-v3-engine-i-d       3458\n",
      "True        chirp-v3-engine-i      1181677\n",
      "            chirp-v3-engine-d       291244\n",
      "            chirp-v3-engine-v0      165322\n",
      "            chirp-v3-engine-s        19234\n",
      "            chirp-v3-engine-i-d       3008\n",
      "Name: count, dtype: int64\n",
      "(3320972, 74)\n",
      "(810029, 74)\n"
     ]
    }
   ],
   "source": [
    "# # let's also kick out the ... ipo and ipo-dpoed model for now?\n",
    "print(df.groupby([\"preference\"])[\"model_name\"].value_counts())\n",
    "print(df.shape)\n",
    "df = df[df[\"model_name\"].isin([\"chirp-v3-engine-d\", \"chirp-v3-engine-v0\"])]\n",
    "df = df[(df[\"model_name\"].isin([\"chirp-v3-engine-d\", \"chirp-v3-engine-v0\"])) & (df[\"created_at\"] <= date_cut)]\n",
    "# only IPO\n",
    "# df = df[(df[\"model_name\"].isin([\"chirp-v3-engine-i\"]) ) & (df[\"created_at\"] >= date_cut)]\n",
    "# for some reason...we can't train dpo on the ipo data...it just doesn't follow lyrics...X.x\n",
    "# df = df[\n",
    "#     (\n",
    "#         (df[\"model_name\"].isin([\"chirp-v3-engine-d\", \"chirp-v3-engine-v0\"]))\n",
    "#         & (df[\"preference\"] == True)\n",
    "#     )\n",
    "#     | (df[\"preference\"] == False)\n",
    "# ]\n",
    "print(df.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 8,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T20:57:54.155952Z",
     "start_time": "2024-04-12T20:57:51.835784Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "(810029, 74)\n",
      "(765750, 74)\n",
      "preference  model_name        \n",
      "False       chirp-v3-engine-v0    212475\n",
      "            chirp-v3-engine-d     170400\n",
      "True        chirp-v3-engine-d     242662\n",
      "            chirp-v3-engine-v0    140213\n",
      "Name: count, dtype: int64\n"
     ]
    }
   ],
   "source": [
    "print(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(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": 9,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T20:57:54.160556Z",
     "start_time": "2024-04-12T20:57:54.158723Z"
    }
   },
   "outputs": [],
   "source": [
    "# Let's use the old selection for now -- for quality assurance"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 35,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-13T03:19:45.401665Z",
     "start_time": "2024-04-13T03:19:42.963994Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "negative 314133 positive 133765\n",
      "total pair requests 382875 selected pair requests 101344 frac 0.265\n"
     ]
    }
   ],
   "source": [
    "normal_pos_play_count = 7\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 = 5\n",
    "\n",
    "neg_filter_selection_mask = (\n",
    "    (df[\"preference\"] == False)  # get basics aligned\n",
    "    & (df[\"reaction_play_count\"] >= 1)  # has to be played once\n",
    "    & (df[\"reaction_play_count\"] <= 3)  # if it is actually bad, shouldn't be listened often\n",
    "    & (df[\"duration\"] >= 10)  # can't be too short, otherwise it is obvious\n",
    "    & (df[\"duration\"] <= 120)  # can't be badly long\n",
    "    & (df[\"has_continue_and_start_continue_at\"].isna())  # won't have any continues\n",
    "    # & (df[\"dislike_count\"] >= 1) # this is kinda strict\n",
    "    #     & (\n",
    "    #         (df_slice[\"is_in_playlist\"] == False)\n",
    "    #         & (df_slice[\"concat_in_playlist\"] == False)\n",
    "    #     )  # can't be part of a playlist -- otherwise there are some like signal in it?\n",
    ")\n",
    "pos_filter_selectin_mask = (\n",
    "    (df[\"preference\"] == True)  # get basics aligned\n",
    "    & (\n",
    "        df[\"good_continue_at\"] == True\n",
    "    )  # if continue, needs to continue off a certain percentage\n",
    "    & (\n",
    "        ((df[\"reaction_play_count\"] >= 1) & (df[\"is_7b\"] == True))\n",
    "        | ((df[\"reaction_play_count\"] >= 10) & (df[\"is_7b\"] == False))\n",
    "    )\n",
    "    & (\n",
    "        (\n",
    "            (df[\"part_of_concat\"] == True)\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\"] == False)\n",
    "            & (df[\"reaction_play_count\"] >= normal_pos_play_count)\n",
    "        )\n",
    "    )\n",
    "    & (df[\"play_rel_diff\"] >= 0)  # this is more like quality assurance\n",
    "    & (df[\"duration\"] >= 10)  # can't be too short, otherwise it is obvious\n",
    "    & (df[\"duration\"] <= 120)  # can't be badly long\n",
    "    & (df[\"dislike_count\"] == 0)  # can't have dislikes\n",
    "    & (df[\"flag_count\"] == 0)  # can't have issues\n",
    "    & (df[\"user_n_clips\"] >= 20)  # user needs to have genereated at least 20\n",
    "    # & (df[\"duration_rel_diff\"] < 10) # positive isn't just longer\n",
    "    # & (df[\"upvote_count\"] >= 1)\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",
    "    \"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",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 36,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-13T03:19:50.927654Z",
     "start_time": "2024-04-13T03:19:50.064473Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "requests 101344 clips 202688 total khrs 3.871; N gpus for 1000 iters 6.334; n unique users 16463\n"
     ]
    }
   ],
   "source": [
    "df_slice = df[df[\"request_id\"].isin(set(unique_requests))].copy()\n",
    "print(\n",
    "    \"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 / 4 / 1000:.3f};\",\n",
    "    f\"n unique users {df_slice['user_id'].nunique()}\",\n",
    ")\n",
    "# 76171 152342 total khrs 2.880 n gpus for 1250 iters 3.809"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 12,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T20:57:59.887341Z",
     "start_time": "2024-04-12T20:57:59.826920Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "positive in playlist (29782, 74)\n"
     ]
    }
   ],
   "source": [
    "test_mask = (df_slice[\"preference\"] == True) & (\n",
    "    (df_slice[\"is_in_playlist\"] == True) | (df_slice[\"concat_in_playlist\"] == True)\n",
    ")\n",
    "print(\"positive in playlist\", df_slice[test_mask].shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 13,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T20:57:59.961609Z",
     "start_time": "2024-04-12T20:57:59.889034Z"
    }
   },
   "outputs": [],
   "source": [
    "interesting_clips_must_be_positive_mask = (\n",
    "    (df_slice[\"upvoted\"] == True)\n",
    "    | (df_slice[\"has_action\"] == True)\n",
    "    | (df_slice[\"part_of_concat\"] == True)\n",
    ")\n",
    "interesting_clips_must_be_not_negative_mask = (df_slice[\"downvoted\"] == False) # & (df_slice[\"dislike_count\"] < 1)\n",
    "interesting_clips_mask = interesting_clips_must_be_positive_mask & interesting_clips_must_be_not_negative_mask\n",
    "assert interesting_clips_mask.eq(df_slice[\"preference\"]).all()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 14,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T20:58:00.003666Z",
     "start_time": "2024-04-12T20:57:59.964472Z"
    }
   },
   "outputs": [],
   "source": [
    "# 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": 37,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-13T03:20:14.363524Z",
     "start_time": "2024-04-13T03:20:14.251837Z"
    }
   },
   "outputs": [
    {
     "data": {
      "text/plain": [
       "49849"
      ]
     },
     "execution_count": 37,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "# don't have continue at\n",
    "df_slice[df_slice[\"continue_at\"].isna()][\"request_id\"].nunique()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 38,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-13T03:20:14.982334Z",
     "start_time": "2024-04-13T03:20:14.935677Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "101344\n"
     ]
    }
   ],
   "source": [
    "final_filtered_requests = df_slice[\"request_id\"].unique()\n",
    "print(len(final_filtered_requests))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 17,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T20:58:01.106411Z",
     "start_time": "2024-04-12T20:58:00.223593Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "100302 1014\n",
      "(200604, 75) (2028, 75)\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\"].isin(set(train_requests))].copy()\n",
    "val_df = df_slice[df_slice[\"request_id\"].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",
    "\n",
    "print(train_df.shape, val_df.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 18,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T20:58:01.124193Z",
     "start_time": "2024-04-12T20:58:01.107750Z"
    }
   },
   "outputs": [],
   "source": [
    "def reshift(arr):\n",
    "    sem_start_idx = 0\n",
    "    sem_end_idx = len(arr) - 1\n",
    "    semantic_arr = arr[:, :SEMANTIC_N_CODEBOOKS]\n",
    "    coarse_arr = arr[:, SEMANTIC_N_CODEBOOKS:]\n",
    "\n",
    "    coarse_start_idx = int(round(sem_start_idx * COARSE_RATE_HZ / SEMANTIC_RATE_HZ))\n",
    "    coarse_end_idx = int(round(sem_end_idx * COARSE_RATE_HZ / SEMANTIC_RATE_HZ))\n",
    "    assert sem_end_idx >= 0 and coarse_start_idx >= 0\n",
    "    assert not (sem_end_idx > len(semantic_arr) or coarse_end_idx > len(coarse_arr))\n",
    "\n",
    "    # get array segments\n",
    "    arr_s = semantic_arr[sem_start_idx:sem_end_idx, :].copy()\n",
    "    arr_c = coarse_arr[coarse_start_idx:coarse_end_idx, :].copy()\n",
    "    assert arr_s.max() <= SEMANTIC_PAD_TOKEN\n",
    "    assert arr_c.max() <= COARSE_PAD_TOKEN\n",
    "    assert len(arr_s) == len(arr_c)\n",
    "    # concat and stack\n",
    "    if len(arr_c) < N_TOKENS_AUDIO:\n",
    "        arr_c = np.pad(\n",
    "            arr_c,\n",
    "            ((0, N_TOKENS_AUDIO - len(arr_c)), (0, 0)),\n",
    "            constant_values=COARSE_PAD_TOKEN,\n",
    "            mode=\"constant\",\n",
    "        )\n",
    "        arr_s = np.pad(\n",
    "            arr_s,\n",
    "            ((0, N_TOKENS_AUDIO - len(arr_s)), (0, 0)),\n",
    "            constant_values=SEMANTIC_PAD_TOKEN,\n",
    "            mode=\"constant\",\n",
    "        )\n",
    "    arr = np.concatenate([arr_s, arr_c], axis=-1)\n",
    "    arr = arr.astype(np.uint16)\n",
    "    assert arr.shape == (N_TOKENS_AUDIO, SEMANTIC_N_CODEBOOKS + COARSE_N_CODEBOOKS)\n",
    "    return arr\n",
    "\n",
    "\n",
    "def make_dataset(input_df, is_val=False):\n",
    "    dset_type = \"val\" if is_val else \"tr\"\n",
    "    out_mmap_path = os.path.join(OUT_DATA_DIR, f\"data_{dset_type}.bin\")\n",
    "    out_metas_path = os.path.join(OUT_DATA_DIR, f\"meta_{dset_type}.jsonl\")\n",
    "    out_info_filepath = os.path.join(OUT_DATA_DIR, f\"info_{dset_type}.json\")\n",
    "\n",
    "    # gather the data\n",
    "    _ = np.memmap(out_mmap_path, dtype=np.uint16, mode=\"w+\", shape=(1,))\n",
    "    n_offs = 0\n",
    "    tot_duration_dict = defaultdict(int)\n",
    "    datasets_info = defaultdict(dict)\n",
    "    n = 0\n",
    "    for i, row in tqdm.tqdm(input_df.iterrows(), total=len(input_df)):\n",
    "        # we need to alternate between preference: neg, pos\n",
    "        # print(i, row)\n",
    "        assert row[\"preference\"] == (i % 2 == 1)\n",
    "        # make mmap -- two different paths\n",
    "        local_path = (\n",
    "            f\"{NPZ_DIR}/{row['s3_id']}.npz\"\n",
    "            if row[\"is_7b\"]\n",
    "            else f\"/app/suno/data/dpo/7b_npz/{row['s3_id']}.npz\"\n",
    "        )\n",
    "        if not os.path.exists(local_path):\n",
    "            # print(row, local_path)\n",
    "            raise ValueError()\n",
    "        # print(local_path)\n",
    "        # print( np.load(local_path))\n",
    "        try:\n",
    "            arr = (\n",
    "                np.load(local_path)[\"v3.0_raw\"]\n",
    "                if row[\"is_7b\"]\n",
    "                else np.load(local_path)[\"v2_raw\"]\n",
    "            )\n",
    "        except Exception as e:\n",
    "            print(local_path)\n",
    "            raise e\n",
    "        assert arr.shape[0] <= 3000\n",
    "        assert arr.shape[1] == 13\n",
    "        arr_duration = arr.shape[0] / 25\n",
    "        # print(arr.shape)\n",
    "        arr = reshift(arr)\n",
    "        # print(\"after shift and pad\", arr.shape)\n",
    "        arr = arr.reshape(\n",
    "            -1,\n",
    "        )\n",
    "        # print(arr.shape)\n",
    "        out_mm = np.memmap(\n",
    "            out_mmap_path,\n",
    "            dtype=np.uint16,\n",
    "            mode=\"r+\",\n",
    "            shape=(n_offs + arr.size,),\n",
    "        )\n",
    "        out_mm[n_offs : n_offs + arr.size] = arr\n",
    "        # print(f\"offset is: {n_offs}\")\n",
    "        # break\n",
    "        # write it once\n",
    "        out_mm.flush()\n",
    "        del out_mm\n",
    "\n",
    "        add_metas = []\n",
    "        add_meta = {\n",
    "            \"dataset\": f\"perference_{int(row['preference'])}\",\n",
    "            \"id\": row[\"s3_id\"],  # this is the row s3_id\n",
    "            \"start_s\": row[\"total_start_s\"] if row[\"total_start_s\"] >= 0 else None,\n",
    "            \"end_s\": (\n",
    "                row[\"total_clip_s\"] if row[\"total_clip_s\"] >= 0 else None\n",
    "            ),  # for full clips we do know it has an edding, other wise, we don't know\n",
    "            \"original_duration_s\": (\n",
    "                row[\"original_duration_s\"]\n",
    "                if row[\"original_duration_s\"] >= 0\n",
    "                else arr_duration\n",
    "            ),  # this nees to be... a bit more complicated, only works with concat!\n",
    "            \"vocal_start_s\": None,  # these are unfortunately missing for now\n",
    "            \"vocal_end_s\": None,  # these are unfortunately missing for now\n",
    "            \"tags\": [\n",
    "                row[\"tags\"] if not pd.isna(row[\"tags\"]) else \"\"\n",
    "            ],  # tags is a list, do you know :)\n",
    "            \"text\": row[\"prompt\"] if not pd.isna(row[\"prompt\"]) else \"\",\n",
    "        }\n",
    "        add_metas.append(add_meta)\n",
    "        tot_duration_dict[row[\"preference\"]] += arr_duration\n",
    "        write_jsonl(\n",
    "            add_metas,\n",
    "            os.path.join(out_metas_path),\n",
    "            do_append=bool(n_offs != 0),\n",
    "        )\n",
    "        if \"idx_list\" not in datasets_info[add_meta[\"dataset\"]]:\n",
    "            datasets_info[add_meta[\"dataset\"]][\"idx_list\"] = [n]\n",
    "        else:\n",
    "            datasets_info[add_meta[\"dataset\"]][\"idx_list\"].append(n)\n",
    "        n += 1\n",
    "        n_offs += arr.size\n",
    "\n",
    "    write_json(datasets_info, out_info_filepath)\n",
    "    print(f\"Total {n} clips\")\n",
    "    for k, v in tot_duration_dict.items():\n",
    "        print(f\"{round(v / 60 / 60):,} hours of {k}\")\n",
    "    print(f\"Done\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Actually make"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 19,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T20:58:01.184248Z",
     "start_time": "2024-04-12T20:58:01.125356Z"
    }
   },
   "outputs": [],
   "source": [
    "# val_df[[\"request_id\", \"metadata\", \"updated_at\", \"user_id\", \"preference\"]].head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 20,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T20:58:08.621142Z",
     "start_time": "2024-04-12T20:58:01.185328Z"
    }
   },
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "100%|█████████████████████████████████████████████████████████████████████████████████████████████████████| 200604/200604 [00:07<00:00, 27206.77it/s]"
     ]
    },
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "3,831 hours of 200604 clips, 6.268875 nodes\n"
     ]
    },
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "\n"
     ]
    }
   ],
   "source": [
    "total_duration = 0\n",
    "for i, row in tqdm.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",
    "    except:\n",
    "        print(i, row)\n",
    "    total_duration += row[\"duration\"]\n",
    "print(f\"{round(total_duration / 60 / 60):,} hours of {train_df.shape[0]} clips, {train_df.shape[0] / 8 / 4 / 1000} nodes\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 21,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T20:58:38.241463Z",
     "start_time": "2024-04-12T20:58:08.622546Z"
    }
   },
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████| 2028/2028 [00:29<00:00, 68.54it/s]"
     ]
    },
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Total 2028 clips\n",
      "20 hours of False\n",
      "19 hours of True\n",
      "Done\n"
     ]
    },
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "\n"
     ]
    }
   ],
   "source": [
    "make_dataset(val_df, is_val=True)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 22,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T21:50:07.161228Z",
     "start_time": "2024-04-12T20:58:38.244440Z"
    }
   },
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      " 56%|██████████████████████████████████████████████████████████                                              | 111911/200604 [29:05<20:37, 71.65it/s]IOPub message rate exceeded.\n",
      "The notebook server will temporarily stop sending output\n",
      "to the client in order to avoid crashing it.\n",
      "To change this limit, set the config variable\n",
      "`--NotebookApp.iopub_msg_rate_limit`.\n",
      "\n",
      "Current values:\n",
      "NotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\n",
      "NotebookApp.rate_limit_window=3.0 (secs)\n",
      "\n"
     ]
    }
   ],
   "source": [
    "make_dataset(train_df, is_val=False)"
   ]
  },
  {
   "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": 23,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T21:50:07.188633Z",
     "start_time": "2024-04-12T21:50:07.162904Z"
    }
   },
   "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, 3008, 13)\n",
    "assert len(mm) == len(test_metas)\n",
    "assert mm[:100, :, 0].min() >= 0\n",
    "assert mm[:100, :, 0].max() <= 4000\n",
    "assert mm[:100, :, 1:].min() >= 0\n",
    "assert mm[:100, :, 1:].max() <= 2048"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 24,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T21:50:07.236535Z",
     "start_time": "2024-04-12T21:50:07.190158Z"
    }
   },
   "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\"] = \"3\"\n",
    "# _ = preload_codec_models(\"/app/suno/tony/v3/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": 25,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T21:50:07.301875Z",
     "start_time": "2024-04-12T21:50:07.237575Z"
    }
   },
   "outputs": [],
   "source": [
    "# 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(\"negative example\", test_metas[idx])\n",
    "# a.play(compress=False)\n",
    "# pos_a = decode(pos_arr)\n",
    "# print(\"positive example\", 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": 26,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T21:50:07.364819Z",
     "start_time": "2024-04-12T21:50:07.302965Z"
    }
   },
   "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": 27,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T21:50:07.427760Z",
     "start_time": "2024-04-12T21:50:07.367630Z"
    }
   },
   "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": 28,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T21:50:07.490935Z",
     "start_time": "2024-04-12T21:50:07.428732Z"
    }
   },
   "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": 29,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T21:50:07.557179Z",
     "start_time": "2024-04-12T21:50:07.491953Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "1014 0\n"
     ]
    }
   ],
   "source": [
    "def validation_on_metas(input_metas):\n",
    "\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(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": 30,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T21:50:07.661593Z",
     "start_time": "2024-04-12T21:50:07.558331Z"
    }
   },
   "outputs": [],
   "source": [
    "train_info = read_json(os.path.join(OUT_DATA_DIR, f\"info_tr.json\"))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 31,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T21:50:07.715783Z",
     "start_time": "2024-04-12T21:50:07.662785Z"
    }
   },
   "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)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 32,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T21:50:07.782764Z",
     "start_time": "2024-04-12T21:50:07.716948Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "total samples 200604 (200604, 75)\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": 33,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T21:50:07.863293Z",
     "start_time": "2024-04-12T21:50:07.783893Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "1 epoch per batch 4, total 1567.21875\n"
     ]
    }
   ],
   "source": [
    "print(\"1 epoch per batch 4, total\", total_iters / 8 / 4 / 4)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 34,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T21:50:08.726150Z",
     "start_time": "2024-04-12T21:50:07.864431Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Submitted batch job 1368\r\n"
     ]
    }
   ],
   "source": [
    "!cd /home/tony/Work/tony/slurm && sbatch sbatch_dpo"
   ]
  },
  {
   "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.13"
  },
  "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": 2
}
