{
 "cells": [
  {
   "cell_type": "markdown",
   "id": "e36d8a2c",
   "metadata": {},
   "source": [
    "### prod mert, prod codec, peaq vae"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "250e8fae",
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"\"\n",
    "os.environ[\"PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION\"] = \"python\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "ed2662f0",
   "metadata": {},
   "outputs": [],
   "source": [
    "# Memmaps:\n",
    "#   Nx9x3584 for audio tokens\n",
    "# Jsons:\n",
    "#   N*Dict with meta keys \n",
    "#     \"dataset\"\n",
    "#     \"original_id\", \"original_duration_s\",\n",
    "#     \"start_s\", \"end_s\", \n",
    "#     \"text_segments\", \"private_text_segments\",\n",
    "#     \"text\", \"private_text\",\n",
    "#     \"tags\", \"private_tags\",\n",
    "#     \"views\",\n",
    "#   Dict with meta keys {\"dataset\": [\"idx_list\"]}\n",
    "\n",
    "# Bundles (mert_25_2x4k & dac_2c_25_12):\n",
    "# s3://suno-data/datasets/bundles/\n",
    "#  v1/youtube_music\n",
    "#  v1/genius_hq"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "9c4df11a",
   "metadata": {},
   "outputs": [],
   "source": [
    "%matplotlib inline\n",
    "from matplotlib import pyplot as plt"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "6b35b54b",
   "metadata": {},
   "outputs": [],
   "source": [
    "import math\n",
    "import numpy as np\n",
    "import tqdm\n",
    "import time\n",
    "import torch\n",
    "import funcy\n",
    "import json\n",
    "import gc\n",
    "import re\n",
    "import random\n",
    "import tempfile\n",
    "import collections\n",
    "from collections import defaultdict\n",
    "from joblib import Parallel, delayed\n",
    "from transformers import BertTokenizer\n",
    "\n",
    "from suno_utils.utils.text import write_jsonl, read_jsonl, write_json, read_json, normalize_whitespace\n",
    "from suno_utils.utils.s3 import read_from_s3, check_s3_file_exists\n",
    "\n",
    "RATE_HZ = 25\n",
    "\n",
    "SEMANTIC_CODEBOOK_SIZE = 4000\n",
    "SEMANTIC_N_CODEBOOKS = 1\n",
    "\n",
    "CODEC_CODEBOOK_SIZE = 2048\n",
    "CODEC_N_CODEBOOKS = 12\n",
    "\n",
    "VAE_DIM = 128\n",
    "\n",
    "N_TOKENS_MEMMAP = 250\n",
    "\n",
    "SEMANTIC_EMBED_DIR = \"mert_25_2x4k\"\n",
    "CODEC_EMBED_DIR = \"dac_2c_25_12\"\n",
    "VAE_EMBED_DIR = \"dac_vae_128\"\n",
    "\n",
    "METAS_DIR = \"/app/suno/data/chirp_v4/metadata\"\n",
    "OUT_DATA_DIR = \"/app/suno/data/chirp_v4/vae\"\n",
    "os.makedirs(OUT_DATA_DIR, exist_ok=True)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "8a4257e2",
   "metadata": {},
   "outputs": [],
   "source": [
    "# load manifests of IDs and text and tags etc\n",
    "meta_info_map = {\n",
    "    \"genius_hq\": {m[\"id\"]: m for m in read_jsonl(os.path.join(METAS_DIR, \"genius_hq_v5.jsonl\"))},\n",
    "    \"youtube_music\": {m[\"id\"]: m for m in read_jsonl(os.path.join(METAS_DIR, \"youtube_music.jsonl\"))},\n",
    "}"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "1db468de",
   "metadata": {},
   "outputs": [],
   "source": [
    "def _trim_to_common(arr_1, arr_2, arr_3):\n",
    "    common_len = min(len(arr_1), len(arr_2), len(arr_3))\n",
    "    arr_1 = arr_1[:common_len]\n",
    "    arr_2 = arr_2[:common_len]\n",
    "    arr_3 = arr_3[:common_len]\n",
    "    return arr_1, arr_2, arr_3\n",
    "\n",
    "\n",
    "def _parse_arrays(dset_name, meta_info, semantic_arr, codec_arr, vae_arr):\n",
    "\n",
    "    if task_name is None:\n",
    "            task_name = \"default\"\n",
    "        # prep segment metas (use semantic for timekeeping)\n",
    "        if \"lyrics\" in meta_info and \"text\" not in meta_info:\n",
    "            # hotfix incase it's called differently\n",
    "            meta_info[\"text\"] = meta_info[\"lyrics\"]\n",
    "        segments_info = []\n",
    "        # each segment infomration is:\n",
    "        # start_idx, end_idx, text, vocal_start_idx, vocal_end_idx, line_start_times\n",
    "        # first check if we have known segments\n",
    "        if \"text_segments\" in meta_info:\n",
    "            for m in meta_info[\"text_segments\"]:\n",
    "                # we keep the relative start time\n",
    "                selected_line_start_s = []\n",
    "                for line_start_s in m.get(\"line_start_s\", []):\n",
    "                    if line_start_s is not None:\n",
    "                        selected_line_start_s.append(round(line_start_s - m[\"start_s\"], 2))\n",
    "                segments_info.append(\n",
    "                    (\n",
    "                        int(round(m[\"start_s\"] * SEMANTIC_RATE_HZ)),\n",
    "                        min(len(semantic_arr), int(round(m[\"end_s\"] * SEMANTIC_RATE_HZ))),\n",
    "                        m[\"text\"],\n",
    "                        (\n",
    "                            int(round(m[\"vocal_start_s\"] * SEMANTIC_RATE_HZ))\n",
    "                            if m[\"vocal_start_s\"] is not None\n",
    "                            else None\n",
    "                        ),\n",
    "                        (\n",
    "                            min(\n",
    "                                len(semantic_arr),\n",
    "                                int(round(m[\"vocal_end_s\"] * SEMANTIC_RATE_HZ)),\n",
    "                            )\n",
    "                            if m[\"vocal_end_s\"] is not None\n",
    "                            else None\n",
    "                        ),\n",
    "                        selected_line_start_s,\n",
    "                    )\n",
    "                )\n",
    "                # if \"text\" in segmented text then add twice\n",
    "                if \"text\" in meta_info:\n",
    "                    segments_info.append(\n",
    "                        (\n",
    "                            0,\n",
    "                            min(len(semantic_arr), N_TOKENS_MEMMAP),\n",
    "                            meta_info[\"text\"],  # add text only to first piece\n",
    "                            None,\n",
    "                            None,\n",
    "                            None,\n",
    "                        )\n",
    "                    )\n",
    "    else:\n",
    "        # randomize offset to not get only multiples if no text available\n",
    "        offs = 0\n",
    "        if task_name == \"default\" and \"text\" not in meta_info and random.random() > 0.5:\n",
    "            # randomize if we don't have lyrics and not special task\n",
    "            offs = random.randint(10 * SEMANTIC_RATE_HZ, N_TOKENS_MEMMAP - 1)\n",
    "            segments_info.append(\n",
    "                (0, min(len(semantic_arr), offs), None, None, None, None)\n",
    "            )\n",
    "        total_steps = int(np.ceil((len(semantic_arr) - offs) / N_TOKENS_MEMMAP))\n",
    "        for n in range(total_steps):\n",
    "            start_idx = offs + n * N_TOKENS_MEMMAP\n",
    "            end_idx = min(len(semantic_arr), offs + (n + 1) * N_TOKENS_MEMMAP)\n",
    "            if end_idx - start_idx < SEMANTIC_RATE_HZ:\n",
    "                # might as well skip mini ones\n",
    "                continue\n",
    "            segments_info.append(\n",
    "                (\n",
    "                    start_idx,\n",
    "                    end_idx,\n",
    "                    (\n",
    "                        meta_info.get(\"text\") if n == 0 else None\n",
    "                    ),  # add text only to first piece\n",
    "                    None,\n",
    "                    None,\n",
    "                    None,\n",
    "                )\n",
    "            )\n",
    "\n",
    "    segments_info = []\n",
    "    # randomize offset to not get only multiples if no text available\n",
    "    offs = 0\n",
    "    if random.random() > 0.75:\n",
    "        # randomize every once in a while\n",
    "        offs = random.randint(1, N_TOKENS_MEMMAP-1)\n",
    "    for n in range(int(np.ceil((len(semantic_arr)-offs)/N_TOKENS_MEMMAP))):\n",
    "        start_idx = offs + n * N_TOKENS_MEMMAP\n",
    "        end_idx = min(len(semantic_arr), offs+(n+1)*N_TOKENS_MEMMAP)\n",
    "        if end_idx - start_idx < N_TOKENS_MEMMAP:\n",
    "            # might as well skip mini ones\n",
    "            continue\n",
    "        segments_info.append((start_idx, end_idx))\n",
    "\n",
    "    arr_list = []\n",
    "    for start_idx, end_idx in segments_info:\n",
    "        # get array segments\n",
    "        arr_s = semantic_arr[start_idx:end_idx, :SEMANTIC_N_CODEBOOKS].copy()\n",
    "        arr_c = codec_arr[start_idx:end_idx, :CODEC_N_CODEBOOKS].copy()\n",
    "        \n",
    "        arr_v = vae_arr[start_idx:end_idx, :].copy()\n",
    "        assert(arr_v.shape[-1] == VAE_DIM)\n",
    "        # fix any alignment mistakes\n",
    "#         arr_s, arr_c, arr_v = _trim_to_common(arr_s, arr_c, arr_v)\n",
    "        assert len(arr_s) == len(arr_c) == len(arr_v) == N_TOKENS_MEMMAP\n",
    "        arr_s = arr_s.astype(np.uint16)\n",
    "        arr_c = arr_c.astype(np.uint16)\n",
    "        arr_v = arr_v.astype(np.float32)\n",
    "        new_meta = {\n",
    "            \"id\": meta_info[\"id\"],\n",
    "            \"start_s\": round(start_idx / RATE_HZ, 2),\n",
    "            \"end_s\": round(end_idx / RATE_HZ, 2),\n",
    "            \"original_duration_s\": round(len(semantic_arr) / RATE_HZ, 2),\n",
    "        }\n",
    "        if \"original_id\" in meta_info:\n",
    "            new_meta[\"original_id\"] = meta_info[\"original_id\"]\n",
    "\n",
    "        if text is not None:\n",
    "            new_meta[\"text\"] = text\n",
    "            new_meta[\"text_lang\"] = meta_info.get(\"lang\")\n",
    "            if task_name == \"default\":\n",
    "                new_meta[\"dset_suffix\"] = (\n",
    "                    \"lyrics\" if meta_info[\"lang\"] == \"en\" else \"lyrics_foreign\"\n",
    "                )\n",
    "        if \"tags\" in meta_info:\n",
    "            new_meta[\"tags\"] = meta_info[\"tags\"]\n",
    "\n",
    "        arr_list.append((arr_s, arr_c, arr_v, new_meta))\n",
    "        del arr_s, arr_c, arr_v\n",
    "    return arr_list\n",
    "\n",
    "\n",
    "def _process_archives(\n",
    "    dset_name,\n",
    "    s3_semantic_archive_filepaths,\n",
    "    s3_codec_archive_filepaths,\n",
    "    s3_vae_archive_filepaths,\n",
    "    relevant_metas,\n",
    "):\n",
    "#     print(len(relevant_metas))\n",
    "    semantic_archive = {}\n",
    "    s3_semantic_archive_filepaths = set(s3_semantic_archive_filepaths)\n",
    "    # print(len(s3_semantic_archive_filepaths))\n",
    "    for s3_semantic_archive_filepath in s3_semantic_archive_filepaths:\n",
    "        if not check_s3_file_exists(s3_semantic_archive_filepath):\n",
    "            print(f\"missing {s3_semantic_archive_filepath}\")\n",
    "            continue\n",
    "        try:\n",
    "            archive = {k: v for k, v in read_from_s3(s3_semantic_archive_filepath, read_f=np.load).items()}\n",
    "        except:\n",
    "            # corrupt archive\n",
    "            print(f\"corrupt {s3_semantic_archive_filepath}\")\n",
    "            continue\n",
    "        for k, v in archive.items():\n",
    "            semantic_archive[k] = v\n",
    "            \n",
    "    codec_archive = {}\n",
    "    s3_codec_archive_filepaths = set(s3_codec_archive_filepaths)\n",
    "    for s3_codec_archive_filepath in s3_codec_archive_filepaths:\n",
    "        if not check_s3_file_exists(s3_codec_archive_filepath):\n",
    "            print(f\"missing {s3_codec_archive_filepath}\")\n",
    "            continue\n",
    "        try:\n",
    "            archive = {k: v for k, v in read_from_s3(s3_codec_archive_filepath, read_f=np.load).items()}\n",
    "        except:\n",
    "            # corrupt archive\n",
    "            print(f\"corrupt {s3_codec_archive_filepath}\")\n",
    "            continue\n",
    "        for k, v in archive.items():\n",
    "            codec_archive[k] = v\n",
    "            \n",
    "    vae_archive = {}\n",
    "    s3_vae_archive_filepaths = set(s3_vae_archive_filepaths)\n",
    "    for s3_vae_archive_filepath in s3_vae_archive_filepaths:\n",
    "        if not check_s3_file_exists(s3_vae_archive_filepath):\n",
    "            print(f\"missing {s3_vae_archive_filepath}\")\n",
    "            continue\n",
    "        try:\n",
    "            archive = {k: v for k, v in read_from_s3(s3_vae_archive_filepath, read_f=np.load).items()}\n",
    "        except:\n",
    "            # corrupt archive\n",
    "            print(f\"corrupt {s3_vae_archive_filepath}\")\n",
    "            continue\n",
    "        for k, v in archive.items():\n",
    "            vae_archive[k] = v\n",
    "\n",
    "    semantic_uids, codec_uids, vae_uids = (\n",
    "        set(semantic_archive.keys()), \n",
    "        set(codec_archive.keys()),\n",
    "        set(vae_archive.keys()),\n",
    "    )\n",
    "    # removing this for now just incase (eg imslp)\n",
    "    # assert (len(semantic_uids) < 10 and len(coarse_uids) < 10) or (\n",
    "    #     len(semantic_uids & coarse_uids) / (len(semantic_uids) + len(coarse_uids)) > 0.1\n",
    "    # )\n",
    "    # assert len(semantic_uids & coarse_uids) > 0\n",
    "    arr_list = []\n",
    "    for uid in semantic_uids & codec_uids & vae_uids:\n",
    "        if uid not in relevant_metas:\n",
    "            continue\n",
    "        semantic_arr = semantic_archive[uid]\n",
    "        codec_arr = codec_archive[uid]\n",
    "        vae_arr = vae_archive[uid]\n",
    "        if (\n",
    "            np.abs(len(codec_arr) / RATE_HZ - len(semantic_arr) / RATE_HZ) > 0.1 or\n",
    "            np.abs(len(codec_arr) / RATE_HZ - len(vae_arr) / RATE_HZ) > 0.1\n",
    "        ):\n",
    "            # skip if embeddings not roughly the same duration\n",
    "            continue\n",
    "        semantic_arr, codec_arr, vae_arr = _trim_to_common(semantic_arr, codec_arr, vae_arr)\n",
    "        assert len(codec_arr) == len(semantic_arr) == len(vae_arr)\n",
    "        arr_list.extend(\n",
    "            _parse_arrays(dset_name, relevant_metas[uid], semantic_arr, codec_arr, vae_arr)\n",
    "        )\n",
    "    del semantic_archive, codec_archive, vae_archive\n",
    "    gc.collect()\n",
    "    return arr_list\n",
    "\n",
    "\n",
    "def _collect_uids(\n",
    "    s3_semantic_metas_filepaths,\n",
    "    s3_codec_metas_filepaths,\n",
    "    s3_vae_metas_filepaths,\n",
    "):\n",
    "    semantic_uids = []\n",
    "    for fp in s3_semantic_metas_filepaths:\n",
    "        try:\n",
    "            metas = read_from_s3(fp, read_f=read_jsonl)\n",
    "        except:\n",
    "            print(f\"failed on metas for fp: {fp}\")\n",
    "            continue\n",
    "        semantic_uids.extend([m[\"id\"] for m in metas])\n",
    "    codec_uids = []\n",
    "    for fp in s3_codec_metas_filepaths:\n",
    "        try:\n",
    "            metas = read_from_s3(fp, read_f=read_jsonl)\n",
    "        except:\n",
    "            print(f\"failed on metas for fp: {fp}\")\n",
    "            continue\n",
    "        codec_uids.extend([m[\"id\"] for m in metas])\n",
    "    vae_uids = []\n",
    "    for fp in s3_vae_metas_filepaths:\n",
    "        try:\n",
    "            metas = read_from_s3(fp, read_f=read_jsonl)\n",
    "        except:\n",
    "            print(f\"failed on metas for fp: {fp}\")\n",
    "            continue\n",
    "        vae_uids.extend([m[\"id\"] for m in metas])\n",
    "    return set(semantic_uids) & set(codec_uids) & set(vae_uids)\n",
    "\n",
    "\n",
    "def _prep_data(\n",
    "    dataset,\n",
    "    njobs=5,\n",
    "    chunksize=10,\n",
    "    is_val=False,\n",
    "    n_offs_s=0,\n",
    "    n_offs_c=0,\n",
    "    n_offs_v=0,\n",
    "):  \n",
    "    dset_name, dset_version, (start_idx, end_idx), n_sem, n_codec, n_vae = dataset\n",
    "    dset_type = \"val\" if is_val else \"tr\"\n",
    "    out_mm_semantic_filepath = os.path.join(OUT_DATA_DIR, f\"data_semantic_{dset_type}.bin\")\n",
    "    out_mm_codec_filepath = os.path.join(OUT_DATA_DIR, f\"data_codec_{dset_type}.bin\")\n",
    "    out_mm_vae_filepath = os.path.join(OUT_DATA_DIR, f\"data_vae_{dset_type}.bin\")\n",
    "    out_metas_filepath = os.path.join(OUT_DATA_DIR, f\"metas_{dset_type}.jsonl\")\n",
    "    tot_duration_dict = defaultdict(int)\n",
    "    n_chunks = int(np.ceil((end_idx - start_idx) / chunksize))\n",
    "    for idx_chunk in tqdm.tqdm(\n",
    "        funcy.chunks(chunksize, list(range(start_idx, end_idx))), total=n_chunks\n",
    "    ):\n",
    "        n_jobs = np.min([njobs, chunksize, len(idx_chunk)])\n",
    "        # collect relevant parts of meta file to avoid copying all to subprocesses\n",
    "        tmp_uid_chunks = Parallel(n_jobs=n_jobs, prefer=\"threads\")(\n",
    "            delayed(_collect_uids)(\n",
    "                [\n",
    "                    f\"s3://suno-data/datasets/bundles/{dset_version}/{dset_name}/{SEMANTIC_EMBED_DIR}/\"\n",
    "                    + f\"metas/part_{idx_idx}.jsonl\"\n",
    "                    for idx_idx in range(idx * n_sem, (idx + 1) * n_sem)\n",
    "                ],\n",
    "                [\n",
    "                    f\"s3://suno-data/datasets/bundles/{dset_version}/{dset_name}/{CODEC_EMBED_DIR}/\"\n",
    "                    + f\"metas/part_{idx_idx}.jsonl\"\n",
    "                    for idx_idx in range(idx * n_codec, (idx + 1) * n_codec)\n",
    "                ],\n",
    "                [\n",
    "                    f\"s3://suno-data/datasets/bundles/{dset_version}/{dset_name}/{VAE_EMBED_DIR}/\"\n",
    "                    + f\"metas/part_{idx_idx}.jsonl\"\n",
    "                    for idx_idx in range(idx * n_vae, (idx + 1) * n_vae)\n",
    "                ],\n",
    "            )\n",
    "            for idx in idx_chunk\n",
    "        )\n",
    "        ## PART A: takes ~40% of loop time\n",
    "        uids_per_part = {idx: tmp_uid_chunks[n] for n, idx in enumerate(idx_chunk)}\n",
    "        # print(len(uids_per_part))\n",
    "        # collect data\n",
    "        encoded_arrays_list = Parallel(n_jobs=n_jobs, prefer=\"processes\")(\n",
    "            delayed(_process_archives)(\n",
    "                dset_name,\n",
    "                [\n",
    "                    f\"s3://suno-data/datasets/bundles/{dset_version}/{dset_name}/{SEMANTIC_EMBED_DIR}/\"\n",
    "                    + f\"part_{idx_idx}.npz\"\n",
    "                    for idx_idx in range(idx * n_sem, (idx + 1) * n_sem)\n",
    "                ],\n",
    "                [\n",
    "                    f\"s3://suno-data/datasets/bundles/{dset_version}/{dset_name}/{CODEC_EMBED_DIR}/\"\n",
    "                    + f\"part_{idx_idx}.npz\"\n",
    "                    for idx_idx in range(idx * n_codec, (idx + 1) * n_codec)\n",
    "                ],\n",
    "                [\n",
    "                    f\"s3://suno-data/datasets/bundles/{dset_version}/{dset_name}/{VAE_EMBED_DIR}/\"\n",
    "                    + f\"part_{idx_idx}.npz\"\n",
    "                    for idx_idx in range(idx * n_vae, (idx + 1) * n_vae)\n",
    "                ],\n",
    "                {\n",
    "                    uid: meta_info_map[dset_name][uid]\n",
    "                    for uid in uids_per_part[idx]\n",
    "                    if uid in meta_info_map[dset_name]\n",
    "                },\n",
    "            )\n",
    "            for idx in idx_chunk\n",
    "        )\n",
    "        ## end Part A\n",
    "        ## PART B: takes ~40% of loop time\n",
    "        add_metas = []\n",
    "        for encoded_arrays in encoded_arrays_list:\n",
    "            to_write_len_s = np.sum([arr.size for arr, _, _, _ in encoded_arrays])\n",
    "            to_write_len_c = np.sum([arr.size for _, arr, _, _ in encoded_arrays])\n",
    "            to_write_len_v = np.sum([arr.size for _, _, arr, _ in encoded_arrays])\n",
    "            if to_write_len_s == 0 or to_write_len_c == 0 or to_write_len_v == 0:\n",
    "                continue\n",
    "            out_mm_semantic = np.memmap(\n",
    "                out_mm_semantic_filepath,\n",
    "                dtype=np.uint16,\n",
    "                mode=\"r+\",\n",
    "                shape=(n_offs_s + to_write_len_s,),\n",
    "            )\n",
    "            out_mm_codec = np.memmap(\n",
    "                out_mm_codec_filepath,\n",
    "                dtype=np.uint16,\n",
    "                mode=\"r+\",\n",
    "                shape=(n_offs_c + to_write_len_c,),\n",
    "            )\n",
    "            out_mm_vae = np.memmap(\n",
    "                out_mm_vae_filepath,\n",
    "                dtype=np.float32,\n",
    "                mode=\"r+\",\n",
    "                shape=(n_offs_v + to_write_len_v,),\n",
    "            )\n",
    "            for arr_s, arr_c, arr_v, arr_meta in encoded_arrays:\n",
    "                out_mm_semantic[n_offs_s : n_offs_s + arr_s.size] = arr_s.reshape(-1,)\n",
    "                out_mm_codec[n_offs_c : n_offs_c + arr_c.size] = arr_c.reshape(-1,)\n",
    "                out_mm_vae[n_offs_v : n_offs_v + arr_v.size] = arr_v.reshape(-1,)\n",
    "                n_offs_s += arr_s.size\n",
    "                n_offs_c += arr_c.size\n",
    "                n_offs_v += arr_v.size\n",
    "                dataset_str = dset_name\n",
    "                add_meta = {\n",
    "                    \"dataset\": dataset_str,\n",
    "                    \"id\": arr_meta[\"id\"],\n",
    "                    \"start_s\": round(arr_meta[\"start_s\"], 2),\n",
    "                    \"end_s\": round(arr_meta[\"end_s\"], 2),\n",
    "                    \"original_duration_s\": arr_meta[\"original_duration_s\"],\n",
    "                }\n",
    "                if \"original_id\" in arr_meta:\n",
    "                    add_meta[\"original_id\"] = arr_meta[\"original_id\"]\n",
    "                tot_duration_dict[dataset_str] += (\n",
    "                    arr_meta[\"end_s\"] - arr_meta[\"start_s\"]\n",
    "                )\n",
    "                add_metas.append(add_meta)\n",
    "            # write it once\n",
    "            out_mm_semantic.flush()\n",
    "            out_mm_codec.flush()\n",
    "            out_mm_vae.flush()\n",
    "            del out_mm_semantic, out_mm_codec, out_mm_vae\n",
    "        ## end Part B\n",
    "        write_jsonl(\n",
    "            add_metas,\n",
    "            os.path.join(out_metas_filepath),\n",
    "            do_append=bool(n_offs_s != 0),\n",
    "        )\n",
    "        del encoded_arrays_list\n",
    "    # TODO: this gc collect takes super long but maybe ok outside of loop. somehow needed sometimes\n",
    "    gc.collect()\n",
    "    for k, v in tot_duration_dict.items():\n",
    "        print(f\"{round(v / 60 / 60):,} hours of {k}\")\n",
    "    return n_offs_s, n_offs_c, n_offs_v\n",
    "\n",
    "\n",
    "def prep_data(\n",
    "    datasets,\n",
    "    is_val=False,\n",
    "    njobs=5,\n",
    "    chunksize=10,\n",
    "):\n",
    "    n_offs_s = 0\n",
    "    n_offs_c = 0\n",
    "    n_offs_v = 0\n",
    "    dset_type = \"val\" if is_val else \"tr\"\n",
    "    out_mm_semantic_filepath = os.path.join(OUT_DATA_DIR, f\"data_semantic_{dset_type}.bin\")\n",
    "    out_mm_codec_filepath = os.path.join(OUT_DATA_DIR, f\"data_codec_{dset_type}.bin\")\n",
    "    out_mm_vae_filepath = os.path.join(OUT_DATA_DIR, f\"data_vae_{dset_type}.bin\")\n",
    "    out_metas_filepath = os.path.join(OUT_DATA_DIR, f\"metas_{dset_type}.jsonl\")\n",
    "    out_mm_semantic = np.memmap(out_mm_semantic_filepath, dtype=np.uint16, mode=\"w+\", shape=(1,))\n",
    "    out_mm_codec = np.memmap(out_mm_codec_filepath, dtype=np.uint16, mode=\"w+\", shape=(1,))\n",
    "    out_mm_vae = np.memmap(out_mm_vae_filepath, dtype=np.float32, mode=\"w+\", shape=(1,))\n",
    "    with open(out_metas_filepath, \"w\") as f:\n",
    "        f.write(\"\")\n",
    "    print(\"start prepare data\")\n",
    "    for dataset in datasets:\n",
    "        n_offs_s, n_offs_c, n_offs_v = _prep_data(\n",
    "            dataset,\n",
    "            njobs=njobs,\n",
    "            chunksize=chunksize,\n",
    "            is_val=is_val,\n",
    "            n_offs_s=n_offs_s,\n",
    "            n_offs_c=n_offs_c,\n",
    "            n_offs_v=n_offs_v,\n",
    "        )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2da8574b",
   "metadata": {},
   "outputs": [],
   "source": [
    "# # TODO: not sure yet what's needed but vae is large so let's use small for now\n",
    "# NJOBS = 32\n",
    "# CHUNKSIZE = 32\n",
    "NJOBS = 8\n",
    "CHUNKSIZE = 8"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "aa8c2cda",
   "metadata": {},
   "outputs": [],
   "source": [
    "# TODO: continuous batching data?"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "d726c704",
   "metadata": {},
   "outputs": [],
   "source": [
    "# (start_idx, end_idx), n_archives_semantic, n_archives_codec, n_archives_vae\n",
    "datasets = [\n",
    "    (\"youtube_music\", \"v1\", (0, 1), 1, 1, 1),\n",
    "    (\"genius_hq\", \"v1\", (0, 1), 1, 1, 1),\n",
    "]\n",
    "prep_data(\n",
    "    datasets,\n",
    "    is_val=True,\n",
    "    njobs=NJOBS,\n",
    "    chunksize=CHUNKSIZE,\n",
    ")\n",
    "# 28 hours of youtube_music\n",
    "# 18 hours of genius_hq"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "d45b342e",
   "metadata": {},
   "outputs": [],
   "source": [
    "# (start_idx, end_idx), n_archives_semantic, n_archives_codec, n_archives_vae\n",
    "datasets = [\n",
    "    (\"youtube_music\", \"v1\", (1, 4204), 1, 1, 1),\n",
    "    (\"genius_hq\", \"v1\", (1, 4302), 1, 1, 1),\n",
    "]\n",
    "prep_data(\n",
    "    datasets,\n",
    "    is_val=False,\n",
    "    njobs=NJOBS,\n",
    "    chunksize=CHUNKSIZE,\n",
    ")\n",
    "# youtube_music: ~45m prep\n",
    "#     88,538 hours of youtube_music\n",
    "#     21,141 hours of youtube_music_lyrics\n",
    "#     18,351 hours of youtube_music_lyrics_foreign\n",
    "# genius_hq_lyrics: ~50m prep\n",
    "#     90,948 hours of genius_hq_lyrics\n",
    "#     49,259 hours of genius_hq_lyrics_foreign"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "afb827ee",
   "metadata": {},
   "outputs": [],
   "source": [
    "!ls -lah /app/suno/data/chirp_v4/vae"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "e5355c5d",
   "metadata": {},
   "outputs": [],
   "source": [
    "# verify\n",
    "mm_semantic = np.memmap(os.path.join(OUT_DATA_DIR, \"data_semantic_val.bin\"), dtype=np.uint16, mode=\"r\")\n",
    "mm_codec = np.memmap(os.path.join(OUT_DATA_DIR, \"data_codec_val.bin\"), dtype=np.uint16, mode=\"r\")\n",
    "mm_vae = np.memmap(os.path.join(OUT_DATA_DIR, \"data_vae_val.bin\"), dtype=np.float32, mode=\"r\")\n",
    "metas = read_jsonl(os.path.join(OUT_DATA_DIR,\"metas_val.jsonl\"))\n",
    "mm_semantic = mm_semantic.reshape(-1, N_TOKENS_MEMMAP, SEMANTIC_N_CODEBOOKS)\n",
    "mm_codec = mm_codec.reshape(-1, N_TOKENS_MEMMAP, CODEC_N_CODEBOOKS)\n",
    "mm_vae = mm_vae.reshape(-1, N_TOKENS_MEMMAP, VAE_DIM)\n",
    "assert(len(mm_semantic) == len(mm_codec) == len(mm_vae) == len(metas))\n",
    "assert(mm_semantic[:100,:,0].min() >= 0)\n",
    "assert(mm_semantic[:100,:,0].max() <= SEMANTIC_CODEBOOK_SIZE)\n",
    "assert(mm_codec[:100,:,1:].min() >= 0)\n",
    "assert(mm_codec[:100,:,1:].max() <= CODEC_CODEBOOK_SIZE)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "846797cd",
   "metadata": {},
   "outputs": [],
   "source": [
    "# randomly listen to some stuff\n",
    "from suno_utils.tasks.dac_2c_12cb import (\n",
    "    preload_models as preload_codec_models,\n",
    "    encode as codec_encode,\n",
    "    decode as codec_decode,\n",
    ")\n",
    "from suno_utils.tasks.dac_vae_peaq import (\n",
    "    preload_models as preload_vae_models,\n",
    "    encode as vae_encode,\n",
    "    decode as vae_decode,\n",
    ")\n",
    "_ = preload_codec_models(\"s3://suno-data/georg/models/codec/dac_2c_25x12.pt\")\n",
    "_ = preload_vae_models(\"s3://suno-data/georg/models/codec/dac_vae_128.pth\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "1ce69c73",
   "metadata": {},
   "outputs": [],
   "source": [
    "idx = random.choice(range(len(metas)))\n",
    "print(mm_semantic[idx].shape)\n",
    "print(mm_codec[idx].shape)\n",
    "print(mm_vae[idx].shape)\n",
    "print(metas[idx])\n",
    "codec_decode(mm_codec[idx][:25*60]).play()\n",
    "vae_decode(mm_vae[idx][:25*60]).play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "11715eb2",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "6a1db89d",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "894aaec6",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "88e3db5a",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "ea2c4a03",
   "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.12"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
