{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "fed7c348",
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"5\""
   ]
  },
  {
   "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",
    "N_BATCH = 2\n",
    "\n",
    "# gpt_model_path_midi = \"/app2/suno/checkpoints/2025-06-16_18-35-59/last_ckpt_infer.pt\"  # synth small\n",
    "# gpt_model_path_midi = \"/app2/suno/checkpoints/2025-06-17_01-16-37/last_ckpt_infer.pt\"  # larger\n",
    "# gpt_model_path_midi_vocals = \"/app2/suno/checkpoints/2025-06-16_21-53-06/last_ckpt_infer.pt\"  # noncausal\n",
    "# causal remi\n",
    "# gpt_model_path_midi_vocals = \"/app2/suno/checkpoints/2025-06-17_16-32-50/last_ckpt_infer.pt\"  # all\n",
    "# gpt_model_path_midi_vocals = (\n",
    "#     \"/app2/suno/checkpoints/2025-06-18_02-40-45/last_ckpt_infer.pt\"  # synthetic only\n",
    "# )\n",
    "# gpt_model_path_midi_vocals = (\n",
    "#     \"/app2/suno/checkpoints/2025-06-18_02-51-12/last_ckpt_infer.pt\"  # linear annealing\n",
    "# )\n",
    "\n",
    "gpt_model_path_midi_vocals = \"/app2/suno/checkpoints/2025-06-18_14-55-00/last_ckpt_infer.pt\"  # bpe 50%\n",
    "# gpt_model_path_midi_vocals = \"/app2/suno/checkpoints/2025-06-18_14-54-24/last_ckpt_infer.pt\"  # bpe 95%\n",
    "gpt_model_path_midi_vocals = \"/app2/suno/checkpoints/2025-06-19_00-12-49/last_ckpt_infer.pt\"  # synthetic\n",
    "# gpt_model_path_midi_vocals = \"/app2/suno/checkpoints/2025-06-19_02-50-42/last_ckpt_infer.pt\"  # longer\n",
    "gpt_model_path_midi_vocals = \"/app2/suno/checkpoints/2025-06-19_14-24-30/last_ckpt_infer.pt\"  # no fret\n",
    "gpt_model_path_midi_vocals = \"/app2/suno/checkpoints/2025-06-20_12-29-53/last_ckpt_infer.pt\"  # no repeat\n",
    "gpt_model_path_midi_vocals = \"/app2/suno/checkpoints/2025-06-20_21-08-10/last_ckpt_infer.pt\"  # no repeat\n",
    "gpt_model_path_midi_vocals = \"/app2/suno/checkpoints/2025-06-21_03-27-02/last_ckpt_infer.pt\"  # sustain\n",
    "gpt_model_path_midi_vocals = \"/app2/suno/checkpoints/2025-06-26_12-04-03/last_ckpt_infer.pt\"  # long\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(gpt_model_path_midi_vocals),\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\n",
    "cfg"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "538a92a0",
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.tasks.mert_25 import preload_models as preload_semantic_models, encode\n",
    "\n",
    "preload_semantic_models(\n",
    "    checkpoint_filepath=\"s3://suno-data/georg/models/semantic/mert_25.pt\",\n",
    "    centroids_filepath=\"s3://suno-data/georg/models/semantic/mert_25_2x4k.npy\",\n",
    ")\n"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "7f4f8960",
   "metadata": {},
   "source": [
    "## MIDI transcription"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "c6ba81da",
   "metadata": {},
   "outputs": [],
   "source": [
    "# load semantic codes\n",
    "\n",
    "from suno_utils.utils.clip import SunoClip\n",
    "from suno_utils.audio.midi import Midi\n",
    "from suno_utils.audio import Audio\n",
    "\n",
    "gen_ids = {\n",
    "    # \"techno\": \"2f327d3b-6c89-4687-bd5e-2cb0d3695884\",\n",
    "    # \"edm\": \"1c6371ab-b335-419d-b746-ae7610d07f8c\",\n",
    "    # \"violin_metal\": \"15c9fb05-4eee-4c1f-b97b-2ef94bf251f1\",\n",
    "    # \"r&b jazz\": \"082db449-660f-4083-aab9-aa0facb87be9\",\n",
    "    # \"blues\": \"1f3c211d-c721-469f-ad07-4ba677174c0d\",\n",
    "    # \"final lie\": \"f1018ab1-8976-437b-940c-8fb10139ca9a\",\n",
    "    # \"pop rock\": \"e3ccb02b-696e-4335-a674-9f7821b202f4\",\n",
    "    \"halo_vocals\": \"f6ed55f9-ba2b-4a74-8a83-76debff7c2b7\",\n",
    "    \"piano\": \"7dd3f0e7-9b78-492f-ab93-3f0a9a30b0e0\",\n",
    "    \"stone\": \"a5e2198a-f352-4abb-9a24-7f81b143ded3\",\n",
    "    # \"tony_piano\": \"e3f1f507-5dee-42b6-b084-e3c23f5e87c7\",\n",
    "    # \"v3_guitar\": \"825f9b8f-2e6d-4f5e-9b81-224a4d7609a5\",\n",
    "    # \"brass_mix\": \"e6714897-8723-4fe0-8f79-070212aa4dde\",\n",
    "    # \"sugar_rush\": \"82e75a1c-d9a7-4ba2-b1e0-028de2fa02f1\",\n",
    "    # \"metal_vocals\": \"3a86fc66-7574-482a-a1bd-622726d92159\",\n",
    "    # \"chi_piano\": \"8f541071-b6a7-49fe-a835-78ce6152e6ea\",\n",
    "    # \"chi_guitar\": \"300d7882-c736-4fa5-ba37-f39385e9c528\",\n",
    "    # \"sam_vocals_chef_tony\": \"62b9ab24-bf0c-4d8f-ab34-2f9c6558f9d5\",\n",
    "    # \"sam_guitar\": \"98ee63ca-bfa1-4524-b03b-9baa631835bc\",\n",
    "    # \"sam_keys\": \"62a72386-95fa-4e24-8053-67060d7c1eb7\",\n",
    "}\n",
    "\n",
    "# gen_id = gen_ids[\"piano\"]\n",
    "# clip = SunoClip(gen_id)\n",
    "# audio = clip.audio()\n",
    "# audio.play()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "c750af12",
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch\n",
    "\n",
    "# from suno_utils.gpt.bct.bct_generation_simple import BCTGenerationConfig, BlockSequence, generate_block\n",
    "from suno_utils.gpt.bct.bct_engine import BCTEngine, BlockSequence, Config\n",
    "from suno_utils.gpt.bct.bct import Block, BlockType, TensorDict, SamplingParams\n",
    "\n",
    "\n",
    "engine_config = Config()\n",
    "engine = BCTEngine(engine_config, model, compile=True)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "f565e87b",
   "metadata": {},
   "outputs": [],
   "source": [
    "def remove_silent_sections(audio: Audio) -> list[tuple[float, float]]:\n",
    "    \"\"\"\n",
    "    Removes silent sections from audio.\n",
    "    Returns a list of tuples containing the start and end times of the active sections.\n",
    "    \"\"\"\n",
    "\n",
    "    from suno_utils.tasks.ss_vad import short_time_energy\n",
    "\n",
    "    step_duration = 0.1\n",
    "    energy = short_time_energy(\n",
    "        audio.array_float, int(step_duration * audio.sample_rate), int(step_duration * audio.sample_rate)\n",
    "    )\n",
    "    is_active = energy > 0.1\n",
    "    # Group frames into 5-second bins and mark as silent if >90% of frames are inactive\n",
    "    frames_per_bin = int(5.0 / step_duration)\n",
    "    silent_sections = []\n",
    "\n",
    "    for i in range(0, len(is_active), frames_per_bin):\n",
    "        bin_frames = is_active[i : i + frames_per_bin]\n",
    "        if len(bin_frames) == 0:\n",
    "            continue\n",
    "\n",
    "        inactive_ratio = (bin_frames == False).sum() / len(bin_frames)\n",
    "        if inactive_ratio > 0.9:\n",
    "            start_time = i * step_duration\n",
    "            end_time = min((i + frames_per_bin) * step_duration, len(is_active) * step_duration)\n",
    "            silent_sections.append((start_time, end_time))\n",
    "\n",
    "    # Convert silent sections to active sections\n",
    "    active_sections = []\n",
    "    current_start = 0.0\n",
    "\n",
    "    for silent_start, silent_end in silent_sections:\n",
    "        if silent_start > current_start:\n",
    "            active_sections.append((current_start, silent_start))\n",
    "        current_start = silent_end\n",
    "\n",
    "    # Add final active section if audio doesn't end with silence\n",
    "    total_duration = len(is_active) * step_duration\n",
    "    if current_start < total_duration:\n",
    "        active_sections.append((current_start, total_duration))\n",
    "\n",
    "    return active_sections\n",
    "\n",
    "\n",
    "audio = SunoClip(gen_ids[\"halo_vocals\"]).audio()\n",
    "sections = remove_silent_sections(audio)\n",
    "print(sections)\n",
    "for start, end in sections:\n",
    "    audio.get_segment(start, end).play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "55b282ab",
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.gpt.bct.bct import Block, BlockType, BlockSequence\n",
    "from suno_utils.gpt.bct.bct_generation_simple import BCTGenerationConfig, generate_block\n",
    "from tqdm import tqdm\n",
    "import torch\n",
    "\n",
    "\n",
    "NonCausalSemanticBlockType = BlockType(name=\"semantic\", is_causal=True)\n",
    "MidiBlockType = BlockType(name=\"midi\", is_causal=True)\n",
    "TextBlockType = BlockType(name=\"text\", is_causal=True)\n",
    "\n",
    "\n",
    "def audio_to_midi(\n",
    "    audio,\n",
    "    text=\"[golden]\",\n",
    "    max_new_tokens=2000,\n",
    "    audio_cfg=1.0,\n",
    "    temperature=0.3,\n",
    "    top_k=1,\n",
    "    midi_prompt=None,\n",
    "    midi_tokenizer=Midi.load_tokenizer(),\n",
    "):\n",
    "    \"\"\"Convert audio to MIDI by transcribing it.\n",
    "\n",
    "    Args:\n",
    "        audio: Audio segment to transcribe\n",
    "        text: Text prompt for the model (default: \"[synthetic]\")\n",
    "        max_new_tokens: Maximum number of tokens to generate (default: 2000)\n",
    "\n",
    "    Returns:\n",
    "        Midi object containing the transcribed MIDI\n",
    "    \"\"\"\n",
    "\n",
    "    # Encode audio to semantic features\n",
    "    mert_latents = encode(audio, pad_to_chunksize=True, batch_size=48, do_clustering=False)\n",
    "    mert_latents = torch.tensor(mert_latents).bfloat16()\n",
    "\n",
    "    # Prepare text tokens\n",
    "    text_tokens = model_container[\"tokenizer\"].encode(text)\n",
    "\n",
    "    text_block = Block(\n",
    "        TextBlockType,\n",
    "        inputs=TensorDict({\"text_input\": torch.tensor(text_tokens).reshape(1, 1, -1)}),\n",
    "    )\n",
    "\n",
    "    sem_block = Block(\n",
    "        NonCausalSemanticBlockType,\n",
    "        inputs=TensorDict({\"continuous_semantic_input\": mert_latents.T.unsqueeze(0)}),\n",
    "    )\n",
    "\n",
    "    midi_tokens = [cfg.midi_infer_token]\n",
    "    if midi_prompt is not None:\n",
    "        midi_tokens += midi_prompt.to_tokens(midi_tokenizer).ids\n",
    "\n",
    "    midi_block = Block(\n",
    "        MidiBlockType,\n",
    "        inputs=TensorDict({\"midi_input\": torch.tensor(midi_tokens).reshape(1, 1, -1)}),\n",
    "    )\n",
    "\n",
    "    sampling_params = SamplingParams(\n",
    "        temperature=temperature,\n",
    "        top_k=top_k,\n",
    "        max_tokens=max_new_tokens,\n",
    "        stop_token_ids=[cfg.midi_pad_token],\n",
    "    )\n",
    "    block_sequence = BlockSequence([text_block, sem_block, midi_block], sampling_params=sampling_params)\n",
    "    neg_seq = BlockSequence([midi_block], sampling_params=sampling_params)\n",
    "\n",
    "    block_sequence.weight = 1 + audio_cfg\n",
    "    neg_seq.weight = -audio_cfg\n",
    "    block_sequence.main_cfg_stream_id = block_sequence.id\n",
    "    neg_seq.main_cfg_stream_id = block_sequence.id\n",
    "    # Generate MIDI tokens\n",
    "    # generated_block_sequence = generate_block(model_container[\"model\"], generation_config)\n",
    "    engine.generate([block_sequence, neg_seq], use_tqdm=True)\n",
    "    midi_tokens = block_sequence.all_tokens()[\"midi_input\"][0, 0]\n",
    "    midi_tokens = midi_tokens.cpu().tolist()[1:-1]\n",
    "    print(midi_tokens)\n",
    "\n",
    "    # Convert tokens to MIDI\n",
    "    midi = Midi.from_tokens(midi_tokens, midi_tokenizer)\n",
    "    midi = filter_silent_sections(midi, audio)\n",
    "\n",
    "    return midi\n",
    "\n",
    "\n",
    "def filter_silent_sections(midi: Midi, audio: Audio):\n",
    "    \"\"\"When audio is silent, remove the corresponding midi notes\"\"\"\n",
    "    from suno_utils.tasks.ss_vad import short_time_energy\n",
    "\n",
    "    step_duration = 0.1\n",
    "    energy = short_time_energy(\n",
    "        audio.array_float, int(step_duration * audio.sample_rate), int(step_duration * audio.sample_rate)\n",
    "    )\n",
    "\n",
    "    def note_is_silent(note, threshold: float = 0.6):\n",
    "        s = int(note.start / step_duration)\n",
    "        e = int(note.end / step_duration)\n",
    "        e = max(e, s + 1)\n",
    "        pct_silent = sum(energy[s:e] < 0.1) / (e - s)\n",
    "        # print(s, e, pct_silent)\n",
    "        return pct_silent > threshold\n",
    "\n",
    "    for instrument in midi.pmidi.instruments:\n",
    "        # print number of notes before and after filtering\n",
    "        # print(f\"Number of notes before filtering: {len(instrument.notes)}\")\n",
    "        instrument.notes = [note for note in instrument.notes if not note_is_silent(note)]\n",
    "        # print(f\"Number of notes after filtering: {len(instrument.notes)}\")\n",
    "    return midi.clean_pmidi()\n",
    "\n",
    "\n",
    "def transcribe(audio, text_prompt=\"[golden]\"):\n",
    "    # process in chunks. input 30s of audio at a time. shift 20s each time.\n",
    "    # feed midi from 20-25s of last chunk to prompt the next chunk\n",
    "\n",
    "    active_sections = remove_silent_sections(audio)\n",
    "    print(active_sections)\n",
    "    audio_sections = []\n",
    "    for start, end in active_sections:\n",
    "        audio_sections.append(audio.get_segment(start, end))\n",
    "        audio_sections.append(Audio.from_silence(3, audio.sample_rate, n_channels=audio.n_channels))\n",
    "    active_audio = Audio.concatenate(audio_sections)\n",
    "\n",
    "    midi_tokenizer = Midi.load_tokenizer()\n",
    "\n",
    "    stride = 30\n",
    "    history_size = 30\n",
    "    future_buffer = 5  # this will get cropped\n",
    "    chunk_size = stride + history_size + future_buffer\n",
    "    midi_chunks = []\n",
    "    programs_seen = set()\n",
    "    for i in tqdm(range(0, int(active_audio.duration_s), stride), desc=\"Transcribing chunks\"):\n",
    "        audio_chunk = active_audio.get_segment(i, i + chunk_size)\n",
    "        if i > 0:\n",
    "            midi_prompt = midi_chunks[-1].get_segment(stride, stride + history_size)\n",
    "        else:\n",
    "            midi_prompt = None\n",
    "        text = f\"[programs in history: {list(programs_seen)}]\"\n",
    "        # print(text)\n",
    "        midi_chunk = audio_to_midi(\n",
    "            audio_chunk,\n",
    "            max_new_tokens=8000,\n",
    "            midi_prompt=midi_prompt,\n",
    "            text=text,\n",
    "            audio_cfg=0.5,\n",
    "            midi_tokenizer=midi_tokenizer,\n",
    "        )\n",
    "        midi_chunks.append(midi_chunk)\n",
    "        programs_seen.update(midi_chunk.program_ids)\n",
    "\n",
    "    # concatenate midi chunks\n",
    "    midi = midi_chunks[0]\n",
    "    for i, midi_chunk in enumerate(midi_chunks[1:]):\n",
    "        midi += midi_chunk.filter_notes_by_onset(0, stride).shift_time(stride * (i + 1))\n",
    "    midi = midi.consolidate_programs()\n",
    "\n",
    "    # boost vocals volume (trombone champ is always 63)\n",
    "    for instrument in midi.pmidi.instruments:\n",
    "        if instrument.program == 53:\n",
    "            for note in instrument.notes:\n",
    "                note.velocity = min(int(note.velocity * 1.5), 127)\n",
    "\n",
    "    # reconstruct midi from active sections\n",
    "    midi_sections = []\n",
    "    midi_pos_s = 0\n",
    "    for i, (start, end) in enumerate(active_sections):\n",
    "        midi_sections.append(midi.get_segment(midi_pos_s, midi_pos_s + end - start).shift_time(start))\n",
    "        midi_pos_s += end - start + 3\n",
    "\n",
    "    return sum(midi_sections).clean_pmidi()\n",
    "\n",
    "\n",
    "# midis = transcribe(audio)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "44aa0173",
   "metadata": {},
   "outputs": [],
   "source": [
    "for name, gen_id in gen_ids.items():\n",
    "    # if name != \"sam_guitar\":\n",
    "    #     continue\n",
    "    clip = SunoClip(gen_id)\n",
    "    audio = clip.audio()\n",
    "    print(name)\n",
    "    midi = transcribe(audio, text_prompt=\"[golden]\")\n",
    "    print(midi.pmidi.instruments)\n",
    "    midi.show()\n",
    "    midi.make_stereo_comparison(audio).play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3d206075",
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
