{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 1,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T04:01:19.099112Z",
     "start_time": "2024-04-12T04:01:17.860891Z"
    }
   },
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "/tmp/ipykernel_2745515/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-12T04:01:19.135094Z",
     "start_time": "2024-04-12T04:01:19.100407Z"
    }
   },
   "outputs": [],
   "source": [
    "OUT_DATA_DIR = \"/app/suno/data/dpo/7v_r2_v0_mix/\"\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-12T04:01:19.172428Z",
     "start_time": "2024-04-12T04:01:19.136738Z"
    }
   },
   "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-12T04:03:16.988164Z",
     "start_time": "2024-04-12T04:01:19.173500Z"
    }
   },
   "outputs": [
    {
     "data": {
      "text/plain": [
       "(2322354, 71)"
      ]
     },
     "execution_count": 4,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "df = pd.read_csv(\"/home/tony/Data/Preference/7b_v2/r0_pre_filter.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-12T04:03:17.416849Z",
     "start_time": "2024-04-12T04:03:16.989438Z"
    }
   },
   "outputs": [
    {
     "data": {
      "text/plain": [
       "is_7b\n",
       "True     2322352\n",
       "False          2\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-12T04:03:17.419925Z",
     "start_time": "2024-04-12T04:03:17.418173Z"
    }
   },
   "outputs": [],
   "source": [
    "date_cut = '2024-03-22 04:30:00'"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 7,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T04:03:19.254135Z",
     "start_time": "2024-04-12T04:03:17.420911Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "preference  model_name         \n",
      "False       chirp-v3-engine-i      882207\n",
      "            chirp-v3-engine-v0     147486\n",
      "            chirp-v3-engine-d      111256\n",
      "            chirp-v3-engine-s       18239\n",
      "            chirp-v3-engine-i-d      1988\n",
      "True        chirp-v3-engine-i      886806\n",
      "            chirp-v3-engine-d      184570\n",
      "            chirp-v3-engine-v0      74197\n",
      "            chirp-v3-engine-s       14008\n",
      "            chirp-v3-engine-i-d      1595\n",
      "Name: count, dtype: int64\n",
      "(2322354, 71)\n",
      "(517509, 71)\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-i\"])) & (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-12T04:03:21.812969Z",
     "start_time": "2024-04-12T04:03:19.256545Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "(517509, 71)\n",
      "(443610, 71)\n",
      "preference  model_name        \n",
      "False       chirp-v3-engine-v0    135871\n",
      "            chirp-v3-engine-d      85934\n",
      "True        chirp-v3-engine-d     151841\n",
      "            chirp-v3-engine-v0     69964\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-12T04:03:22.619148Z",
     "start_time": "2024-04-12T04:03:21.815692Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "negative 221805 positive 136133\n",
      "total pair requests 221805 selected pair requests 136133 frac 0.614\n"
     ]
    }
   ],
   "source": [
    "normal_pos_play_count = 5\n",
    "# this is lower, cause a concat is probably already ensuring that it is good\n",
    "concat_pos_play_count = 1\n",
    "# this is a filter on the concated clip\n",
    "concat_total_play_count = 3\n",
    "\n",
    "neg_filter_selection_mask = df[\"preference\"] == False\n",
    "pos_filter_selectin_mask = (\n",
    "    (df[\"preference\"] == True)  # get basics aligned\n",
    "    & (df[\"user_n_clips\"] >= 50)  # user needs to have genereated at least 40\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[\"is_in_playlist\"] == True) | (df[\"concat_in_playlist\"] == True))\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": 10,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T04:03:23.352901Z",
     "start_time": "2024-04-12T04:03:22.620601Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "requests 136133 clips 272266 total khrs 6.874; N gpus for 1000 iters 8.508; n unique users 28555\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": 11,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T04:03:23.410075Z",
     "start_time": "2024-04-12T04:03:23.354672Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "positive in playlist (24443, 71)\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": 12,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T04:03:23.501237Z",
     "start_time": "2024-04-12T04:03:23.411473Z"
    }
   },
   "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\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": 13,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T04:03:23.543146Z",
     "start_time": "2024-04-12T04:03:23.502941Z"
    }
   },
   "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": 14,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T04:03:23.864887Z",
     "start_time": "2024-04-12T04:03:23.544309Z"
    }
   },
   "outputs": [
    {
     "data": {
      "text/plain": [
       "135759"
      ]
     },
     "execution_count": 14,
     "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": 15,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T04:03:23.915737Z",
     "start_time": "2024-04-12T04:03:23.866172Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "136133\n"
     ]
    }
   ],
   "source": [
    "final_filtered_requests = df_slice[\"request_id\"].unique()\n",
    "print(len(final_filtered_requests))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 16,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T04:03:24.848596Z",
     "start_time": "2024-04-12T04:03:23.917081Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "134771 1362\n",
      "(269542, 72) (2724, 72)\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": 17,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T04:03:24.866841Z",
     "start_time": "2024-04-12T04:03:24.850135Z"
    }
   },
   "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": 18,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T04:03:24.933041Z",
     "start_time": "2024-04-12T04:03:24.868046Z"
    }
   },
   "outputs": [],
   "source": [
    "# val_df[[\"request_id\", \"metadata\", \"updated_at\", \"user_id\", \"preference\"]].head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 19,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T04:03:34.726677Z",
     "start_time": "2024-04-12T04:03:24.934339Z"
    }
   },
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "100%|█████████████████████████████████████████████████████████████████████████████████████████████████████| 269542/269542 [00:09<00:00, 27725.35it/s]"
     ]
    },
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "6,805 hours of 269542 clips, 8.4231875 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": 20,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T04:04:12.819463Z",
     "start_time": "2024-04-12T04:03:34.728363Z"
    }
   },
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████| 2724/2724 [00:37<00:00, 72.05it/s]"
     ]
    },
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Total 2724 clips\n",
      "35 hours of False\n",
      "34 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": 21,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T05:13:18.435229Z",
     "start_time": "2024-04-12T04:04:12.820735Z"
    }
   },
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      " 65%|███████████████████████████████████████████████████████████████████▎                                    | 174344/269542 [44:36<25:26, 62.38it/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",
      "100%|██████████████████████████████████████████████████████████████████████████████████████████████████████| 269542/269542 [1:09:04<00:00, 65.04it/s]\n"
     ]
    },
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Total 269542 clips\n",
      "3,469 hours of False\n",
      "3,339 hours of True\n",
      "Done\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": 22,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T05:13:18.470601Z",
     "start_time": "2024-04-12T05:13:18.436921Z"
    }
   },
   "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": 23,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T05:13:18.510526Z",
     "start_time": "2024-04-12T05:13:18.471918Z"
    }
   },
   "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": 24,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T05:13:18.578773Z",
     "start_time": "2024-04-12T05:13:18.511820Z"
    }
   },
   "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": 25,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T05:13:18.658676Z",
     "start_time": "2024-04-12T05:13:18.580139Z"
    }
   },
   "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": 26,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T05:13:18.749987Z",
     "start_time": "2024-04-12T05:13:18.660121Z"
    }
   },
   "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": 27,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T05:13:18.876012Z",
     "start_time": "2024-04-12T05:13:18.753226Z"
    }
   },
   "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": 28,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T05:13:19.220519Z",
     "start_time": "2024-04-12T05:13:18.877405Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "1362 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": 29,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T05:13:19.321472Z",
     "start_time": "2024-04-12T05:13:19.221903Z"
    }
   },
   "outputs": [],
   "source": [
    "train_info = read_json(os.path.join(OUT_DATA_DIR, f\"info_tr.json\"))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 30,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T05:13:19.373463Z",
     "start_time": "2024-04-12T05:13:19.322933Z"
    }
   },
   "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": 31,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T05:13:19.452080Z",
     "start_time": "2024-04-12T05:13:19.374772Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "total samples 269542 (269542, 72)\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": 32,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T05:13:19.524941Z",
     "start_time": "2024-04-12T05:13:19.453208Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "1 epoch per batch 4, total 2105.796875\n"
     ]
    }
   ],
   "source": [
    "print(\"1 epoch per batch 4, total\", total_iters / 8 / 4 / 4)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 33,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-12T05:13:20.338850Z",
     "start_time": "2024-04-12T05:13:19.526066Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Submitted batch job 1362\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
}
