{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from sunodata.dataset_maker_utils import Bundle, DatasetConfig, MemmapMaker\n",
    "from sunodata.datasets import KaraokeStemsDataset, WetDryDataset, RandomMixDataset\n",
    "from tqdm import tqdm\n"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## make memmap"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "OUT_DATA_DIR = \"/app/suno/data/diffusion_mix/dac_vae_tuned_25hz_stems_grouped\""
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Make karaoke dataset"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "bundle = Bundle(name=\"karaoke_stems_grouped\")\n",
    "bundle.get_part(0)\n",
    "bundle"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "stem_metas = bundle.get_metas(n_lines=None)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "stem_metas[0]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# randomly come up with offsets to get 30s slices and avoid fadein fadeout if possible\n",
    "import copy\n",
    "import random\n",
    "import numpy as np\n",
    "import pandas as pd\n",
    "\n",
    "random.seed(1234)\n",
    "\n",
    "print(round(np.mean([m[\"duration_s\"] >= 30 for m in stem_metas]), 3), \"fraction >= 30s\")\n",
    "offset_s_map = {}\n",
    "for m in tqdm(stem_metas):\n",
    "    if m[\"bundle_id\"] not in offset_s_map:\n",
    "        if m[\"duration_s\"] >= 40:\n",
    "            offset_s_map[m[\"bundle_id\"]] = 5\n",
    "        elif m[\"duration_s\"] <= 30.1:\n",
    "            continue\n",
    "        else:\n",
    "            offset_s_map[m[\"bundle_id\"]] = 0\n",
    "stem_types = []\n",
    "for m in stem_metas:\n",
    "    if \"stem_type\" in m:\n",
    "        stem_types.append(m[\"stem_type\"])\n",
    "\n",
    "\n",
    "s = pd.Series(stem_types).value_counts()\n",
    "stem_type_map = {ss: n for n, ss in enumerate(s[s >= 0].index.tolist())}\n",
    "print(f\"{len(stem_type_map)} stem types in vocab\")\n",
    "retained_stem_metas = []\n",
    "for m in tqdm(stem_metas):\n",
    "    if m[\"bundle_id\"] not in offset_s_map:\n",
    "        continue\n",
    "    if \"stem_type\" in m and m[\"stem_type\"] not in stem_type_map:\n",
    "        continue\n",
    "    new_m = copy.deepcopy(m)\n",
    "    new_m[\"offset_s\"] = offset_s_map[m[\"bundle_id\"]]\n",
    "    if \"stem_type\" in m:\n",
    "        new_m[\"stem_type_id\"] = stem_type_map[m[\"stem_type\"]]\n",
    "    retained_stem_metas.append(new_m)\n",
    "\n",
    "meta_info_map = {m[\"id\"]: m for m in retained_stem_metas}\n",
    "print(len(meta_info_map))\n",
    "# 0.958 fraction >= 30s\n",
    "# 128 stem types in vocab"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def count_missing_stems(metas):\n",
    "    ids = set([m[\"id\"] for m in metas])\n",
    "    bundle_n_stems = {m[\"bundle_id\"]: 0 for m in metas}\n",
    "    for m in metas:\n",
    "        if m.get(\"type\") == \"stem\":\n",
    "            bundle_n_stems[m[\"bundle_id\"]] += 1\n",
    "    num_missing = 0\n",
    "    for m in metas:\n",
    "        if m.get(\"type\") in [\"prefix\", \"suffix\"]:\n",
    "            stem_ids = m[\"stem_ids\"]\n",
    "            assert len(stem_ids) <= bundle_n_stems[m[\"bundle_id\"]], (\n",
    "                f\"{m['bundle_id']} {len(stem_ids)} {bundle_n_stems[m['bundle_id']]}\"\n",
    "            )\n",
    "            for stem_id in stem_ids:\n",
    "                if stem_id not in ids:\n",
    "                    num_missing += 1\n",
    "                    break\n",
    "    return num_missing\n",
    "\n",
    "\n",
    "print(count_missing_stems(stem_metas))\n",
    "print(count_missing_stems(retained_stem_metas))\n",
    "\n",
    "# check"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "meta_info_map[\"9980_6\"]"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Make splice wetdry dataset"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "wetdry_bundle = Bundle(name=\"splice_wetdry\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "wetdry_metas = wetdry_bundle.get_metas(n_lines=None)\n",
    "wetdry_metas[0]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "wetdry_meta_info_map = {m[\"id\"]: m for m in wetdry_metas}"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Make memmap"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "NJOBS = 32\n",
    "CHUNKSIZE = 4\n",
    "\n",
    "# note that this can split the data at an awkward place\n",
    "# stems may be split across val and train\n",
    "VAL_CHUNKS = 10\n",
    "\n",
    "karaoke_dataset = KaraokeStemsDataset(\n",
    "    bundle=bundle,\n",
    "    start_idx=0,\n",
    "    end_idx=VAL_CHUNKS,\n",
    "    n_vae=1,\n",
    "    sort_key=\"bundle_id\",\n",
    ")\n",
    "wetdry_dataset = WetDryDataset(\n",
    "    bundle=wetdry_bundle,\n",
    "    start_idx=0,\n",
    "    end_idx=VAL_CHUNKS,\n",
    "    n_vae=1,\n",
    "    # sort_key=\"bundle_id\",\n",
    ")\n",
    "\n",
    "\n",
    "memmap_maker = MemmapMaker(out_data_dir=OUT_DATA_DIR)\n",
    "memmap_maker.prep_data(\n",
    "    [\n",
    "        (karaoke_dataset, meta_info_map),\n",
    "        (wetdry_dataset, wetdry_meta_info_map),\n",
    "        # (random_mix_dataset, random_mix_meta_info_map),\n",
    "    ],\n",
    "    is_val=True,\n",
    "    njobs=NJOBS,\n",
    "    chunksize=CHUNKSIZE,\n",
    ")\n",
    "# 18 hours of karaoke_stems"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "!ls \"/app/suno/data/diffusion_mix/test\"\n",
    "!head \"/app/suno/data/diffusion_mix/test/metas_val.jsonl\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "n_karaoke_parts = bundle.num_parts()\n",
    "n_wetdry_parts = wetdry_bundle.num_parts()\n",
    "\n",
    "print(n_karaoke_parts, n_wetdry_parts)\n",
    "# n_karaoke_parts = VAL_CHUNKS + 10\n",
    "# n_wetdry_parts = VAL_CHUNKS + 10\n",
    "\n",
    "karaoke_dataset = KaraokeStemsDataset(\n",
    "    bundle=bundle,\n",
    "    start_idx=VAL_CHUNKS,\n",
    "    end_idx=n_karaoke_parts,\n",
    "    n_vae=1,\n",
    "    sort_key=\"bundle_id\",\n",
    ")\n",
    "wetdry_dataset = WetDryDataset(\n",
    "    bundle=wetdry_bundle,\n",
    "    start_idx=VAL_CHUNKS,\n",
    "    end_idx=n_wetdry_parts,\n",
    "    n_vae=1,\n",
    "    # sort_key=\"bundle_id\",\n",
    ")\n",
    "\n",
    "\n",
    "memmap_maker = MemmapMaker(out_data_dir=OUT_DATA_DIR)\n",
    "memmap_maker.prep_data(\n",
    "    [\n",
    "        (karaoke_dataset, meta_info_map),\n",
    "        (wetdry_dataset, wetdry_meta_info_map),\n",
    "        # (random_mix_dataset, random_mix_meta_info_map),\n",
    "    ],\n",
    "    is_val=False,\n",
    "    njobs=NJOBS,\n",
    "    chunksize=10,\n",
    ")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Check"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "!ls /app/suno/data/diffusion_mix/dac_vae_tuned_25hz_stems_pre_suf/\n",
    "!wc -l /app/suno/data/diffusion_mix/dac_vae_tuned_25hz_stems_pre_suf/metas_val.jsonl\n",
    "!wc -l /app/suno/data/diffusion_mix/dac_vae_tuned_25hz_stems_pre_suf/metas_tr.jsonl\n",
    "!head /app/suno/data/diffusion_mix/dac_vae_tuned_25hz_stems_pre_suf/metas_val.jsonl\n",
    "!tail /app/suno/data/diffusion_mix/dac_vae_tuned_25hz_stems_pre_suf/metas_val.jsonl"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.tasks.dac_vae_fixed_25hz import decode, preload_models, encode\n",
    "\n",
    "preload_models(\"s3://suno-data/minz/models/dac_vae_tuned_25hz.pth\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "from suno_utils.utils.text import read_jsonl\n",
    "\n",
    "n_stem_types = 128\n",
    "vae_memmap_filepath = \"/app/suno/data/diffusion_mix/dac_vae_tuned_25hz_stems_pre_suf/data_vae_val.bin\"\n",
    "metas_filepath = \"/app/suno/data/diffusion_mix/dac_vae_tuned_25hz_stems_pre_suf/metas_val.jsonl\"\n",
    "\n",
    "vae_data = np.memmap(\n",
    "    vae_memmap_filepath,\n",
    "    dtype=np.float16,\n",
    "    mode=\"r\",\n",
    ")\n",
    "vae_data = vae_data.reshape(-1, 750, 128)\n",
    "\n",
    "metas = read_jsonl(metas_filepath)\n",
    "assert len(metas) == vae_data.shape[0]\n",
    "metas_dict = {m.get(\"uuid\"): m for m in metas}\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.audio import Audio\n",
    "\n",
    "# find a pair of wetdry\n",
    "a = metas[-2]\n",
    "print(a)\n",
    "b = metas_dict[a[\"dry_uuid\"]]\n",
    "\n",
    "print(a[\"s3_filepath\"])\n",
    "print(b[\"s3_filepath\"])\n",
    "\n",
    "Audio.from_s3(a[\"s3_filepath\"]).play()\n",
    "Audio.from_s3(b[\"s3_filepath\"]).play()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.audio import Audio\n",
    "\n",
    "Audio.from_s3(\n",
    "    \"s3://suno-data/datasets/harvest/splice/wetdry/52f46c3996aea22549374f03577b3179fd63fcbd39532fea774ef7dc232fe192.mp3\"\n",
    ").play()\n",
    "decode(vae_data[-1]).play()\n",
    "decode(vae_data[-2]).play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "from suno_utils.utils.text import read_jsonl\n",
    "\n",
    "n_stem_types = 128\n",
    "vae_memmap_filepath = \"/app/suno/data/diffusion_mix/dac_vae_tuned_25hz_stems_pre_suf/data_vae_tr.bin\"\n",
    "metas_filepath = \"/app/suno/data/diffusion_mix/dac_vae_tuned_25hz_stems_pre_suf/metas_tr.jsonl\"\n",
    "\n",
    "metas = read_jsonl(metas_filepath)\n",
    "print(len(metas))\n",
    "vae_data = np.memmap(\n",
    "    vae_memmap_filepath,\n",
    "    dtype=np.float16,\n",
    "    mode=\"r\",\n",
    ")\n",
    "vae_data = vae_data.reshape(-1, 750, 128)\n",
    "\n",
    "assert len(metas) == vae_data.shape[0], f\"{len(metas)} {vae_data.shape[0]}\"\n",
    "metas_dict = {m.get(\"uuid\"): m for m in metas}\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# stem prep\n",
    "from collections import defaultdict\n",
    "\n",
    "metas_filepath = \"/app/suno/data/diffusion_mix/dac_vae_tuned_25hz_stems_pre_suf/metas_tr.jsonl\"\n",
    "metas = read_jsonl(metas_filepath)\n",
    "\n",
    "metas_by_id = {m[\"id\"]: m for m in metas}\n",
    "metas_stems = [m for m in metas if m.get(\"type\") == \"stem\"]\n",
    "metas_prefixes = [m for m in metas if m.get(\"type\") == \"prefix\"]\n",
    "metas_suffixes = [m for m in metas if m.get(\"type\") == \"suffix\"]\n",
    "\n",
    "bundle_to_stems = defaultdict(list)\n",
    "for stem in metas_stems:\n",
    "    bundle_to_stems[stem[\"bundle_id\"]].append(stem)\n",
    "bundle_to_prefixes = defaultdict(list)\n",
    "for prefix in metas_prefixes:\n",
    "    bundle_to_prefixes[prefix[\"bundle_id\"]].append(prefix)\n",
    "bundle_to_suffixes = defaultdict(list)\n",
    "for suffix in metas_suffixes:\n",
    "    bundle_to_suffixes[suffix[\"bundle_id\"]].append(suffix)\n",
    "\n",
    "print(f\"len(metas_prefixes): {len(metas_prefixes)}\")\n",
    "print(f\"len(metas_suffixes): {len(metas_suffixes)}\")\n",
    "print(\"len(metas_stems):\", len(metas_stems))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "num_missing = 0\n",
    "for prefix in metas_prefixes:\n",
    "    # check all stems are in stems\n",
    "    stem_ids = prefix[\"stem_ids\"]\n",
    "    for stem_id in stem_ids:\n",
    "        if stem_id not in metas_by_id:\n",
    "            print(f\"stem_id: {stem_id} not in metas_by_id\")\n",
    "            num_missing += 1\n",
    "            break\n",
    "print(f\"num_missing: {num_missing}\")\n",
    "metas_by_id[\"karaoke_stems_pre_suf__10105_6\"]\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 4
}
