{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.audio import Audio\n",
    "import os\n",
    "from tqdm import tqdm\n",
    "import random"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "S3_AUDIO_DIR = \"s3://suno-data/datasets/harvest/karaoke_versions/stems/audio/\"\n",
    "\n",
    "splice_metas_file = \"s3://suno-data/datasets/harvest/splice/splice_all_samples_data_cleaned.jsonl\"\n",
    "karaoke_metas_file = \"s3://suno-data/datasets/harvest/karaoke_versions/stems/karaoke_versions.jsonl\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import polars as pl\n",
    "from suno_utils.utils.s3 import download_s3_file_if_needed\n",
    "\n",
    "splice_metas_file = download_s3_file_if_needed(splice_metas_file)\n",
    "karaoke_metas_file = download_s3_file_if_needed(karaoke_metas_file)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "!head $karaoke_metas_file"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "max_n_rows = None  # TODO: remove this\n",
    "splice_df = pl.read_ndjson(splice_metas_file, n_rows=max_n_rows)\n",
    "karaoke_df = pl.read_ndjson(karaoke_metas_file, n_rows=max_n_rows)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Remove the urls column from the splice dataframe\n",
    "splice_df = splice_df.drop(\"urls\")\n",
    "\n",
    "# Cleanup tags column\n",
    "splice_df = splice_df.with_columns(\n",
    "    pl.col(\"tags\").map_elements(lambda x: [tag[\"label\"] for tag in x], return_dtype=pl.List(pl.Utf8))\n",
    ")\n",
    "# Add s3_filepath column to splice_df\n",
    "splice_df = splice_df.with_columns(\n",
    "    pl.lit(\"s3://suno-data/datasets/harvest/splice/audio/\").alias(\"s3_base_path\"),\n",
    "    pl.col(\"uuid\").alias(\"file_id\"),\n",
    ")\n",
    "\n",
    "# Construct the full S3 filepath by combining the base path and file ID\n",
    "splice_df = splice_df.with_columns(\n",
    "    pl.concat_str([pl.col(\"s3_base_path\"), pl.col(\"file_id\"), pl.lit(\".mp3\")]).alias(\"s3_filepath\")\n",
    ")\n",
    "\n",
    "# Clean up temporary columns if needed\n",
    "splice_df = splice_df.drop([\"s3_base_path\", \"file_id\"])\n",
    "\n",
    "\n",
    "splice_df"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "karaoke_df"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Expand karaoke_df to have one row per track\n",
    "# The tracks column contains a list of structs with 6 elements:\n",
    "# [track_name, track_color, track_image_url, is_muted, segments, track_audio_path]\n",
    "\n",
    "# First, explode the tracks column to create one row per track\n",
    "karaoke_tracks_df = karaoke_df.explode(\"tracks\")\n",
    "\n",
    "# # Then extract each field from the struct\n",
    "karaoke_tracks_df = karaoke_tracks_df.with_columns(\n",
    "    [\n",
    "        pl.col(\"tracks\").struct.field(\"description\").alias(\"track_name\"),\n",
    "        pl.col(\"tracks\").struct.field(\"file_path\").alias(\"track_audio_path\"),\n",
    "    ]\n",
    ")\n",
    "\n",
    "# Remove click tracks from the karaoke dataset\n",
    "# First, let's see how many click tracks we have\n",
    "click_tracks_count = karaoke_tracks_df.filter(pl.col(\"track_name\") == \"Click\").shape[0]\n",
    "print(f\"Number of Click tracks in Karaoke dataset: {click_tracks_count}\")\n",
    "\n",
    "# Now remove the click tracks\n",
    "karaoke_tracks_df = karaoke_tracks_df.filter(pl.col(\"track_name\") != \"Click\")\n",
    "\n",
    "# Verify the removal\n",
    "print(f\"Number of tracks after removing Click tracks: {karaoke_tracks_df.shape[0]}\")\n",
    "\n",
    "\n",
    "# Add the s3_filepath column\n",
    "karaoke_tracks_df = karaoke_tracks_df.with_columns(\n",
    "    pl.concat_str(\n",
    "        [\n",
    "            pl.lit(\"s3://suno-data/datasets/harvest/karaoke_versions/stems/audio/\"),\n",
    "            pl.col(\"track_audio_path\"),\n",
    "        ]\n",
    "    ).alias(\"s3_filepath\")\n",
    ")\n",
    "\n",
    "# Verify the s3_filepath column was added correctly\n",
    "print(f\"Sample s3_filepath: {karaoke_tracks_df.select('s3_filepath').head(1)[0, 0]}\")\n",
    "\n",
    "\n",
    "karaoke_tracks_df.head()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Count the number of tracks in each dataset\n",
    "\n",
    "# For splice_df, just count the number of rows\n",
    "print(\"Number of samples in Splice dataset:\", len(splice_df))\n",
    "\n",
    "# For karaoke_df, sum the length of tracks column\n",
    "total_tracks_karaoke = karaoke_df.select(pl.col(\"tracks\").list.len()).sum().item()\n",
    "\n",
    "print(f\"Number of tracks in Karaoke dataset: {total_tracks_karaoke}\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "CATEGORIES = {\n",
    "    \"Drums\": [\n",
    "        \"kicks\",\n",
    "        \"hats\",\n",
    "        \"toms\",\n",
    "        \"breaks\",\n",
    "        \"acoustic\",\n",
    "        \"snares\",\n",
    "        \"claps\",\n",
    "        \"cymbals\",\n",
    "        \"fills\",\n",
    "        \"808\",\n",
    "        \"drum\",\n",
    "        \"drums\",\n",
    "    ],\n",
    "    \"Vocals\": [\n",
    "        \"female vocals\",\n",
    "        \"vocal fx\",\n",
    "        \"vocoder\",\n",
    "        \"screams\",\n",
    "        \"whisper vocals\",\n",
    "        \"male vocals\",\n",
    "        \"spoken word\",\n",
    "        \"vocal phrases\",\n",
    "        \"vocal shouts\",\n",
    "        \"dialogue\",\n",
    "        \"vocal\",\n",
    "        \"vocals\",\n",
    "        \"voice\",\n",
    "    ],\n",
    "    \"Percussion\": [\n",
    "        \"shakers\",\n",
    "        \"grooves\",\n",
    "        \"bongos\",\n",
    "        \"woodblock\",\n",
    "        \"djembe\",\n",
    "        \"conga\",\n",
    "        \"tambourine\",\n",
    "        \"cowbells\",\n",
    "        \"bells\",\n",
    "        \"timbales\",\n",
    "        \"percussion\",\n",
    "    ],\n",
    "    \"Synth\": [\n",
    "        \"synth bass\",\n",
    "        \"pads\",\n",
    "        \"stabs\",\n",
    "        \"plucks\",\n",
    "        \"fx\",\n",
    "        \"leads\",\n",
    "        \"arp\",\n",
    "        \"chords\",\n",
    "        \"analog\",\n",
    "        \"synth melody\",\n",
    "        \"synthesizer\",\n",
    "        \"synth\",\n",
    "    ],\n",
    "    \"Brass & Woodwinds\": [\n",
    "        \"saxophone\",\n",
    "        \"trombone\",\n",
    "        \"ensemble\",\n",
    "        \"riffs\",\n",
    "        \"trumpet\",\n",
    "        \"flute\",\n",
    "        \"harmonica\",\n",
    "        \"brass\",\n",
    "        \"woodwind\",\n",
    "        \"sax\",\n",
    "        \"horn\",\n",
    "    ],\n",
    "    \"Keys\": [\n",
    "        \"piano\",\n",
    "        \"wurlitzer\",\n",
    "        \"chords\",\n",
    "        \"stabs\",\n",
    "        \"electric piano\",\n",
    "        \"organ\",\n",
    "        \"clavinet\",\n",
    "        \"keys melody\",\n",
    "        \"classical\",\n",
    "        \"keys\",\n",
    "        \"keyboard\",\n",
    "    ],\n",
    "    \"Guitar\": [\n",
    "        \"electric\",\n",
    "        \"clean\",\n",
    "        \"leads\",\n",
    "        \"riffs\",\n",
    "        \"acoustic\",\n",
    "        \"distorted\",\n",
    "        \"chords\",\n",
    "        \"guitar melody\",\n",
    "        \"rhythm\",\n",
    "        \"guitar\",\n",
    "    ],\n",
    "    \"Bass\": [\n",
    "        \"synth bass\",\n",
    "        \"analog\",\n",
    "        \"electric\",\n",
    "        \"saw\",\n",
    "        \"wobble\",\n",
    "        \"sub\",\n",
    "        \"acoustic\",\n",
    "        \"acid\",\n",
    "        \"distorted\",\n",
    "        \"pulse\",\n",
    "        \"bass\",\n",
    "    ],\n",
    "    \"FX\": [\n",
    "        \"noise\",\n",
    "        \"downers\",\n",
    "        \"impacts\",\n",
    "        \"textures\",\n",
    "        \"field recordings\",\n",
    "        \"risers\",\n",
    "        \"sweeps\",\n",
    "        \"atmospheres\",\n",
    "        \"reverse\",\n",
    "        \"fx vocals\",\n",
    "        \"effects\",\n",
    "    ],\n",
    "    \"Strings\": [\n",
    "        \"violin\",\n",
    "        \"viola\",\n",
    "        \"ensemble\",\n",
    "        \"staccato\",\n",
    "        \"cello\",\n",
    "        \"bass\",\n",
    "        \"orchestral\",\n",
    "        \"pads\",\n",
    "        \"strings melody\",\n",
    "        \"strings\",\n",
    "    ],\n",
    "}\n",
    "\n",
    "\n",
    "def get_tag_category(tag):\n",
    "    \"\"\"\n",
    "    Determines which category a tag belongs to based on predefined groups.\n",
    "\n",
    "    Args:\n",
    "        tag (str): The tag to categorize\n",
    "\n",
    "    Returns:\n",
    "        str: The category name, or \"Other\" if not found\n",
    "    \"\"\"\n",
    "    tag = tag.lower()\n",
    "\n",
    "    # Check which category the tag belongs to\n",
    "    for category, tags in CATEGORIES.items():\n",
    "        if any(t in tag for t in tags):\n",
    "            return category\n",
    "\n",
    "    return \"Other\"\n",
    "\n",
    "\n",
    "# Add tag category column to the splice dataframe\n",
    "def add_tag_category_to_splice(df):\n",
    "    # Define a function to extract the first tag and determine its category\n",
    "    def extract_first_tag_category(tags_struct):\n",
    "        # Extract the first tag name from the struct\n",
    "        # The structure appears to be a list of structs based on the output\n",
    "        if len(tags_struct) > 0:\n",
    "            first_tag = tags_struct[0]\n",
    "            if first_tag:\n",
    "                return get_tag_category(first_tag)\n",
    "        return \"Unknown\"\n",
    "\n",
    "    # Apply the function to create a new column with specified return_dtype\n",
    "    return df.with_columns(\n",
    "        pl.col(\"tags\")\n",
    "        .map_elements(extract_first_tag_category, return_dtype=pl.Utf8)\n",
    "        .alias(\"tag_category\")\n",
    "    )\n",
    "\n",
    "\n",
    "splice_with_categories = add_tag_category_to_splice(splice_df)\n",
    "splice_with_categories\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Plot histogram of tag categories\n",
    "import matplotlib.pyplot as plt\n",
    "import seaborn as sns\n",
    "\n",
    "# Count the occurrences of each category\n",
    "category_counts = (\n",
    "    splice_with_categories.group_by(\"tag_category\")\n",
    "    .agg(pl.count().alias(\"count\"))\n",
    "    .sort(\"count\", descending=True)\n",
    ")\n",
    "\n",
    "# Convert to pandas for easier plotting with matplotlib\n",
    "category_counts_pd = category_counts.to_pandas()\n",
    "\n",
    "# Set up the plot style\n",
    "plt.figure(figsize=(12, 8))\n",
    "sns.set_style(\"whitegrid\")\n",
    "\n",
    "# Create the bar plot\n",
    "ax = sns.barplot(x=\"tag_category\", y=\"count\", data=category_counts_pd)\n",
    "\n",
    "# Add labels and title\n",
    "plt.title(\"Distribution of Tag Categories\", fontsize=16)\n",
    "plt.xlabel(\"Category\", fontsize=14)\n",
    "plt.ylabel(\"Count\", fontsize=14)\n",
    "plt.xticks(rotation=45, ha=\"right\")\n",
    "\n",
    "# Add count labels on top of each bar\n",
    "for i, v in enumerate(category_counts_pd[\"count\"]):\n",
    "    ax.text(i, v + 0.1, str(v), ha=\"center\", fontsize=10)\n",
    "\n",
    "plt.tight_layout()\n",
    "plt.show()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Create the histogram\n",
    "# Count the occurrences of each track name\n",
    "track_counts = karaoke_tracks_df.group_by(\"track_name\").agg(pl.count().alias(\"count\"))\n",
    "# sort by count\n",
    "track_counts = track_counts.sort(\"count\", descending=True)\n",
    "# get top 50\n",
    "top_50_tracks = track_counts.head(20).to_pandas()\n",
    "# Create the bar plot for top 50 tracks\n",
    "ax = sns.barplot(x=\"track_name\", y=\"count\", data=top_50_tracks)\n",
    "plt.xticks(rotation=45, ha=\"right\")\n",
    "plt.title(\"Top 50 Track Names\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import random\n",
    "import numpy as np\n",
    "\n",
    "\n",
    "def random_lengthen_positions(audio: Audio, duration_s: float):\n",
    "    max_repeat = 1 + (duration_s // audio.duration_s)\n",
    "    n_repeats = random.randint(1, max_repeat)\n",
    "\n",
    "    positions = []\n",
    "    for i in range(n_repeats):\n",
    "        positions.append(random.randint(-5, duration_s))\n",
    "    return positions\n",
    "\n",
    "\n",
    "def random_lengthen(audio: Audio, duration_s: float, positions: list[int] = None):\n",
    "    if positions is None:\n",
    "        positions = random_lengthen_positions(audio, duration_s)\n",
    "    # print(positions)\n",
    "\n",
    "    arr = audio.array_float\n",
    "    out_wav = np.zeros(\n",
    "        (\n",
    "            arr.shape[0],\n",
    "            int(duration_s * audio.sample_rate),\n",
    "        )\n",
    "    )\n",
    "\n",
    "    for i, pos in enumerate(positions):\n",
    "        start_idx = pos * audio.sample_rate\n",
    "\n",
    "        if start_idx < 0:\n",
    "            a = arr[:, start_idx:]\n",
    "            start_idx = 0\n",
    "        else:\n",
    "            a = arr\n",
    "        dur = min(a.shape[-1], out_wav.shape[-1] - start_idx)\n",
    "        out_wav[:, start_idx : start_idx + dur] += a[:, :dur]\n",
    "    return Audio.from_array_float(out_wav, sample_rate=audio.sample_rate, max_allowed_val=100)\n",
    "\n",
    "\n",
    "mp3_path = f\"s3://suno-data/datasets/harvest/splice/audio/{splice_df['uuid'][0]}.mp3\"\n",
    "print(mp3_path)\n",
    "audio = Audio.from_s3(mp3_path, n_channels=2)\n",
    "# print(audio.duration_s)\n",
    "random_lengthen(audio, 60).play()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(CATEGORIES.keys())\n",
    "\n",
    "CATEGORY_WEIGHTS = {\n",
    "    \"Drums\": 0.2,\n",
    "    \"Vocals\": 1,\n",
    "    \"Percussion\": 0.5,\n",
    "    \"Synth\": 0.2,\n",
    "    \"Brass & Woodwinds\": 1,\n",
    "    \"Keys\": 1,\n",
    "    \"Guitar\": 1,\n",
    "    \"Bass\": 1,\n",
    "    \"FX\": 1,\n",
    "    \"Strings\": 1,\n",
    "    \"Other\": 1,\n",
    "}"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import tempfile\n",
    "import multiprocessing as mp\n",
    "from functools import partial\n",
    "\n",
    "from suno_utils.utils.s3 import download_s3_files, upload_s3_files\n",
    "\n",
    "\n",
    "def process_stem(sample):\n",
    "    audio_path = sample[\"s3_filepath\"][0]\n",
    "\n",
    "    # print(audio_path)\n",
    "    # random lengthen\n",
    "    audio = Audio.from_s3(audio_path, sample_rate=48000, n_channels=2)\n",
    "    audio = random_lengthen(audio, 45)\n",
    "\n",
    "    sample = sample.to_dicts()[0]\n",
    "    if \"graphs\" in sample:\n",
    "        del sample[\"graphs\"]\n",
    "    return audio, sample\n",
    "\n",
    "\n",
    "def process_random_song(samples):\n",
    "    stems = [process_stem(sample) for sample in samples]\n",
    "    stem_metas = [stem[1] for stem in stems]\n",
    "    stems = [stem[0] for stem in stems]\n",
    "\n",
    "    # mix stems\n",
    "    mixed_audio = Audio.sum([stem for stem in stems])\n",
    "\n",
    "    complements = []\n",
    "    for i, stem in enumerate(stems):\n",
    "        audios_to_mix = [stem for j, stem in enumerate(stems) if j != i]\n",
    "        complements.append(Audio.sum(audios_to_mix))\n",
    "    return mixed_audio, stem_metas, stems, complements\n",
    "\n",
    "\n",
    "def create_and_upload(samples):\n",
    "    import uuid\n",
    "\n",
    "    mixed_audio, stem_metas, stems, complements = process_random_song(samples)\n",
    "\n",
    "    local_fps = []\n",
    "    s3_filepaths = []\n",
    "    with tempfile.TemporaryDirectory() as tempdir:\n",
    "        new_id = str(uuid.uuid4())\n",
    "\n",
    "        # full mix\n",
    "        fp = os.path.join(tempdir, \"full_mix.mp3\")\n",
    "        mixed_audio.to_hq_mp3(fp)\n",
    "        local_fps.append(fp)\n",
    "        s3_filepaths.append(\n",
    "            f\"s3://suno-data/datasets/harvest/splice/random_mixed_45s/{new_id}/full_mix.mp3\"\n",
    "        )\n",
    "\n",
    "        # stems\n",
    "        for i, stem in enumerate(stems):\n",
    "            fp = os.path.join(tempdir, f\"stem_{i}.mp3\")\n",
    "            stem.to_hq_mp3(fp)\n",
    "            local_fps.append(fp)\n",
    "            s3_filepaths.append(\n",
    "                f\"s3://suno-data/datasets/harvest/splice/random_mixed_45s/{new_id}/stem_{i}.mp3\"\n",
    "            )\n",
    "            stem_metas[i][\"stem_s3_filepath\"] = s3_filepaths[-1]\n",
    "\n",
    "        # complements\n",
    "        for i, complement in enumerate(complements):\n",
    "            fp = os.path.join(tempdir, f\"complement_{i}.mp3\")\n",
    "            complement.to_hq_mp3(fp)\n",
    "            local_fps.append(fp)\n",
    "            s3_filepaths.append(\n",
    "                f\"s3://suno-data/datasets/harvest/splice/random_mixed_45s/{new_id}/complement_{i}.mp3\"\n",
    "            )\n",
    "            stem_metas[i][\"complement_s3_filepath\"] = s3_filepaths[-1]\n",
    "\n",
    "        upload_s3_files(\n",
    "            local_fps,\n",
    "            s3_filepaths,\n",
    "            chunksize=1000,\n",
    "            n_cores=20,\n",
    "            joblib_backend=\"threads\",\n",
    "            silent=True,\n",
    "        )\n",
    "    # create meta\n",
    "    import uuid\n",
    "\n",
    "    meta = {\n",
    "        \"id\": str(uuid.uuid4()),\n",
    "        \"full_mix_s3_filepath\": s3_filepaths[0],\n",
    "        \"duration_s\": 45,\n",
    "        \"stems\": stem_metas,\n",
    "    }\n",
    "    return meta\n",
    "    return mixed_audio, stems, meta\n",
    "\n",
    "\n",
    "def safe_create_and_upload(samples):\n",
    "    try:\n",
    "        return create_and_upload(samples)\n",
    "    except Exception as e:\n",
    "        print(e)\n",
    "        return None\n",
    "\n",
    "\n",
    "# Pre-filter datasets by category for faster sampling\n",
    "category_datasets = {}\n",
    "for category in tqdm(CATEGORY_WEIGHTS.keys(), desc=\"Pre-filtering datasets\"):\n",
    "    category_datasets[category] = splice_with_categories.filter(pl.col(\"tag_category\") == category)\n",
    "\n",
    "\n",
    "def sample_samples(n_samples):\n",
    "    samples = []\n",
    "    for _ in range(n_samples):\n",
    "        # sample dataset\n",
    "        if random.random() < 0.5:\n",
    "            # first sample a category\n",
    "            category = random.choices(\n",
    "                list(CATEGORY_WEIGHTS.keys()), weights=list(CATEGORY_WEIGHTS.values())\n",
    "            )[0]\n",
    "            # then sample from the pre-filtered dataset for that category\n",
    "            sample = (category, random.randint(0, len(category_datasets[category]) - 1))\n",
    "        else:\n",
    "            sample = (\"karaoke\", random.randint(0, len(karaoke_tracks_df) - 1))\n",
    "        samples.append(sample)\n",
    "    return samples\n",
    "\n",
    "\n",
    "def load_sample(samples):\n",
    "    new_samples = []\n",
    "    for sample in samples:\n",
    "        if sample[0] == \"karaoke\":\n",
    "            new_samples.append(karaoke_tracks_df[sample[1]])\n",
    "        else:\n",
    "            new_samples.append(category_datasets[sample[0]][sample[1]])\n",
    "    return new_samples\n",
    "\n",
    "\n",
    "meta = safe_create_and_upload(load_sample(sample_samples(5)))\n",
    "meta"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# mixed_audio = Audio.from_s3(\n",
    "#     meta[\"full_mix_s3_filepath\"], sample_rate=48000, n_channels=2\n",
    "# )\n",
    "# mixed_audio.play()\n",
    "# for i, stem_meta in enumerate(meta[\"stems\"]):\n",
    "#     print(stem_meta)\n",
    "#     stem = Audio.from_s3(stem_meta[\"stem_s3_filepath\"], sample_rate=48000, n_channels=2)\n",
    "#     complement = Audio.from_s3(\n",
    "#         stem_meta[\"complement_s3_filepath\"], sample_rate=48000, n_channels=2\n",
    "#     )\n",
    "#     stem.play()\n",
    "#     complement.play()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# remove all files in s3://suno-data/datasets/harvest/splice/random_mixed_45s/\n",
    "# !aws s3 rm s3://suno-data/datasets/harvest/splice/random_mixed_45s/ --recursive\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Use multiprocessing to speed up processing\n",
    "num_cores = 80  # Leave one core free\n",
    "N = 40_000\n",
    "\n",
    "\n",
    "# Generate samples without using multiprocessing\n",
    "inputs = []\n",
    "for i in tqdm(range(N), desc=\"Generating samples\"):\n",
    "    inputs.append(load_sample(sample_samples(random.randint(3, 15))))\n",
    "\n",
    "with mp.Pool(num_cores) as pool:\n",
    "    results = list(tqdm(pool.imap(safe_create_and_upload, inputs), total=len(inputs)))\n",
    "\n",
    "# Filter out None results and flatten the list\n",
    "metas = [meta for result in results if result is not None]\n",
    "\n",
    "metas[0]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import json\n",
    "\n",
    "# save metas to file\n",
    "with open(\"/tmp/random_mixed_45s_meta.jsonl\", \"w\") as f:\n",
    "    for meta in results:\n",
    "        f.write(json.dumps(meta) + \"\\n\")\n",
    "\n",
    "# upload file to s3\n",
    "!aws s3 cp /tmp/random_mixed_45s_meta.jsonl s3://suno-data/datasets/harvest/splice/random_mixed_45s/meta.jsonl\n",
    "\n",
    "results[0]\n"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Prep bundle"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.utils.s3 import download_s3_file_if_needed\n",
    "\n",
    "metas_file = download_s3_file_if_needed(\n",
    "    \"s3://suno-data/datasets/harvest/splice/random_mixed_45s/meta.jsonl\"\n",
    ")\n",
    "\n",
    "import polars as pl\n",
    "\n",
    "metas = pl.read_ndjson(metas_file).to_dicts()\n",
    "metas[0]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.utils.text import get_filename\n",
    "from tqdm import tqdm\n",
    "\n",
    "flat_metas = []\n",
    "for m in tqdm(metas, desc=\"Expanding metas\"):\n",
    "    if m[\"id\"] is None:\n",
    "        continue\n",
    "    flat_metas.append(\n",
    "        {\n",
    "            \"id\": f\"{m['id']}_main\",\n",
    "            \"bundle_id\": str(m[\"id\"]),\n",
    "            \"s3_filepath\": m[\"full_mix_s3_filepath\"],\n",
    "            \"duration_s\": m[\"duration_s\"],\n",
    "            \"type\": \"full_mix\",\n",
    "        }\n",
    "    )\n",
    "    for i, mm in enumerate(m[\"stems\"]):\n",
    "        flat_metas.append(\n",
    "            {\n",
    "                \"id\": f\"{m['id']}_stem_{i}\",\n",
    "                \"bundle_id\": str(m[\"id\"]),\n",
    "                \"s3_filepath\": mm[\"stem_s3_filepath\"],\n",
    "                \"duration_s\": m[\"duration_s\"],\n",
    "                \"stem_type\": mm.get(\"track_name\", \"\"),\n",
    "                \"stem_tags\": mm.get(\"tags\", []),\n",
    "                \"type\": \"stem\",\n",
    "            }\n",
    "        )\n",
    "        flat_metas.append(\n",
    "            {\n",
    "                \"id\": f\"{m['id']}_complement_{i}\",\n",
    "                \"bundle_id\": str(m[\"id\"]),\n",
    "                \"s3_filepath\": mm[\"complement_s3_filepath\"],\n",
    "                \"duration_s\": m[\"duration_s\"],\n",
    "                \"stem_type\": mm.get(\"title\", \"\"),\n",
    "                \"stem_tags\": mm.get(\"tags\", []),\n",
    "                \"type\": \"complement\",\n",
    "            }\n",
    "        )\n",
    "\n",
    "\n",
    "print(f\"{len([m for m in flat_metas if m['type'] == 'full_mix']):,} main metas\")\n",
    "print(f\"{len([m for m in flat_metas if m['type'] == 'stem']):,} stem metas\")\n",
    "print(f\"{len([m for m in flat_metas if m['type'] == 'complement']):,} complement metas\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import random\n",
    "\n",
    "random.sample(flat_metas, 10)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.utils.text import write_jsonl\n",
    "\n",
    "write_jsonl(flat_metas, \"/tmp/random_mixed_45s_flat_metas.jsonl\")\n",
    "!aws s3 cp /tmp/random_mixed_45s_flat_metas.jsonl s3://suno-data/datasets/bundles/v4/random_mixed_45s/metas.jsonl"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# modal run /home/victor/glockenspiel/suno_utils/suno_utils/scripts/gpt/modal_encode.py \\\n",
    "#     --embed-type='dac_vae_tuned_25hz' \\\n",
    "#     --base-s3-dir='s3://suno-data/datasets/bundles/v4/random_mixed_45s/' \\\n",
    "#     --chunksize=10 \\\n",
    "#     --min-duration-s=5 \\\n",
    "#     --max-duration-s=480 \\\n",
    "#     --output-name='dac_vae_tuned_25hz' \\\n",
    "#     --normalize-volume=False \\\n",
    "#     --first-only=True \\\n",
    "#     --force-overwrite=True"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## check embeds"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "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",
    "    Audio,\n",
    ")\n",
    "\n",
    "_ = preload_codec_models(\"s3://suno-data/minz/models/dac_vae_tuned_25hz.pth\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from sunodata.dataset_maker_utils import Bundle, DatasetConfig, MemmapMaker\n",
    "\n",
    "bundle = Bundle(name=\"random_mixed_45s\")\n",
    "npz = bundle.get_part(0)\n",
    "npz"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "part_metas = bundle.get_part_metas(0)\n",
    "Audio.from_s3(part_metas[0][\"s3_filepath\"]).play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(list(npz.keys()))\n",
    "codec_decode(npz[\"9263be0a-a3e2-42a5-ae8c-9263d13fc1a6_stem_0\"]).play()\n",
    "\n",
    "bundle.get_part_metas(0)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 4
}
