# %% [markdown] # # Consolidate midi datasets # # This notebook consolidates the midi datasets into a single dataset. It also serves as documentation for the midi datasets. # # It produces a metadata file pointing to the midi and audio files. # # ## Datasets # # - Infinite synthetic MIDI # - Mew # - `/app2/suno/data/mew/scoretube` # - `/app2/suno/data/mew/mmd_chunks` # - Match midi instruments to stems. # - A dataset of loosely paired real music with arrangements # - [Slack](https://suno-main.slack.com/archives/C03A7347CSK/p1719682717375969) # - 20k songs matched with high confidence # - [The AI workspace that works for you. | Notion](https://www.notion.so/suno-ai/MooTube-187b01573ccf8065915ecb04af262d6c#187b01573ccf80de941bc4412c900eca) # - Hooktheory melodys, 50k hooks only # - [The AI workspace that works for you. | Notion](https://www.notion.so/suno-ai/hooktheory-91267f75401f4087b506f086abefc65f) # - Trombone champ melodys. 4k full songs # - [The AI workspace that works for you. | Notion](https://www.notion.so/suno-ai/Trombone-Champ-206b01573ccf8009aa1fe22f3440687c) # - Piano Maestro # # # # %% [markdown] # # mew # # Victors's personal stash # %% import os import glob import random from suno_utils.audio import Audio from suno_utils.audio.midi import Midi class MidiPair: def __init__(self, midi_file: str, audio_file: str): self.midi_file = midi_file self.audio_file = audio_file self.midi = None self.audio = None def load_midi(self): if self.midi is not None: return self.midi = Midi.from_path(self.midi_file) def load_audio(self): if self.audio is not None: return self.audio = Audio.from_file(self.audio_file) def load_all(self): self.load_midi() self.load_audio() def __str__(self): return f"MidiPair(midi_file={self.midi_file}, audio_file={self.audio_file})" def __repr__(self): return self.__str__() def play(self): self.load_all() stereo_audio = self.midi.make_stereo_comparison(self.audio) stereo_audio.play() scoretube_dir = "/app2/suno/data/victor/mew/scoretube" def load_pairs(dir: str, audio_ext: str = "mp3", midi_ext: str = "mid"): midi_files = glob.glob(os.path.join(dir, "**", f"*.{midi_ext}"), recursive=True) audio_files = glob.glob(os.path.join(dir, "**", f"*.{audio_ext}"), recursive=True) # Create a mapping of base filenames to audio files audio_map = {} for audio_file in audio_files: base_name = os.path.splitext(os.path.basename(audio_file))[0] audio_map[base_name] = audio_file pairs = [] for midi_file in midi_files: base_name = os.path.splitext(os.path.basename(midi_file))[0] if base_name in audio_map: pairs.append(MidiPair(midi_file, audio_map[base_name])) return pairs scoretube_pairs = load_pairs(scoretube_dir) print(f"Loaded {len(scoretube_pairs)} scoretube pairs") # play random pair random_pair = random.choice(scoretube_pairs) print(random_pair) # random_pair.play() # %% mmd_dir = "/app2/suno/data/victor/mew/mmd_chunks" mmd_pairs = load_pairs(mmd_dir) print(f"Loaded {len(mmd_pairs)} mmd pairs") # %% [markdown] # ## mootube # %% import json mootube_dir = "/app2/suno/data/victor/mootube" mootube_metas = [] for line in open("/app2/suno/data/victor/mootube/score_ytm_clean.jsonl"): meta = json.loads(line) mootube_metas.append(meta) print(f"Loaded {len(mootube_metas)} moo tube metas") mootube_metas[0] # %% mootube_metas[0]["matches"][0]["videoId"] # %% mootube_pairs = [] for line in open("/app2/suno/data/victor/mootube/score_ytm_clean.jsonl"): meta = json.loads(line) midi_file = f"/app2/suno/data/victor/mootube/midi/{meta['id']}.mid" audio_file = f"/app2/suno/data/victor/mootube/audio/{meta['matches'][0]['videoId']}.webm" mootube_pairs.append(MidiPair(midi_file, audio_file)) print(f"Loaded {len(mootube_pairs)} moo tube pairs") # play random pair for _ in range(1): random_pair = random.choice(mootube_pairs) print(random_pair) # random_pair.play() # %% [markdown] # ## maestro # 200 hours of piano # %% maestro_dir = "/app2/suno/data/victor/maestro/maestro-v3.0.0" maestro_pairs = load_pairs(maestro_dir, audio_ext="wav", midi_ext="midi") print(f"Loaded {len(maestro_pairs)} maestro pairs") # play random pair random_pair = random.choice(maestro_pairs) print(random_pair) # random_pair.play() # %% [markdown] # ## trombone champ # %% trombone_champ_dir = "/app2/suno/data/victor/trombone_champ" trombone_champ_pairs = load_pairs(trombone_champ_dir, audio_ext="opus") print(f"Loaded {len(trombone_champ_pairs)} trombone champ pairs") # play random pair random_pair = random.choice(trombone_champ_pairs) print(random_pair) # random_pair.play() # %% [markdown] # ## hook theory # %% hook_theory_dir = "/app2/suno/data/victor/hooktheory" hooktheory_pairs = load_pairs(hook_theory_dir, audio_ext="opus") print(f"Loaded {len(hooktheory_pairs)} hook theory pairs") # play random pair random_pair = random.choice(hooktheory_pairs) print(random_pair) # random_pair.play() # %% [markdown] # ## EDA # %% # Calculate duration statistics for each dataset import random # %% [markdown] # # Extract stems # %% import os procid = int(os.environ.get("SLURM_PROCID", 0)) localid = int(os.environ.get("SLURM_LOCALID", 0)) world_size = int(os.environ.get("SLURM_JOB_NUM_NODES", 1)) * int( os.environ.get("SLURM_NTASKS_PER_NODE", 1) ) assert world_size > 0, "WORLD_SIZE is 0" print(f"PROCID: {procid}, LOCALID: {localid}, WORLD_SIZE: {world_size}") # %% from suno_utils.diffusion import generation as diffusion_gen from suno_utils.tasks.upsample_engine import UpsampleEngine, Request from suno_utils.tasks.dac_vae_fixed_25hz import decode_stream_to_full_audio, encode, decode import torch import numpy as np from tqdm import tqdm from suno_utils.audio import Audio torch.cuda.set_device(localid) # %% out_stem_dir = "/app2/suno/data/victor/midi_stems" os.makedirs(out_stem_dir, exist_ok=True) pair_audio_paths = [] pairs_to_process = { "scoretube": scoretube_pairs, "trombone_champ": trombone_champ_pairs, "hooktheory": hooktheory_pairs, "mootube": mootube_pairs, "mmd": mmd_pairs, } # ignore maestro for dataset, pairs in pairs_to_process.items(): for pair in pairs: name = pair.audio_file.split("/")[-1].split(".")[0] pair_audio_paths.append((pair.audio_file, f"{out_stem_dir}/{dataset}/{name}")) print(len(pair_audio_paths)) pair_audio_paths = pair_audio_paths[procid::world_size] print(len(pair_audio_paths)) # %% import numpy as np diffusion_gen.preload_models( dit_model_filepath="/app/suno/checkpoints/2025-05-26_22-49-34_s8646/last_ckpt_infer.pt", # 12 stems codec_filepath="s3://suno-data/minz/models/dac_vae_tuned_25hz.pth", compile=True, ) engine = UpsampleEngine(min_chunk_size=25 * 30) # %% def gen_stem( audio: Audio, stem_type_cfg_scale=1.0, tags="extract [split_karaoke]", steps=4, seed=3, codec_scale_factor=0.4, scale_ctx_vector=True, noise_ctx_level=0.0, infill_prefix_latents=None, infill_suffix_latents=None, ): vae = encode(audio) gen_cfg = diffusion_gen.DiffusionGenerationConfig( lyrics=tags, steps=steps, seed=seed, codec_scale_factor=codec_scale_factor, scale_ctx_vector=scale_ctx_vector, noise_ctx_level=noise_ctx_level, text_cfg_coef=stem_type_cfg_scale, infill_prefix_latents=infill_prefix_latents, infill_suffix_latents=infill_suffix_latents, drop_semantic_tokens=True, ) # print(gen_cfg) request = Request( id="dummy", generation_config=gen_cfg, tokens=np.zeros((vae.shape[0], 1)), input_tokens_finished=True, stem_ctx_latents=vae, ) result = engine.run_request(request, tqdm_enabled=False) vae_latents = torch.concat(result.vae_latents) # print(f"vae_latents: {vae_latents.shape}") audios = [] for i in tqdm(range(vae_latents.shape[1]), desc="Decoding stems", disable=True): # audios.append(decode(vae_latents[:, i])) audios.append(decode_stream_to_full_audio(vae_latents[:, i], n_stride_tokens=25 * 10)) return audios # %% def write_stem(args): out_root, category, stem = args if stem.loudness < -45: return None out_path = f"{out_root}/{category}.opus" os.makedirs(os.path.dirname(out_path), exist_ok=True) # print(f"Writing {out_path}") stem.write_opus(out_path) return category import threading categories = [ "Vocals", "Backing_Vocals", "Drums", "Bass", "Guitar", "Keyboard", "Percussion", "Strings", "Synth", "FX", "Brass", "Woodwinds", ] def process_id(pair): input_path, output_path = pair audio = Audio.from_file(input_path, n_channels=2) stems = gen_stem(audio, steps=8) # Prepare arguments for parallel processing write_args = [(output_path, category, stem) for category, stem in zip(categories, stems)] # Use threading to write stems in parallel results = [None] * len(write_args) threads = [] def write_stem_thread(i, args): results[i] = write_stem(args) for i, args in enumerate(write_args): thread = threading.Thread(target=write_stem_thread, args=(i, args)) threads.append(thread) thread.start() for thread in threads: thread.join() # Filter out None results found_categories = [result for result in results if result is not None] return found_categories # process_id(random.choice(pair_audio_paths)) # %% for pair in tqdm(pair_audio_paths, mininterval=600): try: process_id(pair) except Exception as e: print(f"Error processing {pair}: {e}") continue # %%