{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 1,
   "id": "5ff0baea",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Populating the interactive namespace from numpy and matplotlib\n"
     ]
    }
   ],
   "source": [
    "%pylab inline"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 2,
   "id": "77817f3e",
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"0\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 3,
   "id": "ec8e0bb6",
   "metadata": {},
   "outputs": [],
   "source": [
    "import time\n",
    "import os\n",
    "import tqdm\n",
    "import shutil\n",
    "import random\n",
    "import funcy\n",
    "import json\n",
    "import numpy as np\n",
    "import multiprocessing\n",
    "import pandas as pd\n",
    "from suno_utils.audio import Audio\n",
    "from suno_utils.audio.conversion import play_audio\n",
    "from suno_utils.utils.podcasts import (\n",
    "    load_podcast_db, find_in_raw_feeds, load_rss_feed, multicore_apply, _load_rss_text, get_file_name,\n",
    "    _clean_episode_url\n",
    ")\n",
    "\n",
    "PODCAST_DATA_DIR = \"/mnt/data-ssd-1/data/podcasts/\"\n",
    "\n",
    "# podcast_df = load_podcast_db(os.path.join(PODCAST_DATA_DIR, \"meta/podcastindex_feeds.db\"), anchor_only=True)\n",
    "# summaries_df = pd.read_csv(os.path.join(PODCAST_DATA_DIR, \"meta/summaries.csv\"))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 4,
   "id": "df62e2ee",
   "metadata": {},
   "outputs": [
    {
     "data": {
      "text/plain": [
       "uid                                    19130554-3955-5aea-a964-41103f863f4a\n",
       "title                                               Transcending La Familia\n",
       "summary                   We are a virtual vocational academy that focus...\n",
       "tags                                                    business;non-profit\n",
       "author                                                       Learning Idiom\n",
       "author_email                                      LearningIdiomar@gmail.com\n",
       "rss_url                            https://anchor.fm/s/62762d40/podcast/rss\n",
       "link                                        https://linktr.ee/Learningidiom\n",
       "language                                                                 es\n",
       "n_episodes                                                               14\n",
       "newest_episode_pubdate                                  2021-10-26 18:32:32\n",
       "oldest_episode_pubdate                                  2021-07-02 09:48:55\n",
       "episode_url_0             https://d3ctxlq1ktw2nl.cloudfront.net/staging/...\n",
       "episode_url_1             https://d3ctxlq1ktw2nl.cloudfront.net/staging/...\n",
       "episode_url_2             https://d3ctxlq1ktw2nl.cloudfront.net/staging/...\n",
       "Name: 0, dtype: object"
      ]
     },
     "execution_count": 4,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "info_df = pd.read_csv(\"tmp/batch_1.csv\")\n",
    "info_df.iloc[0]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 14,
   "id": "2c254a67",
   "metadata": {},
   "outputs": [],
   "source": [
    "for _, row in info_df.iterrows():\n",
    "    fd = os.path.join(PODCAST_DATA_DIR, f\"audio/{row['uid']}\")\n",
    "    if not os.path.exists(fd):\n",
    "        continue\n",
    "    fns = os.listdir(fd)\n",
    "    fps = [os.path.join(fd, fn) for fn in fns if len(fn.split(\".\")[0]) == 36]\n",
    "    if len(fps) == 0:\n",
    "        continue\n",
    "    for fp in fps:\n",
    "        audio = Audio.from_file(fp, sample_rate=16_000, byte_width=2)\n",
    "    break"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 19,
   "id": "f6fe841b",
   "metadata": {},
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "/home/georg/venvs/ml/lib/python3.8/site-packages/huggingface_hub/utils/_deprecation.py:39: FutureWarning: Pass library_name=False as keyword args. From version 0.8 passing these as positional arguments will result in an error\n",
      "  warnings.warn(\n"
     ]
    }
   ],
   "source": [
    "from speechbrain.pretrained import EncoderClassifier\n",
    "\n",
    "language_id = EncoderClassifier.from_hparams(source=\"speechbrain/lang-id-voxlingua107-ecapa\", savedir=\"sb_model_1\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 345,
   "id": "0a072a08",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "tensor([-16.7989,  -8.9914, -13.4097]) tensor([-0.1595]) tensor([96]) ['tl: Tagalog']\n"
     ]
    }
   ],
   "source": [
    "logits, a, b, lang_pred = language_id.classify_batch(torch.Tensor(audio_array))\n",
    "print(logits[0,:3], a, b, lang_pred)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 346,
   "id": "2f7e7922",
   "metadata": {},
   "outputs": [],
   "source": [
    "audio = Audio.from_file(\"data/segment_1.mp3\").convert(16_000, 2, 1)\n",
    "array_segments = [] \n",
    "for s, e in zip([0] + p_boundaries, p_boundaries + [audio.duration_s]):\n",
    "    array_segment = audio.get_slice(s, e).array.astype(np.float32) / np.iinfo(np.int16).max\n",
    "    array_segments.append(torch.Tensor(array_segment))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 347,
   "id": "0a762f69",
   "metadata": {},
   "outputs": [],
   "source": [
    "# Audio.from_array_float(array_segments[1], 16_000).play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 348,
   "id": "95e96bc7",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "En -- max: 0.443 ['la: Latin'] -- p_En: 0.0 -- p_Fi: 0.004\n",
      "Fi -- max: 0.472 ['ko: Korean'] -- p_En: 0.0 -- p_Fi: 0.005\n",
      "En -- max: 0.491 ['tl: Tagalog'] -- p_En: 0.073 -- p_Fi: 0.491\n",
      "Fi -- max: 0.916 ['ar: Arabic'] -- p_En: 0.0 -- p_Fi: 0.032\n",
      "En -- max: 0.304 ['lo: Lao'] -- p_En: 0.0 -- p_Fi: 0.003\n",
      "Fi -- max: 0.973 ['tl: Tagalog'] -- p_En: 0.0 -- p_Fi: 0.973\n",
      "En -- max: 0.472 ['tl: Tagalog'] -- p_En: 0.008 -- p_Fi: 0.472\n",
      "Fi -- max: 0.647 ['tk: Turkmen'] -- p_En: 0.001 -- p_Fi: 0.018\n",
      "En -- max: 0.488 ['tl: Tagalog'] -- p_En: 0.315 -- p_Fi: 0.488\n",
      "Fi -- max: 0.391 ['yi: Yiddish'] -- p_En: 0.037 -- p_Fi: 0.029\n",
      "En -- max: 0.25 ['tl: Tagalog'] -- p_En: 0.148 -- p_Fi: 0.25\n",
      "Fi -- max: 0.958 ['tl: Tagalog'] -- p_En: 0.0 -- p_Fi: 0.958\n",
      "En -- max: 0.656 ['en: English'] -- p_En: 0.656 -- p_Fi: 0.049\n"
     ]
    }
   ],
   "source": [
    "from suno_utils.audio.conversion import _collapse_to_numpy_array\n",
    "\n",
    "assert(len(p_text) == len(array_segments))\n",
    "\n",
    "for (_, lang), array_segment in zip(p_text, array_segments):\n",
    "    logits, a, b, lang_pred = language_id.classify_batch(array_segment)\n",
    "    logits = _collapse_to_numpy_array(logits)\n",
    "    probs = scipy.special.softmax(logits, axis=0)\n",
    "    print(lang, \"--\", \"max:\", round(probs[b], 3), lang_pred, \"--\", \"p_En:\", round(probs[20], 3), \"--\",\"p_Fi:\", round(probs[96], 3))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "809a3f08",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "f98b9bbb",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "df957d09",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "e1da3461",
   "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.8.10"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
