{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "0",
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "import re\n",
    "import torch\n",
    "import math\n",
    "import shutil\n",
    "import numpy as np\n",
    "import matplotlib.pyplot as plt\n",
    "import subprocess\n",
    "\n",
    "os.environ['PATH'] += \":/home/m4burns/\"\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"1\"\n",
    "\n",
    "from suno_utils.utils.text import write_json\n",
    "from suno_utils.tasks.audio_features.beat_this_downbeat import BeatThisDownbeatExtractor\n",
    "from suno_utils.tasks.audio_features.downbeats_data_prep.augment import stretch_audio\n",
    "from suno_utils.audio import Audio\n",
    "\n",
    "from suno_utils.gpt import chirp_v2_5 as chirp_v3\n",
    "from suno_utils.tasks.ditto_v2 import preload_models\n",
    "from suno_utils.tasks.ditto_v2 import encode_overlap as ditto_encode, SAMPLE_RATE, load_model\n",
    "\n",
    "from sklearn.cluster import KMeans\n",
    "from sklearn.metrics import silhouette_score\n",
    "from sklearn.decomposition import PCA\n",
    "from sklearn.metrics.pairwise import euclidean_distances\n",
    "\n",
    "\n",
    "from suno_utils.diffusion import generation as diffusion_gen\n",
    "from suno_utils.tasks.upsample_engine import UpsampleEngine, Request\n",
    "from suno_utils.utils.s3 import upload_s3_files\n",
    "\n",
    "from suno_utils.tasks.dac_vae_fixed_25hz import decode, encode as diff_encode\n"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "1",
   "metadata": {},
   "source": [
    "## Load Models"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2",
   "metadata": {},
   "outputs": [],
   "source": [
    "import logging\n",
    "logging.basicConfig(level=logging.WARN)\n",
    "\n",
    "# load downbeat detector\n",
    "extractor = BeatThisDownbeatExtractor(device=\"cuda\", model_path=\"s3://suno-data/m4burns/beat_this_rc_12l.pt\")\n",
    "\n",
    "# load ditto for embeddings\n",
    "model_filepath = \"s3://suno-data/minz/models/ditto_v2_epoch_57.pt\"\n",
    "model_path = chirp_v3._get_model_if_needed(model_filepath)\n",
    "_ = preload_models(\n",
    "    model_filepath=model_path,\n",
    ")\n",
    "ditto = load_model()[\"model\"]\n",
    "\n",
    "# load stem diffusion\n",
    "diffusion_gen.preload_models(\n",
    "    dit_model_filepath=\"/app2/suno/modal/models//victor/checkpoints/diffusion/stems_12_stem.pt\",  # 12 stems\n",
    "    codec_filepath=\"s3://suno-data/minz/models/dac_vae_tuned_25hz.pth\",\n",
    ")\n",
    "\n",
    "engine = UpsampleEngine(min_chunk_size=25 * 10)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "3",
   "metadata": {},
   "source": [
    "## Setup Parameters"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "4",
   "metadata": {},
   "outputs": [],
   "source": [
    "# which song to decompose\n",
    "\n",
    "uuid = \"75af6dcd-07b2-4f01-9f65-46747e639ed8\" # no dance\n",
    "uuid = \"a5e2198a-f352-4abb-9a24-7f81b143ded3\" # stone\n",
    "uuid = \"cbea9ff8-01e3-42e1-90b4-12eb45f73e98\" # trust in it\n",
    "uuid = \"05f4a926-9df2-4faa-997b-dc6eddc04ca6\" # a stone remix\n",
    "uuid = \"c1f89abb-28a8-4fad-b687-0d2147cb3385\" # tech house example\n",
    "uuid = \"0ffbfe03-4e01-4750-91a2-1209eaef2731\" # a glitchy example\n",
    "uuid = \"aa8eb53f-4a3e-4a45-a694-5dae54922ba9\" # take me back\n",
    "#uuid = \"05ec148f-fd1f-40cf-afee-8b68863ef76c\" #zwolfton pulse\n",
    "#uuid = \"aba17039-59d7-4ee5-9c21-fe241e243292\" # godot-livery\n",
    "#uuid = \"9473efc0-9d4a-44b5-9a1f-11c26a7cb054\" #godot remix orchestral\n",
    "#uuid = \"daf7a527-4c97-4f71-b51f-fcc0db8e9612\" # godot rock\n",
    "uuid = \"a1936d21-d47a-44ad-8551-25feb86a58e9\" # ???\n",
    "\n",
    "# loop length limits\n",
    "MIN_BARS = 8\n",
    "MAX_BARS = 8\n",
    "\n",
    "# minimum confidence value for valid loop\n",
    "CONFIDENCE_CUTOFF = 0.7\n",
    "\n",
    "# hardcode a tempo\n",
    "TARGET_TEMPO = None #120"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "5",
   "metadata": {},
   "source": [
    "## Normalize Tempo"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "6",
   "metadata": {},
   "outputs": [],
   "source": [
    "def extract_beats(audio: np.ndarray) -> np.ndarray:\n",
    "    audio_mono = Audio.convert(audio, n_channels=1, sample_rate=audio.sample_rate, byte_width=audio.byte_width)\n",
    "    out = extractor.extract(audio_mono)\n",
    "    beats_refined = np.array(out[\"downbeats\"])\n",
    "    print(beats_refined[0])\n",
    "    return beats_refined[:, 0]\n",
    "\n",
    "def beats_to_tempo(beats: np.ndarray) -> float:\n",
    "    return 60 / np.diff(beats, axis=-1)\n",
    "\n",
    "def est_target_tempo(beats: np.ndarray) -> int:\n",
    "    tempos = 60 / np.diff(beats, axis=-1)\n",
    "    return round(np.median(tempos))\n",
    "\n",
    "def get_target_beats(beat_times: np.ndarray, target_tempo: int) -> list[float]:\n",
    "    target_beat_time = 60. / target_tempo\n",
    "    print(f\"Beat length: {target_beat_time}\")\n",
    "    target_beats = []\n",
    "\n",
    "    for idx, beat in enumerate(beat_times):\n",
    "        if idx == 0: # first beat\n",
    "            next_beat = beat\n",
    "        else:\n",
    "            last_beat = target_beats[idx - 1]\n",
    "            next_beat = last_beat + target_beat_time\n",
    "        target_beats.append(next_beat)\n",
    "\n",
    "    return target_beats\n",
    "\n",
    "def get_warp_markers(beat_times: np.ndarray, target_tempo) -> list[(float, float)]:\n",
    "    target_beats = get_target_beats(beat_times, target_tempo)\n",
    "    return list(zip(beat_times.tolist(), target_beats))\n",
    "\n",
    "def overlay_metronome(audio: Audio, beat_times: list) -> Audio:\n",
    "    import librosa\n",
    "    audio_beat = librosa.clicks(times = np.array(beat_times), sr=audio.sample_rate, click_freq=1000, length=audio.array_float.shape[-1])\n",
    "    if audio.array_float.ndim > 1:\n",
    "        audio_beat = audio_beat[None, :]\n",
    "    return Audio.from_array_float(audio.array_float + audio_beat.reshape(1, -1), audio.sample_rate, max_allowed_val=10.0)\n",
    "\n",
    "def preprocess(uuid: str) -> tuple[Audio, list[(float, float)], int]:\n",
    "    audio = Audio.from_s3(f\"s3://suno-data-uploads/studio/uploads/{uuid}.mp3\")\n",
    "    beat_times = extract_beats(audio)\n",
    "    target_tempo = est_target_tempo(beat_times) if TARGET_TEMPO is None else TARGET_TEMPO\n",
    "    print(f\"Normalizing to {target_tempo} bpm\")\n",
    "    warp_markers = get_warp_markers(beat_times, target_tempo)\n",
    "    out_audio = stretch_audio(audio, warp_markers)\n",
    "\n",
    "    metronome = overlay_metronome(out_audio, [wm[1] for wm in warp_markers])\n",
    "    metronome.write_mp3(\"test_metronome.mp3\")   \n",
    "\n",
    "    return out_audio, warp_markers, target_tempo"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "7",
   "metadata": {},
   "outputs": [],
   "source": [
    "out_audio, warp_markers, target_tempo = preprocess(uuid)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "8",
   "metadata": {},
   "outputs": [],
   "source": [
    "plt.figure(figsize=(10, 3))\n",
    "plt.title(\"Beat-to-Beat Tempo\")\n",
    "plt.xlabel(\"Beat Index\")\n",
    "plt.ylabel(\"Tempo (BPM)\")\n",
    "plt.plot(beats_to_tempo(np.asarray([wm[0] for wm in warp_markers])), label=\"Original\")\n",
    "plt.plot(beats_to_tempo(np.asarray([wm[1] for wm in warp_markers])), label=\"Normalized\")\n",
    "plt.legend()\n",
    "plt.show()\n"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "9",
   "metadata": {},
   "source": [
    "## Filter Loops"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "10",
   "metadata": {},
   "outputs": [],
   "source": [
    "def get_bar_len_s(target_tempo: int, num_bars: int) -> float:\n",
    "    beat_len_s = (60. / target_tempo)\n",
    "    bar_len_s = 4 * beat_len_s\n",
    "    return bar_len_s * num_bars\n",
    "\n",
    "def get_bar_len_samples(target_tempo: int, num_bars: int, sr: int) -> float:\n",
    "    beat_len_s = (60. / target_tempo)\n",
    "    bar_len_s = 4 * beat_len_s\n",
    "    return bar_len_s * num_bars * sr\n",
    "\n",
    "def find_closest(ordered_list, x):\n",
    "    left, right = 0, len(ordered_list) - 1\n",
    "    \n",
    "    # Binary search for insertion point\n",
    "    while left <= right:\n",
    "        mid = (left + right) // 2\n",
    "        if ordered_list[mid] == x:\n",
    "            return x\n",
    "        elif ordered_list[mid] < x:\n",
    "            left = mid + 1\n",
    "        else:\n",
    "            right = mid - 1\n",
    "    \n",
    "    # Now left is the insertion point\n",
    "    if left == 0:\n",
    "        return ordered_list[0]\n",
    "    if left == len(ordered_list):\n",
    "        return ordered_list[-1]\n",
    "    \n",
    "    # Compare neighbors\n",
    "    before = ordered_list[left - 1]\n",
    "    after = ordered_list[left]\n",
    "    \n",
    "    if abs(x - before) <= abs(x - after):\n",
    "        return before\n",
    "    else:\n",
    "        return after\n",
    "    \n",
    "def find_loops_via_cli(audio_path, min_dur, max_dur):\n",
    "    \"\"\"Use CLI interface from Python\"\"\"\n",
    "    cmd = [\n",
    "        'pymusiclooper', \n",
    "        'export-points', \n",
    "        '--path', audio_path,\n",
    "        '--alt-export-top', '-1',\n",
    "        '--min-loop-duration', str(min_dur),\n",
    "        '--max-loop-duration', str(max_dur)\n",
    "    ]\n",
    "    \n",
    "    try:\n",
    "        result = subprocess.run(cmd, capture_output=True, text=True)\n",
    "        \n",
    "        # Strip ANSI escape codes from output\n",
    "        def strip_ansi_codes(text):\n",
    "            ansi_escape = re.compile(r'\\x1B(?:[@-Z\\\\-_]|\\[[0-?]*[ -/]*[@-~])')\n",
    "            return ansi_escape.sub('', text)\n",
    "        \n",
    "        clean_output = strip_ansi_codes(result.stdout)\n",
    "        \n",
    "        # Parse output (format: start end note_diff loudness_diff score)\n",
    "        loops = []\n",
    "        for line in clean_output.strip().split('\\n'):\n",
    "            # Skip empty lines, comments, and lines without numeric data\n",
    "            if (line.strip() and \n",
    "                not line.startswith('#') and \n",
    "                any(c.isdigit() for c in line)):\n",
    "                \n",
    "                parts = line.split()\n",
    "                \n",
    "                # Make sure we have at least 2 numeric values\n",
    "                if len(parts) == 5:\n",
    "                    try:\n",
    "                        start, end, confidence = int(parts[0]), int(parts[1]), float(parts[4])\n",
    "                        if confidence > CONFIDENCE_CUTOFF:\n",
    "                            loops.append((start, end, confidence))\n",
    "                    except ValueError:\n",
    "                        continue\n",
    "        \n",
    "        return loops\n",
    "    \n",
    "    except Exception as e:\n",
    "        print(f\"Error running PyMusicLooper CLI: {e}\")\n",
    "        return None\n",
    "    \n",
    "def fade_audio(audio, fade_samples):\n",
    "    faded = audio.copy()\n",
    "    \n",
    "    # Fade in (linear)\n",
    "    fade_in = np.linspace(0, 1, fade_samples)\n",
    "    faded[:fade_samples] *= fade_in\n",
    "    \n",
    "    # Fade out (linear)\n",
    "    fade_out = np.linspace(1, 0, fade_samples)\n",
    "    faded[-fade_samples:] *= fade_out\n",
    "    \n",
    "    return faded\n",
    "\n",
    "def get_loop_candidates(audio_fp: str, min_bars: int = 2, max_bars: int = 8) -> list[tuple[float, float, float]]:\n",
    "    min_sec_bars = math.floor(get_bar_len_s(target_tempo=target_tempo, num_bars=min_bars))\n",
    "    max_sec_bars = math.ceil(get_bar_len_s(target_tempo=target_tempo, num_bars=max_bars))\n",
    "    loop_points = find_loops_via_cli(audio_fp, min_dur=min_sec_bars, max_dur=max_sec_bars)\n",
    "    print(f\"Found {len(loop_points)} loops\")\n",
    "\n",
    "    return loop_points\n",
    "\n",
    "def filter_loops_on_endpoints(loop_candidates: list[tuple[float, float, float]], sample_rate: int, target_tempo: int) -> list[tuple[float, float, float]]:\n",
    "    filtered_loops = []\n",
    "    beat_len_s = 60. / target_tempo\n",
    "\n",
    "    for start,end, conf in loop_candidates:\n",
    "        beat_sample = [int(wm[1] * sample_rate) for wm in warp_markers]\n",
    "        downbeat_sample = [b for idx,b in enumerate(beat_sample) if idx % 4 == 0]\n",
    "        closest_start = find_closest(downbeat_sample, start)\n",
    "        closest_end = find_closest(downbeat_sample, end)\n",
    "        if abs(closest_start - start) > (beat_len_s * sample_rate) or abs(closest_end - end) > (beat_len_s * sample_rate):\n",
    "            continue # not a bar length loop\n",
    "        else:\n",
    "            filtered_loops.append((closest_start, closest_end, conf))\n",
    "\n",
    "    return filtered_loops\n",
    "\n",
    "def embed_loops(audio: Audio, loops: list[tuple[float, float, float]], task: str = \"self_sim\") -> np.ndarray:\n",
    "    embeddings = []\n",
    "    source_arr = audio.array_float\n",
    "    for s,e,_ in loops:\n",
    "        loop_segment = source_arr[s:e]\n",
    "        audio = Audio.from_array_float(loop_segment, sample_rate=out_audio.sample_rate, auto_compress=False)\n",
    "        audio = audio.convert(sample_rate=SAMPLE_RATE, n_channels=1, byte_width=2)\n",
    "        encoding = ditto_encode([audio], task=task)[0]\n",
    "        embeddings.append(np.mean(encoding, axis=0))\n",
    "\n",
    "    return np.array(embeddings)\n",
    "\n",
    "def visualize_clusters(embeddings: np.ndarray, cluster_labels: list[int]):\n",
    "    pca = PCA(n_components=2)\n",
    "    embeddings_2d = pca.fit_transform(embeddings)\n",
    "\n",
    "    # Plot clusters\n",
    "    plt.figure(figsize=(10, 8))\n",
    "    scatter = plt.scatter(embeddings_2d[:, 0], embeddings_2d[:, 1], \n",
    "                        c=cluster_labels, cmap='viridis', alpha=0.7)\n",
    "    plt.colorbar(scatter)\n",
    "    plt.title('Embedding Clusters (PCA)')\n",
    "    plt.xlabel('First Principal Component')\n",
    "    plt.ylabel('Second Principal Component')\n",
    "    plt.show()\n",
    "\n",
    "def get_diverse_loops(embeddings: np.ndarray, loops: list[tuple[float, float, float]]) -> list[tuple[float, float, float]]:\n",
    "    silhouette_scores = []\n",
    "    k_range = range(3, 10)\n",
    "\n",
    "    for k in k_range:\n",
    "        kmeans = KMeans(n_clusters=k, random_state=42)\n",
    "        cluster_labels = kmeans.fit_predict(embeddings)\n",
    "        silhouette_scores.append(silhouette_score(embeddings, cluster_labels))\n",
    "\n",
    "    plt.plot(k_range, silhouette_scores, 'ro-')\n",
    "    plt.title('Silhouette Score')\n",
    "    plt.xlabel('Number of clusters')\n",
    "    plt.ylabel('Silhouette Score')\n",
    "    plt.show()\n",
    "\n",
    "    # Cluster with optimal k\n",
    "    optimal_k = k_range[np.argmax(silhouette_scores)]\n",
    "    kmeans = KMeans(n_clusters=optimal_k, random_state=42)\n",
    "    cluster_labels = kmeans.fit_predict(embeddings)\n",
    "    visualize_clusters(embeddings, cluster_labels)\n",
    "\n",
    "    # get best looop from each cluster\n",
    "    closest_indices = []\n",
    "    for cluster_id in range(optimal_k):\n",
    "        cluster_mask = cluster_labels == cluster_id\n",
    "        if not cluster_mask.any():\n",
    "            continue\n",
    "            \n",
    "        cluster_points = np.array(embeddings)[cluster_mask]\n",
    "        cluster_point_indices = np.where(cluster_mask)[0]\n",
    "        \n",
    "        # Distance from center to each point in cluster\n",
    "        distances = euclidean_distances(\n",
    "            kmeans.cluster_centers_[cluster_id].reshape(1, -1), \n",
    "            cluster_points\n",
    "        )[0]\n",
    "        \n",
    "        # Original index of closest point\n",
    "        closest_original_idx = cluster_point_indices[np.argmin(distances)]\n",
    "        closest_indices.append(closest_original_idx)\n",
    "        \n",
    "        print(f\"Cluster {cluster_id}: representative loop {loops[closest_original_idx]}\")\n",
    "\n",
    "    # Get the representative loops\n",
    "    representative_loops = [loops[idx] for idx in closest_indices]\n",
    "    return representative_loops\n",
    "    \n",
    "def get_loop_audio(start: int, end: int, audio: Audio, target_tempo: int) -> tuple[int, Audio]:\n",
    "    bar_len_s = 4. * 60. / target_tempo\n",
    "    float_arr = audio.array_float\n",
    "    loop_segment = float_arr[start:end]\n",
    "    fade_segment = fade_audio(loop_segment, 100)\n",
    "    loop_audio = Audio.from_array_float(audio_arr=fade_segment, sample_rate=out_audio.sample_rate, auto_compress=False)\n",
    "    bar_length = round((end - start) / (bar_len_s * audio.sample_rate))\n",
    "    return bar_length, loop_audio"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "11",
   "metadata": {},
   "outputs": [],
   "source": [
    "min_bars = MIN_BARS\n",
    "max_bars = MAX_BARS\n",
    "\n",
    "loop_candidates = get_loop_candidates(audio_fp=\"/home/sara/sara/sbbq_hackathon/test_metronome.mp3\", min_bars=min_bars, max_bars=max_bars)\n",
    "sample_rate = out_audio.sample_rate\n",
    "filtered_loops = filter_loops_on_endpoints(loop_candidates=loop_candidates, sample_rate=sample_rate, target_tempo=target_tempo)\n",
    "embeddings = embed_loops(out_audio, filtered_loops)\n",
    "representative_loops = get_diverse_loops(embeddings=embeddings, loops=filtered_loops)\n",
    "\n",
    "audio_loops = []\n",
    "for start,end,conf in representative_loops:\n",
    "    _, loop_audio = get_loop_audio(start=start, end=end, audio=out_audio, target_tempo=target_tempo)\n",
    "    audio_loops.append(loop_audio)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "12",
   "metadata": {},
   "source": [
    "## Stem Out the Loops"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "13",
   "metadata": {},
   "outputs": [],
   "source": [
    "stem_group_names = [\n",
    "    \"Vocals\",\n",
    "    \"Backing_Vocals\",\n",
    "    \"Drums\",\n",
    "    \"Bass\",\n",
    "    \"Guitar\",\n",
    "    \"Keyboard\",\n",
    "    \"Percussion\",\n",
    "    \"Strings\",\n",
    "    \"Synth\",\n",
    "    \"FX\",\n",
    "    \"Brass\",\n",
    "    \"Woodwinds\",\n",
    "]\n",
    "\n",
    "def gen_stem(\n",
    "    audio: Audio,\n",
    "    stem_type_cfg_scale=1.0,\n",
    "    tags=\"extract [group_Vocals]\",\n",
    "    steps=8,\n",
    "    seed=3,\n",
    "    codec_scale_factor=0.4,\n",
    "    scale_ctx_vector=True,\n",
    "    noise_ctx_level=0.0,\n",
    "    infill_prefix_latents=None,\n",
    "    infill_suffix_latents=None,\n",
    "):\n",
    "    vae = diff_encode(audio, normalize_volume=False)\n",
    "    gen_cfg = diffusion_gen.DiffusionGenerationConfig(\n",
    "        lyrics=tags,\n",
    "        steps=steps,\n",
    "        seed=seed,\n",
    "        codec_scale_factor=codec_scale_factor,\n",
    "        scale_ctx_vector=scale_ctx_vector,\n",
    "        noise_ctx_level=noise_ctx_level,\n",
    "        text_cfg_coef=stem_type_cfg_scale,\n",
    "        infill_prefix_latents=infill_prefix_latents,\n",
    "        infill_suffix_latents=infill_suffix_latents,\n",
    "        drop_semantic_tokens=True,\n",
    "    )\n",
    "    print(gen_cfg)\n",
    "\n",
    "    request = Request(\n",
    "        id=\"dummy\",\n",
    "        generation_config=gen_cfg,\n",
    "        tokens=np.zeros((vae.shape[0], 1)),\n",
    "        input_tokens_finished=True,\n",
    "        stem_ctx_latents=vae,\n",
    "    )\n",
    "\n",
    "    result = engine.run_request(request)\n",
    "    vae_latents = torch.concat(result.vae_latents)\n",
    "    print(f\"vae_latents: {vae_latents.shape}\")\n",
    "    audios = []\n",
    "    for i in range(vae_latents.shape[1]):\n",
    "        audios.append(decode(vae_latents[:, i]))\n",
    "    return audios\n",
    "\n",
    "\n",
    "def get_stems(audio):\n",
    "    audios = gen_stem(\n",
    "        audio,\n",
    "        tags=\"extract [split_karaoke]\",\n",
    "    )\n",
    "\n",
    "    return audios\n",
    "\n",
    "def get_active_stems(audio_loops):\n",
    "    final_stems = []\n",
    "    for final_loop in audio_loops:\n",
    "        stem_loops = get_stems(final_loop)\n",
    "        final_stems.append(stem_loops)\n",
    "    \n",
    "    metadata = {\"tempo\": target_tempo, \"loops\": {}}\n",
    "    non_silent_stems = metadata[\"loops\"]\n",
    "\n",
    "    for loop_idx, loop in enumerate(final_stems):\n",
    "        loop_stems = []\n",
    "        duration_s = loop[0].duration_s\n",
    "        duration_bar = 4. * 60. / target_tempo\n",
    "        num_bars = round(duration_s / duration_bar)\n",
    "        for stem_idx, stem in enumerate(loop):\n",
    "            arr = stem.mono().array_float\n",
    "            max_amplitude = np.max(np.abs(arr))\n",
    "            if max_amplitude > 0.1:\n",
    "                loop_stems.append(\n",
    "                    {\n",
    "                        \"stem_type\": stem_group_names[stem_idx],\n",
    "                        \"max_amplitude\": float(max_amplitude),\n",
    "                        \"audio_data\": stem,\n",
    "                        \"duration_bars\": num_bars\n",
    "                    }\n",
    "                )\n",
    "        non_silent_stems[f\"loop_{loop_idx}\"] = loop_stems\n",
    "\n",
    "    return metadata\n",
    "\n",
    "def audition_stems(stems):\n",
    "    for _, loops in stems.items():\n",
    "        combined = Audio.sum([loop[\"audio_data\"] for loop in loops])\n",
    "        repeats = Audio.concatenate([combined, combined, combined, combined])\n",
    "        repeats.play()\n",
    "\n",
    "def save_stems_locally(stem_metadata):\n",
    "    output_dir = uuid\n",
    "    if os.path.exists(output_dir):\n",
    "        shutil.rmtree(output_dir)\n",
    "        \n",
    "    metadata_path = os.path.join(output_dir, \"metadata.json\")\n",
    "    flat_local = [metadata_path]\n",
    "    if not os.path.exists(output_dir):\n",
    "        os.makedirs(output_dir)\n",
    "\n",
    "    for loop_name, loops in stem_metadata[\"loops\"].items():\n",
    "        loop_path = os.path.join(output_dir, loop_name)\n",
    "        if not os.path.exists(loop_path):\n",
    "            os.makedirs(loop_path)\n",
    "        for stem in loops:\n",
    "            stem_type = stem[\"stem_type\"]\n",
    "            stem_fp = os.path.join(loop_path, f\"{stem_type}.wav\")\n",
    "            stem_audio = stem[\"audio_data\"]\n",
    "            stem_audio.write_wav(stem_fp)\n",
    "            stem[\"audio_path\"] = stem_fp\n",
    "            flat_local.append(stem_fp)\n",
    "            stem.pop(\"audio_data\")\n",
    "\n",
    "    write_json(stem_metadata, metadata_path)\n",
    "    return flat_local"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "14",
   "metadata": {},
   "outputs": [],
   "source": [
    "stem_data = get_active_stems(audio_loops)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "15",
   "metadata": {},
   "source": [
    "## Listen to Loop Mixes"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "16",
   "metadata": {},
   "outputs": [],
   "source": [
    "audition_stems(stem_data[\"loops\"])"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "17",
   "metadata": {},
   "source": [
    "## Save Outputs"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "18",
   "metadata": {},
   "outputs": [],
   "source": [
    "flat_local = save_stems_locally(stem_data)\n",
    "print(f\"Saving {len(flat_local) - 1} stems\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "19",
   "metadata": {},
   "outputs": [],
   "source": [
    "all_the_files = flat_local[1:]\n",
    "stem_paths = {stem: [] for stem in stem_group_names}\n",
    "for loop_fp in all_the_files:\n",
    "    for stem in stem_group_names:\n",
    "        if stem in loop_fp:\n",
    "            stem_paths[stem].append(loop_fp)\n",
    "\n",
    "stem_paths = {stem:data for stem,data in stem_paths.items() if len(data) > 0}\n",
    "\n",
    "condensed = {\"Drums\": [], \"Bass\": [], \"Synth\": [], \"Other\": []}\n",
    "for stem_name, loops in stem_paths.items():\n",
    "    if stem_name in condensed:\n",
    "        for loop in loops:\n",
    "            condensed[stem_name].append(loop)\n",
    "    else:\n",
    "        for loop in loops:\n",
    "            condensed[\"Other\"].append(loop)\n",
    "\n",
    "outdir = f\"{uuid}_flat\"\n",
    "if os.path.exists(outdir):\n",
    "    shutil.rmtree(outdir)\n",
    "os.makedirs(outdir)\n",
    "\n",
    "for stem_name, loop_paths in condensed.items():\n",
    "    for idx, loop_path in enumerate(loop_paths):\n",
    "        new_path = os.path.join(outdir, f\"{stem_name}_{idx}.wav\")\n",
    "        shutil.copy(loop_path, new_path)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "20",
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "suno_clean",
   "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": 5
}
