# %% [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

# %%
