{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# reload module\n",
    "%load_ext autoreload\n",
    "%autoreload 2\n",
    "\n",
    "import os\n",
    "import torch\n",
    "import funcy\n",
    "import IPython\n",
    "import numpy as np\n",
    "\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"4\"\n",
    "\n",
    "from suno_utils.utils.text import (    \n",
    "    write_jsonl,\n",
    "    read_jsonl,\n",
    "    write_json,\n",
    "    read_json,\n",
    "    normalize_whitespace,\n",
    ")\n",
    "\n",
    "from dac.model.dac4 import DAC\n",
    "from suno_utils.utils.s3 import read_from_s3"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "ytm_alignments = read_jsonl(\"/home/tony/Work/tony/hoot/tmp/ytm_hq_alignments_t30_v1.jsonl\", progress=False)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "genius_alignments = read_jsonl(\"/home/tony/Work/tony/hoot/tmp/genius_hq_alignments_t30_v1.jsonl\", progress=False)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# check if any ids are in both ytm and genius\n",
    "ytm_ids = set([a[0] for a in ytm_alignments])\n",
    "genius_ids = set([a[0] for a in genius_alignments])\n",
    "print(len(ytm_ids), len(genius_ids))\n",
    "print(len(ytm_ids.intersection(genius_ids)))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(len(ytm_ids) + len(genius_ids))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "ytm_alignments_map = {a[0]: a[1] for a in ytm_alignments}\n",
    "genius_alignments_map = {a[0]: a[1] for a in genius_alignments}\n",
    "# merge alignments into single map\n",
    "alignments_map = {**ytm_alignments_map, **genius_alignments_map}\n",
    "print(len(alignments_map))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "#vae_name = \"vae_100hz_30s\"\n",
    "vae_name = \"vae_25hz_64_30s\"\n",
    "subset = \"tr\"\n",
    "base_metas = read_jsonl(f\"/app/suno/data/diffusion_mix/{vae_name}/metas_{subset}.jsonl\")\n",
    "#base_metas = read_jsonl(\"/app/suno/data/diffusion_mix/metadata/metas.jsonl\") # main meats"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from tqdm.auto import tqdm"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# add context to metas\n",
    "\n",
    "new_metas = []\n",
    "\n",
    "for idx, meta in tqdm(enumerate(base_metas)):\n",
    "    new_meta = meta.copy()\n",
    "    meta_id = meta[\"id\"]\n",
    "    # check start_s\n",
    "    start_s = meta[\"start_s\"]\n",
    "    # this is the first block so we don't have any context\n",
    "    new_meta[\"prev_context_id\"] = []\n",
    "\n",
    "    if float(start_s) != 0:\n",
    "        # look at previous indicies from the current index and find the first index where start_s is 0\n",
    "        for i in range(idx - 1, -1, -1):\n",
    "            # check if the id is the same\n",
    "            if base_metas[i][\"id\"] == meta_id:\n",
    "                new_meta[\"prev_context_id\"].append(i)\n",
    "            else:\n",
    "                break\n",
    "    new_metas.append(new_meta)\n",
    "\n",
    "print(len(new_metas), len(base_metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "base_metas = new_metas"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# merge alignments into metas\n",
    "new_metas = []\n",
    "alignment_count = 0\n",
    "\n",
    "for base_meta_idx, base_meta in enumerate(base_metas):\n",
    "    if base_meta[\"id\"] in alignments_map:\n",
    "        alignments = alignments_map[base_meta[\"id\"]]\n",
    "        start_s = base_meta[\"start_s\"]\n",
    "        end_s = base_meta[\"end_s\"]\n",
    "        found_alignment = False\n",
    "        for alignment in alignments:\n",
    "            if start_s == alignment[\"start_s\"] and end_s == alignment[\"end_s\"]:     \n",
    "                base_meta[\"text_aligned\"] = alignment[\"text\"]\n",
    "                found_alignment = True\n",
    "                alignment_count += 1\n",
    "                new_metas.append(base_meta)\n",
    "        if not found_alignment:\n",
    "            new_metas.append(base_meta)\n",
    "    else:\n",
    "        new_metas.append(base_meta)\n",
    "\n",
    "print(len(new_metas))\n",
    "print(len(base_metas))\n",
    "print(alignment_count, alignment_count / len(base_metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(new_metas[110][\"text_aligned\"])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# save new metas file\n",
    "write_jsonl(new_metas, f\"/app/suno/data/diffusion_mix/{vae_name}/metas_context_aligned_{subset}.jsonl\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# load VAE\n",
    "device = \"cuda:0\"\n",
    "#device = \"cpu\"\n",
    "# checkpoint_filepath = \"/app/suno/christian/checkpoints/dac/100hz_vae_peaq_kl_0.005/best/dac/weights.pth\"\n",
    "if vae_name == \"vae_100hz_30s\":\n",
    "    checkpoint_filepath = \"s3://suno-data/christian/100hz_vae_peaq_kl_0.005.pth\"\n",
    "elif vae_name == \"vae_25hz_30s\":\n",
    "    checkpoint_filepath = \"s3://suno-data/christian/25hz_vae_peaq_kl_0.005.pth\"\n",
    "elif vae_name == \"vae_25hz_64_30s\":\n",
    "    checkpoint_filepath = \"s3://suno-data/christian/25hz_vae_peaq_64_kl_0.005.pth\"\n",
    "else:\n",
    "    raise ValueError(f\"Unknown vae name: {vae_name}\")\n",
    "\n",
    "load_f = funcy.partial(torch.load, map_location=\"cpu\")\n",
    "\n",
    "if checkpoint_filepath.startswith(\"s3://\"):\n",
    "    sd = read_from_s3(checkpoint_filepath, read_f=load_f)\n",
    "else:\n",
    "    sd = load_f(checkpoint_filepath)\n",
    "\n",
    "sd[\"metadata\"][\"kwargs\"] = {\n",
    "    k: v\n",
    "    for k, v in sd[\"metadata\"][\"kwargs\"].items()\n",
    "    if k in DAC.__init__.__code__.co_varnames\n",
    "}\n",
    "vae_model = DAC(**sd[\"metadata\"][\"kwargs\"])\n",
    "vae_model.load_state_dict(sd[\"state_dict\"])\n",
    "vae_model.eval()\n",
    "vae_model.to(device)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "VAE_DIM = 64\n",
    "\n",
    "if vae_name == \"vae_100hz_30s\":\n",
    "    VAE_N_MEMMAP_TOKENS = 3000\n",
    "elif vae_name == \"vae_25hz_30s\" or vae_name == \"vae_25hz_64_30s\":\n",
    "    VAE_N_MEMMAP_TOKENS = 750\n",
    "else:\n",
    "    raise ValueError(f\"Unknown vae name: {vae_name}\")\n",
    "\n",
    "SEMANTIC_N_CODEBOOKS = 1\n",
    "SEMANTIC_N_MEMMAP_TOKENS = 750\n",
    "\n",
    "base_dir = f\"/app/suno/data/diffusion_mix/{vae_name}\"\n",
    "#base_dir = \"/mnt/localdisk/cjs_shards/\"\n",
    "read_new_metas = read_jsonl(f\"/app/suno/data/diffusion_mix/{vae_name}/metas_context_aligned_{subset}.jsonl\")\n",
    "vae_memmap_filepath = f\"{base_dir}/data_vae_{subset}.bin\"\n",
    "semantic_memmap_filepath = f\"{base_dir}/data_semantic_{subset}.bin\"\n",
    "\n",
    "# load memmaps\n",
    "vae_memmap = np.memmap(vae_memmap_filepath, dtype=np.float16, mode=\"r\")\n",
    "semantic_memmap = np.memmap(semantic_memmap_filepath, dtype=np.uint16, mode=\"r\")\n",
    "\n",
    "# reshape memmaps\n",
    "vae_data = vae_memmap.reshape(-1, VAE_N_MEMMAP_TOKENS, VAE_DIM)\n",
    "semantic_data = semantic_memmap.reshape(-1, SEMANTIC_N_MEMMAP_TOKENS, SEMANTIC_N_CODEBOOKS)\n",
    "\n",
    "print(vae_data.shape, semantic_data.shape, len(read_new_metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from tqdm.auto import tqdm\n",
    "counter = 0\n",
    "for meta in tqdm(read_new_metas):\n",
    "    aligned_text = meta.get(\"text_aligned\", None)\n",
    "    if aligned_text is not None:\n",
    "        #meta[\"text\"] = aligned_text\n",
    "        counter += 1\n",
    "print(counter)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "read_new_metas = read_jsonl(f\"/app/suno/data/diffusion_mix/{vae_name}/metas_context_aligned_{subset}.jsonl\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "rand_idx = np.random.randint(0, len(read_new_metas))\n",
    "rand_idx = 115\n",
    "print(rand_idx)\n",
    "if \"text\" in read_new_metas[rand_idx]:\n",
    "    print(read_new_metas[rand_idx][\"text\"])\n",
    "    print(\"-\"*100)\n",
    "    print(read_new_metas[rand_idx][\"text_aligned\"])    \n",
    "    print(read_new_metas[rand_idx][\"start_s\"])\n",
    "    print(read_new_metas[rand_idx][\"end_s\"])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "vae_seq = vae_data[rand_idx]\n",
    "print(vae_seq.shape)\n",
    "# crop out vae pad tokens before decode\n",
    "vae_seq = vae_seq[:base_metas[rand_idx][\"n_vae_tokens\"]]\n",
    "print(vae_seq.shape)\n",
    "vae_seq = torch.from_numpy(vae_seq.copy()).to(device).unsqueeze(0).float()\n",
    "\n",
    "with torch.no_grad():\n",
    "    audio = vae_model.decode(vae_seq.permute(0, 2, 1))[0].detach().cpu()         \n",
    "    audio /= audio.abs().max().clamp(1e-8)\n",
    "    print(audio.mean())\n",
    "\n",
    "IPython.display.display(IPython.display.Audio(audio.numpy(), rate=48000))"
   ]
  },
  {
   "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
}
