{
 "cells": [
  {
   "cell_type": "markdown",
   "id": "0",
   "metadata": {},
   "source": [
    "## Setup"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "1",
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "import glob\n",
    "from suno_utils.audio import Audio\n",
    "from suno_utils.audio.midi import Midi\n",
    "from pathlib import Path\n",
    "from suno_utils.utils.s3 import upload_s3_files\n",
    "\n",
    "class MidiPair:\n",
    "    def __init__(self, midi_file: str, audio_file: str):\n",
    "        self.midi_file = midi_file\n",
    "        self.audio_file = audio_file\n",
    "        self.midi = None\n",
    "        self.audio = None\n",
    "\n",
    "    def load_midi(self):\n",
    "        if self.midi is not None:\n",
    "            return\n",
    "        self.midi = Midi.from_path(self.midi_file)\n",
    "\n",
    "    def load_audio(self):\n",
    "        if self.audio is not None:\n",
    "            return\n",
    "        self.audio = Audio.from_file(self.audio_file)\n",
    "\n",
    "    def load_all(self):\n",
    "        self.load_midi()\n",
    "        self.load_audio()\n",
    "\n",
    "    def __str__(self):\n",
    "        return f\"MidiPair(midi_file={self.midi_file}, audio_file={self.audio_file})\"\n",
    "\n",
    "    def __repr__(self):\n",
    "        return self.__str__()\n",
    "\n",
    "    def play(self):\n",
    "        self.load_all()\n",
    "        stereo_audio = self.midi.make_stereo_comparison(self.audio)\n",
    "        stereo_audio.play()\n",
    "\n",
    "def load_pairs(\n",
    "    dir: str, audio_ext: str = \"mp3\", midi_ext: str = \"mid\", include_parent_dir: bool = False\n",
    "):\n",
    "    midi_files = glob.glob(os.path.join(dir, \"**\", f\"*.{midi_ext}\"), recursive=True)\n",
    "    audio_files = glob.glob(os.path.join(dir, \"**\", f\"*.{audio_ext}\"), recursive=True)\n",
    "\n",
    "    # Create a mapping of base filenames to audio files\n",
    "    audio_map = {}\n",
    "    for audio_file in audio_files:\n",
    "        base_name = os.path.splitext(os.path.basename(audio_file))[0]\n",
    "        if include_parent_dir:\n",
    "            parent_dir = os.path.basename(os.path.dirname(audio_file))\n",
    "            base_name = os.path.join(parent_dir, base_name)\n",
    "        audio_map[base_name] = audio_file\n",
    "\n",
    "    pairs = []\n",
    "    for midi_file in midi_files:\n",
    "        base_name = os.path.splitext(os.path.basename(midi_file))[0]\n",
    "        if include_parent_dir:\n",
    "            parent_dir = os.path.basename(os.path.dirname(midi_file))\n",
    "            base_name = os.path.join(parent_dir, base_name)\n",
    "        if base_name in audio_map:\n",
    "            pairs.append(MidiPair(midi_file, audio_map[base_name]))\n",
    "\n",
    "    return pairs"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2",
   "metadata": {},
   "outputs": [],
   "source": [
    "from tqdm import tqdm\n",
    "import random\n",
    "\n",
    "def load_and_shuffle_pairs():\n",
    "    synthetic_dir = \"/app2/suno/data/victor/gmd/MIDIs/\"\n",
    "    synthetic_pairs = []\n",
    "    for i in tqdm([\"0\", \"1\", \"2\", \"3\", \"4\", \"5\", \"6\", \"7\", \"8\", \"9\", \"a\", \"b\", \"c\", \"d\", \"e\", \"f\"]):\n",
    "        synthetic_pairs.extend(load_pairs(synthetic_dir + i, audio_ext=\"opus\"))\n",
    "        print(f\"Loaded {len(synthetic_pairs)} synthetic pairs\")\n",
    "\n",
    "    random.shuffle(synthetic_pairs)\n",
    "    return synthetic_pairs\n",
    "    \n",
    "synthetic_pairs = load_and_shuffle_pairs()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3",
   "metadata": {},
   "outputs": [],
   "source": [
    "test_set = synthetic_pairs[:1000]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "4",
   "metadata": {},
   "outputs": [],
   "source": [
    "local_files = [pair.midi_file for pair in test_set]\n",
    "s3_files = [f\"s3://suno-data/sara/midi/midi_test_set/{Path(x).name}\" for x in local_files]\n",
    "upload_s3_files(from_local_filepaths=local_files, to_s3_filepaths=s3_files)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "5",
   "metadata": {},
   "source": [
    "## Fixed sample of gmd dataset with marc's comments"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "6",
   "metadata": {},
   "outputs": [],
   "source": [
    "def gmd_from_basename(basename: str):\n",
    "    return MidiPair(\n",
    "        midi_file=f\"/app2/suno/data/victor/gmd/MIDIs/{basename[0]}/{basename}.mid\",\n",
    "        audio_file=f\"/app2/suno/data/victor/gmd/MIDIs/{basename[0]}/{basename}.opus\",\n",
    "    )\n",
    "\n",
    "basenames_with_comments = [\n",
    "    (\"a6f035c45845d66eb3ebf0c3d61d086a\", \"ok\"),\n",
    "    (\"af1b5c2180b19148161a3cbfe5eebe5a\", \"ok\"),\n",
    "    (\"272670b0ef33e65da97b47e5a0562e31\", \"ok\"),\n",
    "    (\"a96df1cea93a25470f87eb2b0e1a2eac\", \"ok\"),\n",
    "    (\"a753d3f11a56f02fa65203c73dac94ab\", \"different soundfont?\"),\n",
    "    (\"a06d5f7a7d7057a0c02907740451c454\", \"in stored: some very short and some inappropriately sustained notes, buggy-sounding transposition of bass starting around 37s\"),\n",
    "    (\"01c17e7e621c8e4d71091dbda28b735b\", \"in stored: way too much sustain, sounds like all the piano notes are sustained for a whole bar at a time, missing instruments\"),\n",
    "    (\"f17699aed9d6443ae69d80d3d25604c4\", \"in stored: instruments transposed incorrectly, unharmonic\"),\n",
    "    (\"abc7b35355f46579749555f3dbede21d\", \"in stored: more transpositions, missing instruments?\"),\n",
    "    (\"a6c2140f51b857ebc8c233f03abb48a4\", \"ok\")\n",
    "]\n",
    "\n",
    "for basename, comment in basenames_with_comments[:10]:\n",
    "    print(basename, comment)\n",
    "    # stereo test: left channel = just synthesized from midi, right channel = stored audio file for model trainin\n",
    "    gmd_from_basename(basename).play()"
   ]
  }
 ],
 "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
}
