{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.audio import Audio\n",
    "import os\n",
    "from tqdm import tqdm\n",
    "import random\n",
    "\n",
    "\n",
    "import os\n",
    "import pandas as pd\n",
    "from suno_utils.utils.text import read_jsonl, write_jsonl\n",
    "from suno_utils.utils.s3 import read_from_s3\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "S3_AUDIO_DIR = \"s3://suno-data/datasets/harvest/karaoke_versions/stems/audio/\"\n",
    "\n",
    "raw_metas = read_from_s3(\n",
    "    \"s3://suno-data/datasets/harvest/karaoke_versions/stems/karaoke_versions.jsonl\",\n",
    "    read_f=read_jsonl,\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "track_descriptions = []\n",
    "for m in raw_metas:\n",
    "    for mm in m[\"tracks\"]:\n",
    "        track_descriptions.append(mm[\"description\"])\n",
    "s = pd.Series(track_descriptions).value_counts()\n",
    "print(f\"{s.shape[0]} types of stems\")\n",
    "print(f\"{s[s >= 1000].shape[0]} with >=1000\")\n",
    "print(f\"{s[s >= 100].shape[0]} with >=100\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### parse categories"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import json\n",
    "\n",
    "# Create a mapping dictionary from category to instrument list\n",
    "with open(\"instrument_categories.json\", \"r\") as f:\n",
    "    category_to_instruments = json.load(f)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "write_jsonl(category_to_instruments, \"instrument_categories.jsonl\")\n",
    "instrument_to_category = {}\n",
    "for category, instruments in category_to_instruments.items():\n",
    "    for instrument in instruments:\n",
    "        instrument_to_category[instrument] = category\n",
    "\n",
    "\n",
    "def get_category(instrument):\n",
    "    instrument = instrument.replace(\" \", \"_\")\n",
    "    return instrument_to_category.get(instrument, \"Other\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "metas = []\n",
    "for m in raw_metas:\n",
    "    if any([\"file_path\" not in mm for mm in m[\"tracks\"]]):\n",
    "        # only very few are missing this\n",
    "        continue\n",
    "    # excluding click (metronome) here\n",
    "    metas.append(\n",
    "        {\n",
    "            \"id\": m[\"song_id\"],\n",
    "            \"duration_s\": round((m[\"preview_end\"] - m[\"preview_start\"]) / 1000, 1),\n",
    "            \"stems\": [\n",
    "                {\n",
    "                    \"id\": f\"{m['song_id']}_{i}\",\n",
    "                    \"title\": mm[\"description\"],\n",
    "                    \"s3_filepath\": os.path.join(S3_AUDIO_DIR, mm[\"file_path\"]),\n",
    "                    \"category\": get_category(mm[\"description\"]),\n",
    "                }\n",
    "                for i, mm in enumerate(m[\"tracks\"])\n",
    "                if mm[\"description\"] != \"Click\"\n",
    "            ],\n",
    "        }\n",
    "    )\n",
    "assert len([m[\"id\"] for m in metas]) == len(set([m[\"id\"] for m in metas]))\n",
    "print(f\"{len(metas):,} tracks\")\n",
    "n_stems = 0\n",
    "tot_track_duration_s = 0\n",
    "tot_stem_duration_s = 0\n",
    "for m in metas:\n",
    "    n_stems += len(m[\"stems\"])\n",
    "    tot_track_duration_s += m[\"duration_s\"]\n",
    "    tot_stem_duration_s += m[\"duration_s\"] * len(m[\"stems\"])\n",
    "print(f\"{n_stems:,} stems\")\n",
    "print(f\"{round(tot_track_duration_s / 60 / 60):,} hours total tracks\")\n",
    "print(f\"{round(tot_stem_duration_s / 60 / 60):,} hours total stems\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "metas[0]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import random\n",
    "import uuid\n",
    "from tqdm import tqdm\n",
    "import numpy as np\n",
    "from collections import defaultdict\n",
    "\n",
    "\n",
    "def mix_stack(bundle_metas, verbose=False):\n",
    "    stem_metas = bundle_metas[\"stems\"]\n",
    "\n",
    "    stems_by_category = defaultdict(list)\n",
    "    for stem in stem_metas:\n",
    "        stems_by_category[stem[\"category\"]].append(stem)\n",
    "    for category in stems_by_category:\n",
    "        random.shuffle(stems_by_category[category])\n",
    "\n",
    "    category_order = list(stems_by_category.keys())\n",
    "    random.shuffle(category_order)\n",
    "\n",
    "    stem_metas = []\n",
    "    for category in category_order:\n",
    "        stem_metas.extend(stems_by_category[category])\n",
    "\n",
    "    stem_audios = [\n",
    "        Audio.from_s3(mm[\"s3_filepath\"], sample_rate=48_000, n_channels=2)\n",
    "        for mm in tqdm(stem_metas, desc=\"Loading stems\", disable=not verbose)\n",
    "    ]\n",
    "\n",
    "    # make prefix's\n",
    "    prefix_audios = []\n",
    "    prefix_metas = []\n",
    "    for i, stem_audio in tqdm(enumerate(stem_audios), desc=\"Mixing prefix\", disable=not verbose):\n",
    "        prefix_meta = {\n",
    "            \"id\": str(uuid.uuid4()),\n",
    "            \"stems\": [m[\"id\"] for m in stem_metas[: i + 1]],\n",
    "        }\n",
    "        prefix_metas.append(prefix_meta)\n",
    "        if len(prefix_audios) == 0:\n",
    "            prefix_audios.append(stem_audio)\n",
    "        else:\n",
    "            prefix_audios.append(Audio.sum((prefix_audios[-1], stem_audio)))\n",
    "    assert len(prefix_audios) == len(prefix_metas)\n",
    "    assert len(stem_audios) == len(stem_metas)\n",
    "\n",
    "    # make suffix's\n",
    "    suffix_audios = []\n",
    "    suffix_metas = []\n",
    "    reversed_stem_audios = list(reversed(stem_audios))\n",
    "    reversed_stem_metas = list(reversed(stem_metas))\n",
    "    for i, stem_audio in tqdm(\n",
    "        enumerate(reversed_stem_audios), desc=\"Mixing suffix\", disable=not verbose\n",
    "    ):\n",
    "        suffix_meta = {\n",
    "            \"id\": str(uuid.uuid4()),\n",
    "            \"stems\": [m[\"id\"] for m in reversed_stem_metas[: i + 1]],\n",
    "        }\n",
    "        suffix_metas.append(suffix_meta)\n",
    "        if len(suffix_audios) == 0:\n",
    "            suffix_audios.append(stem_audio)\n",
    "        else:\n",
    "            suffix_audios.append(Audio.sum((suffix_audios[-1], stem_audio)))\n",
    "    assert len(suffix_audios) == len(suffix_metas)\n",
    "    assert len(stem_audios) == len(stem_metas)\n",
    "\n",
    "    # make groups\n",
    "    group_audios = []\n",
    "    group_metas = []\n",
    "    groups_to_mix = [category for category in category_order if len(stems_by_category[category]) > 1]\n",
    "    for category in tqdm(groups_to_mix, desc=\"Mixing groups\", disable=not verbose):\n",
    "        group_meta = {\n",
    "            \"id\": str(uuid.uuid4()),\n",
    "            \"stems\": [m[\"id\"] for m in stem_metas if m[\"category\"] == category],\n",
    "            \"category\": category,\n",
    "        }\n",
    "        group_metas.append(group_meta)\n",
    "        group_audios.append(\n",
    "            Audio.sum(\n",
    "                [stem_audios[i] for i in range(len(stem_metas)) if stem_metas[i][\"category\"] == category]\n",
    "            )\n",
    "        )\n",
    "\n",
    "    # # compute energy for stems\n",
    "    # for i, stem_audio in enumerate(stem_audios):\n",
    "    #     energy = stem_audio.get_energy(bin_size_s=1).astype(np.int32).tolist()\n",
    "    #     stem_metas[i][\"energy\"] = energy\n",
    "\n",
    "    return (\n",
    "        stem_metas,\n",
    "        stem_audios,\n",
    "        prefix_metas,\n",
    "        prefix_audios,\n",
    "        suffix_metas,\n",
    "        suffix_audios,\n",
    "        group_metas,\n",
    "        group_audios,\n",
    "    )\n",
    "\n",
    "\n",
    "(\n",
    "    stem_metas,\n",
    "    stem_audios,\n",
    "    prefix_metas,\n",
    "    prefix_audios,\n",
    "    suffix_metas,\n",
    "    suffix_audios,\n",
    "    group_metas,\n",
    "    group_audios,\n",
    ") = mix_stack(metas[0], verbose=True)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "\n",
    "audio = stem_audios[0]\n",
    "audio.play()\n",
    "stem_metas\n",
    "group_metas"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# for i, (stem_meta, stem_audio, prefix_meta, prefix_audio, suffix_meta, suffix_audio) in enumerate(\n",
    "#     zip(stem_metas, stem_audios, prefix_metas, prefix_audios, suffix_metas, suffix_audios)\n",
    "# ):\n",
    "#     print(stem_meta)\n",
    "#     stem_audio.play()\n",
    "#     print(suffix_meta)\n",
    "#     suffix_audio.play()\n",
    "#     print(prefix_meta)\n",
    "#     prefix_audio.play()\n",
    "#     if i > 3:\n",
    "#         break"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import uuid\n",
    "import os\n",
    "import tempfile\n",
    "from suno_utils.utils.s3 import download_s3_files, upload_s3_files\n",
    "\n",
    "\n",
    "def create_and_upload(bundle_metas, verbose=False):\n",
    "    (\n",
    "        stem_metas,\n",
    "        stem_audios,\n",
    "        prefix_metas,\n",
    "        prefix_audios,\n",
    "        suffix_metas,\n",
    "        suffix_audios,\n",
    "        group_metas,\n",
    "        group_audios,\n",
    "    ) = mix_stack(bundle_metas, verbose=verbose)\n",
    "    bundle_id = bundle_metas[\"id\"]\n",
    "\n",
    "    local_fps = []\n",
    "    s3_filepaths = []\n",
    "    with tempfile.TemporaryDirectory() as tempdir:\n",
    "        for prefix_meta, prefix_audio in zip(prefix_metas, prefix_audios):\n",
    "            fp = os.path.join(tempdir, f\"{prefix_meta['id']}.opus\")\n",
    "            prefix_audio.to_opus(fp)\n",
    "            local_fps.append(fp)\n",
    "            s3_path = f\"s3://suno-data/datasets/harvest/karaoke_versions/stems_v2/prefix_mix/{bundle_id}/{prefix_meta['id']}.opus\"\n",
    "            s3_filepaths.append(s3_path)\n",
    "            prefix_meta[\"s3_filepath\"] = s3_path\n",
    "\n",
    "        for suffix_meta, suffix_audio in zip(suffix_metas, suffix_audios):\n",
    "            fp = os.path.join(tempdir, f\"{suffix_meta['id']}.opus\")\n",
    "            suffix_audio.to_opus(fp)\n",
    "            local_fps.append(fp)\n",
    "            s3_path = f\"s3://suno-data/datasets/harvest/karaoke_versions/stems_v2/suffix_mix/{bundle_id}/{suffix_meta['id']}.opus\"\n",
    "            s3_filepaths.append(s3_path)\n",
    "            suffix_meta[\"s3_filepath\"] = s3_path\n",
    "\n",
    "        for group_meta, group_audio in zip(group_metas, group_audios):\n",
    "            fp = os.path.join(tempdir, f\"{group_meta['id']}.opus\")\n",
    "            group_audio.to_opus(fp)\n",
    "            local_fps.append(fp)\n",
    "            s3_path = f\"s3://suno-data/datasets/harvest/karaoke_versions/stems_v2/group_mix/{bundle_id}/{group_meta['id']}.opus\"\n",
    "            s3_filepaths.append(s3_path)\n",
    "            group_meta[\"s3_filepath\"] = s3_path\n",
    "\n",
    "        upload_s3_files(\n",
    "            local_fps,\n",
    "            s3_filepaths,\n",
    "            chunksize=1000,\n",
    "            n_cores=20,\n",
    "            joblib_backend=\"threads\",\n",
    "            silent=True,\n",
    "        )\n",
    "\n",
    "    meta = {\n",
    "        \"id\": str(uuid.uuid4()),\n",
    "        \"duration_s\": bundle_metas[\"duration_s\"],\n",
    "        \"stems\": stem_metas,\n",
    "        \"prefix\": prefix_metas,\n",
    "        \"suffix\": suffix_metas,\n",
    "        \"group\": group_metas,\n",
    "    }\n",
    "    return meta\n",
    "\n",
    "\n",
    "def safe_create_and_upload(bundle_metas):\n",
    "    try:\n",
    "        return create_and_upload(bundle_metas)\n",
    "    except Exception as e:\n",
    "        print(e)\n",
    "        return None"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "meta = safe_create_and_upload(metas[0])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "meta"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# check\n",
    "\n",
    "# for i, prefix_meta in enumerate(meta[\"prefix\"]):\n",
    "#     print(prefix_meta)\n",
    "#     prefix = Audio.from_s3(prefix_meta[\"s3_filepath\"], sample_rate=48000, n_channels=2)\n",
    "#     prefix.play()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# break"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import multiprocessing as mp\n",
    "\n",
    "# Use multiprocessing to speed up processing\n",
    "num_cores = 80  # Leave one core free\n",
    "N = 40_000\n",
    "\n",
    "\n",
    "# Generate samples without using multiprocessing\n",
    "\n",
    "with mp.Pool(num_cores) as pool:\n",
    "    results = list(tqdm(pool.imap(safe_create_and_upload, metas), total=len(metas)))\n",
    "\n",
    "# Filter out None results and flatten the list\n",
    "metas = [meta for result in results if result is not None]\n",
    "\n",
    "metas[0]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import json\n",
    "\n",
    "# save metas to file\n",
    "with open(\"/tmp/grouped_metas.jsonl\", \"w\") as f:\n",
    "    for meta in results:\n",
    "        f.write(json.dumps(meta) + \"\\n\")\n",
    "\n",
    "# upload file to s3\n",
    "!aws s3 cp /tmp/grouped_metas.jsonl s3://suno-data/datasets/harvest/karaoke_versions/grouped_metas.jsonl\n",
    "\n",
    "results[0]\n"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Prep bundle"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.utils.s3 import download_s3_file_if_needed\n",
    "\n",
    "metas_file = download_s3_file_if_needed(\n",
    "    \"s3://suno-data/datasets/harvest/karaoke_versions/grouped_metas.jsonl\"\n",
    ")\n",
    "\n",
    "import polars as pl\n",
    "\n",
    "metas = pl.read_ndjson(metas_file).to_dicts()\n",
    "metas[0]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.utils.text import get_filename\n",
    "from tqdm import tqdm\n",
    "\n",
    "flat_metas = []\n",
    "for m in tqdm(metas, desc=\"Expanding metas\"):\n",
    "    if m[\"id\"] is None:\n",
    "        continue\n",
    "    for mm in m[\"stems\"]:\n",
    "        flat_metas.append(\n",
    "            {\n",
    "                \"id\": f\"{mm['id']}\",\n",
    "                \"bundle_id\": str(m[\"id\"]),\n",
    "                \"s3_filepath\": mm[\"s3_filepath\"],\n",
    "                \"duration_s\": m[\"duration_s\"],\n",
    "                \"stem_type\": mm.get(\"title\", \"\"),\n",
    "                \"stem_tags\": mm.get(\"tags\", []),\n",
    "                \"category\": mm.get(\"category\", \"\"),\n",
    "                # \"energy\": mm.get(\"energy\", []),\n",
    "                \"type\": \"stem\",\n",
    "            }\n",
    "        )\n",
    "    for mm in m[\"prefix\"]:\n",
    "        flat_metas.append(\n",
    "            {\n",
    "                \"id\": f\"{mm['id']}\",\n",
    "                \"bundle_id\": str(m[\"id\"]),\n",
    "                \"s3_filepath\": mm[\"s3_filepath\"],\n",
    "                \"duration_s\": m[\"duration_s\"],\n",
    "                \"stem_type\": mm.get(\"title\", \"\"),\n",
    "                \"stem_tags\": mm.get(\"tags\", []),\n",
    "                \"stem_ids\": mm.get(\"stems\", []),\n",
    "                \"type\": \"prefix\",\n",
    "            }\n",
    "        )\n",
    "    for mm in m[\"suffix\"]:\n",
    "        flat_metas.append(\n",
    "            {\n",
    "                \"id\": f\"{mm['id']}\",\n",
    "                \"bundle_id\": str(m[\"id\"]),\n",
    "                \"s3_filepath\": mm[\"s3_filepath\"],\n",
    "                \"duration_s\": m[\"duration_s\"],\n",
    "                \"stem_type\": mm.get(\"title\", \"\"),\n",
    "                \"stem_tags\": mm.get(\"tags\", []),\n",
    "                \"stem_ids\": mm.get(\"stems\", []),\n",
    "                \"type\": \"suffix\",\n",
    "            }\n",
    "        )\n",
    "    for mm in m[\"group\"]:\n",
    "        flat_metas.append(\n",
    "            {\n",
    "                \"id\": f\"{mm['id']}\",\n",
    "                \"bundle_id\": str(m[\"id\"]),\n",
    "                \"s3_filepath\": mm[\"s3_filepath\"],\n",
    "                \"duration_s\": m[\"duration_s\"],\n",
    "                \"stem_type\": mm.get(\"title\", \"\"),\n",
    "                \"stem_tags\": mm.get(\"tags\", []),\n",
    "                \"stem_ids\": mm.get(\"stems\", []),\n",
    "                \"category\": mm.get(\"category\", \"\"),\n",
    "                \"type\": \"group\",\n",
    "            }\n",
    "        )\n",
    "\n",
    "\n",
    "print(f\"{len([m for m in flat_metas if m['type'] == 'prefix']):,} prefix metas\")\n",
    "print(f\"{len([m for m in flat_metas if m['type'] == 'suffix']):,} suffix metas\")\n",
    "print(f\"{len([m for m in flat_metas if m['type'] == 'stem']):,} stem metas\")\n",
    "print(f\"{len([m for m in flat_metas if m['type'] == 'group']):,} group metas\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import random\n",
    "\n",
    "random.sample(flat_metas, 10)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.utils.text import write_jsonl\n",
    "\n",
    "write_jsonl(flat_metas, \"/tmp/karaoke_stems_grouped_flat_metas.jsonl\")\n",
    "!aws s3 cp /tmp/karaoke_stems_grouped_flat_metas.jsonl s3://suno-data/datasets/bundles/v4/karaoke_stems_grouped/metas.jsonl"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# modal run /home/victor/glockenspiel/suno_utils/suno_utils/scripts/gpt/modal_encode.py \\\n",
    "#     --embed-type='dac_vae_tuned_25hz' \\\n",
    "#     --base-s3-dir='s3://suno-data/datasets/bundles/v4/karaoke_stems_grouped/' \\\n",
    "#     --chunksize=10 \\\n",
    "#     --min-duration-s=5 \\\n",
    "#     --max-duration-s=480 \\\n",
    "#     --output-name='dac_vae_tuned_25hz' \\\n",
    "#     --normalize-volume=False \\\n",
    "#     --first-only=True \\\n",
    "#     --force-overwrite=True"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## check embeds"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.tasks.dac_vae_fixed_25hz import (\n",
    "    preload_models as preload_codec_models,\n",
    "    decode as codec_decode,\n",
    "    encode as codec_encode,\n",
    "    Audio,\n",
    ")\n",
    "\n",
    "_ = preload_codec_models(\"s3://suno-data/minz/models/dac_vae_tuned_25hz.pth\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from sunodata.dataset_maker_utils import Bundle, DatasetConfig, MemmapMaker\n",
    "\n",
    "bundle = Bundle(name=\"karaoke_stems_grouped\")\n",
    "npz = bundle.get_part(0)\n",
    "npz"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "part_metas = bundle.get_part_metas(0)\n",
    "Audio.from_s3(part_metas[0][\"s3_filepath\"]).play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(list(npz.keys()))\n",
    "codec_decode(npz[\"9980_7\"]).play()\n",
    "\n",
    "bundle.get_part_metas(0)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 4
}
