{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "fed7c348",
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"4\""
   ]
  },
  {
   "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",
    "from suno_utils.audio import Audio"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "e9c62b25",
   "metadata": {},
   "source": [
    "## GPT engine"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3c7d0311",
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch\n",
    "from suno_utils.gpt.generation import GenerationConfig\n",
    "from suno_utils.gpt.engine import Engine\n",
    "from suno_utils.gpt.generation_engine import make_request\n",
    "\n",
    "N_BATCH = 1\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "06f57a07",
   "metadata": {},
   "outputs": [],
   "source": [
    "dit_model_filepath = \"/app/suno/checkpoints/2025-04-27_00-05-49_s1686/last_ckpt.pt\"  # base\n",
    "diffusion_gen.preload_models(\n",
    "    dit_model_filepath=dit_model_filepath,\n",
    "    codec_filepath=\"/app/suno/data/dpo/models/dac_vae_tuned_25hz.pth\",\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "1b21561e",
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.tasks.dac_vae_fixed_25hz import preload_models as preload_codec_models, decode, encode\n",
    "\n",
    "preload_codec_models(\"/app/suno/data/dpo/models/dac_vae_tuned_25hz.pth\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "6971612f",
   "metadata": {},
   "outputs": [],
   "source": [
    "engine = Engine(\n",
    "    # \"/app/suno/checkpoints/2025-04-09_22-40-29/last_ckpt_infer.pt\",  # 6b baseline\n",
    "    # \"/app/suno/checkpoints/2025-04-02_04-49-15/last_ckpt_infer.pt\",  # 6b oracle dl\n",
    "    # \"/app/suno/checkpoints/2025-04-10_21-20-14/last_ckpt_infer.pt\",  # 3b baseline\n",
    "    # \"/app/suno/checkpoints/2025-04-23_16-23-43/last_ckpt_infer.pt\",  # 6b ft playlist\n",
    "    # \"/app/suno/checkpoints/2025-05-30_01-41-55/last_ckpt_infer.pt\",  # 6b ft playlist\n",
    "    \"/app2/suno/checkpoints/2025-08-24_08-31-23/last_ckpt_infer.pt\",  # dpo\n",
    "    # \"/app2/suno/checkpoints/2025-08-22_19-13-34/last_ckpt_infer.pt\",  # sft\n",
    "    \"s3://suno-data/georg/models/tokenizers/tokenizer_60k.json\",\n",
    "    max_sequences=1 * N_BATCH,\n",
    "    compile=False,\n",
    ")\n",
    "cfg = engine.model.config"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "df37a180",
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.utils.clip import SunoClip\n",
    "\n",
    "# gen_id = \"99bffa17-7e59-47b4-a048-5528cbda05d5\" # sister\n",
    "# gen_id = \"081d73c4-7805-4212-9c80-8db1137ca3c4\" # friends\n",
    "# gen_id = \"562f762d-6ced-4080-9af1-910ee3d0a5dc\" # something real\n",
    "# gen_id = \"23c15c62-494d-422d-8a60-8b0454044322\" # rubber duck\n",
    "# gen_id = \"4b140a9e-964b-422c-85b5-5861ad1a9d38\" # once\n",
    "# gen_id = \"7b214347-fa38-4e9b-96f4-f7ec65adea45\"  # rock n roll, 3.5\n",
    "gen_id = \"a5e2198a-f352-4abb-9a24-7f81b143ded3\"  # stone\n",
    "gen_id_1 = \"2d38216b-ee07-44cc-b1e3-dfcf25cfdda1\"  # bluegrass\n",
    "gen_id_2 = \"80ed085f-c036-4da0-b56f-d1f3ecd50b46\"  # bluegrass\n",
    "cond_codes = []\n",
    "for gen_id in [gen_id_1, gen_id_2]:\n",
    "    clip = SunoClip(gen_id)\n",
    "    audio = clip.audio()\n",
    "    audio.play()\n",
    "    gen_codes = clip._get_npz()[\"v5.0_raw\"][: 25 * 60, :1]\n",
    "    print(gen_codes.shape)\n",
    "    # gen_codes = gen_codes.reshape(1, -1)\n",
    "    cond_codes.append(gen_codes)\n",
    "print(len(cond_codes))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "f2e7e428",
   "metadata": {},
   "outputs": [],
   "source": [
    "lyrics = \"\"\"\n",
    "[verse]\n",
    "oh, my love\n",
    "My friend you know\n",
    "it's been a while\n",
    "Without thinking of you\n",
    "but the thought makes me smile\n",
    "\n",
    "[chorus]\n",
    "I'm so tired of wanting\n",
    "wanting more than this\n",
    "i know it but what am i to do\n",
    "i need some space to breathe,\n",
    "so give me some room\n",
    "\n",
    "[verse]\n",
    "oh, my love\n",
    "you have a heart of stone\n",
    "cause since i've come home\n",
    "i've never felt so alone\n",
    "but the thought makes me smile\n",
    "\"\"\"\n",
    "gconf = GenerationConfig(\n",
    "    text=lyrics,\n",
    "    text_tags=\"\",\n",
    "    cfg_coef=1.0,\n",
    "    # cfg_coef_tags=1.5,\n",
    "    n_batch=1,\n",
    "    min_text_offset=0,\n",
    "    eos_pad_duration_s=0,\n",
    "    max_gen_duration_s=120,\n",
    "    playlist_arr=cond_codes,\n",
    ")\n",
    "requests = [make_request(f\"{i}\", gconf, engine.model.config, engine.tokenizer) for i in range(1)]\n",
    "jobs = engine.run_request(requests, tqdm_enabled=True)\n",
    "out_gpt = []\n",
    "for n, job in enumerate(jobs):\n",
    "    stream = engine.token_generator(job)\n",
    "    arr = torch.stack(list(stream))[:, 1]\n",
    "    if arr[-1] == 4000:\n",
    "        arr = arr[:-1]\n",
    "    out_gpt.append(arr)\n",
    "out_gpt[0].shape"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "df4ebc89",
   "metadata": {},
   "source": [
    "## Diffusion engine"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "c0c34cd3",
   "metadata": {},
   "outputs": [],
   "source": [
    "diffusion_engine = UpsampleEngine(compile=False)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "37b3fcef",
   "metadata": {},
   "outputs": [],
   "source": [
    "gen_cfg = diffusion_gen.DiffusionGenerationConfig(\n",
    "    # audio=vae_latents,\n",
    "    # tags=\"extract Lead Vocal\",\n",
    "    lyrics=lyrics,\n",
    "    text_cfg_coef=2.0,\n",
    "    ctx_cfg_coef=1.0,\n",
    "    steps=10,\n",
    "    codec_scale_factor=0.4,\n",
    "    scale_ctx_vector=True,\n",
    "    noise_ctx_level=0.75,\n",
    "    noise_ctx_pad_len=0,\n",
    ")\n",
    "request = Request(\n",
    "    id=\"dummy\",\n",
    "    generation_config=gen_cfg,\n",
    "    tokens=out_gpt[0],\n",
    "    input_tokens_finished=True,\n",
    ")\n",
    "result = diffusion_engine.run_request(request)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "9b601aff",
   "metadata": {},
   "source": [
    "## Codec"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "d0aeca0a",
   "metadata": {},
   "outputs": [],
   "source": [
    "vae_latents = torch.concat(result.vae_latents)\n",
    "audio = decode(vae_latents)\n",
    "audio.play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "8628f293",
   "metadata": {},
   "outputs": [],
   "source": [
    "# audio.write_mp3(\"output.mp3\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3021a427",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "markdown",
   "id": "wl94e1pm1r",
   "metadata": {},
   "source": [
    "## Test Stem Conditioning\n",
    "\n",
    "Let's test the new stem conditioning functionality by using one of the existing audio clips as the stem."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "43522cfe",
   "metadata": {},
   "outputs": [],
   "source": [
    "gen_id = \"a5e2198a-f352-4abb-9a24-7f81b143ded3\"  # stone\n",
    "\n",
    "clip = SunoClip(gen_id)\n",
    "print(list(clip._get_npz().keys()))\n",
    "stem_codes = clip._get_npz()[\"v3.5_raw\"][:, :1][600:800]\n",
    "stem_codes.shape"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "jqxwhntoou",
   "metadata": {},
   "outputs": [],
   "source": [
    "# Test stem conditioning with one of our existing clips\n",
    "import numpy as np\n",
    "from suno_utils.gpt.prompt import Prompt\n",
    "\n",
    "# Test stem conditioning configuration\n",
    "stem_gconf = GenerationConfig(\n",
    "    text=\"{add drums;activity:80%}\",\n",
    "    # text_tags=\"folk rock\",\n",
    "    # text_start_control_tags=[\"{add drums;activity:80%}\"],\n",
    "    cfg_coef=1.0,\n",
    "    cfg_coef_tags=0.0,\n",
    "    n_batch=1,\n",
    "    min_text_offset=0,\n",
    "    eos_pad_duration_s=0,\n",
    "    max_gen_duration_s=300,  # Shorter duration for testing\n",
    "    sample_arr=stem_codes,  # This is our new stem conditioning field\n",
    "    # stem_arr=stem_codes,  # This is our new stem conditioning field\n",
    ")\n",
    "\n",
    "print(\"Creating stem conditioning request...\")\n",
    "stem_requests = [\n",
    "    make_request(f\"stem_test_{i}\", stem_gconf, engine.model.config, engine.tokenizer) for i in range(1)\n",
    "]\n",
    "# Prompt(\"\", engine.model.config).visualize(stem_requests[0].streams[0].prompt, compress=True)\n",
    "# Prompt(\"\", engine.model.config).visualize(stem_requests[0].streams[0].prompt, compress=False)\n",
    "print(\"Running stem conditioning generation...\")\n",
    "stem_jobs = engine.run_request(stem_requests, tqdm_enabled=True)\n",
    "\n",
    "stem_out_gpt = []\n",
    "for n, job in enumerate(stem_jobs):\n",
    "    stream = engine.token_generator(job)\n",
    "    arr = torch.stack(list(stream))[:, 1]\n",
    "    if arr[-1] == 4000:  # Remove EOS token if present\n",
    "        arr = arr[:-1]\n",
    "    stem_out_gpt.append(arr)\n",
    "\n",
    "print(f\"Generated stem conditioning output shape: {stem_out_gpt[0].shape}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "ote5r988hmq",
   "metadata": {},
   "outputs": [],
   "source": [
    "# Run the stem conditioning output through the diffusion engine\n",
    "print(\"Running stem conditioning output through diffusion engine...\")\n",
    "\n",
    "stem_diffusion_cfg = diffusion_gen.DiffusionGenerationConfig(\n",
    "    ctx_cfg_coef=1.0,\n",
    "    steps=10,\n",
    "    codec_scale_factor=0.4,\n",
    "    scale_ctx_vector=True,\n",
    "    noise_ctx_level=0.75,\n",
    "    noise_ctx_pad_len=0,\n",
    ")\n",
    "\n",
    "stem_request = Request(\n",
    "    id=\"stem_test_diffusion\",\n",
    "    generation_config=stem_diffusion_cfg,\n",
    "    tokens=stem_out_gpt[0],\n",
    "    input_tokens_finished=True,\n",
    ")\n",
    "\n",
    "stem_result = diffusion_engine.run_request(stem_request)\n",
    "\n",
    "# Decode the audio\n",
    "stem_vae_latents = torch.concat(stem_result.vae_latents)\n",
    "stem_audio = decode(stem_vae_latents)\n",
    "\n",
    "print(\"Stem conditioning test complete!\")\n",
    "print(f\"Generated audio duration: {stem_audio.duration_s:.2f} seconds\")\n",
    "\n",
    "# Play the original stem conditioning source\n",
    "print(\"\\nPlaying original stem conditioning source:\")\n",
    "clip_1 = SunoClip(gen_id)\n",
    "# clip_1.audio().play()\n",
    "\n",
    "print(\"\\nPlaying stem conditioned generation:\")\n",
    "stem_audio.play()\n",
    "\n",
    "\n",
    "mixed_audio = clip_1.audio() + stem_audio\n",
    "mixed_audio.play()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "527d7746",
   "metadata": {},
   "source": [
    "## Test sample conditioning"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "zyhxmkds4rm",
   "metadata": {},
   "outputs": [],
   "source": [
    "gen_id = \"a5e2198a-f352-4abb-9a24-7f81b143ded3\"  # stone\n",
    "\n",
    "clip = SunoClip(gen_id)\n",
    "print(list(clip._get_npz().keys()))\n",
    "stem_codes = clip._get_npz()[\"v3.5_raw\"][:, :1][600:800]\n",
    "clip.audio().get_segment(600 / 25, 800 / 25).play()\n",
    "stem_codes.shape"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a59ce906",
   "metadata": {},
   "outputs": [],
   "source": [
    "# Test stem conditioning with one of our existing clips\n",
    "import numpy as np\n",
    "from suno_utils.gpt.prompt import Prompt\n",
    "\n",
    "# Test stem conditioning configuration\n",
    "stem_gconf = GenerationConfig(\n",
    "    text=\"[mumble mode]}\",\n",
    "    text_tags=\"folk rock\",\n",
    "    cfg_coef=1.0,\n",
    "    cfg_coef_tags=0.0,\n",
    "    n_batch=1,\n",
    "    min_text_offset=0,\n",
    "    eos_pad_duration_s=0,\n",
    "    max_gen_duration_s=300,  # Shorter duration for testing\n",
    "    sample_arr=stem_codes,  # This is our new stem conditioning field\n",
    "    # stem_arr=stem_codes,  # This is our new stem conditioning field\n",
    ")\n",
    "\n",
    "print(\"Creating stem conditioning request...\")\n",
    "stem_requests = [\n",
    "    make_request(f\"stem_test_{i}\", stem_gconf, engine.model.config, engine.tokenizer) for i in range(1)\n",
    "]\n",
    "# Prompt(\"\", engine.model.config).visualize(stem_requests[0].streams[0].prompt, compress=True)\n",
    "# Prompt(\"\", engine.model.config).visualize(stem_requests[0].streams[0].prompt, compress=False)\n",
    "print(\"Running stem conditioning generation...\")\n",
    "stem_jobs = engine.run_request(stem_requests, tqdm_enabled=True)\n",
    "\n",
    "stem_out_gpt = []\n",
    "for n, job in enumerate(stem_jobs):\n",
    "    stream = engine.token_generator(job)\n",
    "    arr = torch.stack(list(stream))[:, 1]\n",
    "    if arr[-1] == 4000:  # Remove EOS token if present\n",
    "        arr = arr[:-1]\n",
    "    stem_out_gpt.append(arr)\n",
    "\n",
    "print(f\"Generated stem conditioning output shape: {stem_out_gpt[0].shape}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "9634755c",
   "metadata": {},
   "outputs": [],
   "source": [
    "# Run the stem conditioning output through the diffusion engine\n",
    "print(\"Running stem conditioning output through diffusion engine...\")\n",
    "\n",
    "stem_diffusion_cfg = diffusion_gen.DiffusionGenerationConfig(\n",
    "    ctx_cfg_coef=1.0,\n",
    "    steps=10,\n",
    "    codec_scale_factor=0.4,\n",
    "    scale_ctx_vector=True,\n",
    "    noise_ctx_level=0.75,\n",
    "    noise_ctx_pad_len=0,\n",
    ")\n",
    "\n",
    "stem_request = Request(\n",
    "    id=\"stem_test_diffusion\",\n",
    "    generation_config=stem_diffusion_cfg,\n",
    "    tokens=stem_out_gpt[0],\n",
    "    input_tokens_finished=True,\n",
    ")\n",
    "\n",
    "stem_result = diffusion_engine.run_request(stem_request)\n",
    "\n",
    "# Decode the audio\n",
    "stem_vae_latents = torch.concat(stem_result.vae_latents)\n",
    "stem_audio = decode(stem_vae_latents)\n",
    "\n",
    "print(\"Stem conditioning test complete!\")\n",
    "print(f\"Generated audio duration: {stem_audio.duration_s:.2f} seconds\")\n",
    "\n",
    "# Play the original stem conditioning source\n",
    "print(\"\\nPlaying original stem conditioning source:\")\n",
    "clip_1 = SunoClip(gen_id)\n",
    "# clip_1.audio().play()\n",
    "\n",
    "print(\"\\nPlaying stem conditioned generation:\")\n",
    "stem_audio.play()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a96d311e",
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
