{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 23,
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch\n",
    "import torchaudio\n",
    "import whisper\n",
    "import numpy as np\n",
    "import os\n",
    "import IPython\n",
    "import glob\n",
    "\n",
    "from tokenizers import AddedToken\n",
    "from transformers import PreTrainedTokenizerFast\n",
    "\n",
    "from suno_utils.utils.text import read_jsonl"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "whisper_model = whisper.load_model(\"base\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 35,
   "metadata": {},
   "outputs": [
    {
     "data": {
      "text/plain": [
       "1"
      ]
     },
     "execution_count": 35,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "tokenizer = PreTrainedTokenizerFast(\n",
    "    tokenizer_file=\"/app/suno/data/chirp_v4/multi/tokenizer_60k.json\",\n",
    "    unk_token=\"[UNK]\",\n",
    "    pad_token=\"[PAD]\",\n",
    "    padding=\"max_length\",\n",
    "    truncation=True,\n",
    "    max_length=16,\n",
    ")\n",
    "tokenizer.add_special_tokens({\"additional_special_tokens\": [AddedToken(\"\\n\")]})"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 38,
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "60000\n",
      "tensor([[25430, 19221, 19170,  ...,     1,     1,     1],\n",
      "        [19221, 19170,    66,  ...,     1,     1,     1]])\n",
      "torch.Size([2, 2560, 512])\n"
     ]
    }
   ],
   "source": [
    "texts = [\"hello this is a test of the t5 tokenizer\", \"this is a different test of the t5 tokenizer\"]\n",
    "\n",
    "inputs = tokenizer(\n",
    "    texts,\n",
    "    return_tensors=\"pt\",\n",
    "    padding=\"max_length\",\n",
    "    truncation=True,\n",
    "    max_length=2560,\n",
    ")\n",
    "print(tokenizer.vocab_size)\n",
    "\n",
    "print(inputs[\"input_ids\"])\n",
    "\n",
    "embedding = torch.nn.Embedding(tokenizer.vocab_size, 512)\n",
    "embeds = embedding(inputs[\"input_ids\"])\n",
    "print(embeds.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 34,
   "metadata": {},
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "100%|██████████| 5879/5879 [00:00<00:00, 119690.67it/s]"
     ]
    },
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "\n"
     ]
    },
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "count 1590 mean 508.84150943396224 median 462.0 max 1811 min 4\n"
     ]
    },
    {
     "data": {
      "image/png": "iVBORw0KGgoAAAANSUhEUgAAAh8AAAGfCAYAAAD/BbCUAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjkuMCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy80BEi2AAAACXBIWXMAAA9hAAAPYQGoP6dpAAAhw0lEQVR4nO3df3DUxeH/8VcC5BKEJCTAhQgB/AUoP6pRwlVFC6khZRBLZqqUKWgpVhuoELWYtoowbZORGbC2AZ0OhHYUUTqCoygOBAF/JAiRqGjJABNMLCRUaC4BzSWQ/f7Rb+7jQYBcctnLXZ6PmfdMbt9779tl8+PF3u2+I4wxRgAAAJZEBrsBAACgeyF8AAAAqwgfAADAKsIHAACwivABAACsInwAAACrCB8AAMAqwgcAALCK8AEAAKwifAAAAKt6+lP56aef1tKlS33KRowYoYMHD0qSGhoa9Oijj2rDhg3yeDzKyMjQqlWr5HQ62/wazc3NOnbsmPr27auIiAh/mgcAAILEGKP6+nolJycrMvLScxt+hQ9JuuGGG7R9+/b/u0DP/7vEokWLtGXLFm3cuFFxcXGaP3++ZsyYoQ8++KDN1z927JiGDBnib7MAAEAXUFVVpcGDB1+yjt/ho2fPnkpKSrqg3O12a82aNVq/fr0mTZokSSosLNSoUaNUUlKiCRMmtOn6ffv2lfS/xsfGxvrbPAAAEAR1dXUaMmSI9+/4pfgdPg4dOqTk5GRFR0fL5XIpLy9PKSkpKi0tVVNTk9LT0711R44cqZSUFBUXF180fHg8Hnk8Hu/j+vp6SVJsbCzhAwCAENOWj0z49YHTtLQ0rVu3Tlu3btXq1atVUVGh22+/XfX19aqurlZUVJTi4+N9nuN0OlVdXX3Ra+bl5SkuLs578JYLAADhza+Zj8zMTO/XY8eOVVpamoYOHapXX31VMTEx7WpAbm6ucnJyvI9bpm0AAEB46tBS2/j4eF133XU6fPiwkpKS1NjYqNraWp86NTU1rX5GpIXD4fC+xcJbLQAAhL8OhY/Tp0/ryJEjGjRokFJTU9WrVy8VFRV5z5eXl6uyslIul6vDDQUAAOHBr7ddHnvsMU2bNk1Dhw7VsWPHtGTJEvXo0UMzZ85UXFyc5s6dq5ycHCUkJCg2NlYLFiyQy+Vq80oXAAAQ/vwKH1999ZVmzpypkydPasCAAbrttttUUlKiAQMGSJJWrlypyMhIZWVl+WwyBgAA0CLCGGOC3YjvqqurU1xcnNxuN5//AAAgRPjz95t7uwAAAKsIHwAAwCrCBwAAsIrwAQAArCJ8AAAAqwgfAADAKsIHAACwyq9NxoCLGfbEFp/HR/Ondsp1A3ltAEBwMPMBAACsInwAAACrCB8AAMAqwgcAALCK8AEAAKwifAAAAKsIHwAAwCrCBwAAsIrwAQAArCJ8AAAAqwgfAADAKsIHAACwivABAACsInwAAACrega7AbCnvbenP/957b2lfaCuAwAIbcx8AAAAqwgfAADAKsIHAACwivABAACsInwAAACrCB8AAMAqltqGidaW0QbzOt1Ze5c0A0B3wcwHAACwivABAACsInwAAACrCB8AAMAqwgcAALCK8AEAAKwifAAAAKvY56ObC4d9PdhXAwBCCzMfAADAKsIHAACwivABAACsInwAAACrCB8AAMAqwgcAALCKpbYh4PylpCwjBQCEMmY+AACAVYQPAABgFeEDAABYRfgAAABWET4AAIBVhA8AAGAVS23RKcLhbrkAgM7BzAcAALCK8AEAAKwifAAAAKsIHwAAwCrCBwAAsIrwAQAArCJ8AAAAq9jnAyGHPUQAILQx8wEAAKwifAAAAKs6FD7y8/MVERGhhQsXessaGhqUnZ2txMRE9enTR1lZWaqpqeloOwEAQJhod/jYu3evXnjhBY0dO9anfNGiRXrjjTe0ceNG7dq1S8eOHdOMGTM63FAAABAe2hU+Tp8+rVmzZulvf/ub+vXr5y13u91as2aNVqxYoUmTJik1NVWFhYX68MMPVVJSErBGAwCA0NWu8JGdna2pU6cqPT3dp7y0tFRNTU0+5SNHjlRKSoqKi4tbvZbH41FdXZ3PAQAAwpffS203bNigjz/+WHv37r3gXHV1taKiohQfH+9T7nQ6VV1d3er18vLytHTpUn+bAVzS+ctxj+ZPDVJLAADn82vmo6qqSo888oheeuklRUdHB6QBubm5crvd3qOqqiog1wUAAF2TX+GjtLRUJ06c0E033aSePXuqZ8+e2rVrl5577jn17NlTTqdTjY2Nqq2t9XleTU2NkpKSWr2mw+FQbGyszwEAAMKXX2+7TJ48WZ999plP2QMPPKCRI0dq8eLFGjJkiHr16qWioiJlZWVJksrLy1VZWSmXyxW4VgMAgJDlV/jo27evRo8e7VN2xRVXKDEx0Vs+d+5c5eTkKCEhQbGxsVqwYIFcLpcmTJgQuFYDAICQFfB7u6xcuVKRkZHKysqSx+NRRkaGVq1aFeiXAQAAIarD4WPnzp0+j6Ojo1VQUKCCgoKOXhoAAIQh7u0CAACsInwAAACrCB8AAMAqwgcAALCK8AEAAKwifAAAAKsIHwAAwCrCBwAAsIrwAQAArCJ8AAAAqwgfAADAKsIHAACwivABAACsInwAAACrCB8AAMCqnsFuAPBdw57YEuwmAAA6GTMfAADAKsIHAACwivABAACsInwAAACrCB8AAMAqwgcAALCKpbbA/9eWZb5H86daaAkAhDdmPgAAgFWEDwAAYBXhAwAAWEX4AAAAVhE+AACAVYQPAABgFeEDAABYxT4f6Lbasq8HACDwmPkAAABWET4AAIBVhA8AAGAV4QMAAFhF+AAAAFYRPgAAgFWEDwAAYBXhAwAAWEX4AAAAVhE+AACAVYQPAABgFeEDAABYRfgAAABWcVfbLoY7rXYO/l0BoOtg5gMAAFhF+AAAAFYRPgAAgFWEDwAAYBXhAwAAWEX4AAAAVhE+AACAVezzgaBh7w0A6J6Y+QAAAFYRPgAAgFWEDwAAYBXhAwAAWEX4AAAAVhE+AACAVYQPAABgFeEDAABYRfgAAABWET4AAIBVfoWP1atXa+zYsYqNjVVsbKxcLpfefvtt7/mGhgZlZ2crMTFRffr0UVZWlmpqagLeaAAAELr8Ch+DBw9Wfn6+SktLtW/fPk2aNEnTp0/X559/LklatGiR3njjDW3cuFG7du3SsWPHNGPGjE5pOAAACE1+3Vhu2rRpPo//+Mc/avXq1SopKdHgwYO1Zs0arV+/XpMmTZIkFRYWatSoUSopKdGECRNavabH45HH4/E+rqur87cPAAAghLT7rrbnzp3Txo0bdebMGblcLpWWlqqpqUnp6eneOiNHjlRKSoqKi4svGj7y8vK0dOnS9jYDCDruzgsA/vH7A6efffaZ+vTpI4fDoYceekibNm3S9ddfr+rqakVFRSk+Pt6nvtPpVHV19UWvl5ubK7fb7T2qqqr87gQAAAgdfs98jBgxQmVlZXK73frnP/+pOXPmaNeuXe1ugMPhkMPhaPfzAQBAaPE7fERFRemaa66RJKWmpmrv3r3685//rHvvvVeNjY2qra31mf2oqalRUlJSwBoMAABCW4f3+WhubpbH41Fqaqp69eqloqIi77ny8nJVVlbK5XJ19GUAAECY8GvmIzc3V5mZmUpJSVF9fb3Wr1+vnTt36p133lFcXJzmzp2rnJwcJSQkKDY2VgsWLJDL5broh00BAED341f4OHHihGbPnq3jx48rLi5OY8eO1TvvvKMf/vCHkqSVK1cqMjJSWVlZ8ng8ysjI0KpVqzql4QAAIDT5FT7WrFlzyfPR0dEqKChQQUFBhxoFAADCV7v3+QAQWOfvF3I0f2qQWgIAnYsbywEAAKsIHwAAwCrCBwAAsIrwAQAArCJ8AAAAqwgfAADAKpbahqBg38I92K8PAAhtzHwAAACrCB8AAMAqwgcAALCK8AEAAKwifAAAAKsIHwAAwCrCBwAAsIrwAQAArCJ8AAAAqwgfAADAKsIHAACwivABAACsInwAAACrCB8AAMCqnsFuABBKhj2xJdhNAICQx8wHAACwivABAACsInwAAACrCB8AAMAqwgcAALCK8AEAAKxiqS0QBG1ZsttanaP5UzujOQBgFTMfAADAKsIHAACwivABAACsInwAAACrCB8AAMAqwgcAALCKpbaABdwNFwD+DzMfAADAKsIHAACwivABAACsInwAAACrCB8AAMAqwgcAALCK8AEAAKxinw8ghJy/X8jR/KntqgMAwcTMBwAAsIrwAQAArCJ8AAAAqwgfAADAKsIHAACwivABAACsInwAAACrCB8AAMAqwgcAALCK8AEAAKwifAAAAKsIHwAAwCrCBwAAsIq72lp0/t1GJe44io5p7XuqPc/j+xCATcx8AAAAqwgfAADAKsIHAACwyq/wkZeXp1tuuUV9+/bVwIEDdc8996i8vNynTkNDg7Kzs5WYmKg+ffooKytLNTU1AW00AAAIXX6Fj127dik7O1slJSXatm2bmpqadNddd+nMmTPeOosWLdIbb7yhjRs3ateuXTp27JhmzJgR8IYDAIDQ5Ndql61bt/o8XrdunQYOHKjS0lJNnDhRbrdba9as0fr16zVp0iRJUmFhoUaNGqWSkhJNmDAhcC0HAAAhqUOf+XC73ZKkhIQESVJpaamampqUnp7urTNy5EilpKSouLi41Wt4PB7V1dX5HAAAIHy1O3w0Nzdr4cKFuvXWWzV69GhJUnV1taKiohQfH+9T1+l0qrq6utXr5OXlKS4uznsMGTKkvU0CAAAhoN3hIzs7WwcOHNCGDRs61IDc3Fy53W7vUVVV1aHrAQCArq1dO5zOnz9fb775pnbv3q3Bgwd7y5OSktTY2Kja2lqf2Y+amholJSW1ei2HwyGHw9GeZgAAgBDk18yHMUbz58/Xpk2btGPHDg0fPtznfGpqqnr16qWioiJvWXl5uSorK+VyuQLTYgAAENL8mvnIzs7W+vXr9frrr6tv377ez3HExcUpJiZGcXFxmjt3rnJycpSQkKDY2FgtWLBALpeLlS4AAECSn+Fj9erVkqQ777zTp7ywsFD333+/JGnlypWKjIxUVlaWPB6PMjIytGrVqoA0FgAAhD6/wocx5rJ1oqOjVVBQoIKCgnY3CgAAhC/u7QIAAKwifAAAAKsIHwAAwCrCBwAAsIrwAQAArCJ8AAAAqwgfAADAqnbd2wWBM+yJLcFuAgAAVjHzAQAArCJ8AAAAqwgfAADAKsIHAACwivABAACsInwAAACrCB8AAMAq9vkA0Cat7UlzNH9qEFoCINQx8wEAAKwifAAAAKsIHwAAwCrCBwAAsIrwAQAArCJ8AAAAqwgfAADAKsIHAACwivABAACsInwAAACrCB8AAMAqwgcAALCK8AEAAKzq9ne1be+dOs9/Hnf3BACgbZj5AAAAVhE+AACAVYQPAABgFeEDAABYRfgAAABWET4AAIBVhA8AAGBVt9/nozWt7f0BoPOwbw7QvTDzAQAArCJ8AAAAqwgfAADAKsIHAACwivABAACsInwAAACrWGoLoN1YIgugPZj5AAAAVhE+AACAVYQPAABgFeEDAABYRfgAAABWET4AAIBVLLUFwlxb7tLcWp1QXDbL0l8gNDDzAQAArCJ8AAAAqwgfAADAKsIHAACwivABAACsInwAAACrCB8AAMCqbrfPR1v2PADQ9X5WgrmHR7jsgwJ0Fcx8AAAAqwgfAADAKr/Dx+7duzVt2jQlJycrIiJCmzdv9jlvjNFTTz2lQYMGKSYmRunp6Tp06FCg2gsAAEKc3+HjzJkzGjdunAoKClo9/8wzz+i5557T888/rz179uiKK65QRkaGGhoaOtxYAAAQ+vz+wGlmZqYyMzNbPWeM0bPPPqvf//73mj59uiTpH//4h5xOpzZv3qz77ruvY60FAAAhL6Cf+aioqFB1dbXS09O9ZXFxcUpLS1NxcXGrz/F4PKqrq/M5AABA+AroUtvq6mpJktPp9Cl3Op3ec+fLy8vT0qVLA9kMAEHSluW5bVmi2tWW+QIIrKCvdsnNzZXb7fYeVVVVwW4SAADoRAENH0lJSZKkmpoan/KamhrvufM5HA7Fxsb6HAAAIHwFNHwMHz5cSUlJKioq8pbV1dVpz549crlcgXwpAAAQovz+zMfp06d1+PBh7+OKigqVlZUpISFBKSkpWrhwof7whz/o2muv1fDhw/Xkk08qOTlZ99xzTyDbDQAAQpTf4WPfvn36wQ9+4H2ck5MjSZozZ47WrVun3/zmNzpz5owefPBB1dbW6rbbbtPWrVsVHR0duFYDAICQ5Xf4uPPOO2WMuej5iIgILVu2TMuWLetQwwAAQHjqdne1tYnlgkDXE8y74wL4n6AvtQUAAN0L4QMAAFhF+AAAAFYRPgAAgFWEDwAAYBXhAwAAWEX4AAAAVrHPB4CQ1JZ9dNq7105n7dHT2nXZZwTdETMfAADAKsIHAACwivABAACsInwAAACrCB8AAMAqwgcAALCKpbYB0llL8wB0Tef/zAdqySzLcdEdMPMBAACsInwAAACrCB8AAMAqwgcAALCK8AEAAKwifAAAAKtYagvAKpalA2DmAwAAWEX4AAAAVhE+AACAVYQPAABgFeEDAABYRfgAAABWET4AAIBV7PMBoFsLl31HWuvH0fypQWgJcHnMfAAAAKsIHwAAwCrCBwAAsIrwAQAArCJ8AAAAqwgfAADAKpbaAkCYOn/5LUtv0VUw8wEAAKwifAAAAKsIHwAAwCrCBwAAsIrwAQAArCJ8AAAAq1hqCwCdpC13zG1vnfYsm23vdYK5ZJe79YYnZj4AAIBVhA8AAGAV4QMAAFhF+AAAAFYRPgAAgFWEDwAAYBXhAwAAWMU+HwAQAG3ZryOUX89f7M+BS2HmAwAAWEX4AAAAVhE+AACAVYQPAABgFeEDAABYRfgAAABWsdQWALqx9izZbe8y3/Of19rS285aQtze6wZzeXCglit3xWXPzHwAAACrCB8AAMAqwgcAALCq08JHQUGBhg0bpujoaKWlpemjjz7qrJcCAAAhpFPCxyuvvKKcnBwtWbJEH3/8scaNG6eMjAydOHGiM14OAACEkE5Z7bJixQrNmzdPDzzwgCTp+eef15YtW7R27Vo98cQTPnU9Ho88Ho/3sdvtliTV1dV1RtPU7PmmU64LAPBPa7/n2/I7uj1/H9r7u7+z/ha1RWttDlTfO6NfLdc0xly+sgkwj8djevToYTZt2uRTPnv2bHP33XdfUH/JkiVGEgcHBwcHB0cYHFVVVZfNCgGf+fj666917tw5OZ1On3Kn06mDBw9eUD83N1c5OTnex83NzTp16pQSExMVERERsHbV1dVpyJAhqqqqUmxsbMCu29V1135L3bfv9Lt79Vvqvn3vrv2WumbfjTGqr69XcnLyZesGfZMxh8Mhh8PhUxYfH99prxcbG9tlBsqm7tpvqfv2nX53P921792131LX63tcXFyb6gX8A6f9+/dXjx49VFNT41NeU1OjpKSkQL8cAAAIMQEPH1FRUUpNTVVRUZG3rLm5WUVFRXK5XIF+OQAAEGI65W2XnJwczZkzRzfffLPGjx+vZ599VmfOnPGufgkGh8OhJUuWXPAWT7jrrv2Wum/f6Xf36rfUffveXfsthX7fI4xpy5oY//31r3/V8uXLVV1dre9973t67rnnlJaW1hkvBQAAQkinhQ8AAIDWcG8XAABgFeEDAABYRfgAAABWET4AAIBV3SJ8FBQUaNiwYYqOjlZaWpo++uijYDepQ/Ly8nTLLbeob9++GjhwoO655x6Vl5f71LnzzjsVERHhczz00EM+dSorKzV16lT17t1bAwcO1OOPP66zZ8/a7Irfnn766Qv6NXLkSO/5hoYGZWdnKzExUX369FFWVtYFG96FYr+HDRt2Qb8jIiKUnZ0tKXzGe/fu3Zo2bZqSk5MVERGhzZs3+5w3xuipp57SoEGDFBMTo/T0dB06dMinzqlTpzRr1izFxsYqPj5ec+fO1enTp33qfPrpp7r99tsVHR2tIUOG6Jlnnunsrl3Wpfre1NSkxYsXa8yYMbriiiuUnJys2bNn69ixYz7XaO37JD8/36dOV+v75cb8/vvvv6BPU6ZM8akTjmMuqdWf+YiICC1fvtxbJxTHXJICfmO5rmbDhg0mKirKrF271nz++edm3rx5Jj4+3tTU1AS7ae2WkZFhCgsLzYEDB0xZWZn50Y9+ZFJSUszp06e9de644w4zb948c/z4ce/hdru958+ePWtGjx5t0tPTzf79+81bb71l+vfvb3Jzc4PRpTZbsmSJueGGG3z69Z///Md7/qGHHjJDhgwxRUVFZt++fWbChAnm+9//vvd8qPb7xIkTPn3etm2bkWTeffddY0z4jPdbb71lfve735nXXnvNSLrgBpX5+fkmLi7ObN682XzyySfm7rvvNsOHDzfffvutt86UKVPMuHHjTElJiXnvvffMNddcY2bOnOk973a7jdPpNLNmzTIHDhwwL7/8somJiTEvvPCCrW626lJ9r62tNenp6eaVV14xBw8eNMXFxWb8+PEmNTXV5xpDhw41y5Yt8/k++O7vha7Y98uN+Zw5c8yUKVN8+nTq1CmfOuE45sYYnz4fP37crF271kRERJgjR45464TimBtjTNiHj/Hjx5vs7Gzv43Pnzpnk5GSTl5cXxFYF1okTJ4wks2vXLm/ZHXfcYR555JGLPuett94ykZGRprq62lu2evVqExsbazweT2c2t0OWLFlixo0b1+q52tpa06tXL7Nx40Zv2b/+9S8jyRQXFxtjQrff53vkkUfM1VdfbZqbm40x4Tne5/8ybm5uNklJSWb58uXestraWuNwOMzLL79sjDHmiy++MJLM3r17vXXefvttExERYf79738bY4xZtWqV6devn0+/Fy9ebEaMGNHJPWq71v4Qne+jjz4yksyXX37pLRs6dKhZuXLlRZ/T1ft+sfAxffr0iz6nO4359OnTzaRJk3zKQnXMw/ptl8bGRpWWlio9Pd1bFhkZqfT0dBUXFwexZYHldrslSQkJCT7lL730kvr376/Ro0crNzdX33zzjfdccXGxxowZ43P34YyMDNXV1enzzz+30/B2OnTokJKTk3XVVVdp1qxZqqyslCSVlpaqqanJZ7xHjhyplJQU73iHcr9bNDY26sUXX9TPf/5znzs/h+t4t6ioqFB1dbXP+MbFxSktLc1nfOPj43XzzTd766SnpysyMlJ79uzx1pk4caKioqK8dTIyMlReXq7//ve/lnrTcW63WxERERfciDM/P1+JiYm68cYbtXz5cp+31kK17zt37tTAgQM1YsQIPfzwwzp58qT3XHcZ85qaGm3ZskVz58694FwojnnQ72rbmb7++mudO3fO5xeuJDmdTh08eDBIrQqs5uZmLVy4ULfeeqtGjx7tLf/pT3+qoUOHKjk5WZ9++qkWL16s8vJyvfbaa5Kk6urqVv9dWs51VWlpaVq3bp1GjBih48ePa+nSpbr99tt14MABVVdXKyoq6oJfxk6n09unUO33d23evFm1tbW6//77vWXhOt7f1dLO1vrx3fEdOHCgz/mePXsqISHBp87w4cMvuEbLuX79+nVK+wOpoaFBixcv1syZM33uaPrrX/9aN910kxISEvThhx8qNzdXx48f14oVKySFZt+nTJmiGTNmaPjw4Tpy5Ih++9vfKjMzU8XFxerRo0e3GfO///3v6tu3r2bMmOFTHqpjHtbhozvIzs7WgQMH9P777/uUP/jgg96vx4wZo0GDBmny5Mk6cuSIrr76atvNDJjMzEzv12PHjlVaWpqGDh2qV199VTExMUFsmT1r1qxRZmamkpOTvWXhOt64UFNTk37yk5/IGKPVq1f7nMvJyfF+PXbsWEVFRemXv/yl8vLyQvYeIPfdd5/36zFjxmjs2LG6+uqrtXPnTk2ePDmILbNr7dq1mjVrlqKjo33KQ3XMw/ptl/79+6tHjx4XrHaoqalRUlJSkFoVOPPnz9ebb76pd999V4MHD75k3Zb76hw+fFiSlJSU1Oq/S8u5UBEfH6/rrrtOhw8fVlJSkhobG1VbW+tT57vjHer9/vLLL7V9+3b94he/uGS9cBzvlnZe6uc5KSlJJ06c8Dl/9uxZnTp1Kiy+B1qCx5dffqlt27b5zHq0Ji0tTWfPntXRo0clhXbfW1x11VXq37+/z/d2OI+5JL333nsqLy+/7M+9FDpjHtbhIyoqSqmpqSoqKvKWNTc3q6ioSC6XK4gt6xhjjObPn69NmzZpx44dF0yptaasrEySNGjQIEmSy+XSZ5995vND2/LL7Prrr++UdneG06dP68iRIxo0aJBSU1PVq1cvn/EuLy9XZWWld7xDvd+FhYUaOHCgpk6desl64Tjew4cPV1JSks/41tXVac+ePT7jW1tbq9LSUm+dHTt2qLm52RvIXC6Xdu/eraamJm+dbdu2acSIEV16+r0leBw6dEjbt29XYmLiZZ9TVlamyMhI79sSodr37/rqq6908uRJn+/tcB3zFmvWrFFqaqrGjRt32bohM+ZB/birBRs2bDAOh8OsW7fOfPHFF+bBBx808fHxPp/6DzUPP/ywiYuLMzt37vRZXvXNN98YY4w5fPiwWbZsmdm3b5+pqKgwr7/+urnqqqvMxIkTvddoWXp51113mbKyMrN161YzYMCALrf08nyPPvqo2blzp6moqDAffPCBSU9PN/379zcnTpwwxvxvqW1KSorZsWOH2bdvn3G5XMblcnmfH6r9NuZ/K7VSUlLM4sWLfcrDabzr6+vN/v37zf79+40ks2LFCrN//37vio78/HwTHx9vXn/9dfPpp5+a6dOnt7rU9sYbbzR79uwx77//vrn22mt9ll3W1tYap9Npfvazn5kDBw6YDRs2mN69ewd96eGl+t7Y2GjuvvtuM3jwYFNWVubzc9+yiuHDDz80K1euNGVlZebIkSPmxRdfNAMGDDCzZ8/2vkZX7Pul+l1fX28ee+wxU1xcbCoqKsz27dvNTTfdZK699lrT0NDgvUY4jnkLt9ttevfubVavXn3B80N1zI3pBkttjTHmL3/5i0lJSTFRUVFm/PjxpqSkJNhN6hBJrR6FhYXGGGMqKyvNxIkTTUJCgnE4HOaaa64xjz/+uM++D8YYc/ToUZOZmWliYmJM//79zaOPPmqampqC0KO2u/fee82gQYNMVFSUufLKK829995rDh8+7D3/7bffml/96lemX79+pnfv3ubHP/6xOX78uM81QrHfxhjzzjvvGEmmvLzcpzycxvvdd99t9Xt7zpw5xpj/Lbd98sknjdPpNA6Hw0yePPmCf4+TJ0+amTNnmj59+pjY2FjzwAMPmPr6ep86n3zyibntttuMw+EwV155pcnPz7fVxYu6VN8rKiou+nPfstdLaWmpSUtLM3FxcSY6OtqMGjXK/OlPf/L5I21M1+v7pfr9zTffmLvuussMGDDA9OrVywwdOtTMmzfvgv88huOYt3jhhRdMTEyMqa2tveD5oTrmxhgTYYwxnTq1AgAA8B1h/ZkPAADQ9RA+AACAVYQPAABgFeEDAABYRfgAAABWET4AAIBVhA8AAGAV4QMAAFhF+AAAAFYRPgAAgFWEDwAAYNX/A3QR9koMpyvlAAAAAElFTkSuQmCC",
      "text/plain": [
       "<Figure size 640x480 with 1 Axes>"
      ]
     },
     "metadata": {},
     "output_type": "display_data"
    }
   ],
   "source": [
    "# open metadata, iterate over lyrics to tokenize and then measure the length of the tokenized lyrics\n",
    "metas = read_jsonl(\"/app/suno/data/chirp_v4/multi/metas_val.jsonl\", progress=True)\n",
    "\n",
    "counts = []\n",
    "\n",
    "for meta in metas:\n",
    "    if \"text\" not in meta:\n",
    "        continue\n",
    "    text = meta[\"text\"]\n",
    "    tokenized_text = tokenizer(text, return_tensors=\"pt\")\n",
    "    counts.append(tokenized_text[\"input_ids\"].shape[1])\n",
    "\n",
    "print(\"count\", len(counts), \"mean\", np.mean(counts), \"median\", np.median(counts), \"max\", np.max(counts), \"min\", np.min(counts))\n",
    "\n",
    "# histogram of tokenized lengths\n",
    "import matplotlib.pyplot as plt\n",
    "plt.hist(counts, bins=100)\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "root_dir = \"/app/suno/christian/data/tiktok_covers_48khz/train\"\n",
    "filepaths = glob.glob(os.path.join(root_dir, \"*.wav\"))\n",
    "print(len(filepaths))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "\n",
    "random_filepath = np.random.choice(filepaths)\n",
    "print(random_filepath)\n",
    "result = whisper_model.transcribe(random_filepath)\n",
    "print(result[\"text\"], len(result[\"text\"]), result)\n",
    "IPython.display.display(IPython.display.Audio(random_filepath))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# test T5 on top of whisper transcript\n",
    "\n",
    "from transformers import T5Tokenizer, T5ForConditionalGeneration\n",
    "\n",
    "tokenizer = T5Tokenizer.from_pretrained(\"google-t5/t5-small\")\n",
    "model = T5ForConditionalGeneration.from_pretrained(\"google-t5/t5-small\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Tokenize the input text\n",
    "from tqdm import tqdm\n",
    "\n",
    "texts = []\n",
    "for n in tqdm(range(4)):\n",
    "\n",
    "    random_filepath = np.random.choice(filepaths)\n",
    "    #print(random_filepath)\n",
    "    result = whisper_model.transcribe(random_filepath)\n",
    "    #print(result[\"text\"], len(result[\"text\"]))\n",
    "\n",
    "    text = result[\"text\"]\n",
    "    texts.append(text)\n",
    "\n",
    "inputs = tokenizer(texts, return_tensors='pt', padding='max_length', truncation=True, max_length=256)\n",
    "#inputs = tokenizer(texts, return_tensors='pt', padding=True, truncation=True)\n",
    "print(inputs)\n",
    "\n",
    "# Get the encoder outputs\n",
    "with torch.no_grad():\n",
    "    encoder_outputs = model.encoder(**inputs)\n",
    "\n",
    "# The hidden states are the embeddings (last hidden state)\n",
    "embeddings = encoder_outputs.last_hidden_state\n",
    "\n",
    "# To get a single vector per text, you can average the token embeddings\n",
    "#embeddings = embeddings.mean(dim=1)\n",
    "\n",
    "print(embeddings.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "import torch\n",
    "import IPython\n",
    "import numpy as np\n",
    "import glob\n",
    "import torchaudio\n",
    "from transformers import WhisperProcessor, WhisperForConditionalGeneration\n",
    "\n",
    "\n",
    "root_dir = \"/app/suno/christian/data/tiktok_covers_48khz/train\"\n",
    "root_dir = \"/app/suno/data/audio_2ch_48khz_lg/val/genius_hq\"\n",
    "filepaths = glob.glob(os.path.join(root_dir, \"*.wav\"))\n",
    "print(len(filepaths))\n",
    "\n",
    "\n",
    "# load whisper model for transcription\n",
    "whisper_processor = WhisperProcessor.from_pretrained(\"openai/whisper-base\")\n",
    "whisper_model = WhisperForConditionalGeneration.from_pretrained(\n",
    "    \"openai/whisper-base\"\n",
    ")\n",
    "whisper_model = whisper_model.cuda()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "audios = []\n",
    "indices = np.random.choice(len(filepaths), 5)\n",
    "print(indices)\n",
    "for n in indices):\n",
    "    print(n)\n",
    "    audio, sr = torchaudio.load(filepaths[n])\n",
    "    audio_16k = torchaudio.functional.resample(audio.mean(dim=0, keepdim=False), sr, 16000)\n",
    "    audios.append(audio_16k.numpy())\n",
    "\n",
    "input_features = whisper_processor(\n",
    "    audios, # audios must be list of numpy arrays\n",
    "    sampling_rate=16000,\n",
    "    return_tensors=\"pt\",\n",
    ").input_features\n",
    "\n",
    "print(input_features.shape)\n",
    "input_features = input_features.cuda()\n",
    "\n",
    "# Generate token ids\n",
    "predicted_ids = whisper_model.generate(input_features)\n",
    "\n",
    "# Decode token ids to text\n",
    "transcriptions = whisper_processor.batch_decode(\n",
    "    predicted_ids, skip_special_tokens=True\n",
    ")\n",
    "\n",
    "for transcription, index in zip(transcriptions, indices):\n",
    "    print(f\"Transcription {index}: {transcription}\")\n",
    "    IPython.display.display(IPython.display.Audio(filepaths[index]))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "suno_env",
   "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.9"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
