{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import re\n",
    "import funcy\n",
    "import numpy as np\n",
    "from tqdm import tqdm\n",
    "import multiprocessing\n",
    "from phonemizer import phonemize\n",
    "from suno_utils.utils.text import read_jsonl, write_jsonl\n",
    "\n",
    "# you will need espeak installed\n",
    "# sudo apt-get install espeak-ng "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "use_val = True\n",
    "metas_dir = \"/app/suno/data/diffusion_mix/vae_2min_v2\"\n",
    "#metas_dir = \"/app/suno/data/diffusion_mix/metadata_2min/\"\n",
    "\n",
    "if use_val:\n",
    "    name = \"val\"\n",
    "else:\n",
    "    name = \"tr\"\n",
    "\n",
    "# load metadata \n",
    "metas = read_jsonl(f\"{metas_dir}/metas_tr.jsonl\", progress=True)\n",
    "print(len(metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "def clean_text(text):\n",
    "    # Remove meta tags like [verse] and [chorus]\n",
    "    text = re.sub(r'\\[.*?\\]', '', text)\n",
    "    \n",
    "    # Remove punctuation and special characters\n",
    "    # This keeps alphanumeric characters and spaces\n",
    "    text = re.sub(r'[^a-zA-Z0-9\\s]', '', text)\n",
    "    \n",
    "    # Convert to lowercase\n",
    "    text = text.lower()\n",
    "    \n",
    "    # Remove extra whitespace\n",
    "    text = ' '.join(text.split())\n",
    "    \n",
    "    return text\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def phonemize_chunk(texts, language=\"en-us\", backend=\"espeak\"):\n",
    "    return phonemize(\n",
    "        texts,\n",
    "        language=language,\n",
    "        backend=backend,\n",
    "        strip=True,\n",
    "        preserve_punctuation=False,\n",
    "    )\n",
    "\n",
    "def chunker(seq, size):\n",
    "    return (seq[pos:pos + size] for pos in range(0, len(seq), size))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "inputs = []\n",
    "for idx, meta in enumerate(metas):\n",
    "    if \"text\" in meta:\n",
    "        # clean the text\n",
    "        cleaned_text = clean_text(meta[\"text\"])\n",
    "        if len(cleaned_text) > 0:\n",
    "            inputs.append((idx, cleaned_text))\n",
    "\n",
    "print(len(inputs))\n",
    "\n",
    "input_ids = [input_[0] for input_ in inputs]\n",
    "input_texts = [input_[1] for input_ in inputs]\n",
    "\n",
    "chunksize = 10_000\n",
    "num_processes = 32 #multiprocessing.cpu_count()   # Use all available CPU cores\n",
    "\n",
    "chunks = list(chunker(input_texts, chunksize))\n",
    "\n",
    "with multiprocessing.Pool(processes=num_processes) as pool:\n",
    "    phonemized_chunks = list(tqdm(\n",
    "        pool.imap(phonemize_chunk, chunks),\n",
    "        total=len(chunks),\n",
    "        desc=\"Phonemizing\"\n",
    "    ))\n",
    "\n",
    "phonemized_texts = [item for sublist in phonemized_chunks for item in sublist]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "inputs = []\n",
    "for idx, meta in enumerate(metas):\n",
    "    if \"text\" in meta:\n",
    "        # clean the text\n",
    "        cleaned_text = clean_text(meta[\"text\"])\n",
    "        if len(cleaned_text) > 0:\n",
    "            inputs.append((idx, cleaned_text))\n",
    "\n",
    "print(len(inputs))\n",
    "\n",
    "input_ids = [input_[0] for input_ in inputs]\n",
    "input_texts = [input_[1] for input_ in inputs]\n",
    "\n",
    "num_processes = multiprocessing.cpu_count() - 4\n",
    "\n",
    "phonemized_texts = []\n",
    "#chunks = list(chunker(input_texts, chunksize))\n",
    "\n",
    "\n",
    "chunksize = 10_000\n",
    "for items_chunk in tqdm(\n",
    "    funcy.chunks(chunksize, input_texts), \n",
    "    total=int(np.ceil(len(input_texts) / chunksize))\n",
    "):\n",
    "    print(len(items_chunk))\n",
    "    tmp_phonemized_texts = phonemize(\n",
    "        items_chunk,\n",
    "        language=\"en-us\",\n",
    "        backend=\"espeak\",\n",
    "        strip=True,\n",
    "    #     with_stress=True,\n",
    "        preserve_punctuation=False,\n",
    "        njobs=1,\n",
    "        prepend_text=True\n",
    "    )\n",
    "    print(len(tmp_phonemized_texts))\n",
    "    phonemized_texts.extend(tmp_phonemized_texts)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(len(phonemized_texts), len(input_ids))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# slow version\n",
    "new_metas = []\n",
    "for idx, meta in enumerate(metas):\n",
    "    if idx in input_ids:\n",
    "        meta[\"phonemized_text\"] = phonemized_texts[input_ids.index(idx)]\n",
    "    new_metas.append(meta)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# fast version\n",
    "## Create a dictionary for faster lookups\n",
    "phonemized_dict = {idx: text for idx, text in zip(input_ids, phonemized_texts)}\n",
    "\n",
    "# Use a list comprehension with dictionary lookup\n",
    "new_metas = [\n",
    "    {**meta, \"phonemized_text\": phonemized_dict.get(idx, meta.get(\"phonemized_text\"))}\n",
    "    for idx, meta in enumerate(metas)\n",
    "]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "for i, j in enumerate(input_ids[100:110]):\n",
    "    print(i, j)\n",
    "    print(metas[j][\"text\"][:100])\n",
    "    print(\"----\")\n",
    "    print(phonemized_texts[i][:100])\n",
    "    print(\"----\")\n",
    "    print()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(len(metas), len(new_metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "write_jsonl(new_metas, f\"{metas_dir}/metas_tr_phonemes.jsonl\") # careful this will overwrite the original file"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def get_unique_phonemes(phoneme_strings):\n",
    "    return set(''.join(phoneme_strings))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Get all phonemes\n",
    "phoneme_strings = [meta[\"phonemized_text\"] for meta in new_metas if meta[\"phonemized_text\"] is not None]\n",
    "print(len(phoneme_strings))\n",
    "phoneme_set = get_unique_phonemes(phoneme_strings)\n",
    "phoneme_set = sorted(phoneme_set)\n",
    "# create a dict mapping phonemes to indices\n",
    "phoneme_to_idx = {phoneme: idx for idx, phoneme in enumerate(phoneme_set)}\n",
    "print(phoneme_to_idx)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# download speed of .5 gb/s\n",
    "# need to transfer 8tb of data\n",
    "# 8 * 1024 / .5 / 60 / 60 = 45 hours"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# find all unique phonemes in the dataset\n",
    "example = \"\"\"\n",
    "You need a love that will not change\n",
    "You want a lover to remain, forever yours\n",
    "Oh, but you don't have to worry, you never have to fear\n",
    "Through thick and thin I'll always be here\n",
    "\n",
    "I'll be your bridge over and through troubled waters\n",
    "You never have to face it alone\n",
    "And when the world seems to treat you unfair\n",
    "Baby, for you I'll always be there\n",
    "\n",
    "I won't be no fairweather friend\n",
    "I'll be there till the end\n",
    "Even through stormy weather\n",
    "Time and time again\n",
    "\n",
    "I won't be no fairweather friend\"\"\".strip()\n",
    "\n",
    "phonemized_example = phonemize(\n",
    "    clean_text(example),\n",
    "    language=\"en-us\",\n",
    "    backend=\"espeak\",\n",
    "    strip=True,\n",
    "    preserve_punctuation=False,\n",
    ")\n",
    "\n",
    "print(phonemized_example)\n"
   ]
  },
  {
   "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
}
