{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 2,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "import boto3\n",
    "import numpy as np\n",
    "from tqdm import tqdm\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": [
    "# load the base metas\n",
    "VAE_DIM = 128\n",
    "VAE_RATE_HZ = 100\n",
    "VAE_MEMMAP_SIZE = 3000\n",
    "SEMANTIC_VOCAB_SIZE = 4000\n",
    "SEMANTIC_MEMMAP_SIZE = 750\n",
    "\n",
    "OUT_DATA_DIR = \"/app/suno/data/diffusion_ft/upsample_100hz_v4_t_5_20241018_pairs_text_cfg/\"\n",
    "os.makedirs(OUT_DATA_DIR, exist_ok=True)\n",
    "\n",
    "# christian/data/upsample_100hz_v4_t_5_20241018\n",
    "\n",
    "bucket_name = \"suno-data\"\n",
    "base_dir = \"christian/data/upsample_100z_v1\"\n",
    "output_name = \"v2\"\n",
    "base_metas_path = os.path.join(\"s3://\", bucket_name, base_dir, \"metas.jsonl\")\n",
    "base_metas = read_from_s3(base_metas_path, read_f=read_jsonl)\n",
    "print(len(base_metas))\n",
    "\n",
    "# s3 client\n",
    "s3 = boto3.client('s3')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 73,
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": 8,
   "metadata": {},
   "outputs": [],
   "source": [
    "def get_s3_files(bucket_name, prefix, max_keys:int=100000):\n",
    "    all_files = []\n",
    "    continuation_token = None\n",
    "\n",
    "    while True:\n",
    "        # Prepare the arguments for the request\n",
    "        list_kwargs = {\n",
    "            'Bucket': bucket_name,\n",
    "            'Prefix': prefix  # List objects under this prefix, or leave blank for all objects\n",
    "        }\n",
    "        \n",
    "        if continuation_token:\n",
    "            list_kwargs['ContinuationToken'] = continuation_token\n",
    "        \n",
    "        # Make the request to list objects\n",
    "        response = s3.list_objects_v2(**list_kwargs)\n",
    "        \n",
    "        # Collect the file keys\n",
    "        all_files += [obj['Key'] for obj in response.get('Contents', [])]\n",
    "        \n",
    "        # Check if more results are available\n",
    "        if response.get('IsTruncated'):  # True if there are more results to fetch\n",
    "            continuation_token = response['NextContinuationToken']\n",
    "        else:\n",
    "            break  # No more results to fetch\n",
    "\n",
    "    return all_files"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "val_size = 100\n",
    "train_size = len(base_metas) - val_size\n",
    "train_metas = base_metas[:train_size]\n",
    "val_metas = base_metas[train_size:]\n",
    "print(len(train_metas), len(val_metas))\n",
    "\n",
    "# find all files on s3 with the pattern\n",
    "filepaths = get_s3_files(bucket_name, f\"{base_dir}/{output_name}\")\n",
    "print(len(filepaths))\n",
    "id_to_s3_paths = {}\n",
    "\n",
    "# create a dict with the id as key and the s3 paths a list of values\n",
    "for filepath in filepaths:\n",
    "    meta_id = filepath.split(\"/\")[-1].split(\".\")[0].split(\"-\")[:-1]\n",
    "    meta_id = \"-\".join(meta_id)\n",
    "    if meta_id not in id_to_s3_paths:\n",
    "        id_to_s3_paths[meta_id] = set()\n",
    "\n",
    "    s3_filepath_basename = filepath.split(\".\")[0]\n",
    "    id_to_s3_paths[meta_id].add(s3_filepath_basename)\n",
    "\n",
    "print(len(id_to_s3_paths))\n",
    "\n",
    "# now iterate over the val, then train metas\n",
    "for dset_type in [\"val\", \"train\"]:\n",
    "    metas = val_metas if dset_type == \"val\" else train_metas\n",
    "\n",
    "    valid_metas = []\n",
    "    for meta in metas:\n",
    "        if meta[\"id\"] in id_to_s3_paths:\n",
    "            valid_metas.append(meta)\n",
    "\n",
    "    print(len(valid_metas))\n",
    "\n",
    "    # get a list of all the ids in the id_to_s3_paths\n",
    "    to_write_len_s = SEMANTIC_MEMMAP_SIZE * len(valid_metas)\n",
    "    to_write_len_v = VAE_MEMMAP_SIZE * VAE_DIM * len(valid_metas)\n",
    "\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_filepath = os.path.join(OUT_DATA_DIR, f\"data_semantic_{dset_type}.bin\")\n",
    "\n",
    "    n_offs_s = 0\n",
    "    n_offs_v = 0\n",
    "\n",
    "    out_mm_semantic = np.memmap(\n",
    "        out_mm_semantic_filepath,\n",
    "        dtype=np.uint16,\n",
    "        mode=\"w+\",\n",
    "        shape=(n_offs_s + to_write_len_s,),\n",
    "    )\n",
    "\n",
    "    out_mm_vae = np.memmap(\n",
    "        out_mm_vae_filepath,\n",
    "        dtype=np.float16,\n",
    "        mode=\"w+\",\n",
    "        shape=(n_offs_v + to_write_len_v,),\n",
    "    )\n",
    "\n",
    "    new_metas = []\n",
    "\n",
    "    for meta in tqdm(valid_metas):\n",
    "\n",
    "        # get the s3 paths for the current meta\n",
    "        if meta[\"id\"] not in id_to_s3_paths:\n",
    "            continue\n",
    "\n",
    "        s3_paths = list(id_to_s3_paths[meta[\"id\"]])\n",
    "        s3_path = s3_paths[0]\n",
    "\n",
    "        filepath = f\"s3://{bucket_name}/{s3_path}.npz\"\n",
    "        data = read_from_s3(filepath, read_f=np.load)\n",
    "\n",
    "        arr_s = data[\"semantic_codes\"]\n",
    "        arr_v = data[\"upsampled_latents\"]\n",
    "\n",
    "        # write to memmap\n",
    "        out_mm_semantic[n_offs_s : n_offs_s + arr_s.size] = arr_s.reshape(\n",
    "            -1,\n",
    "        )\n",
    "        out_mm_vae[n_offs_v : n_offs_v + arr_v.size] = arr_v.reshape(\n",
    "            -1,\n",
    "        )\n",
    "        n_offs_s += arr_s.size\n",
    "        n_offs_v += arr_v.size\n",
    "        \n",
    "        # add the meta to the new metas\n",
    "        new_metas.append(meta)\n",
    "\n",
    "\n",
    "    # write it once\n",
    "    out_mm_semantic.flush()\n",
    "    out_mm_vae.flush()\n",
    "    del out_mm_semantic, out_mm_vae\n",
    "\n",
    "    write_jsonl(new_metas, out_metas_filepath)\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import funcy\n",
    "import torch\n",
    "from dac.model.dac4 import DAC\n",
    "\n",
    "# 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",
    "#checkpoint_filepath = \"s3://suno-data/christian/100hz_vae_peaq_kl_0.005.pth\"\n",
    "checkpoint_filepath = \"s3://suno-data/christian/25hz_vae_peaq_kl_0.005.pth\"\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",
    "model_100hz = DAC(**sd[\"metadata\"][\"kwargs\"])\n",
    "model_100hz.load_state_dict(sd[\"state_dict\"])\n",
    "model_100hz.eval()\n",
    "#model_100hz.to(device)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "VAE_DIM = 128\n",
    "VAE_RATE_HZ = 25\n",
    "VAE_MEMMAP_SIZE = 750\n",
    "SEMANTIC_VOCAB_SIZE = 4000\n",
    "SEMANTIC_MEMMAP_SIZE = 750\n",
    "#OUT_DATA_DIR = \"/home/christian/data/dpo/genius_hq_corrupt_dpo_25hz_30_v2/\"\n",
    "OUT_DATA_DIR = \"/app/suno/data/diffusion_ft/upsample_v4_t_5_20241018_25hz_20241031_v1\"\n",
    "\n",
    "\n",
    "dset_type = \"train\"\n",
    "vae_memmap_filepath = os.path.join(OUT_DATA_DIR, f\"data_vae_{dset_type}.bin\")\n",
    "semantic_memmap_filepath = os.path.join(OUT_DATA_DIR, f\"data_semantic_{dset_type}.bin\")\n",
    "metas_filepath = os.path.join(OUT_DATA_DIR, f\"metas_{dset_type}.jsonl\")\n",
    "metas = read_jsonl(metas_filepath)\n",
    "print(len(metas))\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_MEMMAP_SIZE, VAE_DIM)\n",
    "semantic_data = semantic_memmap.reshape(-1, SEMANTIC_MEMMAP_SIZE)\n",
    "\n",
    "print(vae_data.shape, semantic_data.shape, len(metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import IPython\n",
    "\n",
    "# select a random idx on 0, 2, 4, 6, ...\n",
    "rand_idx = np.arange(0, len(metas), 2)\n",
    "rand_idx = np.random.choice(rand_idx, replace=False)\n",
    "\n",
    "\n",
    "print(rand_idx)\n",
    "print(metas[rand_idx])\n",
    "vae_seq = vae_data[rand_idx]\n",
    "vae_seq = torch.from_numpy(vae_seq.copy()).unsqueeze(0).float()\n",
    "\n",
    "audio = model_100hz.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": [
    "print(metas[rand_idx + 1])\n",
    "vae_seq = vae_data[rand_idx + 1]\n",
    "vae_seq = torch.from_numpy(vae_seq.copy()).unsqueeze(0).float()\n",
    "\n",
    "audio = model_100hz.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_env2",
   "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.15"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
