{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "fed7c348",
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"7\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "c046f7b4",
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.diffusion import generation as diffusion_gen\n",
    "from suno_utils.tasks.upsample_engine import UpsampleEngine, Request\n",
    "import torch\n",
    "from suno_utils.audio import Audio"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "35f40f0f",
   "metadata": {},
   "outputs": [],
   "source": [
    "!pip freeze | grep miditok\n",
    "!pip freeze | grep tokeni\n",
    "!pip freeze | grep transfor"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "6fc385c0",
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.diffusion import generation as diffusion_gen\n",
    "from suno_utils.tasks.upsample_engine import UpsampleEngine, Request\n",
    "import torch\n",
    "from suno_utils.audio import Audio\n",
    "from suno_utils.audio.midi import Midi\n",
    "from suno_utils.utils.s3 import download_s3_file_if_needed\n",
    "\n",
    "Midi.load_tokenizer(download_s3_file_if_needed(\"s3://suno-data/victor/midi_tokenizer_v1.json\"))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a718fb2b",
   "metadata": {},
   "outputs": [],
   "source": [
    "dit_model_filepath = \"/app/suno/checkpoints/2025-04-29_15-30-01_s1948/last_ckpt_infer.pt\"  # complements\n",
    "diffusion_gen.preload_models(\n",
    "    dit_model_filepath=\"/app/suno/checkpoints/2025-05-26_22-49-34_s8646/last_ckpt_infer.pt\",  # 12 stems\n",
    "    codec_filepath=\"s3://suno-data/minz/models/dac_vae_tuned_25hz.pth\",\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "4e6e3359",
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.tasks.mert_25 import preload_models as preload_semantic_models, encode as encode_semantic\n",
    "\n",
    "preload_semantic_models(\n",
    "    checkpoint_filepath=\"s3://suno-data/georg/models/semantic/mert_25.pt\",\n",
    "    centroids_filepath=\"s3://suno-data/georg/models/semantic/mert_25_2x4k.npy\",\n",
    ")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cf696991",
   "metadata": {},
   "outputs": [],
   "source": [
    "import json\n",
    "import numpy as np\n",
    "from suno_utils.utils.s3 import read_from_s3\n",
    "from suno_utils.utils.clip import SunoClip\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "fa9ef970",
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.tasks.dac_vae_fixed_25hz import (\n",
    "    preload_models as preload_codec_models,\n",
    "    decode,\n",
    "    encode,\n",
    "    get_embedding_rate,\n",
    "    load_model as load_codec_model,\n",
    ")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "c0c34cd3",
   "metadata": {},
   "outputs": [],
   "source": [
    "upsample_engine = UpsampleEngine(min_chunk_size=25 * 30)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "7098d2e2",
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.tasks.midi_transcription import MidiTranscriber\n",
    "\n",
    "midi_transcriber = MidiTranscriber()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "dda4f313",
   "metadata": {},
   "source": [
    "## Stems\n",
    "\n",
    "The latest model supports the following prompt types:\n",
    "\n",
    "- extract [group_Vocals]\n",
    "    - A category\n",
    "- extract [Backing Vocals]\n",
    "    - A specific instrument\n",
    "- extract [binary]\n",
    "    - Split into halves with random stems in each half"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2197af65",
   "metadata": {},
   "outputs": [],
   "source": [
    "def gen_stem(\n",
    "    audio: Audio,\n",
    "    stem_type_cfg_scale=1.0,\n",
    "    tags=\"extract [group_Vocals]\",\n",
    "    steps=8,\n",
    "    seed=3,\n",
    "    codec_scale_factor=0.4,\n",
    "    scale_ctx_vector=True,\n",
    "    noise_ctx_level=0.0,\n",
    "    infill_prefix_latents=None,\n",
    "    infill_suffix_latents=None,\n",
    "):\n",
    "    vae = encode(audio)\n",
    "    gen_cfg = diffusion_gen.DiffusionGenerationConfig(\n",
    "        lyrics=tags,\n",
    "        steps=steps,\n",
    "        seed=seed,\n",
    "        codec_scale_factor=codec_scale_factor,\n",
    "        scale_ctx_vector=scale_ctx_vector,\n",
    "        noise_ctx_level=noise_ctx_level,\n",
    "        text_cfg_coef=stem_type_cfg_scale,\n",
    "        infill_prefix_latents=infill_prefix_latents,\n",
    "        infill_suffix_latents=infill_suffix_latents,\n",
    "        drop_semantic_tokens=True,\n",
    "    )\n",
    "    print(gen_cfg)\n",
    "\n",
    "    request = Request(\n",
    "        id=\"dummy\",\n",
    "        generation_config=gen_cfg,\n",
    "        tokens=np.zeros((vae.shape[0], 1)),\n",
    "        input_tokens_finished=True,\n",
    "        stem_ctx_latents=vae,\n",
    "    )\n",
    "\n",
    "    result = upsample_engine.run_request(request)\n",
    "    vae_latents = torch.concat(result.vae_latents)\n",
    "    print(f\"vae_latents: {vae_latents.shape}\")\n",
    "    audios = []\n",
    "    for i in range(vae_latents.shape[1]):\n",
    "        audios.append(decode(vae_latents[:, i]))\n",
    "    return audios\n"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "0fc7e7f2",
   "metadata": {},
   "source": [
    "## midi"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "28d3e221",
   "metadata": {},
   "outputs": [],
   "source": [
    "def decompose(gen_id):\n",
    "    # v2 codec clips\n",
    "    # gen_id = \"5ed9d01d-6e8b-41a0-968d-70ebbe51ba24\"  #\n",
    "    # gen_id = \"f497991a-ad42-45be-89f9-5fc324252f24\"\n",
    "    clip = SunoClip(gen_id)\n",
    "    audio = clip.audio()\n",
    "\n",
    "    audios = gen_stem(audio.get_segment(from_s=0), tags=\"extract [split_karaoke]\")\n",
    "    mix = Audio.sum(audios)\n",
    "    # mix.play()\n",
    "\n",
    "    # for audio in audios:\n",
    "    #     audio.play()\n",
    "\n",
    "    pairs = []\n",
    "    for audio in audios:\n",
    "        if audio.loudness < -40:\n",
    "            continue\n",
    "        midi_stem = midi_transcriber.transcribe(audio)\n",
    "        # midi_stem.show()\n",
    "        # midi_stem.make_stereo_comparison(audio).play()\n",
    "        pairs.append((midi_stem, audio))\n",
    "\n",
    "    midi_sum = sum([midi for midi, _ in pairs]).consolidate_programs().clean_pmidi()\n",
    "    # midi_sum.show()\n",
    "    # midi_sum.make_stereo_comparison(mix).play()\n",
    "    return midi_sum, mix, pairs\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "93b6e439",
   "metadata": {},
   "outputs": [],
   "source": [
    "ids_to_decompose = {\n",
    "    \"stone\": \"a5e2198a-f352-4abb-9a24-7f81b143ded3\",\n",
    "    \"sara\": \"a1936d21-d47a-44ad-8551-25feb86a58e9\",\n",
    "    \"pop\": \"e3ccb02b-696e-4335-a674-9f7821b202f4\",\n",
    "    \"chill\": \"63774ef1-ab73-472b-84dc-3c9d9a385339\",\n",
    "    \"rainbow\": \"7aa56e05-3e52-4cd8-b7e4-cb009480b0e7\",\n",
    "    \"surf\": \"90a879c4-1124-4c35-8f3a-723faf6130cc\",\n",
    "    \"henry\": \"313c8579-456f-4218-828e-ca165ae2b432\",\n",
    "    \"turkey\": \"0e72b9be-75e7-4d1c-bcb2-091a0d95f293\",\n",
    "}\n",
    "\n",
    "for name, gen_id in ids_to_decompose.items():\n",
    "    midi_sum, mix, pairs = decompose(gen_id)\n",
    "    midi_sum.show()\n",
    "    midi_sum.make_stereo_comparison(mix).play()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "5bf3e26b",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "57fd915a",
   "metadata": {},
   "outputs": [],
   "source": [
    "# midi_sum.write(\"henry.mid\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "14bbc689",
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
