{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "fee795b8",
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"2\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "47377e6a",
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch\n",
    "import torchaudio\n",
    "#import polars as pl\n",
    "from suno_utils.utils.text import read_jsonl\n",
    "from suno_utils.audio import Audio\n",
    "from pathlib import Path\n",
    "from suno_utils.utils.opusfile import OpusFile"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "d6eedadf",
   "metadata": {},
   "outputs": [],
   "source": [
    "# test some of the local data\n",
    "CODEC_FILEPATH = \"s3://suno-data/minz/models/dac_vae_tuned_25hz.pth\"\n",
    "\n",
    "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",
    "    decode_stream_to_full_audio,\n",
    ")\n",
    "\n",
    "_ = preload_codec_models(CODEC_FILEPATH)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2a4b0889",
   "metadata": {},
   "outputs": [],
   "source": [
    "# lets count how many 30s chunks per second we can load\n",
    "import sys\n",
    "sys.path.append(\"/home/christian/code/neon/sunoDiff\")\n",
    "\n",
    "from dynamic_dataset import DynamicDataset"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "06b728f8",
   "metadata": {},
   "outputs": [],
   "source": [
    "run_config = {\n",
    "    \"data\" : {\n",
    "        \"codec_filepath\": \"s3://suno-data/minz/models/dac_vae_tuned_25hz.pth\",\n",
    "        \"semantic_model_filepath\": \"s3://suno-data/georg/models/semantic/mert_25.pt\",\n",
    "        \"semantic_clusters_filepath\": \"s3://suno-data/georg/models/semantic/mert_25_2x4k.npy\",\n",
    "        \"train_metas_filepath\": \"/app2/suno/data/diffusion/v1/metas_diff_v1_val.jsonl\",\n",
    "        \"val_metas_filepath\": \"/app2/suno/data/diffusion/v1/metas_diff_v1_val.jsonl\",\n",
    "        \"audio_chunk_s\": 30.02,\n",
    "        \"audio_ctx_s\": 30.02,\n",
    "        \"audio_vox_s\": 30.02,\n",
    "        \"target_loudness_db\" : -16.0,\n",
    "        \"vae_dim\": 128,\n",
    "        \"semantic_rate_hz\": 25,\n",
    "        \"vae_scale_factor\": 0.4,\n",
    "        \"foreign_weight\" : 0.5,\n",
    "        \"text_aligned_weight\" : 3.0,\n",
    "        \"stem_weight\" : 2.0,\n",
    "        \"scale_vae_ctx\" : True,\n",
    "        \"noise_ctx\" : 1.0,\n",
    "        \"text_drop_prob\": 0.1,\n",
    "        \"semantic_mask_prob\": 0.1,\n",
    "        \"ctx_mask_prob\": 0.2,\n",
    "        \"use_stem_prob\" : 0.1,\n",
    "        \"use_vox_prob\" : 0.8,\n",
    "        \"infill_prob\": 0.1,\n",
    "        \"infill_min_ratio\": 0.2,\n",
    "        \"infill_max_ratio\": 0.8,\n",
    "    }\n",
    "}\n",
    "\n",
    "dataset = DynamicDataset(\n",
    "    metas=run_config[\"data\"][\"train_metas_filepath\"],\n",
    "    audio_chunk_s=run_config[\"data\"][\"audio_chunk_s\"],\n",
    "    audio_ctx_s=run_config[\"data\"][\"audio_ctx_s\"],\n",
    "    audio_vox_s=run_config[\"data\"][\"audio_vox_s\"],\n",
    "    cond_text_len=1536,\n",
    "    text_drop_prob=run_config[\"data\"][\"text_drop_prob\"],\n",
    "    target_loudness_db=run_config[\"data\"][\"target_loudness_db\"],\n",
    "    use_stem_prob=run_config[\"data\"][\"use_stem_prob\"],\n",
    "    is_training=True,\n",
    "    use_vox_prob=run_config[\"data\"][\"use_vox_prob\"],\n",
    "    foreign_weight=run_config[\"data\"][\"foreign_weight\"],\n",
    "    text_aligned_weight=run_config[\"data\"][\"text_aligned_weight\"],\n",
    "    stem_weight=run_config[\"data\"][\"stem_weight\"],\n",
    "\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "d2736318",
   "metadata": {},
   "outputs": [],
   "source": [
    "dataset.is_training = True"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "c15fa022",
   "metadata": {},
   "outputs": [],
   "source": [
    "dataloader = torch.utils.data.DataLoader(\n",
    "    dataset,\n",
    "    batch_size=1,\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "9ba3936a",
   "metadata": {},
   "outputs": [],
   "source": [
    "for idx, batch in enumerate(dataset):\n",
    "    print(idx, batch)\n",
    "    audio_target, audio_ctx, audio_vox, text_codes, raw_text, audio_target_24k = batch\n",
    "    #print(audio_target, audio_ctx, audio_vox, text_codes, audio_target_24k)\n",
    "    print(raw_text)\n",
    "    break"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "97438930",
   "metadata": {},
   "outputs": [],
   "source": [
    "audio_target.duration_s"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "c9268d99",
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "\n",
    "for idx, batch in enumerate(dataset):\n",
    "    #print(idx, batch)\n",
    "    audio_target, audio_ctx, audio_vox, text_codes, raw_text, audio_target_24k = batch\n",
    "\n",
    "    #if audio_vox is None:\n",
    "    #    continue\n",
    "    #else:\n",
    "    print()\n",
    "    print(idx)\n",
    "    print(\"audio_target\")\n",
    "    audio_target.play()\n",
    "    #0print(\"audio_target (cycled)\")\n",
    "    #audio_target_cycled = codec_decode(codec_encode(audio_target))\n",
    "    #audio_target_cycled.play()\n",
    "    print(\"audio_target_24k\")\n",
    "    audio_target_24k.play()\n",
    "    if audio_ctx is not None:\n",
    "        print(\"audio_ctx\")  \n",
    "        audio_ctx.play()\n",
    "        print(\"audio_target+audio_ctx\")\n",
    "        # Concatenate audio_ctx and audio_target and play the result\n",
    "        audio_target_array = audio_target.array_float\n",
    "        audio_ctx_array = audio_ctx.array_float\n",
    "        concat_audio = np.concatenate([audio_ctx_array, audio_target_array], axis=-1)\n",
    "        concat_audio = Audio.from_array_float(concat_audio, sample_rate=48000)\n",
    "        concat_audio.play()\n",
    "    if audio_vox is not None:\n",
    "        print(\"audio_vox\")\n",
    "        audio_vox.play()\n",
    "    print()\n",
    "    print(\"raw_text\")\n",
    "    print(raw_text)\n",
    "    break\n",
    "    "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "fc5e475e",
   "metadata": {},
   "outputs": [],
   "source": [
    "podcast_count = 0\n",
    "for meta in dataset.metas:\n",
    "    if meta[\"id\"].startswith(\"extreme\"):\n",
    "        print(meta[\"id\"])\n",
    "        print(f\"\"\" \"{meta[\"local_filepath\"]}\" \"\"\")\n",
    "        podcast_count += 1\n",
    "        if podcast_count > 10:\n",
    "            break"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2a52f796",
   "metadata": {},
   "outputs": [],
   "source": [
    "audio = Audio.from_file(\"/app2/suno/data/extreme_music/audio/56944/Full Version.mp3\", n_channels=2)\n",
    "audio.play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "59b1f649",
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "suno_env_fa2",
   "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": 5
}
