{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"5\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.gpt.generation import load_model, GPT, GPTConfig\n",
    "from suno_utils.utils.s3 import download_s3_file_if_needed\n",
    "import torch\n",
    "\n",
    "N_BATCH = 2\n",
    "\n",
    "ckpt_path = \"/app2/suno/checkpoints/2025-07-02_15-21-34/last_ckpt_infer.pt\"\n",
    "\n",
    "# load model\n",
    "checkpoint = torch.load(str(ckpt_path), mmap=True)\n",
    "model_args = checkpoint[\"model_args\"]\n",
    "state_dict = checkpoint[\"model\"]\n",
    "\n",
    "# load model\n",
    "print(model_args)\n",
    "gptconf = GPTConfig(**model_args)\n",
    "print(gptconf)\n",
    "model = GPT(gptconf)\n",
    "\n",
    "# set up model\n",
    "model.load_state_dict(state_dict, strict=False, assign=True)\n",
    "print(f\"model loaded: {ckpt_path}\")\n",
    "del checkpoint, state_dict\n",
    "n_params = model.get_num_params()\n",
    "print(f\"model loaded: {round(n_params / 1e6, 1)}M params\")\n",
    "model.eval()\n",
    "\n",
    "model.model_args = model_args\n",
    "\n",
    "device = \"cuda\"\n",
    "if device == \"cuda\" and torch.cuda.is_bf16_supported():\n",
    "    model.to(device, dtype=torch.bfloat16)  # this should be fine since it was trained with AMP\n",
    "else:\n",
    "    model.to(device)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import torchaudio\n",
    "\n",
    "device = \"cuda\"\n",
    "\n",
    "mel_spectrogram = torchaudio.transforms.MelSpectrogram(\n",
    "    sample_rate=16000,\n",
    "    n_fft=1025,\n",
    "    hop_length=160,  # 100Hz, but it will be 25Hz after flattening them.\n",
    "    n_mels=128,\n",
    "    f_min=0,\n",
    "    f_max=8000,\n",
    ").to(device)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.audio import Audio\n",
    "import torchaudio\n",
    "import soxr\n",
    "import numpy as np\n",
    "\n",
    "B = 100\n",
    "\n",
    "\n",
    "def _resample_to_musicfm(arr):\n",
    "    assert arr.ndim == 2\n",
    "    assert arr.shape[0] == 2\n",
    "    out_arr = soxr.resample(arr.T, 48_000, 16_000).astype(np.float32)  # stereo\n",
    "    return out_arr\n",
    "\n",
    "\n",
    "audio_48kHz_arr = Audio.from_silence(5, 48_000, n_channels=2).array_float\n",
    "audio_16kHz_arr = torch.from_numpy(_resample_to_musicfm(audio_48kHz_arr)).T.to(device)\n",
    "fm_melspec_clean = mel_spectrogram(audio_16kHz_arr)\n",
    "from einops import rearrange\n",
    "\n",
    "fm_melspec_clean = rearrange(fm_melspec_clean, \"s c t -> t s c\")  # s: stereo, c: mel, t: time\n",
    "fm_melspec_clean = rearrange(fm_melspec_clean, \"(t n) s c -> (n s c) t \", n=4)  # (25hz time, 4, 2, 128)\n",
    "fm_melspec_clean = fm_melspec_clean.bfloat16().unsqueeze(0).repeat(B, 1, 1)\n",
    "print(fm_melspec_clean.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# %%timeit\n",
    "# fm_melspec_clean = mel_spectrogram(audio_16kHz_arr)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.gpt.bct.bct_generation_simple import (\n",
    "    BCTGenerationConfig,\n",
    "    Block,\n",
    "    BlockSequence,\n",
    "    generate_block,\n",
    "    prefill,\n",
    ")\n",
    "from suno_utils.gpt.bct.bct import TensorDict, BlockType\n",
    "\n",
    "mel_block_type = BlockType(name=\"melspec\", is_causal=False)\n",
    "\n",
    "mel_block = Block(spec=mel_block_type, inputs=TensorDict({\"melspec_input\": fm_melspec_clean}))\n",
    "blocks = BlockSequence([mel_block])\n",
    "\n",
    "bct_config = BCTGenerationConfig(prompts=[(1.0, blocks)])\n",
    "\n",
    "cfg = model.config\n",
    "num_blocks = B * cfg.n_layer * (cfg.block_size // 256)\n",
    "kv_cache = model.setup_caches(block_size=256, num_blocks=num_blocks)\n",
    "block_table = torch.arange(num_blocks, device=model.device, dtype=torch.int32).reshape(\n",
    "    B, cfg.n_layer, -1\n",
    ")\n",
    "\n",
    "model = torch.compile(model, mode=\"reduce-overhead\")\n",
    "prefill(model, blocks, kv_cache, block_table)\n",
    "\n",
    "# generate_block(model, bct_config)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "%%timeit\n",
    "prefill(model, blocks, kv_cache, block_table)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "with torch.profiler.profile(\n",
    "    activities=[\n",
    "        torch.profiler.ProfilerActivity.CPU,\n",
    "        torch.profiler.ProfilerActivity.CUDA,\n",
    "    ],\n",
    "    record_shapes=True,\n",
    "    with_stack=True,\n",
    ") as prof:\n",
    "    prefill(model, blocks, kv_cache, block_table)\n",
    "prof.export_chrome_trace(\"trace.json.gz\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "!aws s3 cp trace.json.gz s3://suno-data/traces/musicfm_inf/trace.json.gz\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
