{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 1,
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch\n",
    "import os\n",
    "import tempfile\n",
    "import re\n",
    "import boto3\n",
    "from tokenizers import Tokenizer\n",
    "from contextlib import contextmanager\n",
    "\n",
    "\n",
    "S3_BUCKET_PATH_RE = r\"s3\\:\\/\\/(.+?)\\/\"\n",
    "\n",
    "def get_filename(filepath, keep_ext=True):\n",
    "    if \"http\" in filepath:\n",
    "        clean_filepath = filepath.split(\"?\")[0]\n",
    "    else:\n",
    "        clean_filepath = filepath\n",
    "    filename = clean_filepath.split(\"/\")[-1]\n",
    "    if \".\" not in filename:\n",
    "        raise ValueError(\"filename does not seem to contain a period.\")\n",
    "    m = re.search(r\"(.+)\\.([^\\.]+)$\", filename)\n",
    "    if not m:\n",
    "        raise ValueError(f\"filename could not be parsed for `{filepath}`\")\n",
    "    filename = m.group(1)\n",
    "    file_ext = m.group(2).lower()\n",
    "    if len(file_ext) > 10:\n",
    "        raise ValueError(f\"file extension suspiciously long for `{filepath}`\")\n",
    "    if keep_ext:\n",
    "        filename = filename + \".\" + file_ext\n",
    "    return filename\n",
    "\n",
    "def _parse_s3_filepath(s3_filepath):\n",
    "    bucket_name = re.search(S3_BUCKET_PATH_RE, s3_filepath).group(1)\n",
    "    rel_s3_filepath = re.sub(S3_BUCKET_PATH_RE, \"\", s3_filepath)\n",
    "    return bucket_name, rel_s3_filepath\n",
    "\n",
    "\n",
    "def download_s3_file(\n",
    "    from_s3_filepath,\n",
    "    to_local_filepath,\n",
    "):\n",
    "    bucket_name, from_rel_s3_filepath = _parse_s3_filepath(from_s3_filepath)\n",
    "    client = boto3.client(\"s3\")\n",
    "    client.download_file(bucket_name, from_rel_s3_filepath, to_local_filepath)\n",
    "\n",
    "@contextmanager\n",
    "def _download_from_s3_if_needed(maybe_s3_filepath):\n",
    "    tmp_filepath = maybe_s3_filepath\n",
    "    if maybe_s3_filepath.startswith(\"s3://\"):\n",
    "        temp_dir = tempfile.TemporaryDirectory()\n",
    "        filename = get_filename(maybe_s3_filepath, keep_ext=True)\n",
    "        tmp_filepath = os.path.join(temp_dir.name, filename)\n",
    "        download_s3_file(maybe_s3_filepath, tmp_filepath)\n",
    "    yield tmp_filepath\n",
    "\n",
    "\n",
    "def load_tokenizer(\n",
    "    tokenizer_filepath=\"s3://suno-data/georg/models/tokenizers/tokenizer_60k.json\",\n",
    "):\n",
    "    with _download_from_s3_if_needed(tokenizer_filepath) as tmp_fp:\n",
    "        tokenizer = Tokenizer.from_file(tmp_fp)\n",
    "    tokenizer.add_special_tokens([\"\\n\"])\n",
    "    tokenizer.pad_idx = tokenizer.token_to_id(\"[PAD]\")\n",
    "    return tokenizer\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 2,
   "metadata": {},
   "outputs": [],
   "source": [
    "tokenizer = load_tokenizer()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# print out the tokens\n",
    "vocab = tokenizer.get_vocab()\n",
    "print(vocab)\n",
    "\n",
    "# sort the vocab by token\n",
    "sorted_vocab = sorted(vocab.items(), key=lambda x: x[1])\n",
    "print(sorted_vocab)\n",
    "\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "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
}
