import json
import os
import random
from collections import defaultdict
import traceback
import copy

import numpy as np
import orjson
import re
import struct
import webrtcvad
from scipy.ndimage.morphology import binary_dilation
from scipy.signal import butter, lfilter
import soxr
import torch
from suno_utils.audio import Audio
from torch.utils.data import IterableDataset
from tqdm import tqdm

from data_types import SampleData, SamplingParams, AudioType
from modules.gpt import GPTConfig, GPTTrainConfig
from oracle_dataset import get_sample_oracle_file_segment
from sample_extraction import create_sample_from_stems, create_sample_from_full_song, apply_audio_effects
from text_utils import load_tokenizer

MOCK_SUBSAMPLE_RATE = 100


class JSONLMemmap:
    def __init__(self, path, verbose=False):
        self.path = path
        self._index = self.get_line_index_map(verbose=verbose)

    def get_line_index_map(self, verbose=False):
        line_positions = [0]
        with open(self.path, "rb") as f:
            for line in tqdm(f, disable=not verbose, desc="Making memmap index"):
                line_positions.append(len(line) + line_positions[-1])
        return np.array(line_positions[:-1])

    def get_line_from_index(self, index: int):
        with open(self.path, "rb") as f:
            f.seek(self._index[index])
            return json.loads(f.readline())

    def __getitem__(self, index: int):
        return self.get_line_from_index(index)

    def __len__(self):
        return len(self._index)

    def __iter__(self):
        for i in range(len(self)):
            yield self[i]


def _resample_to_mert(arr):
    assert arr.ndim == 2
    assert arr.shape[0] == 2
    out_arr = soxr.resample(arr.T, 48_000, 24_000).mean(axis=1).astype(np.float32)
    return out_arr


def _resample_to_musicfm(arr):
    assert arr.ndim == 2
    assert arr.shape[0] == 2
    out_arr = soxr.resample(arr.T, 48_000, 16_000).astype(np.float32).T  # (2, T)
    return out_arr


def trim_long_silences(wav_stereo, sampling_rate):
    """
    Ensures that segments without voice in the waveform remain no longer than a
    threshold determined by the VAD parameters in params.py.

    :param wav: the raw waveform as a numpy array of floats
    :return: the same waveform with silences trimmed away (length <= original wav length)
    """
    vad_moving_average_width = 8
    vad_max_silence_length = 6
    vad_window_length = 30
    int16_max = (2**15) - 1

    # stereo to mono
    wav_mono = wav_stereo.mean(axis=0)

    # Compute the voice detection window size
    samples_per_window = (vad_window_length * sampling_rate) // 1000

    # Trim the end of the audio to have a multiple of the window size
    wav_mono = wav_mono[: len(wav_mono) - (len(wav_mono) % samples_per_window)]
    wav_stereo = wav_stereo[:, : len(wav_mono) - (len(wav_mono) % samples_per_window)]

    # Convert the float waveform to 16-bit mono PCM
    pcm_wave = struct.pack("%dh" % len(wav_mono), *(np.round(wav_mono * int16_max)).astype(np.int16))

    # Perform voice activation detection
    voice_flags = []
    vad = webrtcvad.Vad(mode=3)
    for window_start in range(0, len(wav_mono), samples_per_window):
        window_end = window_start + samples_per_window
        voice_flags.append(
            vad.is_speech(pcm_wave[window_start * 2 : window_end * 2], sample_rate=sampling_rate)
        )
    voice_flags = np.array(voice_flags)

    # Smooth the voice detection with a moving average
    def moving_average(array, width):
        array_padded = np.concatenate((np.zeros((width - 1) // 2), array, np.zeros(width // 2)))
        ret = np.cumsum(array_padded, dtype=float)
        ret[width:] = ret[width:] - ret[:-width]
        return ret[width - 1 :] / width

    audio_mask = moving_average(voice_flags, vad_moving_average_width)
    audio_mask = np.round(audio_mask).astype(bool)

    # Dilate the voiced regions
    audio_mask = binary_dilation(audio_mask, np.ones(vad_max_silence_length + 1))
    audio_mask = np.repeat(audio_mask, samples_per_window)

    return wav_stereo[:, audio_mask == True]


def bandpass_filter(data, sr, lowcut=300.0, highcut=3400.0, order=5):
    nyq = 0.5 * sr
    low = lowcut / nyq
    high = highcut / nyq
    b, a = butter(order, [low, high], btype="band")

    # Handle stereo input
    if len(data.shape) == 2:  # stereo
        return np.array([lfilter(b, a, channel) for channel in data])
    else:  # mono
        return lfilter(b, a, data)


def add_room_noise(wave, sr, snr_db=None):
    """Overlay white noise at a given SNR (dB) per *mix* (not per channel)."""
    if snr_db is None:
        snr_db = random.uniform(10.0, 30.0)
    noise = np.random.normal(0.0, 1.0, size=wave.shape)
    noise *= _snr_scale(wave, noise, snr_db)
    return np.clip(wave + noise, -1.0, 1.0)


def _soft_clip(x, alpha=2.0):
    return np.tanh(alpha * x) / np.tanh(alpha)


def _snr_scale(clean, noise, snr_db):
    power_clean = np.mean(clean**2)
    power_noise = np.mean(noise**2) + 1e-12
    target_noise_power = power_clean / (10.0 ** (snr_db / 10.0))
    return np.sqrt(target_noise_power / power_noise)


def soft_clipping(wave, alpha=None):
    if alpha is None:
        alpha = random.uniform(1.5, 3.0)
    return _soft_clip(wave, alpha)


def heavy_compression(wave, threshold_db=-18.0, ratio=None):
    if ratio is None:
        ratio = random.uniform(4.0, 10.0)
    eps = 1e-12
    mag = np.abs(wave) + eps
    db = 20.0 * np.log10(mag)
    over = db - threshold_db
    gain_db = np.where(over > 0.0, -over * (1.0 - 1.0 / ratio), 0.0)
    gain_lin = 10.0 ** (gain_db / 20.0)
    return wave * gain_lin


def loudness_wobble(wave, sr, depth_db=3.0, lfo_hz=None):
    if lfo_hz is None:
        lfo_hz = random.uniform(0.5, 3.0)
    t = np.arange(wave.shape[0]) / sr
    lfo = (1.0 + (10.0 ** (depth_db / 20.0) - 1.0) * np.sin(2 * np.pi * lfo_hz * t)) / (
        10.0 ** (depth_db / 20.0)
    )
    return np.clip(wave * lfo[:, None] if wave.ndim == 2 else wave * lfo, -1.0, 1.0)


def random_amp(audio, min_amp=0.3, max_amp=1.0):
    return audio * np.random.uniform(min_amp, max_amp)


def interleave_pitch_shift(y_stereo, sr, semitone_range=(-0.6, 0.6)):
    duration = int(y_stereo.shape[1] / sr)
    hopsize = random.sample([0.5, 1], 1)[0]
    indices = [int(sr * hopsize * i) for i in range(int((sr * duration) // (sr * hopsize) + 1))]
    y_stereo = torch.tensor(y_stereo.astype(np.float32))
    for i in range(len(indices) - 1):
        start = indices[i]
        end = min(indices[i + 1], y_stereo.shape[1])
        part = y_stereo[:, start : end + 1]
        if random.random() < 0.1:
            n_semitone = random.uniform(*semitone_range)
            shifted = apply_audio_effects(part, sr, pitch_semitones=n_semitone, rate_factor=1.0)
            shifted = shifted[:, : end - start]
            if (end - start) != shifted.shape[1]:
                shifted = torch.nn.functional.pad(shifted, (0, int(end - start - shifted.shape[1])))
            y_stereo[:, start:end] = shifted
    return y_stereo.numpy()


class AudioLoaderDataset(IterableDataset):
    """Gets SampleData objects from the oracle dataset, which contain raw wav data"""

    def __init__(
        self,
        model_cfg: GPTConfig,
        train_cfg: GPTTrainConfig,
        metas_path: str,
        batch_size_tokens: int,
        tokenizer_fp: str,
        device: str,
        split="train",
        info_path: str | None = None,
        dataset_idx=None,
        sampling_params: SamplingParams | None = None,
        stem_active_sections_weight: float = 4,
    ):
        self.model_cfg = model_cfg
        self.train_cfg = train_cfg
        self.metas_path = metas_path
        self.info_path = info_path
        self.batch_size_tokens = batch_size_tokens
        self.tokenizer_fp = tokenizer_fp
        self.device = device
        self.dataset_idx = dataset_idx
        self.split = split
        self.sampling_params = sampling_params
        self.stem_active_sections_weight = stem_active_sections_weight

        assert model_cfg.semantic_type.startswith("mert") or model_cfg.semantic_type.startswith(
            "musicfm"
        )
        if model_cfg.semantic_type.startswith("mert"):
            self.resample_fn = _resample_to_mert
            self.semantic_sample_rate = 24_000
            self.semantic_n_channels = 1
        elif model_cfg.semantic_type.startswith("musicfm"):
            self.resample_fn = _resample_to_musicfm
            self.semantic_sample_rate = 16_000
            self.semantic_n_channels = 2

        if split == "val":
            assert self.sampling_params.inference

        # load metas as a memmap to save memory
        self.metas = JSONLMemmap(self.metas_path, verbose=False)

        if self.info_path is not None:
            print(f"Filtering metas with {self.info_path}")
            with open(self.info_path, "r") as f:
                self.info = json.load(f)
            self.valid_ids = {}
            for key, val in self.info.items():
                if isinstance(val, list):
                    self.valid_ids[key] = set(val)
                print(f"{key}: {len(val):,}")

        else:
            self.valid_ids = None

        # make essential dicts
        self.weights = []
        self.id_to_index = {}
        self.artist_to_ids = defaultdict(list)
        self.playlist_to_ids = defaultdict(list)
        self.artist_to_vox_ids = defaultdict(list)

        meta_count = 0
        with open(self.metas_path, "r") as f:
            for i, line in tqdm(
                enumerate(f),
                disable=True,
                total=len(self.metas),
                desc="Loading metas",
                mininterval=10,
            ):
                meta = orjson.loads(line)
                if self.valid_ids is not None and meta["id"] not in self.valid_ids["audio_ids"]:
                    weight = 0  # set weight to 0 if not in valid ids
                else:
                    meta_count += 1  # count valid ids
                    weight = meta.get("weight", 0)
                if "stem_active_sections" in meta:
                    weight *= self.stem_active_sections_weight
                self.weights.append(weight)
                self.id_to_index[meta["id"]] = i
                for artist_id in meta.get("artist_ids", []):
                    self.artist_to_ids[artist_id].append(meta["id"])
                for playlist_id in meta.get("playlist_ids", []):
                    self.playlist_to_ids[playlist_id].append(meta["id"])
                if "underpaint_id" in meta and meta.get("artist_ids"):
                    self.artist_to_vox_ids[meta["artist_ids"][0]].append(meta["id"])

        if self.valid_ids is not None:
            prev_len = len(self.metas)
            filtered_pct = meta_count / prev_len
            print(f"Filtered to {meta_count:,} metas from {prev_len:,} ({filtered_pct:.2%})")
        self.tokenizer = load_tokenizer(self.tokenizer_fp)
        self.random_cache = []

    def __iter__(self):
        return self

    def _load_audio_for_semantic(
        self,
        local_filepath,
        s3_filepath=None,
        expected_duration_s=180,
        start_s=0,
        max_duration_s=60 * 30,
        mock=False,
        load_for_semantic=True,
    ):
        model_max_duration_s = self.model_cfg.block_size / self.model_cfg.semantic_rate_hz
        max_duration_s = min(max_duration_s, model_max_duration_s)
        try:
            if mock:
                use_duration_s = min(max_duration_s, expected_duration_s)
                # subsample to avoid ipc limit
                audio_arr = Audio.from_silence(
                    use_duration_s, self.semantic_sample_rate, n_channels=self.semantic_n_channels
                ).array_float[..., ::MOCK_SUBSAMPLE_RATE]
            else:
                audio = get_sample_oracle_file_segment(
                    local_filepath=local_filepath,
                    s3_filepath=s3_filepath,
                    start_s=start_s,
                    max_duration_s=max_duration_s,
                )
                assert audio.sample_rate == 48_000
                assert audio.n_channels == 2
                if load_for_semantic:
                    audio_arr = self.resample_fn(audio.array_float)
                else:
                    return audio
        except Exception as e:
            print(traceback.format_exc())
            print(
                f"Host {os.environ['HOSTNAME']} Error loading audio for {s3_filepath} "
                f"(local_fp: {local_filepath}, start_s: {start_s}, max_duration_s: {max_duration_s}): {e}"
            )
        return audio_arr, audio.array_float

    def _load_cover_audio(self, main_meta, mock=False, load_for_semantic=True):
        covers = main_meta.get("cover_ids", [])
        if self.valid_ids is not None:
            covers = [id for id in covers if id in self.valid_ids]
        if len(covers) > 0:
            cover_id = random.choice(covers)
            cover_meta = self.metas[self.id_to_index[cover_id]]
            data_row_cover, _ = self._load_audio_for_semantic(
                cover_meta["local_filepath"],
                cover_meta["s3_filepath"],
                cover_meta["duration_s"],
                mock=mock,
                load_for_semantic=load_for_semantic,
            )
        else:
            data_row_cover = None
        return data_row_cover

    def _load_artist_audio(self, main_meta, mock=False, load_for_semantic=True):
        artists = main_meta.get("artist_ids", [])
        if len(artists) == 0:
            return None
        artist_id = random.choice(artists)
        artist_song_ids = self.artist_to_ids[artist_id]
        # remove main_meta id from artist_song_ids
        artist_song_ids = [id for id in artist_song_ids if id != main_meta["id"]]
        if len(artist_song_ids) == 0:
            return None

        # sample multiple segments
        n_segments = random.randint(1, 10)
        min_segment_len = 5
        max_segment_len = 120
        artist_audio_arrs = []

        for _ in range(n_segments):
            track_id = random.choice(artist_song_ids)
            track_meta = self.metas[self.id_to_index[track_id]]
            dur_s = random.random() * max_segment_len
            dur_s = min(max(min_segment_len, dur_s), track_meta["duration_s"])
            start_s = random.random() * (track_meta["duration_s"] - dur_s)
            assert start_s >= 0
            data_row_artist, _ = self._load_audio_for_semantic(
                track_meta["local_filepath"],
                track_meta["s3_filepath"],
                track_meta["duration_s"],
                start_s=start_s,
                max_duration_s=dur_s,
                mock=mock,
                load_for_semantic=load_for_semantic,
            )
            artist_audio_arrs.append(data_row_artist)

        return artist_audio_arrs

    def _load_playlist_audio(self, main_meta, mock=False, load_for_semantic=True):
        playlists = main_meta.get("playlist_ids", [])
        if len(playlists) == 0:
            return None
        playlist_id = random.choice(playlists)
        playlist_song_ids = self.playlist_to_ids[playlist_id]
        # remove main_meta id from playlist_song_ids
        playlist_song_ids = [id for id in playlist_song_ids if id != main_meta["id"]]
        if len(playlist_song_ids) == 0:
            return None

        playlist_audio_arrs = []
        n_segments = random.randint(1, 10)
        min_segment_len = 5
        max_segment_len = 120
        for _ in range(n_segments):
            track_id = random.choice(playlist_song_ids)
            track_meta = self.metas[self.id_to_index[track_id]]
            dur_s = random.random() * max_segment_len
            dur_s = min(max(min_segment_len, dur_s), track_meta["duration_s"])
            start_s = random.random() * (track_meta["duration_s"] - dur_s)
            assert start_s >= 0
            data_row_playlist, _ = self._load_audio_for_semantic(
                track_meta["local_filepath"],
                track_meta["s3_filepath"],
                track_meta["duration_s"],
                start_s=start_s,
                max_duration_s=dur_s,
                mock=mock,
                load_for_semantic=load_for_semantic,
            )
            playlist_audio_arrs.append(data_row_playlist)
        return playlist_audio_arrs

    def _load_overpaint_audio(self, main_meta, mock=False, load_for_semantic=True):
        overpaint_id = main_meta.get("overpaint_id", None)
        if overpaint_id is None:
            return None
        overpaint_meta = self.metas[self.id_to_index[overpaint_id]]
        data_row_overpaint, _ = self._load_audio_for_semantic(
            overpaint_meta["local_filepath"],
            overpaint_meta["s3_filepath"],
            overpaint_meta["duration_s"],
            mock=mock,
            load_for_semantic=load_for_semantic,
        )
        return data_row_overpaint

    def _load_underpaint_audio(self, main_meta, mock=False, load_for_semantic=True):
        underpaint_id = main_meta.get("underpaint_id", None)
        if underpaint_id is None:
            return None
        underpaint_meta = self.metas[self.id_to_index[underpaint_id]]
        data_row_underpaint, _ = self._load_audio_for_semantic(
            underpaint_meta["local_filepath"],
            underpaint_meta["s3_filepath"],
            underpaint_meta["duration_s"],
            mock=mock,
            load_for_semantic=load_for_semantic,
        )
        return data_row_underpaint

    def _load_sample_source_audio(self, main_meta, mock=False, load_for_semantic=True):
        sample_source_id = main_meta.get("sample_source_id", None)
        if sample_source_id is None:
            return None
        sample_source_meta = self.metas[self.id_to_index[sample_source_id]]
        data_row_sample_source, _ = self._load_audio_for_semantic(
            sample_source_meta["local_filepath"],
            sample_source_meta["s3_filepath"],
            sample_source_meta["duration_s"],
            mock=mock,
            load_for_semantic=load_for_semantic,
        )
        return data_row_sample_source

    def _load_remix_source_audio(self, main_meta, mock=False, load_for_semantic=True):
        remix_source_id = main_meta.get("remix_source_id", None)
        if remix_source_id is None:
            return None
        remix_source_meta = self.metas[self.id_to_index[remix_source_id]]
        data_row_remix_source, _ = self._load_audio_for_semantic(
            remix_source_meta["local_filepath"],
            remix_source_meta["s3_filepath"],
            remix_source_meta["duration_s"],
            mock=mock,
            load_for_semantic=load_for_semantic,
        )
        return data_row_remix_source

    def _load_mashup_audio(self, main_meta, mock=False, load_for_semantic=True):
        mashup_song_ids = main_meta.get("mashup_source_ids", [])
        if len(mashup_song_ids) == 0:
            return None

        # select up to 4 unique mashup sources
        num_to_select = random.randint(1, min(4, len(mashup_song_ids)))
        mashup_song_ids = random.sample(mashup_song_ids, num_to_select)

        mashup_audio_arrs = []
        min_segment_len = 60
        max_segment_len = 180

        for track_id in mashup_song_ids:
            track_meta = self.metas[self.id_to_index[track_id]]

            dur_s = random.random() * max_segment_len
            dur_s = min(max(min_segment_len, dur_s), track_meta["duration_s"])
            start_s = random.random() * (track_meta["duration_s"] - dur_s)
            assert start_s >= 0

            data_row_mashup, _ = self._load_audio_for_semantic(
                track_meta["local_filepath"],
                track_meta["s3_filepath"],
                track_meta["duration_s"],
                start_s=start_s,
                max_duration_s=dur_s,
                mock=mock,
                load_for_semantic=load_for_semantic,
            )
            mashup_audio_arrs.append(data_row_mashup)

        if len(mashup_audio_arrs) == 0:
            return None
        return mashup_audio_arrs

    def _crop_and_patch(
        self, ref_audio, sample_rate, n_channels, cond_length_s=30.0, patch_length_s=0.5
    ):
        patch_length = int(sample_rate * patch_length_s)
        num_patch = int(sample_rate * cond_length_s / patch_length)
        dtype = ref_audio.dtype
        # Ensure ref_audio is 2D: (n_channels, T) or (T,) → (T, n_channels)
        if ref_audio.ndim == 1:
            ref_audio = ref_audio[:, np.newaxis]  # (T,) → (T, 1)
        else:
            ref_audio = ref_audio.T  # (n_channels, T) → (T, n_channels)

        out_patch = np.zeros((int(sample_rate * cond_length_s), n_channels), dtype=dtype)

        # crop input duration to be 5s to 30s
        if (self.split == "train") and (len(ref_audio) > 5 * sample_rate):
            patch_duration = min(len(ref_audio), random.randint(5 * sample_rate, len(ref_audio)))
            start_ix = random.randint(0, len(ref_audio) - patch_duration)
            ref_audio = ref_audio[start_ix : start_ix + patch_duration]

        if len(ref_audio) < patch_length:
            if ref_audio.ndim == 1:
                ref_audio = np.pad(ref_audio, (0, patch_length - len(ref_audio)))
            elif ref_audio.ndim == 2:
                ref_audio = np.pad(ref_audio, ((0, patch_length - len(ref_audio)), (0, 0)))
        if self.split == "train":
            for i in range(num_patch):
                start_ix = random.randint(0, max(0, len(ref_audio) - patch_length))
                out_patch[i * patch_length : (i + 1) * patch_length] = ref_audio[
                    start_ix : start_ix + patch_length
                ]
            return out_patch.T
        else:
            if len(ref_audio) < int(sample_rate * cond_length_s):
                ref_audio = np.pad(ref_audio, (0, int(sample_rate * cond_length_s) - len(ref_audio)))
            return ref_audio[max(0, len(ref_audio) - int(sample_rate * cond_length_s)) :].T

    def _load_vox_audio(self, main_meta, mock=False, load_for_semantic=True):
        mix_id = main_meta["id"]

        # when it has artist_ids
        if main_meta.get("artist_ids", None) is not None:
            artist_id = main_meta["artist_ids"][0]
            vox_ids = self.artist_to_vox_ids[artist_id]
            vox_ids = [id for id in vox_ids if id != mix_id]
            if len(vox_ids) == 0:
                vox_ids = self.artist_to_vox_ids[artist_id]

            if len(vox_ids) == 0:
                return None

            if self.split == "train":
                vox_id = random.choice(vox_ids)
            else:
                vox_id = vox_ids[0]
            vox_meta = self.metas[self.id_to_index[vox_id]]

            # load vox
            data_row_vox = self._load_audio_for_semantic(
                vox_meta["local_filepath"].replace(".opus", "_vocals.opus"),
                vox_meta["s3_filepath"].replace(".opus", "_vocals.opus"),
                vox_meta["duration_s"],
                mock=mock,
                load_for_semantic=False,
            )
        # when it is cover data without artist_ids
        elif (
            main_meta.get("stems", None) is not None
            and main_meta.get("stems").get("Vocals", None) is not None
        ):
            vox_meta = main_meta
            data_row_vox = self._load_audio_for_semantic(
                vox_meta.get("stems").get("Vocals"),
                vox_meta["s3_filepath"].replace(
                    ".opus", "_vocals.opus"
                ),  # <- this is a placeholder (fake path)
                vox_meta["duration_s"],
                mock=mock,
                load_for_semantic=False,
            )
        else:
            return None

        # remove silence
        trimmed_data_row_vox = trim_long_silences(data_row_vox.array_float, 48000)
        if trimmed_data_row_vox.shape[-1] > 0:
            data_row_vox = trimmed_data_row_vox
        else:
            data_row_vox = data_row_vox.array_float

        # add noise
        prob = 0.5
        if self.split == "train":
            if random.random() < prob:
                data_row_vox = bandpass_filter(data_row_vox, 48000)
            if random.random() < prob:
                data_row_vox = add_room_noise(data_row_vox, 48000)
            if random.random() < prob:
                data_row_vox = soft_clipping(data_row_vox)
            if random.random() < prob:
                data_row_vox = heavy_compression(data_row_vox)
            if random.random() < prob:
                data_row_vox = loudness_wobble(data_row_vox, 48000)
            if random.random() < prob:
                data_row_vox = random_amp(data_row_vox)
            if random.random() < prob:
                data_row_vox = interleave_pitch_shift(data_row_vox, 48000)
            data_row_vox = data_row_vox.astype(np.float32)

        # prevent signal overflow
        if np.max(np.abs(data_row_vox)) > 1.0:
            data_row_vox = data_row_vox / np.max(np.abs(data_row_vox))

        # resample
        data_row_vox = self.resample_fn(data_row_vox)
        data_row_vox = self._crop_and_patch(
            data_row_vox,
            sample_rate=self.semantic_sample_rate,
            n_channels=self.semantic_n_channels,
            patch_length_s=3.0,
        )

        return data_row_vox

    def _load_and_mix_stems(
        self, stem_dict, stem_keys, mock=False, load_for_semantic=True, start_s=0, end_s=None
    ):
        """Load and mix multiple stems together."""
        stem_audios = [
            self._load_audio_for_semantic(
                stem_dict[k],
                mock=mock,
                load_for_semantic=False,
            ).get_segment(start_s, end_s)
            for k in stem_keys
        ]
        try:
            mixed_audio = Audio.sum(stem_audios)
        except Exception as e:
            print(f"Error mixing stems: {e}")
            print(f"Stem keys: {stem_keys}")
            print(f"Stem dict: {stem_dict}")
            print(f"Start s: {start_s}")
            print(f"End s: {end_s}")
            raise e
        if load_for_semantic:
            return self.resample_fn(mixed_audio.array_float)
        else:
            return mixed_audio.array_float

    def _load_stem_audio(self, main_meta, task="add", mock=False, load_for_semantic=True):
        """Load and process stem audio data for training.
        Randomly splits available stems into input and output subsets, then loads
        and mixes the stems in each subset. Returns the stem type description
        and the mixed audio arrays.
        """
        stem_dict = copy.deepcopy(main_meta.get("stems", {}))
        # filter out keys that contain " and " or "&". multi instrument isnt good.
        stem_dict = {
            k: v for k, v in stem_dict.items() if " and " not in k.lower() and "&" not in k.lower()
        }
        if len(stem_dict) < 2:
            return None, None, None, None

        keys = list(stem_dict.keys())
        # Clean up keys by removing trailing single characters/numbers
        cleaned_keys = []
        for key in keys:
            # Remove trailing single character or number (e.g., "synths 3" -> "synths")
            cleaned_key = re.sub(r"\s+[a-zA-Z0-9]$", "", key).strip()
            if cleaned_key:  # Only add non-empty keys
                cleaned_keys.append(cleaned_key)

            stem_dict[cleaned_key] = stem_dict.pop(key)

        # Use cleaned keys and remove duplicates while preserving order
        seen = set()
        unique_keys = []
        for key in cleaned_keys:
            if key not in seen:
                seen.add(key)
                unique_keys.append(key)

        keys = unique_keys
        random.shuffle(keys)

        # Check if we have enough keys after deduplication
        if len(keys) < 2:
            return None, None, None, None

        # Split keys into input and output subsets
        split_idx = random.randint(1, len(keys) - 1)
        input_subset = keys[:split_idx]
        # Randomly choose subsets of each
        input_subset = random.sample(input_subset, random.randint(1, len(input_subset)))
        if task == "add":
            output_subset = keys[split_idx : split_idx + 1]  # only add one stem
            # Randomly set instrument name to "auto" for some percentage of the data
            if random.random() < 0.0:
                stem_type = f"{task} auto"
            else:
                stem_type = f"{task} {', '.join(output_subset)}"
        elif task == "extract":
            output_subset = input_subset
            output_subset = random.sample(output_subset, 1)
            stem_type = f"{task} {', '.join(output_subset)}"
        elif task == "remove":
            output_subset = input_subset
            output_subset = random.sample(
                output_subset, random.randint(1, max(1, len(output_subset) - 1))
            )
            # print(f"remove, {len(input_subset)} -> {len(output_subset)}")
            stems_removed = [k for k in input_subset if k not in output_subset]
            stem_type = f"{task} {', '.join(stems_removed)}"
        else:
            raise ValueError(f"Invalid task: {task}")

        # 50% chance to crop to an active section
        start_s = 0
        end_s = None
        if (
            random.random() < 1
            and "stem_active_sections" in main_meta
            and len(output_subset) == 1
            and len(main_meta["stem_active_sections"].get(output_subset[0], [])) > 0
        ):
            active_sections = main_meta["stem_active_sections"][output_subset[0]]

            # Try to find a random span between section boundaries where the stem is active at least 30% of the time
            max_attempts = 100
            for _ in range(max_attempts):
                # choose a span of sections
                section_start, section_end = sorted(random.choices(active_sections, k=2))
                span_start = section_start[0]
                span_end = section_end[1]
                span_duration = span_end - span_start

                # Calculate how much of this span overlaps with active sections
                active_duration = 0
                for section_start, section_end in active_sections:
                    overlap_start = max(span_start, section_start)
                    overlap_end = min(span_end, section_end)
                    if overlap_start < overlap_end:
                        active_duration += overlap_end - overlap_start

                # Check if at least 30% of the span is active
                activity_ratio = active_duration / span_duration
                # print(
                #     f"Found span for {main_meta['id']}, using {span_start} to {span_end}, activity ratio: {activity_ratio}"
                # )
                if activity_ratio >= 0.3:
                    start_s = span_start
                    end_s = span_end
                    break
            else:
                # Fallback: use a random active section if no good span found
                print(f"No good span found for {main_meta['id']}, using random active section")
                start_s, end_s = random.choice(active_sections)
            if end_s - start_s < 2:  # skip if the active section is too short
                start_s = 0
                end_s = None

        # Load and mix stems
        output_stem_audio = self._load_and_mix_stems(
            stem_dict,
            output_subset,
            mock=mock,
            load_for_semantic=load_for_semantic,
            start_s=start_s,
            end_s=end_s,
        )
        input_subset_audio = self._load_and_mix_stems(
            stem_dict,
            input_subset,
            mock=mock,
            load_for_semantic=load_for_semantic,
            start_s=start_s,
            end_s=end_s,
        )
        # pad to same length
        if self.model_cfg.semantic_type.startswith("mert"):
            input_subset_audio = np.pad(
                input_subset_audio, (0, max(0, len(output_stem_audio) - len(input_subset_audio)))
            )
            output_stem_audio = np.pad(
                output_stem_audio, (0, max(0, len(input_subset_audio) - len(output_stem_audio)))
            )
            assert len(input_subset_audio) == len(output_stem_audio)
        elif self.model_cfg.semantic_type.startswith("musicfm"):
            input_subset_audio = np.pad(
                input_subset_audio,
                ((0, 0), (0, max(0, output_stem_audio.shape[1] - input_subset_audio.shape[1]))),
            )
            output_stem_audio = np.pad(
                output_stem_audio,
                ((0, 0), (0, max(0, input_subset_audio.shape[1] - output_stem_audio.shape[1]))),
            )
            assert input_subset_audio.shape[1] == output_stem_audio.shape[1]

        # Create full mix (input + output stem) for output target
        full_mix_audio = input_subset_audio + output_stem_audio * 3

        return stem_type, input_subset_audio, output_stem_audio, full_mix_audio

    def __next__(self):
        return self._next()

    def _sample_meta(self):
        """
        Randomly sample a meta from the metas list.
        Weighted sampling is slow for large weighted datasets so we sample 10k at a time.
        """
        if len(self.random_cache) == 0:
            choices = random.choices(range(len(self.metas)), weights=self.weights, k=10000)
            self.random_cache.extend(choices)
        idx = self.random_cache.pop()
        return self.metas[idx]

    def _next(self):
        # randomly sample a meta
        main_meta = self._sample_meta()
        main_meta["s3_filepath"] = main_meta.get("s3_filepath", None)  # fine if missing
        audio_type = main_meta.get("audio_type", AudioType.MUSIC)  # fine if missing
        use_raw_audio = self.sampling_params.output_distribution == "vae"

        # load wav
        data_row, data_row_48kHz = self._load_audio_for_semantic(
            main_meta["local_filepath"],
            main_meta.get("s3_filepath", None),
            main_meta["duration_s"],
            mock=self.sampling_params.mock_data,
        )

        # TODO: add cover, artist, playlist, overpaint, underpaint
        if (
            not self.sampling_params.inference
            and self.sampling_params.allow_cover
            and random.random() < self.sampling_params.prob_cover
        ):
            data_row_cover = self._load_cover_audio(
                main_meta, mock=self.sampling_params.mock_data, load_for_semantic=not use_raw_audio
            )
        else:
            data_row_cover = None

        p_use_artist = (
            0.5 if data_row_cover is not None else self.sampling_params.prob_artist
        )  # increase chance for cover cause voice beautifier
        if (
            not self.sampling_params.inference
            and self.sampling_params.allow_artist
            and random.random() < p_use_artist
        ):
            data_row_artist = self._load_artist_audio(
                main_meta, mock=self.sampling_params.mock_data, load_for_semantic=not use_raw_audio
            )
        else:
            data_row_artist = None

        if (
            not self.sampling_params.inference
            and self.sampling_params.use_ditto
            and random.random() < self.sampling_params.prob_use_ditto
        ):
            data_row_ditto = _resample_to_mert(data_row_48kHz)
        else:
            data_row_ditto = None

        if (
            not self.sampling_params.inference
            and self.sampling_params.allow_playlist
            and random.random() < self.sampling_params.prob_playlist
        ):
            data_row_playlist = self._load_playlist_audio(
                main_meta, mock=self.sampling_params.mock_data, load_for_semantic=not use_raw_audio
            )
        else:
            data_row_playlist = None

        if (
            not self.sampling_params.inference
            and self.sampling_params.allow_overpaint
            and random.random() < self.sampling_params.prob_overpaint
        ):
            data_row_overpaint = self._load_overpaint_audio(
                main_meta, mock=self.sampling_params.mock_data, load_for_semantic=not use_raw_audio
            )
            # sometimes add noise to hide imperfect source sep at input
            if data_row_overpaint is not None and random.random() < 0.5:
                noise = np.random.normal(0, random.random() * 0.25, data_row_overpaint.shape)
                data_row_overpaint = np.clip(
                    data_row_overpaint + noise.astype(data_row_overpaint.dtype), -1.1, 1.1
                )
        else:
            data_row_overpaint = None

        if (
            not self.sampling_params.inference
            and self.sampling_params.allow_underpaint
            and random.random() < self.sampling_params.prob_underpaint
        ):
            data_row_underpaint = self._load_underpaint_audio(
                main_meta, mock=self.sampling_params.mock_data, load_for_semantic=not use_raw_audio
            )
            # sometimes add noise to hide imperfect source sep at input
            if data_row_underpaint is not None and random.random() < 0.5:
                noise = np.random.normal(0, random.random() * 0.25, data_row_underpaint.shape)
                data_row_underpaint = np.clip(
                    data_row_underpaint + noise.astype(data_row_underpaint.dtype), -1.1, 1.1
                )
        else:
            data_row_underpaint = None

        if (
            not self.sampling_params.inference
            and self.sampling_params.allow_vox
            and random.random() < self.sampling_params.prob_vox
        ):
            data_row_vox = self._load_vox_audio(
                main_meta, mock=self.sampling_params.mock_data, load_for_semantic=not use_raw_audio
            )
        else:
            data_row_vox = None

        if (
            not self.sampling_params.inference
            and self.sampling_params.allow_remix
            and random.random() < self.sampling_params.prob_remix
        ):
            data_row_remix = self._load_remix_source_audio(
                main_meta, mock=self.sampling_params.mock_data, load_for_semantic=not use_raw_audio
            )
        else:
            data_row_remix = None

        if (
            not self.sampling_params.inference
            and self.sampling_params.allow_sample_source
            and random.random() < self.sampling_params.prob_sample_source
        ):
            data_row_sample_source = self._load_sample_source_audio(
                main_meta, mock=self.sampling_params.mock_data, load_for_semantic=not use_raw_audio
            )
        else:
            data_row_sample_source = None
        if (
            not self.sampling_params.inference
            and self.sampling_params.allow_mashup
            and random.random() < self.sampling_params.prob_mashup
        ):
            mashup_tracks = self._load_mashup_audio(
                main_meta, mock=self.sampling_params.mock_data, load_for_semantic=not use_raw_audio
            )
        else:
            mashup_tracks = None
        # Sample conditioning - moved before stem conditioning to preserve intact data_row
        audio_sample_tracks = None
        audio_sample_start_times_s = None
        audio_sample_sources = None
        if (
            not self.sampling_params.inference
            and self.sampling_params.allow_sample
            and random.random() < self.sampling_params.prob_sample
        ):
            # Sample conditioning mode: extract random number of audio samples
            max_samples = getattr(self.sampling_params, "max_num_audio_samples", 1)
            num_samples = random.randint(1, max_samples)
            prob_stems = getattr(
                self.sampling_params, "prob_sample_from_stems", 0.95
            )  # Default 95% stems, 5% full song

            extracted_samples = []
            extracted_times = []
            extracted_sources = []

            for _ in range(num_samples):
                result = None
                # Get sample permutation probability from sampling params (only for training)
                sample_permutation_prob = 0.0
                if self.split == "train":
                    sample_permutation_prob = getattr(
                        self.sampling_params, "sample_permutation_prob", 0.3
                    )

                if random.random() < prob_stems:
                    # Try stems first (preferably vocal)
                    result = create_sample_from_stems(
                        self, main_meta, data_row, self.semantic_sample_rate, sample_permutation_prob
                    )

                # If stems weren't chosen or failed, use full song approach
                # TODO: also apply smart cropping for full song, instead of the naive approach
                if result is None:
                    result = create_sample_from_full_song(
                        data_row, main_meta, self.semantic_sample_rate, sample_permutation_prob
                    )

                if result is not None:
                    sample_audio, sample_time, source_type = result
                    extracted_samples.append(sample_audio)
                    extracted_times.append(sample_time)
                    extracted_sources.append(source_type)

            if extracted_samples:
                audio_sample_tracks = extracted_samples
                audio_sample_start_times_s = extracted_times
                audio_sample_sources = extracted_sources

        stem_type = None
        data_row_stem = None
        stem_output_mix = None
        if (
            not self.sampling_params.inference
            and self.sampling_params.allow_stem
            and main_meta.get("stems", {})
            and (random.random() < self.sampling_params.prob_stem or "stem_active_sections" in main_meta)
        ):
            task = random.choices(["add", "extract", "remove"], weights=[0.8, 0.0, 0.0])[0]
            stem_type, stem_input_mix, isolated_stem, full_mix = self._load_stem_audio(
                main_meta,
                task=task,
                mock=self.sampling_params.mock_data,
                load_for_semantic=not use_raw_audio,
            )
            if stem_type is not None:
                data_row_stem = stem_input_mix
                data_row = isolated_stem
                stem_output_mix = full_mix

        # passing 48khz raw audio through dataloaders reduces throughput 20%
        use_hoot = self.sampling_params.use_hoot and random.random() < self.sampling_params.prob_use_hoot
        use_repa_hoot = (
            self.sampling_params.repa_hoot and random.random() < self.sampling_params.prob_repa_hoot
        )
        use_repa_midi = (
            self.sampling_params.repa_midi and random.random() < self.sampling_params.prob_repa_midi
        )
        if (
            self.sampling_params.use_vae_input
            or self.sampling_params.output_distribution == "vae"
            or use_hoot
            or use_repa_hoot
        ):
            use_raw_audio = True

        if use_raw_audio:
            raw_audio = self._load_audio_for_semantic(
                main_meta["local_filepath"],
                main_meta["s3_filepath"],
                main_meta["duration_s"],
                mock=self.sampling_params.mock_data,
                load_for_semantic=False,
            )
        else:
            raw_audio = None

        sample_data = SampleData(
            data_row=data_row,
            data_meta=main_meta,
            sampling_params=self.sampling_params,
            data_row_cover=data_row_cover,
            artist_tracks=data_row_artist,
            ditto_track=data_row_ditto,
            playlist_tracks=data_row_playlist,
            overpaint_track=data_row_overpaint,
            underpaint_track=data_row_underpaint,
            vox_track=data_row_vox,
            remix_track=data_row_remix,
            sample_source_track=data_row_sample_source,
            mashup_tracks=mashup_tracks,
            stem_track=data_row_stem,
            stem_output_mix=stem_output_mix,
            audio_sample_tracks=audio_sample_tracks,
            audio_sample_start_times_s=audio_sample_start_times_s,
            audio_sample_sources=audio_sample_sources,
            raw_audio=raw_audio,
            stem_type=stem_type,
            audio_type=audio_type,
            # signals passed from audioloader to data_utils.py
            use_repa_hoot=use_repa_hoot,
            use_repa_midi=use_repa_midi,
            use_hoot=use_hoot,
        )
        return sample_data
