{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"5\"\n",
    "\n",
    "import json\n",
    "from suno_utils.audio import Audio\n",
    "from suno_utils.utils.text import read_jsonl\n",
    "from suno_utils.utils.s3 import list_s3_dir, read_from_s3\n",
    "\n",
    "import torch\n",
    "import numpy as np\n",
    "\n",
    "from suno_utils.diffusion.generation import (\n",
    "    preload_dit_model,\n",
    "    preload_tokenizer,\n",
    "    TOKENIZER_FILEPATH,\n",
    "    SEMANTIC_MODEL_FILEPATH,\n",
    "    SEMANTIC_CLUSTERS_FILEPATH,\n",
    "    _retrieve_models,\n",
    "    #preload_ear_model,\n",
    ")\n",
    "\n",
    "\n",
    "from suno_utils.tasks.mert_25 import (\n",
    "    preload_models as preload_semantic_models,\n",
    "    encode as encode_semantic,\n",
    ")\n",
    "\n",
    "from suno_utils.tasks.dac_vae_fixed_25hz import (\n",
    "    preload_models as preload_codec_models,\n",
    "    decode as codec_decode,\n",
    "    decode_stream_to_full_audio,\n",
    ")\n",
    "\n",
    "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.upsample_engine_old import UpsampleEngine, Request, Job\n",
    "\n",
    "import sys\n",
    "sys.path.insert(0, \"/home/m4burns/glockenspiel/suno_utils/notebooks/shimmerscore/\")\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# load diffusion model\n",
    "num_gpus = torch.cuda.device_count()\n",
    "cuda_device = torch.cuda.current_device()\n",
    "print(f\"Found {num_gpus} GPUs. Using GPU {cuda_device}.\")\n",
    "\n",
    "# dit models\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-04-12_05-22-06_s1150/step_6000_ckpt.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-02-17_16-54-01_s7787/last_ckpt_infer.pt\" # base model\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-04-11_14-14-12_s4538/last_ckpt_infer.pt\" # audio tag finetune\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-04-10_01-05-22_s3784/last_ckpt_infer.pt\" # shared context\n",
    "#dit_model_filepath = \"s3://suno-data/christian/checkpoints/diffusion/v1_t6_5E6_beta100_n16_bt4.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-04-19_20-59-02_s6204/last_ckpt_infer.pt\" # ear finetune ongoing\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-04-21_21-18-08_s4974/last_ckpt_infer.pt\" # sft infill 100k\n",
    "#dit_model_filepath = \"/app/suno/data/dpo/models/diff_v2_2b_2mil_ft_infill_20250421_v2.pt\"\n",
    "#dit_model_filepath = \"/app/suno/modal/models/tony/tmp/diff/v45_2b_step_2mil_ft_8k_infill_apr21_t1_t6_cs0.pt\"\n",
    "dit_model_filepath = \"/app2/suno/modal/models/tony/tmp/diff/v45_2b_step_2mil_ft_8k_infill_apr21_d4_v39.pt\"\n",
    "\n",
    "# marc ckpts testing\n",
    "#dit_model_filepath = \"/app/suno/data/dpo/models/diff_v2_2b_2mil_ft_v0.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-04-10_04-13-39_s12/step_9000_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-04-16_19-50-48_s4352/last_ckpt_infer.pt\"\n",
    "\n",
    "# new infill\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-04-16_11-28-40_s3492/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-04-16_11-28-40_s3492/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-04-15_00-52-47_s9732/last_ckpt_infer.pt\"\n",
    "\n",
    "# dpo ckpts\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-04-25_15-17-24_s7084/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-04-25_16-19-23_s384/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-04-25_17-37-31_s8956/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-04-25_17-55-38_s9857/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-04-25_18-28-57_s5037/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-04-25_19-56-14_s7541/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-04-25_19-26-14_s1282/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-06-30_15-13-54_s814/last_ckpt_infer.pt\" # t28\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-14_04-22-20_s1097/last_ckpt_infer.pt\" # t38\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-15_15-57-35_s4091/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-15_21-22-20_s6636/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-14_15-27-20_s7143/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-16_16-30-58_s665/last_ckpt_infer.pt\"\n",
    "\n",
    "# sft ckpts\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-04-23_15-33-35_s2185/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-04-24_13-31-37_s3763/last_ckpt_infer.pt\"\n",
    "\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-04-11_14-12-15_s544/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-04-03_19-36-00_s7765/last_ckpt_infer.pt\"\n",
    "# 4b ckpts\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-05-09_19-28-35_s3585/step_370000_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-05-13_10-48-23_s3130/last_ckpt_infer.pt\"\n",
    "\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-05-15_17-55-31_s3782/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-06-02_21-00-56_s4958/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-06-03_00-09-44_s9599/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-06-04_16-09-11_s9069/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-06-04_18-05-00_s4398/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-06-04_19-29-45_s8700/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-06-05_13-25-20_s4365/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-06-05_15-53-31_s6230/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-07-02_13-58-09_s2849/last_ckpt_infer.pt\"\n",
    "\n",
    "# rectified flow\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-05-19_15-27-59_s697/last_ckpt_infer.pt\" #2b \n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-05-20_18-01-27_s320/last_ckpt_infer.pt\" #4bd\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-07-04_20-52-26_s300/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-07-24_21-01-48_s3483/last_ckpt_infer.pt\" # 2b sft\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-28_19-42-35_s3828/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-28_19-04-28_s6406/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-29_01-28-34_s2744/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-07-21_16-49-05_s324/last_ckpt_infer.pt\" # 30s -> 60s\n",
    "\n",
    "# distilled\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-05-28_14-55-42_s2329/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-06-06_15-39-38_s4725/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-06-24_11-51-11_s5757/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-07_18-31-04_s9125/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-08_10-31-28_s6751/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-08_18-08-55_s7740/last_ckpt_infer.pt\"#\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-10_10-25-15_s3025/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-10_14-24-29_s1186/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-11_13-59-09_s4679/last_ckpt_infer.pt\" # subtract\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-11_15-10-20_s252/last_ckpt_infer.pt\" # subtract higher LR\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-14_14-58-55_s7635/last_ckpt_infer.pt\" # 16n\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-14_15-23-51_s2636/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-14_19-46-40_s7238/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-15_20-53-46_s7434/last_ckpt_infer.pt\" # working!\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-16_18-03-26_s4029/last_ckpt_infer.pt\" # cfg\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-20_23-40-33_s5649/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-24_21-10-53_s9089/last_ckpt_infer.pt\" # 1-step only\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-24_21-22-42_s5920/last_ckpt_infer.pt\" # higher gen lr\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-28_11-34-51_s4168/last_ckpt_infer.pt\" # no ema 2.0 cfg, 5e-6\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-28_11-34-51_s4168/step_210k_infer.pt\" # no ema 2.0 cfg, 5e-6\n",
    "\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-28_13-08-21_s6706/last_ckpt_infer.pt\" # no ema 2.0 cfg, 5e-7\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-29_01-36-33_s947/last_ckpt_infer.pt\" # no ema 1.5 cfg\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-29_19-56-20_s8294/last_ckpt_infer.pt\" # N = 2\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-30_20-12-21_s1398/last_ckpt_infer.pt\" # 0.1 g_loss\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-28_11-34-51_s4168/step_210k_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-08-04_16-12-15_s8367/last_ckpt_infer.pt\" # residual flow cfg 1\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-08-04_17-33-57_s5109/last_ckpt_infer.pt\" # residual flow cfg 2\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-08-04_17-33-57_s5109/step_7k.pt\" # residual flow cfg 2\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-08-04_20-50-13_s3177/last_ckpt_infer.pt\" # residual flow cfg 1.5\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-08-05_12-58-07_s5517/last_ckpt_infer.pt\" # residual flow cfg 1.25\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-08-05_19-13-00_s1190/last_ckpt_infer.pt\" # residual flow cfg 1.25 aug\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-08-06_16-44-03_s5006/last_ckpt_infer.pt\" # residual flow cfg 1.25 + sft\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-08-11_10-20-26_s1039/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-08-19_20-55-03_s7168/last_ckpt_infer.pt\" # variable cfg distill\n",
    "\n",
    "\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-08-05_19-42-00_s2407/last_ckpt_infer.pt\" # dpo'd\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-08-11_17-42-28_s2024/last_ckpt_infer.pt\" # dpo'd\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-08-11_19-56-44_s4273/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-08-12_11-42-01_s2590/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-08-12_15-44-41_s6728/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-08-18_21-15-45_s9697/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-08-20_19-56-03_s5138/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-08-21_17-05-00_s7014/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-08-25_10-45-03_s3473/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-08-25_19-50-45_s330/last_ckpt_infer.pt\"\n",
    "\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/diffusion/16n_25hz_v45_infill_shared_flow_resume_1_75m.pt\"\n",
    "\n",
    "# sfx\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-07-30_17-01-42_s4725/last_ckpt_infer.pt\" # ft\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-07-30_16-16-20_s6206/last_ckpt_infer.pt\" # scratch\n",
    "\n",
    "\n",
    "#sft\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-08-01_14-11-57_s1109/last_ckpt_infer.pt\" # trad\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-08-01_15-59-37_s4709/last_ckpt_infer.pt\" # syn sft t1\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-08-01_14-11-57_s1109/step_10k.pt\" # syn sft t1, 1e-6\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-08-01_14-11-57_s1109/step_30k.pt\" # syn sft t1, 1e-6\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-08-01_14-11-57_s1109/last_ckpt_infer.pt\" # syn sft t1, 1e-6\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-08-04_10-39-54_s6463/last_ckpt_infer.pt\" # syn sft t2, 1e-6\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-08-04_11-37-48_s7572/last_ckpt_infer.pt\" # syn sft t2, 1e-6\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/diffusion/4n_25hz_2b_flow_5e5_sft_t8_500k.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-09-10_11-02-20_s8021/step_3000_infer.pt\"\n",
    "\n",
    "# adv\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-06-26_14-05-51_s8891/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-06-30_14-31-49_s3703/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-07_18-21-35_s4183/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-07_19-57-38_s6350/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-08_13-02-27_s8817/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-08_20-42-35_s9505/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-07-09_12-38-47_s9647/last_ckpt_infer.pt\"\n",
    "\n",
    "# no semantic flow\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-06-05_20-54-37_s4802/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-06-12_04-28-25_s9058/last_ckpt_infer.pt\"\n",
    "\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-09-18_15-47-55_s7983/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-09-26_00-57-42_s2763/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-09-27_02-43-54_s1279/last_ckpt_infer.pt\" # artist hash \n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-09-28_01-22-02_s4103/last_ckpt_infer.pt\" # artist hash no dropout\n",
    "\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-10-01_14-31-46_s4317/last_ckpt_infer.pt\" # dpo + distill\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-10-01_21-12-10_s7928/last_ckpt_infer.pt\"\n",
    "\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-10-10_18-29-43_s9245/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-10-10_21-32-41_s8045/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-10-13_11-48-58_s6473/last_ckpt_infer.pt\"\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-10-14_00-24-25_s7951/last_ckpt_infer.pt\" # pretrain 16n\n",
    "#dit_model_filepath = \"/app/suno/checkpoints/2025-10-15_01-07-00_s3307/last_ckpt_infer.pt\" # pretrain actnorm 2 nodes\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-10-17_10-44-59_s2569/last_ckpt_infer.pt\" # init 4n\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-10-17_17-14-47_s5386/last_ckpt_infer.pt\" # 16n 8 hr new data\n",
    "#dit_model_filepath = \"/app2/suno/checkpoints/2025-10-19_21-39-31_s7424/last_ckpt_infer.pt\" # 8n new data\n",
    "dit_model_filepath = \"/app2/suno/checkpoints/2025-10-21_00-20-19_s819/last_ckpt_infer.pt\"\n",
    "dit_model_filepath = \"/app2/suno/checkpoints/2025-10-24_01-10-49_s116/last_ckpt_infer.pt\"\n",
    "\n",
    "# other models\n",
    "tokenizer_filepath = \"s3://suno-data/georg/models/tokenizers/tokenizer_60k.json\"\n",
    "semantic_model_filepath = \"s3://suno-data/georg/models/semantic/mert_25.pt\"\n",
    "semantic_clusters_filepath = (\n",
    "    \"s3://suno-data/georg/models/semantic/mert_25_2x4k.npy\"\n",
    ")\n",
    "codec_filepath = \"s3://suno-data/minz/models/dac_vae_tuned_25hz.pth\"\n",
    "CODEC_SCALE_FACTOR = 0.4\n",
    "SCALE_CTX_VECTOR = True\n",
    "\n",
    "_ = diffusion_gen.preload_dit_model(\n",
    "    dit_model_filepath=dit_model_filepath,\n",
    "    use_ema_if_exists=True,\n",
    "    compile=False,\n",
    "    weights_precision=torch.bfloat16,\n",
    ")\n",
    "_ = preload_tokenizer(tokenizer_filepath)\n",
    "_ = preload_semantic_models(semantic_model_filepath, semantic_clusters_filepath)\n",
    "_ = preload_codec_models(codec_filepath)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "#ear_model_filepath = \"/app2/suno/checkpoints/2025-09-05_15-51-22_s7077/last_ckpt.pt\"\n",
    "ear_model_filepath = \"s3://suno-data/christian/checkpoints/ear/ear_v3_s8558.pt\"\n",
    "_ = preload_ear_model(ear_model_filepath)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# load the test clip here\n",
    "#item_id = \"a5e2198a-f352-4abb-9a24-7f81b143ded3\" # stone\n",
    "#item = work_items[10]\n",
    "#item_id = \"c75e6dc3-de5a-4a18-aae4-a24ac25dc850\"\n",
    "#item_id = \"fce5c7de-27db-409c-be2b-8ace15720c2b\"\n",
    "#item_id = \"1ec51ef1-378e-4f33-94a5-05c9b1e1448f\" \n",
    "#item_id = \"1210dc12-03d2-48fc-a51c-87d02b3517c5\" #long\n",
    "#item_id = \"ea86b583-913a-45df-9401-c48f01f4caf3\"\n",
    "#item_id = \"dfcbd035-2b9d-4c43-89ba-39411d0c6e43\"\n",
    "#item_id = \"a0438b2e-c6d2-4119-8bc6-863949b04d76\"\n",
    "#item_id = \"0ccb861e-8fc9-4a58-bc3f-e7c3e0bbd18b\"\n",
    "#item_id = \"b37c1380-7ad5-4a1b-8608-aabc26af42a1\"\n",
    "#item_id = \"44e11913-4b48-4f31-9c26-5bb19105b474\" # dna\n",
    "#item_id = \"0b0268b5-1599-4888-985f-46de37d42cd0\" # feels so good\n",
    "#item_id = \"1025518a-ace8-41c9-b979-13b07ae4d301\"\n",
    "#item_id = \"7868c62c-0d0d-4355-a35f-8a5d8ecc234f\"\n",
    "#item_id = \"1bb8ea30-31a2-43a6-8b94-ba15fe1e4eff\"\n",
    "#item_id = \"ace679d6-0c59-4#d6f-893f-951b80f33490\"\n",
    "#item_id = \"6c425fb1-0194-4995-82e0-dc2764be8bb3\"\n",
    "#item_id = \"e88fc58b-3e7f-4401-aae1-62207d89e140\" # mumble singing\n",
    "#item_id = \"b0d17aa0-8cb6-4db3-89ac-2c89d89f742a\"\n",
    "#item_id = \"26824a36-a837-4cca-95a9-e9458868af05\"\n",
    "#item_id = \"04110bb9-bb25-4afb-9cd6-7b987f3510a4\"\n",
    "#item_id = \"d09b0bb9-20cc-4228-b92a-700f884c84dc\"\n",
    "#item_id = \"a06329ca-9453-4730-82b6-f48576fe4530\"\n",
    "#item_id = \"13792213-e460-4e2f-b1db-2e7894ae6221\"\n",
    "#item_id = \"de5e5ecd-2616-4f3f-a856-392111badc3f\"\n",
    "#item_id = \"eaf614e7-bac5-4746-baa5-f4885bb14080\"\n",
    "#item_id = \"25461a13-9afa-4a87-9094-7a85908d2f34\"\n",
    "#item_id = \"bd8a3225-7753-4a64-88e6-26da0eee35c9\"\n",
    "#item_id = \"520863e7-e5da-4365-97a6-31832e0a9f06\"\n",
    "#item_id = \"adad8d8e-5e9f-4355-a3bb-605152adbe52\" # hot to go\n",
    "#item_id = \"4f23519e-d301-43aa-9f45-287b90abfe8e\" # hot to go 80s\n",
    "#item_id = \"5a4c3600-309a-452b-ba74-6e634a84bf63\" # stone vaporwave\n",
    "#item_id = \"e0ecb0ef-13f3-4dff-b0cb-3ce259719ac3\" # between the two\n",
    "#item_id = \"aae216ae-e05e-477e-b173-b7b42daabfc6\" # beautiful dream\n",
    "#item_id = \"0a4e795d-2373-424e-9c69-d3098f4934a3\" # she cries \n",
    "#item_id = \"ebbddbbc-f0fa-400e-9789-94978fb0bf05\" # did i stutter\n",
    "\n",
    "#item_id = \"f24dfa54-7ee6-499e-b7bc-3ff4c9c8f0af\" # hello\n",
    "#item_id = \"4dcbe31d-a57c-4284-9cb1-10298d5532a1\"\n",
    "#item_id = \"76e54a97-f20b-4707-8ff7-51295d9a598f\"\n",
    "#item_id = \"9f71fed1-4419-4d03-a814-198955a684d3\" # dorado\n",
    "\n",
    "# shimmer tracks\n",
    "#item_id = \"ce01fb47-511a-4114-a1ab-11ca7948bc4b\"\n",
    "#item_id = \"4b403c33-ebf6-452d-a6ea-14b8fcf49f29\"\n",
    "\n",
    "# metal \n",
    "#item_id = \"8d558fea-02d4-4138-8bf6-ea2b7424d923\"\n",
    "\n",
    "# opera\n",
    "item_id = \"44b8d3f6-a5a8-44b1-853e-6f1a2fcf4301\"\n",
    "\n",
    "#item_id = \"7198b4b0-cd68-4576-a33d-f6480c2d03f5\"\n",
    "#item_id = \"3df0dfaf-30d9-44df-ac58-818d26858b67\"\n",
    "#item_id = \"bd6d76c3-43b2-4952-a8ca-2057146ac263\"\n",
    "#item_id = \"5a4c3600-309a-452b-ba74-6e634a84bf63\"\n",
    "#item_id = \"ef5a404c-1582-4d90-a25b-691a51e8a9de\"\n",
    "\n",
    "# auk clips\n",
    "#item_id = \"7f6d8870-69c4-49ba-9530-b4971bcf401d\"\n",
    "#item_id = \"560618f4-58df-4303-b10d-13f5135cc300\"\n",
    "#item_id = \"a44b8af6-de6e-4228-b2fb-e7cec751700e\" # test test fuck you pat\n",
    "\n",
    "#item_id = \"91ae68ac-8a8c-46b0-bc60-b90029ab5383\"\n",
    "#item_id = \"a675a2b4-cc38-4d6a-84e1-5095938e3979\"\n",
    "#item_id = \"6232308c-c294-4d71-97d0-b54385a1647d\" # rings cover\n",
    "#item_id = \"428c0d70-5270-4856-9a43-de05f72f8c1f\" # john mayer persona\n",
    "#item_id = \"d296aeb2-23e3-4618-b394-0c74a782b167\" # brewster vox gen\n",
    "#item_id = \"88ec703b-8d9f-43ab-8b59-0176c71f9338\" # brewster vox gen\n",
    "\n",
    "# vox clone forth wanderers\n",
    "#item_id = \"fc9b27df-48d8-4723-81b9-240a419885e8\" # care for ya by forth wanderers\n",
    "#item_id = \"b48f1adc-7f6e-4712-9304-7246a0276259\" #\n",
    "#item_id = \"2777a4af-b552-4b13-85f3-878065aee218\"  #7months vocal\n",
    "#item_id = \"a20580e6-6250-461b-9b63-f1b62e57b418\" # 7months full track \n",
    "#item_id = \"64afb3de-51b1-41ff-be5a-c14736efe978\" # pinegrove rings\n",
    "#item_id = \"426d9bed-5404-4147-b280-6b03f7c47d11\" # ariana grande - 7 rings\n",
    "#item_id = \"57ad0d51-de89-48b2-a918-15dc7504b067\" # blank space by taylor swift\n",
    "#item_id = \"a1f2dcac-97e6-4b80-bba8-de46b195ffea\" # brewster, legacy lives on#\n",
    "\n",
    "print(item_id)\n",
    "# load the semantic codes from s3\n",
    "s3_filepath = f\"s3://suno-data-uploads/studio/uploads/{item_id}.npz\"\n",
    "\n",
    "try:\n",
    "    data = read_from_s3(s3_filepath, read_f=np.load)\n",
    "except Exception as e:\n",
    "    print(f\"Error loading {s3_filepath}: {e}\")\n",
    "\n",
    "if \"v3.0_raw\" in data:\n",
    "    codes = data[\"v3.0_raw\"]\n",
    "elif \"v3.5_raw\" in data:\n",
    "    codes = data[\"v3.5_raw\"]\n",
    "elif \"v4.0_raw\" in data:\n",
    "    codes = data[\"v4.0_raw\"]\n",
    "elif \"v5.0_raw\" in data:\n",
    "    codes = data[\"v5.0_raw\"]\n",
    "else:\n",
    "    raise ValueError(\"No codes found\")\n",
    "\n",
    "semantic_codes = torch.from_numpy(codes[:, 0]).long()  # .cuda()\n",
    "# semantic_codes = semantic_codes[:3000]\n",
    "print(semantic_codes.shape)\n",
    "\n",
    "# load the vae\n",
    "vae_data_filepath = f\"s3://suno-data-uploads/studio/uploads/{item_id}_vae.npz\"\n",
    "vae_data = read_from_s3(vae_data_filepath, read_f=np.load)\n",
    "print(vae_data.keys())\n",
    "init_vae_latents = vae_data[\"vae_latents\"]\n",
    "print(init_vae_latents.shape)\n",
    "\n",
    "gen_seed = vae_data.get(\"seed\", None)\n",
    "print(gen_seed)\n",
    "#tags_str = item[\"tags\"]\n",
    "\n",
    "#if tags_str is None:\n",
    "#    tags_str = \"\"\n",
    "#lyrics = item[\"text\"]\n",
    "\n",
    "os.system(f\"aws s3 cp s3://suno-data-uploads/studio/uploads/{item_id}_hoot.json text_data.json\")\n",
    "text_data = open(\"text_data.json\", \"r\", encoding=\"utf-8\").read()\n",
    "aligned_lyrics = json.loads(text_data)\n",
    "print(aligned_lyrics)\n",
    "\n",
    "tags = \"\"\n",
    "\n",
    "lyrics = \"\"\n",
    "for elem in aligned_lyrics:\n",
    "    if \"word\" in elem:\n",
    "        lyrics += elem[\"word\"]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "#item_id = \"2777a4af-b552-4b13-85f3-878065aee218\" # forth wanderers vox persona\n",
    "#tem_id = \"cbf58ff5-9a6d-4691-96a7-8c26902c61d2\" # taylor swift vox persona\n",
    "#item_id = \"83ebfdff-2390-4c24-b68c-b4bd04ed4394\" # john mayer\n",
    "#item_id = \"323d7c4e-32fb-4fd9-aed5-87a927e151eb\" # brewster\n",
    "item_id = \"37fe8c7f-6652-4765-8a2e-9626c1868ac6\" # brewster 2\n",
    "\n",
    "# load mp3 file \n",
    "#mp3_filepath = f\"s3://suno-data-uploads/studio/uploads/{item_id}.mp3\"\n",
    "#audio = Audio.from_file(mp3_filepath, n_channels=1)\n",
    "\n",
    "# load the vae\n",
    "vae_data_filepath = f\"s3://suno-data-uploads/studio/uploads/{item_id}_vae.npz\"\n",
    "vae_data = read_from_s3(vae_data_filepath, read_f=np.load)\n",
    "print(vae_data.keys())\n",
    "vox_vae_latents = vae_data[\"vae_latents\"]\n",
    "print(vox_vae_latents.shape)\n",
    "\n",
    "# decode to audio\n",
    "vox_audio = codec_decode(torch.from_numpy(vox_vae_latents))\n",
    "\n",
    "# adjust to mono\n",
    "vox_audio_mono = vox_audio.mono().stereo()\n",
    "\n",
    "# now encode back\n",
    "vox_vae_latents_mono = codec_encode(vox_audio_mono)\n",
    "print(vox_vae_latents_mono.shape)\n",
    "\n",
    "vox_vae_latents_mono = vox_vae_latents\n",
    "\n",
    "#semantic_codes = encode_semantic(audio_segment)\n",
    "\n",
    "start_s = .0\n",
    "end_s = start_s + 30.0\n",
    "\n",
    "start_samp = int(start_s * 25)\n",
    "end_samp = int(end_s * 25)\n",
    "\n",
    "# perform repeat padding and cropping\n",
    "latent_len_needed = 25 * 30  # 750 for 30 seconds at 25Hz\n",
    "\n",
    "# Crop first, then pad via repeat if needed\n",
    "extracted = vox_vae_latents[start_samp:end_samp]\n",
    "cur_len = extracted.shape[0]\n",
    "if cur_len < latent_len_needed:\n",
    "    repeats = (latent_len_needed + cur_len - 1) // cur_len  # repeats needed to cover target length\n",
    "    padded = np.tile(extracted, (repeats, 1)) if extracted.ndim > 1 else np.tile(extracted, repeats)\n",
    "    vox_vae_latents_cropped = padded[:latent_len_needed]\n",
    "else:\n",
    "    vox_vae_latents_cropped = extracted[:latent_len_needed]\n",
    "\n",
    "print(vox_vae_latents_cropped.shape)\n",
    "\n",
    "vox_audio = codec_decode(torch.from_numpy(vox_vae_latents_cropped))\n",
    "vox_audio.play()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "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,
   "metadata": {},
   "outputs": [],
   "source": [
    "# audio from disk\n",
    "\n",
    "lyrics = \"\"\"\n",
    "[Verse 1: Stevie Nicks]\n",
    "Now, here you go again\n",
    "You say you want your freedom\n",
    "Well, who am I to keep you down?\n",
    "It's only right that you should\n",
    "Play the way you feel it\n",
    "But listen carefully to the sound\n",
    "Of your loneliness\n",
    "\n",
    "[Pre-Chorus: Stevie Nicks]\n",
    "Like a heartbeat drives you mad (Heartbeat)\n",
    "In the stillness of rememberin' (Stillness)\n",
    "What you had and what you lost (Lonely, ooh)\n",
    "And what you had and what you lost (Ooh, ooh)\n",
    "\n",
    "[Chorus: Stevie Nicks & Lindsey Buckingham, Christine McVie]\n",
    "Oh, thunder only happens when it's rainin'\n",
    "Players only love you when they're playing\n",
    "Say, \"Women, they will come and they will go\"\n",
    "When the rain washes you clean, you'll know\n",
    "You'll know\n",
    "\n",
    "[Instrumental Break]\n",
    "[Verse 2: Stevie Nicks]\n",
    "Now, here I go again\n",
    "I see the crystal visions\n",
    "I keep my visions to myself\n",
    "It's only me who wants to\n",
    "Wrap around your dreams\n",
    "And have you any dreams you'd like to sell?\n",
    "Dreams of loneliness\n",
    "\n",
    "[Pre-Chorus: Stevie Nicks]\n",
    "Like a heartbeat drives you mad (Heartbeat)\n",
    "In the stillness of rememberin' (Stillness)\n",
    "What you had and what you lost (Lonely, ooh)\n",
    "Oh, what you had, oh, what you lost (Ooh, ah)\n",
    "\n",
    "[Chorus: Stevie Nicks & Lindsey Buckingham, Christine McVie]\n",
    "Thunder only happens when it's rainin'\n",
    "Players only love you when they're playing\n",
    "Women, they will come and they will go\n",
    "When the rain washes you clean, you'll know\n",
    "Oh, thunder only happens when it's rainin'\n",
    "Players only love you when they're playing\n",
    "Say, \"Women, they will come and they will go\"\n",
    "When the rain washes you clean, you'll know\n",
    "\n",
    "[Outro: Stevie Nicks & Lindsey Buckingham, Christine McVie]\n",
    "You'll know\n",
    "You will know\n",
    "Oh, you'll know\n",
    "\"\"\"\n",
    "\n",
    "# load audio file\n",
    "audio = Audio.from_file(\"/home/christian/audio/reference-audio-wav/02 Dreams.wav\", n_channels=2)\n",
    "audio = audio.get_segment(from_s=15, to_s=45.02)\n",
    "audio.play()\n",
    "\n",
    "sample_rate = 24000\n",
    "byte_width = 2\n",
    "n_channels = 1\n",
    "\n",
    "audio_segment = audio.convert(sample_rate, byte_width, n_channels)\n",
    "audio_segment.play()\n",
    "\n",
    "semantic_codes = encode_semantic(audio_segment)\n",
    "semantic_codes = torch.from_numpy(semantic_codes[:, 0]).long()\n",
    "# semantic_codes = semantic_codes[:3000]\n",
    "print(semantic_codes.shape)\n",
    "\n",
    "# also get the vae\n",
    "init_vae_latents = codec_encode(audio)\n",
    "print(init_vae_latents.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "\n",
    "from transformers import AutoTokenizer, AutoModel\n",
    "\n",
    "_mmbert_tokenizer = None\n",
    "_mmbert_encoder = None\n",
    "\n",
    "def load_mmbert_tokenizer_and_encoder():\n",
    "    \"\"\"Load mmBERT tokenizer and encoder\"\"\"\n",
    "    global _mmbert_tokenizer, _mmbert_encoder\n",
    "    if _mmbert_tokenizer is None or _mmbert_encoder is None:\n",
    "        _mmbert_tokenizer = AutoTokenizer.from_pretrained(\"jhu-clsp/mmBERT-base\")\n",
    "        _mmbert_encoder = AutoModel.from_pretrained(\"jhu-clsp/mmBERT-base\")\n",
    "    return _mmbert_tokenizer, _mmbert_encoder"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "mmbert_tokenizer, mmbert_encoder = load_mmbert_tokenizer_and_encoder()\n",
    "\n",
    "# Check if encoder is already on GPU to avoid redundant transfers\n",
    "current_device = next(mmbert_encoder.parameters()).device\n",
    "if current_device.type != \"cuda\":\n",
    "    mmbert_encoder = mmbert_encoder.cuda()\n",
    "\n",
    "# Encode text with mmBERT\n",
    "mmbert_text_arr = mmbert_tokenizer(lyrics, return_tensors=\"pt\").input_ids.cuda()\n",
    "print(mmbert_text_arr.shape)\n",
    "mmbert_text_arr = mmbert_encoder(mmbert_text_arr).last_hidden_state[0]\n",
    "print(mmbert_text_arr.shape)\n",
    "\n",
    "\n",
    "\n",
    "# Truncate if exceeds max tokens\n",
    "if mmbert_text_arr.shape[0] > 1024:\n",
    "    mmbert_text_arr = mmbert_text_arr[:1024]\n",
    "\n",
    "# nohup bash -c \"echo -e '\\n' | ./setup_suno_env_fa2_fixed_commit.sh suno_env_auto\" > setup_suno_env_auto.log 2>&1 &\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(len(lyrics))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "diffusion_engine = UpsampleEngine(min_chunk_size=25*30) #, vae_version=\"v_vae_25_tuned_2\")\n",
    "#diffusion_engine = UpsampleEngine(min_chunk_size=1500, vae_version=\"v_vae_25_tuned_2\") #, vae_version=\"v_vae_25_tuned_2\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print()\n",
    "# Map slider value from 0–1 to a target range\n",
    "def map_slider(slider_value, min_out, max_out):\n",
    "    \"\"\"Map slider from 0-1 range to min_out-max_out range, handling None.\"\"\"\n",
    "    if slider_value is None:\n",
    "        return None\n",
    "    slider_value = max(0, min(1, slider_value))  # clamp to 0–1\n",
    "    return min_out + (max_out - min_out) * slider_value\n",
    "\n",
    "# === Set global slider values from frontend input (0.0 to 1.0) ===\n",
    "slider_strength = 0.95   # example: strength from user input\n",
    "slider_quality = 0.5    # example: quality\n",
    "slider_tone = 0.5      # example: tone\n",
    "slider_stereo = 0.5    # example: stereo\n",
    "\n",
    "# === Map to model ranges ===\n",
    "strength_value = map_slider(slider_strength, 1.0, 4.0)     # cfg\n",
    "quality_value = map_slider(slider_quality, 12, 28)         # ear score\n",
    "tone_value = map_slider(slider_tone, -2.0, 2.0)            # spectral centroid\n",
    "stereo_value = map_slider(slider_stereo, -2.0, 2.0)        # stereo balance\n",
    "\n",
    "\n",
    "\n",
    "# === Format audio tags ===\n",
    "audio_tags_parts = []\n",
    "if quality_value is not None and quality_value != 20.0:\n",
    "    audio_tags_parts.append(f\"quality: {int(quality_value)}\")\n",
    "if tone_value is not None and tone_value != 0.0:\n",
    "    audio_tags_parts.append(f\"spectral_centroid: {tone_value:.1f}\")\n",
    "if stereo_value is not None and stereo_value != 0.0:\n",
    "    #if stereo_value <= -1.5:\n",
    "    #    audio_tags_parts.append(f\"mono\")\n",
    "    #else:/\n",
    "    audio_tags_parts.append(f\"stereo_width: {stereo_value:.1f}\")\n",
    "\n",
    "audio_tags = \", \".join(audio_tags_parts) if audio_tags_parts else \"\"\n",
    "print(f\"Prepared audio tags: {audio_tags}\")\n",
    "\n",
    "# === Example of using values in model logic ===\n",
    "text_cfg_coef = 2.5 if strength_value is None else strength_value\n",
    "tags = \"\"\n",
    "\n",
    "steps = 2\n",
    "\n",
    "print(f\"Using config: text_cfg_coef={text_cfg_coef}, steps={steps}, tags={tags}\")\n",
    "\n",
    "#diffusion_seeds = [0, 0, 0]\n",
    "#diffusion_seeds = np.random.randint(0, 1000000, size=2)\n",
    "#diffusion_seeds = 0\n",
    "diffusion_steps = [4]\n",
    "\n",
    "diffusion_text_cfg_coef = 1.0\n",
    "diffusion_ctx_cfg_coef = 1.0\n",
    "\n",
    "noise_ctx_level = 0.75\n",
    "noise_ctx_pad_len = 0\n",
    "CODEC_SCALE_FACTOR = 0.4\n",
    "\n",
    "t = [0.9999999403953552, 0.6418998837471008, 0.0]\n",
    "\n",
    "\n",
    "upsampled_audios = []\n",
    "use_vox_latents = [False, True, True]\n",
    "ctx_cfgs = [1.0, 1.0, 3.0]\n",
    "cfgs = [1.0, 3.0]\n",
    "for idx, seed in enumerate(np.random.randint(0, 1000000, size=3)):\n",
    "    print(f\"noise_ctx_level: {noise_ctx_level}\")\n",
    "    tags = \"\"\n",
    "    print(tags)\n",
    "    use_vox_latents = True\n",
    "    gen_cfg = diffusion_gen.DiffusionGenerationConfig(\n",
    "        steps=32,\n",
    "        lyrics=lyrics,\n",
    "        tags=\"guitar\",\n",
    "        text_cfg_coef=2.5,\n",
    "        ctx_cfg_coef=1.5,\n",
    "        codec_scale_factor=CODEC_SCALE_FACTOR,\n",
    "        scale_ctx_vector=True,\n",
    "        noise_ctx_level=noise_ctx_level,\n",
    "        noise_ctx_pad_len=noise_ctx_pad_len,\n",
    "        drop_semantic_tokens=False,\n",
    "        semantic_mask_ratio=0.0,\n",
    "        seed=int(seed),#np.random.randint(0, 1000000),\n",
    "        rho=1.0,\n",
    "        sigma_min=0.5,\n",
    "        sigma_max=50.0,\n",
    "        vox_latents=torch.from_numpy(vox_vae_latents_cropped).unsqueeze(0) if use_vox_latents else None,\n",
    "        objective=\"rectified_flow\",\n",
    "        #sampler_type=\"pingpong\",\n",
    "        #variation=variation,\n",
    "        #init_latents=torch.from_numpy(init_vae_latents) if variation != 1.0 else None,\n",
    "        #rank_candidates=rank_candidates,\n",
    "        #variation=variation,\n",
    "        #init_latents=torch.from_numpy(init_vae_latents * 0.4) if variation != 1.0 else None,\n",
    "        #variation_strength=0.2,\n",
    "        #sigmas=torch.tensor(t),\n",
    "        #semantic_mask_ratio=0.2,\n",
    "        #num_retries=0,\n",
    "    )\n",
    "\n",
    "    request = Request(\n",
    "        id=\"dummy\",\n",
    "        generation_config=gen_cfg,\n",
    "        tokens=semantic_codes[0:1500],\n",
    "        input_tokens_finished=True,\n",
    "    )\n",
    "\n",
    "    result = diffusion_engine.run_request(request)\n",
    "\n",
    "    vae_latents = []\n",
    "    for vae_latent in result.vae_latents:\n",
    "        #print(vae_latent.shape)\n",
    "        mean = vae_latent.mean()\n",
    "        std = vae_latent.std()\n",
    "        print(f\"mean: {mean}, std: {std}\")\n",
    "        vae_latents.append(vae_latent)\n",
    "\n",
    "    vae_latents = torch.concat(vae_latents)\n",
    "    #print(vae_latents.shape)\n",
    "    upsampled_audio = decode_stream_to_full_audio(vae_latents)\n",
    "    upsampled_audio.play()\n",
    "    upsampled_audios.append(upsampled_audio)\n",
    "\n",
    "    #v3_flow_distill_1E6_beta100_n8_bt2_acc2_4k"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "stereo_width_deltas = []\n",
    "total_octave_deltas = []\n",
    "\n",
    "for idx, audio in enumerate(upsampled_audios):\n",
    "    results = analyze_audio_first_last(torch.from_numpy(audio.array_float), 48000)\n",
    "    stereo_width_delta = results['stereo_width_delta']\n",
    "    print(f\"{idx+1}: stereo width delta: {stereo_width_delta}\")\n",
    "    total_octave_delta = results['total_delta']\n",
    "    print(f\"{idx+1}: total octave delta: {total_octave_delta}\")\n",
    "    print()\n",
    "    stereo_width_deltas.append(stereo_width_delta)\n",
    "    total_octave_deltas.append(total_octave_delta)\n",
    "\n",
    "combined_score = combine_width_octave(stereo_width_deltas, total_octave_deltas)\n",
    "print(f\"combined score: {combined_score}\")\n",
    "# get the index of the lowest combined score\n",
    "lowest_combined_score_index = np.argmin(combined_score)\n",
    "print(f\"lowest combined score index: {lowest_combined_score_index+1}\")\n",
    "# get the upsampled audio with the lowest combined score\n",
    "lowest_combined_score_audio = upsampled_audios[lowest_combined_score_index]\n",
    "lowest_combined_score_audio.play()\n",
    "\n",
    "# get the index of the highest combined score\n",
    "highest_combined_score_index = np.argmax(combined_score)\n",
    "print(f\"highest combined score index: {highest_combined_score_index+1}\")\n",
    "# get the upsampled audio with the highest combined score\n",
    "highest_combined_score_audio = upsampled_audios[highest_combined_score_index]\n",
    "highest_combined_score_audio.play()\n",
    "\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch\n",
    "import torch.fft as fft\n",
    "import numpy as np\n",
    "\n",
    "import numpy as np\n",
    "\n",
    "def combine_width_octave(\n",
    "    width_deltas,          # array-like of stereo width deltas (can be signed)\n",
    "    octave_totals,         # array-like of total octave deltas (>=0)\n",
    "    method=\"robust\"        # \"robust\" or \"minmax\"\n",
    "):\n",
    "    w = np.abs(np.asarray(width_deltas, dtype=float))\n",
    "    o = np.asarray(octave_totals, dtype=float)\n",
    "\n",
    "    def robust_norm(x):\n",
    "        med = np.median(x)\n",
    "        q1, q3 = np.percentile(x, 25), np.percentile(x, 75)\n",
    "        iqr = max(q3 - q1, 1e-12)\n",
    "        return (x - med) / iqr\n",
    "\n",
    "    def minmax_norm(x):\n",
    "        xmin, xmax = float(np.min(x)), float(np.max(x))\n",
    "        return (x - xmin) / (max(xmax - xmin, 1e-12))\n",
    "\n",
    "    if method == \"robust\":\n",
    "        z_w, z_o = robust_norm(w), robust_norm(o)\n",
    "        combo = 0.5 * (z_w + z_o)\n",
    "        # map to 0–100 for readability\n",
    "        score = 100 * minmax_norm(combo)\n",
    "    elif method == \"minmax\":\n",
    "        n_w, n_o = minmax_norm(w), minmax_norm(o)\n",
    "        score = 100 * (0.5 * (n_w + n_o))\n",
    "    else:\n",
    "        raise ValueError(\"method must be 'robust' or 'minmax'\")\n",
    "\n",
    "    return score  # higher = more change; use (100 - score) for stability\n",
    "\n",
    "\n",
    "def third_octave_bands(sr, fmin=20.0, fmax=None):\n",
    "    \"\"\"\n",
    "    Compute 1/3-octave band center frequencies and edges.\n",
    "    \"\"\"\n",
    "    if fmax is None:\n",
    "        fmax = sr / 2.0\n",
    "\n",
    "    k = np.arange(-30, 30)  # wide enough range\n",
    "    f_center = 1000.0 * (2.0 ** (k / 3.0))  # ISO 1/3 octave centers\n",
    "    f_center = f_center[(f_center >= fmin) & (f_center <= fmax)]\n",
    "    \n",
    "    f_lower = f_center / (2 ** (1/6))\n",
    "    f_upper = f_center * (2 ** (1/6))\n",
    "    return f_center, f_lower, f_upper\n",
    "\n",
    "\n",
    "def third_octave_response_db(waveform: torch.Tensor, sr: int):\n",
    "    \"\"\"\n",
    "    Compute 1/3 octave magnitude response in dB from waveform.\n",
    "    \n",
    "    Args:\n",
    "        waveform (torch.Tensor): shape (n_samples,) or (1, n_samples)\n",
    "        sr (int): sample rate\n",
    "    \n",
    "    Returns:\n",
    "        freqs (np.ndarray): band center frequencies\n",
    "        mags_db (torch.Tensor): band magnitudes in dB\n",
    "    \"\"\"\n",
    "    if waveform.ndim > 1:\n",
    "        waveform = waveform.squeeze(0)\n",
    "    \n",
    "    n = waveform.numel()\n",
    "    spec = fft.rfft(waveform)\n",
    "    mag = torch.abs(spec) / n\n",
    "    freqs = torch.fft.rfftfreq(n, d=1.0/sr)\n",
    "\n",
    "    # Get bands\n",
    "    f_center, f_lower, f_upper = third_octave_bands(sr)\n",
    "    band_mags = []\n",
    "    for fl, fu in zip(f_lower, f_upper):\n",
    "        idx = (freqs >= fl) & (freqs < fu)\n",
    "        if idx.any():\n",
    "            band_mags.append(mag[idx].mean())\n",
    "        else:\n",
    "            band_mags.append(torch.tensor(0.0))\n",
    "\n",
    "    band_mags = torch.stack(band_mags)\n",
    "\n",
    "    # Convert to dB (avoid log(0))\n",
    "    mags_db = 20 * torch.log10(band_mags + 1e-12)\n",
    "    \n",
    "    return f_center, mags_db\n",
    "\n",
    "def stereo_width(waveform: torch.Tensor):\n",
    "    \"\"\"\n",
    "    Compute the stereo width of a waveform.\n",
    "    \"\"\"\n",
    "    # can you implement this?\n",
    "    # Assume waveform shape is (2, seq_len)\n",
    "    if waveform.ndim != 2 or waveform.shape[0] != 2:\n",
    "        raise ValueError(\"waveform must have shape (2, seq_len) for stereo width calculation\")\n",
    "    left = waveform[0]\n",
    "    right = waveform[1]\n",
    "    # Compute correlation coefficient between L and R\n",
    "    left = left - left.mean()\n",
    "    right = right - right.mean()\n",
    "    numerator = (left * right).mean()\n",
    "    denominator = torch.sqrt((left ** 2).mean() * (right ** 2).mean()) + 1e-12\n",
    "    corr = numerator / denominator\n",
    "    # Stereo width: 0 = mono, 1 = fully wide (L and R uncorrelated), -1 = fully out of phase\n",
    "    width = torch.sqrt(1 - corr ** 2)\n",
    "    return width.item()\n",
    "\n",
    "import torch\n",
    "from typing import Optional, Dict, Any\n",
    "\n",
    "# Assumes third_octave_bands, third_octave_response_db, stereo_width\n",
    "# are defined exactly as in your snippet above.\n",
    "\n",
    "def analyze_audio_first_last(\n",
    "    audio: torch.Tensor,\n",
    "    sample_rate: int,\n",
    "    segment_duration: int = 60,\n",
    ") -> Dict[str, Any]:\n",
    "    \"\"\"\n",
    "    Compute 1/3-octave response and stereo width for the first and last segments\n",
    "    of a single audio tensor, and return their deltas.\n",
    "\n",
    "    Args:\n",
    "        audio: Tensor of shape (channels, samples) or (samples,)\n",
    "        sample_rate: sample rate in Hz\n",
    "        segment_duration: segment length in seconds for the first/last comparison\n",
    "\n",
    "    Returns:\n",
    "        {\n",
    "            'center_freqs': np.ndarray,\n",
    "            'third_octave_first': Tensor[dB],\n",
    "            'third_octave_last': Tensor[dB],\n",
    "            'delta_response': Tensor[dB],      # first - last\n",
    "            'total_delta': Tensor[scalar],     # L1 magnitude of delta_response\n",
    "            'stereo_width_first': float or None,\n",
    "            'stereo_width_last': float or None,\n",
    "            'stereo_width_delta': float or None,  # first - last\n",
    "        }\n",
    "    \"\"\"\n",
    "    if not isinstance(audio, torch.Tensor):\n",
    "        raise TypeError(\"audio must be a torch.Tensor\")\n",
    "    if audio.numel() == 0:\n",
    "        raise ValueError(\"audio is empty\")\n",
    "\n",
    "    # Determine segment length in samples (clip to available length)\n",
    "    seg_len = min(audio.shape[-1], segment_duration * sample_rate)\n",
    "\n",
    "    # Helper: mono mixdown for 1/3-octave analysis\n",
    "    mono = audio.mean(dim=0) if audio.ndim > 1 else audio\n",
    "    first_mono = mono[:seg_len]\n",
    "    last_mono  = mono[-seg_len:]\n",
    "\n",
    "    # 1/3-octave responses (dB)\n",
    "    center_freqs_first, oct_first = third_octave_response_db(first_mono, sample_rate)\n",
    "    center_freqs_last,  oct_last  = third_octave_response_db(last_mono,  sample_rate)\n",
    "\n",
    "    # Centers should match; keep the first as canonical\n",
    "    if len(center_freqs_first) != len(center_freqs_last) or (center_freqs_first != center_freqs_last).any():\n",
    "        raise RuntimeError(\"Mismatched third-octave centers between first and last segments.\")\n",
    "\n",
    "    delta = oct_first - oct_last\n",
    "    total_delta = delta.abs().sum()\n",
    "\n",
    "    # Stereo width (if stereo input)\n",
    "    def maybe_width(x: torch.Tensor) -> Optional[float]:\n",
    "        if x.ndim == 2 and x.shape[0] == 2 and x.shape[1] > 0:\n",
    "            return stereo_width(x)\n",
    "        return None\n",
    "\n",
    "    first_full = audio[:, :seg_len] if audio.ndim == 2 else audio\n",
    "    last_full  = audio[:, -seg_len:] if audio.ndim == 2 else audio\n",
    "\n",
    "    width_first = maybe_width(first_full)\n",
    "    width_last  = maybe_width(last_full)\n",
    "    width_delta = None\n",
    "    if (width_first is not None) and (width_last is not None):\n",
    "        width_delta = float(width_first - width_last)\n",
    "\n",
    "    return {\n",
    "        \"center_freqs\": center_freqs_first,     # np.ndarray\n",
    "        \"third_octave_first\": oct_first,        # Tensor[dB]\n",
    "        \"third_octave_last\": oct_last,          # Tensor[dB]\n",
    "        \"delta_response\": delta,                # Tensor[dB]\n",
    "        \"total_delta\": total_delta,             # Tensor[scalar]\n",
    "        \"stereo_width_first\": width_first,      # float or None\n",
    "        \"stereo_width_last\": width_last,        # float or None\n",
    "        \"stereo_width_delta\": width_delta,      # float or None\n",
    "    }\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import matplotlib.pyplot as plt\n",
    "\n",
    "\n",
    "fig, (ax1, ax2) = plt.subplots(2, 1, figsize=(10, 6), sharex=True)\n",
    "\n",
    "for quality_level, upsampled_audio in zip(quality_levels, upsampled_audios):\n",
    "    spectrum_mid, spectrum_side = calculate_average_stereo_spectrum(upsampled_audio.array_float, 48000)\n",
    "\n",
    "    # Frequency axis setup\n",
    "    n_bins = np.array(spectrum_mid).shape[0]\n",
    "    freqs = np.linspace(0, 48000 / 2, n_bins)\n",
    "\n",
    "    ax1.plot(freqs, spectrum_mid, label=quality_level)\n",
    "    ax2.plot(freqs, spectrum_side, label=quality_level)\n",
    "\n",
    "    # Configure Mid plot\n",
    "    ax1.set_xscale('log')\n",
    "    ax1.set_ylabel('Mid Channel (dB)', fontsize=12)\n",
    "    ax1.set_title('Average Mid Spectrum per Model', fontsize=14)\n",
    "    ax1.grid(True, which='both', linestyle='--', alpha=0.5)\n",
    "    ax1.set_ylim(-60, 48)\n",
    "    ax1.legend()\n",
    "\n",
    "    # Configure Side plot\n",
    "    ax2.set_xscale('log')\n",
    "    ax2.set_xlabel('Frequency (Hz)', fontsize=12)\n",
    "    ax2.set_ylabel('Side Channel (dB)', fontsize=12)\n",
    "    ax2.set_title('Average Side Spectrum per Model', fontsize=14)\n",
    "    ax2.grid(True, which='both', linestyle='--', alpha=0.5)\n",
    "    ax2.set_ylim(-60, 48)\n",
    "    ax2.set_xlim(20, 24000)\n",
    "    #ax2.legend()\n",
    "\n",
    "plt.tight_layout()\n",
    "plt.show()\n",
    "#plt.savefig(f\"{plot_dir}/average_mid_side_spectrum_per_model.png\")\n",
    "\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def calculate_average_spectrum_db(waveform, n_fft=16384, hop_length=8192):\n",
    "    # Keep as torch tensor or convert to torch tensor if it's numpy\n",
    "    if not isinstance(waveform, torch.Tensor):\n",
    "        waveform = torch.from_numpy(waveform)\n",
    "\n",
    "    # Calculate the average spectrum using STFT for efficiency\n",
    "    n_fft = 2048  # Choose an appropriate FFT size\n",
    "    hop_length = n_fft // 4  # Standard hop length\n",
    "\n",
    "    # Compute STFT using torch\n",
    "    if waveform.dim() > 1:\n",
    "        # For stereo, compute STFT for each channel\n",
    "        stft_results = []\n",
    "        for channel in range(waveform.shape[0]):\n",
    "            stft = torch.stft(\n",
    "                waveform[channel],\n",
    "                n_fft=n_fft,\n",
    "                hop_length=hop_length,\n",
    "                window=torch.hann_window(n_fft),\n",
    "                return_complex=True,\n",
    "            )\n",
    "            # Get magnitude\n",
    "            stft_magnitude = torch.abs(stft)\n",
    "            stft_results.append(stft_magnitude)\n",
    "\n",
    "        # Average across time frames for each channel\n",
    "        magnitude_spectrum = torch.stack(\n",
    "            [torch.mean(stft, dim=1) for stft in stft_results]\n",
    "        )\n",
    "    else:\n",
    "        # For mono\n",
    "        stft = torch.stft(\n",
    "            waveform,\n",
    "            n_fft=n_fft,\n",
    "            hop_length=hop_length,\n",
    "            window=torch.hann_window(n_fft),\n",
    "            return_complex=True,\n",
    "        )\n",
    "        # Get magnitude\n",
    "        stft_magnitude = torch.abs(stft)\n",
    "        magnitude_spectrum = torch.mean(stft_magnitude, dim=1)\n",
    "\n",
    "    # Convert to dB scale\n",
    "    spectrum_db = 20 * torch.log10(\n",
    "        magnitude_spectrum + 1e-10\n",
    "    )  # Adding small value to avoid log(0)\n",
    "\n",
    "    # Convert to numpy for consistency with the rest of the code\n",
    "    return spectrum_db.numpy()\n",
    "\n",
    "\n",
    "def calculate_average_stereo_spectrum(waveform, sr):\n",
    "    if not isinstance(waveform, torch.Tensor):\n",
    "        waveform = torch.from_numpy(waveform)\n",
    "    assert waveform.dim() == 2 and waveform.shape[0] == 2\n",
    "    # split into left and right channels\n",
    "    left = waveform[0]\n",
    "    right = waveform[1]\n",
    "\n",
    "    # compute mid and side channels\n",
    "    mid = (left + right) / 2\n",
    "    side = (left - right) / 2\n",
    "\n",
    "    # calculate spectrum for mid and side channels\n",
    "    spectrum_mid = calculate_average_spectrum_db(mid, sr)\n",
    "    spectrum_side = calculate_average_spectrum_db(side, sr)\n",
    "\n",
    "    return spectrum_mid, spectrum_side"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import math\n",
    "import torch\n",
    "import torch.nn as nn\n",
    "\n",
    "\n",
    "class SinusoidalPositionalEncoding(nn.Module):\n",
    "    def __init__(self, dim, max_len=2048):\n",
    "        super().__init__()\n",
    "        pe = torch.zeros(max_len, dim)\n",
    "        position = torch.arange(0, max_len).unsqueeze(1)\n",
    "        div_term = torch.exp(torch.arange(0, dim, 2) * -(math.log(10000.0) / dim))\n",
    "        pe[:, 0::2] = torch.sin(position * div_term)\n",
    "        pe[:, 1::2] = torch.cos(position * div_term)\n",
    "        self.register_buffer(\"pe\", pe)\n",
    "\n",
    "    def forward(self, x):\n",
    "        # x: (B, T, D)\n",
    "        seq_len = x.size(1)\n",
    "        return x + self.pe[:seq_len].unsqueeze(0).to(x.dtype)  # (1, T, D)\n",
    "\n",
    "\n",
    "class SimpleTransformerEncoder(nn.Module):\n",
    "    def __init__(self, vae_dim, embed_dim=768, num_layers=6, num_heads=12, ff_dim=2048, dropout=0.1, max_len=2048):\n",
    "        super().__init__()\n",
    "        self.input_proj = nn.Linear(vae_dim, embed_dim)\n",
    "        self.pos_encoding = SinusoidalPositionalEncoding(embed_dim, max_len)\n",
    "\n",
    "        encoder_layer = nn.TransformerEncoderLayer(\n",
    "            d_model=embed_dim,\n",
    "            nhead=num_heads,\n",
    "            dim_feedforward=ff_dim,\n",
    "            dropout=dropout,\n",
    "            batch_first=True\n",
    "        )\n",
    "        self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)\n",
    "        self.norm = nn.LayerNorm(embed_dim)\n",
    "        self.output_head = nn.Linear(embed_dim, 1)\n",
    "\n",
    "    def forward(self, x):\n",
    "        # x: (B, T, vae_dim)\n",
    "        x = self.input_proj(x)                      # (B, T, embed_dim)\n",
    "        x = self.pos_encoding(x)                   # add positional encodings\n",
    "        x = self.transformer(x)                    # (B, T, embed_dim)\n",
    "        x = self.norm(x)\n",
    "        return self.output_head(x)                 # (B, T, 1)\n",
    "\n",
    "def load_checkpoint(checkpoint_path, device):\n",
    "    checkpoint = torch.load(checkpoint_path)\n",
    "    model = SimpleTransformerEncoder(**checkpoint[\"run_config\"][\"model\"])\n",
    "    state_dict = checkpoint[\"model\"]\n",
    "    state_dict = {k.replace(\"_orig_mod.\", \"\"): v for k, v in state_dict.items()}\n",
    "    model.load_state_dict(state_dict)\n",
    "    model.eval()\n",
    "    model.to(device)\n",
    "    return model"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Test-time compute"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "#ear_ckpt = \"/app2/suno/checkpoints/2025-07-26_10-27-27_s7301/last_ckpt.pt\"\n",
    "#ear_model = load_checkpoint(ear_ckpt, device=\"cuda\")\n",
    "\n",
    "from suno_utils.tasks.ear import load_model\n",
    "model = load_model(\"s3://suno-data/christian/checkpoints/ear/ear_v2_s3080.pt\", compile=True)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.tasks.ear_v3 import load_model\n",
    "model = load_model(\"/app2/suno/checkpoints/2025-08-04_19-38-58_s1811/last_ckpt.pt\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "num_runs = 10\n",
    "diffusion_seeds = np.random.randint(0, 1000000, size=num_runs)\n",
    "diffusion_steps = 16\n",
    "\n",
    "diffusion_text_cfg_coef = 2.5\n",
    "diffusion_ctx_cfg_coef = 1.0\n",
    "\n",
    "noise_ctx_level = 0.0\n",
    "noise_ctx_pad_len = 0\n",
    "CODEC_SCALE_FACTOR = 0.4\n",
    "\n",
    "upsampled_audios = []\n",
    "for idx, diffusion_seed in enumerate(diffusion_seeds):\n",
    "    diffusion_seed = np.random.randint(0, 1000000)\n",
    "    diffusion_text_cfg_coef = np.random.uniform(1.0, 3.0)\n",
    "    diffusion_steps = np.random.randint(8, 32)\n",
    "    diffusion_noise_ctx_level = np.random.uniform(0.0, 0.75)\n",
    "\n",
    "    gen_cfg = diffusion_gen.DiffusionGenerationConfig(\n",
    "        steps=diffusion_steps,\n",
    "        lyrics=lyrics,\n",
    "        tags=f\"pop\",\n",
    "        text_cfg_coef=diffusion_text_cfg_coef,\n",
    "        ctx_cfg_coef=diffusion_ctx_cfg_coef,\n",
    "        codec_scale_factor=CODEC_SCALE_FACTOR,\n",
    "        scale_ctx_vector=True,\n",
    "        noise_ctx_level=noise_ctx_level,\n",
    "        noise_ctx_pad_len=noise_ctx_pad_len,\n",
    "        drop_semantic_tokens=False,\n",
    "        seed=diffusion_seed,\n",
    "        rho=1.0,\n",
    "        sigma_min=0.5,\n",
    "        sigma_max=50.0,\n",
    "        objective=\"rectified_flow\",\n",
    "        #sampler_type=\"pingpong\",\n",
    "        #semantic_mask_ratio=0.2,\n",
    "    )\n",
    "\n",
    "    request = Request(\n",
    "        id=\"dummy\",\n",
    "        generation_config=gen_cfg,\n",
    "        tokens=semantic_codes,\n",
    "        input_tokens_finished=True,\n",
    "    )\n",
    "\n",
    "    result = diffusion_engine.run_request(request)\n",
    "    #rint(result.generated_audios)\n",
    "    #concat_audio = Audio.concatenate(result.generated_audios)\n",
    "    #concat_audio.play()\n",
    "\n",
    "    # 34d2011075ab7d281191e5b3857446773116df9d\n",
    "\n",
    "    vae_latents = []\n",
    "    for vae_latent in result.vae_latents:\n",
    "        #print(vae_latent.shape)\n",
    "        mean = vae_latent.mean()\n",
    "        std = vae_latent.std()\n",
    "        #print(f\"mean: {mean}, std: {std}\")\n",
    "        vae_latents.append(vae_latent)\n",
    "\n",
    "    vae_latents = torch.concat(vae_latents)\n",
    "    #with torch.no_grad():\n",
    "        #ear_logits = ear_model(vae_latents.unsqueeze(0)).mean(dim=1).mean(dim=1)\n",
    "        #print(\"ear logits: \", ear_logits)\n",
    "    upsampled_audio = decode_stream_to_full_audio(vae_latents)\n",
    "    ear_logits = model.get_score(torch.from_numpy(upsampled_audio.array_float), sample_rate=48000)\n",
    "\n",
    "    metadata = {\n",
    "        \"diffusion_seed\": diffusion_seed,\n",
    "        \"diffusion_text_cfg_coef\": diffusion_text_cfg_coef,\n",
    "        \"diffusion_ctx_cfg_coef\": diffusion_ctx_cfg_coef,\n",
    "        \"diffusion_steps\": diffusion_steps,\n",
    "        \"ear_logits\": ear_logits,\n",
    "    }\n",
    "    #upsampled_audio.play()\n",
    "    upsampled_audios.append((upsampled_audio, metadata))\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# sort the upsampled_audios by ear_logits\n",
    "upsampled_audios.sort(key=lambda x: x[1][\"ear_logits\"])\n",
    "# play the worst and best upsampled audio\n",
    "print(\"worst upsampled audio\", upsampled_audios[0][1])\n",
    "upsampled_audios[0][0].play()\n",
    "print(\"best upsampled audio\", upsampled_audios[-1][1])\n",
    "upsampled_audios[-1][0].play()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "for i in range(5000,5010):\n",
    "    print(semantic_codes[i])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "BLOCK_SIZE = 25 * 30\n",
    "min_chunk_size = 25 * 10\n",
    "\n",
    "chunk_size_schedule = [\n",
    "    min_chunk_size * (2**i) for i in range(10) if min_chunk_size * (2**i) < BLOCK_SIZE\n",
    "] + [BLOCK_SIZE]\n",
    "\n",
    "print(chunk_size_schedule)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "lyrics_chunk = gen_cfg.lyrics\n",
    "tags = gen_cfg.tags\n",
    "text_pieces = []\n",
    "if len(tags) > 0:\n",
    "    text_pieces.append(f\"[{tags}]\")\n",
    "if len(lyrics_chunk) > 0:\n",
    "    text_pieces.append(lyrics_chunk)\n",
    "text = \"\\n\\n\".join(text_pieces)\n",
    "print(text)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def _retrieve_models():\n",
    "    global models\n",
    "    return models\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(models[\"dit_model\"].cond_text_len)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "text_codes = torch.full(\n",
    "        (1, models[\"dit_model\"].cond_text_len), models[\"tokenizer\"].pad_idx, dtype=torch.long\n",
    ")\n",
    "\n",
    "models = diffusion_gen._retrieve_models()\n",
    "\n",
    "for n, codes_row in enumerate(models[\"tokenizer\"].encode_batch([text])):\n",
    "    print(codes_row)\n",
    "    codes_row = torch.tensor(codes_row.ids[: models[\"dit_model\"].cond_text_len], dtype=torch.long)\n",
    "    print(codes_row.shape)\n",
    "    text_codes[n, : len(codes_row)] = codes_row\n",
    "\n",
    "print(text_codes)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "import numpy as np\n",
    "from shimmerscore_new import shimmer_score\n",
    "from tempfile import NamedTemporaryFile\n",
    "\n",
    "audios_first30 = []\n",
    "scores_first30 = []\n",
    "score_means_first30 = []\n",
    "audios_next30 = []\n",
    "scores_next30 = []\n",
    "score_means_next30 = []\n",
    "audios_first30.append(upsampled_audio.get_segment(from_s=0, to_s=30))\n",
    "audios_next30.append(upsampled_audio.get_segment(from_s=30))\n",
    "\n",
    "for in_arr, out_arr, means_arr in [(audios_first30, scores_first30, score_means_first30), (audios_next30, scores_next30, score_means_next30)]:\n",
    "    for audio in in_arr:\n",
    "        with NamedTemporaryFile(suffix=\".wav\") as f:\n",
    "            audio.write_wav(f.name)\n",
    "            out_arr.append(shimmer_score(f.name))\n",
    "    means_arr.append(np.mean(out_arr))\n",
    "\n",
    "print(\"First 30s: \", [round(score, 2) for score in score_means_first30])\n",
    "print(\"Next 30s: \", [round(score, 2) for score in score_means_next30])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "request = Request(\n",
    "    id=\"dummy\",\n",
    "    generation_config=gen_cfg,\n",
    "    tokens=semantic_codes,\n",
    "    input_tokens_finished=True,\n",
    ")\n",
    "\n",
    "job = diffusion_engine.add_request(request)\n",
    "audios = []\n",
    "for audio in job.audio_generator():\n",
    "    audios.append(audio)\n",
    "    print(len(audio))"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Infill"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "list(vae_data.keys())\n",
    "ctx_vae_latents = vae_data[\"vae_latents\"]\n",
    "ctx_vae_latents.shape\n",
    "print(ctx_vae_latents.shape)\n",
    "decode_stream_to_full_audio(ctx_vae_latents).play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "lyrics"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "lyrics_30s = \"\"\"\n",
    "H-O-T-T-O-G-O\\nSnap and clap and touch your toes\\nRaise your feet, now body roll\\nDance it out, you're hot to go\\nH-O-T-T-O-G-O\n",
    "\"\"\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "SEMANTIC_RATE_HZ = 25\n",
    "\n",
    "start_s = 21\n",
    "end_s = 22\n",
    "infill_dur_s = end_s - start_s\n",
    "context_dur_s = 30.0 - infill_dur_s\n",
    "\n",
    "# count total tokens of the context + infill\n",
    "infill_tokens = int(np.ceil(infill_dur_s * SEMANTIC_RATE_HZ))\n",
    "print(infill_dur_s)\n",
    "print(infill_tokens)\n",
    "\n",
    "prefix_start_s = start_s - (context_dur_s // 2)\n",
    "prefix_end_s = start_s\n",
    "suffix_start_s = end_s\n",
    "suffix_end_s = end_s + (context_dur_s // 2)\n",
    "print(\"prefix_start_s\", prefix_start_s)\n",
    "print(\"suffix_end_s\", suffix_end_s)\n",
    "\n",
    "prefix_start_idx = int(np.ceil(prefix_start_s * SEMANTIC_RATE_HZ))\n",
    "prefix_end_idx = int(np.ceil(prefix_end_s * SEMANTIC_RATE_HZ))\n",
    "prefix_tokens = prefix_end_idx - prefix_start_idx\n",
    "\n",
    "print(\"prefix_start_idx\", prefix_start_idx)\n",
    "print(\"prefix_end_idx\", prefix_end_idx)\n",
    "\n",
    "suffix_start_idx = int(np.ceil(suffix_start_s * SEMANTIC_RATE_HZ))\n",
    "suffix_end_idx = int(np.ceil(suffix_end_s * SEMANTIC_RATE_HZ))\n",
    "suffix_tokens = suffix_end_idx - suffix_start_idx\n",
    "\n",
    "print(\"suffix_start_idx\", suffix_start_idx)\n",
    "print(\"suffix_end_idx\", suffix_end_idx)\n",
    "\n",
    "total_context_tokens = prefix_tokens + suffix_tokens\n",
    "print(\"total_context_tokens\", total_context_tokens)\n",
    "\n",
    "total_tokens = total_context_tokens + infill_tokens\n",
    "if total_tokens < 750: \n",
    "    # add more context to the right side\n",
    "    context_to_add = 750 - total_tokens\n",
    "    print(context_to_add)\n",
    "    suffix_end_idx += context_to_add\n",
    "\n",
    "\n",
    "infill_prefix_latents = torch.tensor(ctx_vae_latents[prefix_start_idx: prefix_end_idx, :])\n",
    "infill_suffix_latents = torch.tensor(ctx_vae_latents[suffix_start_idx: suffix_end_idx, :])\n",
    "print(infill_prefix_latents.shape)\n",
    "print(infill_suffix_latents.shape)\n",
    "\n",
    "\n",
    "gen_cfg = diffusion_gen.DiffusionGenerationConfig(\n",
    "    audio=None,\n",
    "    steps=64,\n",
    "    lyrics=lyrics_30s,\n",
    "    tags=tags,\n",
    "    text_cfg_coef=4.0,\n",
    "    ctx_cfg_coef=1.0,\n",
    "    codec_scale_factor=0.4,\n",
    "    scale_ctx_vector=True,\n",
    "    infill_prefix_latents=infill_prefix_latents,\n",
    "    infill_suffix_latents=infill_suffix_latents,\n",
    "    noise_ctx_level=0.2,\n",
    "    drop_semantic_tokens=False,\n",
    "    semantic_skip_factor=1,\n",
    "    semantic_mask_ratio=1.0,\n",
    "    is_diff_infill=True,\n",
    "    objective=\"rectified_flow\",\n",
    ")\n",
    "input_tokens = codes[prefix_start_idx : suffix_end_idx, 0]\n",
    "print(input_tokens.shape)\n",
    "request = Request(\n",
    "    id=\"dummy\",\n",
    "    generation_config=gen_cfg,\n",
    "    tokens=input_tokens,\n",
    "    input_tokens_finished=True,\n",
    ")\n",
    "result = diffusion_engine.run_request(request)\n",
    "\n",
    "vae_latents = torch.concat(result.vae_latents)\n",
    "print(vae_latents.shape)\n",
    "decode_stream_to_full_audio(vae_latents).play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "decode_stream_to_full_audio(ctx_vae_latents).play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(infill_prefix_latents * 0.4)\n",
    "print(infill_suffix_latents * 0.4)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "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": 2
}
