{
 "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, Job\n",
    "from suno_utils.tasks.dac_vae_fixed_25hz import decode_stream_to_full_audio, decode\n",
    "import torch\n",
    "from suno_utils.audio import Audio"
   ]
  },
  {
   "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/victor/2b_prefix_ft.pt\",\n",
    "    # codec_filepath=\"/app/suno/victor/25hz_vae_peaq_kl_0.005.pth\",\n",
    "    # dit_model_filepath=\"/app/suno/checkpoints/2025-03-24_22-31-58_s4919/last_ckpt_infer.pt\", # finetune infill fixed\n",
    "    # dit_model_filepath=\"/app/suno/checkpoints/2025-03-28_18-46-17_s894/last_ckpt_infer.pt\",  # base model\n",
    "    # dit_model_filepath=\"/app/suno/checkpoints/2025-04-08_23-03-53_s7257/step_30000_infer.pt\",  # stem infill\n",
    "    # dit_model_filepath=\"/app/suno/checkpoints/2025-04-29_15-30-01_s1948/step_80000_infer.pt\",  # complements\n",
    "    # dit_model_filepath=\"/app/suno/checkpoints/2025-05-09_11-32-03_s7649/last_ckpt_infer.pt\",  # complements groups\n",
    "    # dit_model_filepath=\"/app/suno/checkpoints/2025-05-09_11-32-03_s7649/last_ckpt_infer.pt\",  # stereo\n",
    "    # dit_model_filepath=\"/app/suno/checkpoints/2025-05-16_18-19-47_s2957/last_ckpt_infer.pt\",  # 8 stems\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": "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",
    "\n",
    "\n",
    "gen_id = \"a5e2198a-f352-4abb-9a24-7f81b143ded3\"  # stone\n",
    "\n",
    "\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)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "0dc18892",
   "metadata": {},
   "outputs": [],
   "source": [
    "audio = clip.audio()"
   ]
  },
  {
   "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",
    "\n",
    "vae = encode(audio)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "c0c34cd3",
   "metadata": {},
   "outputs": [],
   "source": [
    "engine = UpsampleEngine(min_chunk_size=25 * 30)"
   ]
  },
  {
   "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 = 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": "code",
   "execution_count": null,
   "id": "9f6bc7c4",
   "metadata": {},
   "outputs": [],
   "source": [
    "audios = gen_stem(\n",
    "    audio.get_segment(from_s=0, to_s=320),\n",
    "    # tags=\"extract [group_Vocals]\",\n",
    "    tags=\"extract [split_karaoke]\",\n",
    ")\n",
    "mix = Audio.sum(audios)\n",
    "mix.play()\n",
    "\n",
    "for audio in audios:\n",
    "    audio.play()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "9557d673",
   "metadata": {},
   "outputs": [],
   "source": [
    "break"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "0fc7e7f2",
   "metadata": {},
   "source": [
    "## Modal call"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "bb4413f7",
   "metadata": {},
   "outputs": [],
   "source": [
    "import modal\n",
    "import json\n",
    "from uuid import uuid4\n",
    "\n",
    "model_f_stem = modal.Cls.lookup(\"stems-stems_v1-dev\", \"StemStub\")\n",
    "gen_id = \"a5e2198a-f352-4abb-9a24-7f81b143ded3\"  # stone\n",
    "# can also be a s3 url like \"s3://suno-data-uploads/studio/uploads/a5e2198a-f352-4abb-9a24-7f81b143ded3.mp3\"\n",
    "id = str(uuid4())\n",
    "\n",
    "input = {\n",
    "    \"id\": id,\n",
    "    \"prompt_audio\": gen_id,\n",
    "    \"model_name\": \"stems\",\n",
    "    \"callback_url\": \"https://api-staging.suno.ai/api/generate/finish-clip/\",\n",
    "    \"metadata\": {\n",
    "        \"stem_type_group_name\": \"Vocals\",\n",
    "        \"multi_ids\": [id + \"_stem\", id + \"_complement\"],\n",
    "    },\n",
    "}\n",
    "\n",
    "print(f\"Starting modal call for {input['id']}\")\n",
    "model_f_stem.stem.remote(json.dumps(input))\n",
    "print(\n",
    "    f\"Done with modal call for {input['id']}. See it at https://cdn1.suno.ai/{input['id']}_stem.opus and https://cdn1.suno.ai/{input['id']}_complement.opus\"\n",
    ")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "feafce07",
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
