{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "fed7c348",
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"4\""
   ]
  },
  {
   "cell_type": "markdown",
   "id": "e9c62b25",
   "metadata": {},
   "source": [
    "# BCT Inference\n",
    "\n",
    "This notebook is a hackable place to test BCT model inference."
   ]
  },
  {
   "cell_type": "markdown",
   "id": "23f330f5",
   "metadata": {},
   "source": [
    "## Load model\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3c7d0311",
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.gpt.generation import load_model, GPT\n",
    "from suno_utils.utils.s3 import download_s3_file_if_needed\n",
    "import torch\n",
    "\n",
    "gpt_model_path = \"/app2/suno/checkpoints/2025-08-24_08-31-23/last_ckpt_infer.pt\"\n",
    "more_stems_model_path = \"/app2/suno/checkpoints/2025-10-24_13-54-25/last_ckpt_infer.pt\"\n",
    "crow_model_path = \"/app2/suno/modal/models/tony/sem/model_45_6b_crow_sft0828_t14_r15_d87.pt\"\n",
    "dodo_sft = \"/app2/suno/checkpoints/2025-11-01_13-45-35/last_ckpt_infer.pt\"\n",
    "# stems = \"/app2/suno/checkpoints/2025-11-03_20-46-04/last_ckpt_infer.pt\" # input repa\n",
    "stems = \"/app2/suno/checkpoints/2025-11-04_16-31-51/last_ckpt_infer.pt\"  # mix repa\n",
    "# stems = \"/app2/suno/checkpoints/2025-11-05_21-10-24/last_ckpt_infer.pt\"  # + hoot repa\n",
    "# stems = \"/app2/suno/checkpoints/2025-11-06_04-31-03/last_ckpt_infer.pt\"  # + midi repa\n",
    "stem_dpo = \"/app2/suno/checkpoints/2025-11-11_09-20-50/last_ckpt_infer.pt\"\n",
    "infill = \"/app2/suno/checkpoints/2025-11-08_19-46-16/last_ckpt_infer.pt\"\n",
    "stem_test = \"/app2/suno/checkpoints/2025-11-17_12-56-45/last_ckpt_infer.pt\"\n",
    "tokenizer_path = \"s3://suno-data/georg/models/tokenizers/tokenizer_60k.json\"\n",
    "\n",
    "\n",
    "model_container = load_model(\n",
    "    ckpt_path=download_s3_file_if_needed(stem_dpo),\n",
    "    tokenizer_path=download_s3_file_if_needed(tokenizer_path),\n",
    ")\n",
    "model: GPT = model_container[\"model\"]\n",
    "assert isinstance(model, GPT)\n",
    "\n",
    "\n",
    "cfg = model.config"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3118c1b4",
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.diffusion import generation as diffusion_gen\n",
    "from suno_utils.tasks.upsample_engine import UpsampleEngine, Request\n",
    "from suno_utils.audio import Audio\n",
    "\n",
    "dit_model_filepath = \"/app/suno/checkpoints/2025-02-17_16-54-01_s7787/last_ckpt_infer.pt\"  # base\n",
    "diffusion_gen.preload_models(\n",
    "    dit_model_filepath=dit_model_filepath,\n",
    ")\n",
    "from suno_utils.tasks.dac_vae_fixed_25hz import preload_models as preload_codec_models, decode\n",
    "\n",
    "preload_codec_models(\"s3://suno-data/minz/models/dac_vae_tuned_25hz.pth\")\n",
    "\n",
    "diffusion_engine = UpsampleEngine(compile=False)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "1a1d5952",
   "metadata": {},
   "source": [
    "## Inference"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "8a44dbed",
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.gpt.bct.bct_test_utils import Sequencer, TestCase\n",
    "from suno_utils.gpt.bct.bct_generation_simple import BCTGenerationConfig, BlockSequence, generate_block\n",
    "\n",
    "sequencer = Sequencer(model.config, model_container[\"tokenizer\"])\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3b09b80a",
   "metadata": {},
   "outputs": [],
   "source": [
    "def generate_audio(test_case: TestCase, prepend_token=False):\n",
    "    block = generate_block(model, test_case.generation_config)\n",
    "    sem_codes_full = block.inputs[\"semantic_input_0\"][0, 0, 1 : len(block)]\n",
    "\n",
    "    # keep up to the first invalid token\n",
    "    mask_valid = sem_codes_full < model.config.semantic_codebook_size\n",
    "    if not torch.all(mask_valid):\n",
    "        invalid_indices = (~mask_valid).nonzero(as_tuple=True)[0]\n",
    "        if invalid_indices.numel() == 0:\n",
    "            sem_codes = sem_codes_full\n",
    "        else:\n",
    "            first_invalid = invalid_indices[0].item()\n",
    "            sem_codes = sem_codes_full[:first_invalid]\n",
    "    else:\n",
    "        sem_codes = sem_codes_full\n",
    "    assert torch.all(sem_codes < model.config.semantic_codebook_size)\n",
    "    if prepend_token:\n",
    "        # 813 is a silence token. used to offset the audio\n",
    "        sem_codes = torch.cat(\n",
    "            [torch.tensor([813], device=sem_codes.device, dtype=sem_codes.dtype), sem_codes]\n",
    "        )\n",
    "\n",
    "    gen_cfg = diffusion_gen.DiffusionGenerationConfig(\n",
    "        # audio=vae_latents,\n",
    "        # tags=\"extract Lead Vocal\",\n",
    "        lyrics=test_case.lyrics,\n",
    "        text_cfg_coef=2.0,\n",
    "        ctx_cfg_coef=1.0,\n",
    "        steps=12,\n",
    "        codec_scale_factor=0.4,\n",
    "        scale_ctx_vector=True,\n",
    "    )\n",
    "    request = Request(\n",
    "        id=\"dummy\",\n",
    "        generation_config=gen_cfg,\n",
    "        tokens=sem_codes.cpu(),\n",
    "        input_tokens_finished=True,\n",
    "    )\n",
    "    result = diffusion_engine.run_request(request)\n",
    "    vae_latents = torch.concat(result.vae_latents)\n",
    "    audio = decode(vae_latents)\n",
    "    return audio\n",
    "\n",
    "\n",
    "# test_case = sequencer.stone_cover_blocks()\n",
    "# generate_audio(test_case).play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "e234c811",
   "metadata": {},
   "outputs": [],
   "source": [
    "# test_case = sequencer.infill_blocks(infill_duration_s=10)\n",
    "# infill_audio = generate_audio(test_case)\n",
    "# infill_audio.play()\n",
    "# Audio.concatenate(\n",
    "#     [test_case.extra_data[\"prefix_audio\"], infill_audio, test_case.extra_data[\"suffix_audio\"]]\n",
    "# ).play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "7d54e5e7",
   "metadata": {},
   "outputs": [],
   "source": [
    "# lyrics = \"plastic plastic plastic, I love plastic\\n\" * 20\n",
    "# for i in range(4):\n",
    "#     test_case = sequencer.text_conditional_blocks(lyrics)\n",
    "#     test_case.generation_config.max_autoregressive_steps = 25*240\n",
    "#     generate_audio(test_case).play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "e6acb234",
   "metadata": {},
   "outputs": [],
   "source": [
    "# lyrics = \"[JS bach, JSbach, fugue, piano, baroque]\"\n",
    "# for i in range(4):\n",
    "#     test_case = sequencer.text_conditional_blocks(lyrics)\n",
    "#     test_case.generation_config.max_autoregressive_steps = 25*240\n",
    "#     generate_audio(test_case).play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "f695cdfd",
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.utils.clip import SunoClip\n",
    "import os\n",
    "\n",
    "clip = SunoClip(\"04785fb0-f0f3-4867-aa66-c2158e6054aa\")\n",
    "# clip = SunoClip(\"f58d86ce-357d-4bf5-b1e7-9e3794b54ea0\")\n",
    "for i in range(8):\n",
    "    test_case = sequencer.stem_add_blocks(\n",
    "        stem_tag=f\"add Kick Drums\",\n",
    "        clip=clip,\n",
    "        start_s=0,\n",
    "    )\n",
    "    stem_audio = generate_audio(test_case, prepend_token=True)\n",
    "    original_audio = test_case.extra_data[\"input_clip\"]\n",
    "    stem_audio.play()\n",
    "    mix = original_audio + stem_audio.normalize_volume(-10)\n",
    "    mix.play()\n",
    "    os.makedirs(\"outputs\", exist_ok=True)\n",
    "    mix.write_opus(f\"outputs/mix_{i}.opus\")\n",
    "    stem_audio.write_opus(f\"outputs/stem_{i}.opus\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2cb4d71a",
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.tasks.ditto_v2 import preload_models as preload_ditto_models\n",
    "\n",
    "preload_ditto_models()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "f2fdbd90",
   "metadata": {},
   "outputs": [],
   "source": [
    "test_case = sequencer.ditto_blocks()\n",
    "for i in range(4):\n",
    "    generate_audio(test_case).play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a727e0f1",
   "metadata": {},
   "outputs": [],
   "source": [
    "generate_audio(sequencer.overpaint_blocks()).play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "205fdd9a",
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
