from collections import defaultdict
import json
import math
import os
import random
import re
from dataclasses import dataclass
from tqdm import tqdm

import numpy as np
import orjson
from tokenizers import AddedToken
import torch
import torch.nn.functional as F
from torch.utils.data import IterableDataset
from transformers import PreTrainedTokenizerFast

from oracle_dataset import get_sample_oracle_file_segment
from modules.gpt import GPTConfig, GPTTrainConfig
from utils.helpers import print_with_time_master
from utils.bct import Block, BlockSequence, PackedBlockSequence, BlockType

# to avoid: "The current process just got forked, after parallelism has already been used"
os.environ["TOKENIZERS_PARALLELISM"] = "False"

RELOAD_MEMMAP = False

global tokenizer
g_tokenizer = None

TextBlockType = BlockType(
    name="text",
    is_causal=True,
)

CausalSemanticBlockType = BlockType(
    name="semantic",
    is_causal=True,
)
ArtistBlockType = BlockType(
    name="artist",
    is_causal=True,
)
PlaylistBlockType = BlockType(
    name="playlist",
    is_causal=True,
)
UnderpaintBlockType = BlockType(
    name="underpaint",
    is_causal=True,
)
OverpaintBlockType = BlockType(
    name="overpaint",
    is_causal=True,
)
CoverBlockType = BlockType(
    name="cover",
    is_causal=True,
)
PrefixBlockType = BlockType(
    name="prefix",
    is_causal=True,
)
SuffixBlockType = BlockType(
    name="suffix",
    is_causal=True,
)

NonCausalSemanticBlockType = BlockType(
    name="non_causal_semantic",
    is_causal=False,
)


# ++ Define a dataclass to hold sampling parameters ++
@dataclass
class SamplingParams:
    # Flags controlling data augmentation and task types
    allow_infill: bool = False
    allow_artist: bool = False
    allow_cover: bool = False
    allow_overpaint: bool = False
    allow_underpaint: bool = False
    allow_skip: bool = False
    allow_playlist: bool = False
    text_loss: bool = False  # Renamed from allow_text_loss for consistency if desired

    # probs
    prob_infill: float = 0.5  # Used when random.random() <= 0.25
    prob_artist: float = 0.1  # Used when random.random() < p_use_artist (0.1 or 0.5)
    prob_cover: float = 0.5  # Used when random.random() < 0.5
    prob_overpaint: float = 0.5  # Used when random.random() < 0.5
    prob_underpaint: float = 0.5  # Used when random.random() < 0.5
    prob_skip: float = 0.05  # Used when random.random() <= 0.1
    prob_playlist: float = 0.2  # Used when random.random() < 0.1

    # Flags controlling sampling mode
    inference: bool = False
    dummy_data: bool = False
    suppress_text: bool = False
    dataset_idx: int | None = None

    # Flags/params controlling batching and padding
    pack: bool = False
    mask_padding: bool = False
    min_text_offs: int | None = None

    # Add any other related parameters that are frequently passed together


MAX_TAG_LEN = 1024  # characters, not tokens
MAX_TOT_TAGS_LEN = 2048  # characters, not tokens
MAX_N_TAGS = 10  # to avoid overfitting to an artist


def _load_tokenizer(tokenizer_fp=None):
    global g_tokenizer
    if g_tokenizer is not None:
        return g_tokenizer
    assert os.path.exists(tokenizer_fp)
    g_tokenizer = PreTrainedTokenizerFast(
        tokenizer_file=tokenizer_fp,
        unk_token="[UNK]",
        pad_token="[PAD]",
    )
    g_tokenizer.add_special_tokens({"additional_special_tokens": [AddedToken("\n")]})
    return g_tokenizer


def _space_repl(m):
    s = m.group()
    n_newline = s.count("\n")
    if n_newline >= 2:
        return "\n\n"
    elif n_newline == 1:
        return "\n"
    return " "


def _simplify_whitespace(text, retain_newlines=True):
    """simplify while respecting up to 2 newlines"""
    if retain_newlines:
        text = re.sub(r"\s+", _space_repl, text).strip()
    else:
        text = re.sub(r"\s+", " ", text).strip()
    return text


def tokenize_batch(
    text_list,
    max_tokens=None,
    pad_token_id=0,
    retain_newlines=True,
    tokenizer_fp=None,
):
    tokenizer = _load_tokenizer(tokenizer_fp)
    text_list = [_simplify_whitespace(s, retain_newlines=retain_newlines) for s in text_list]
    text_enc = tokenizer(
        text_list,
        add_special_tokens=False,
        truncation=True,
        max_length=max_tokens,
        padding="longest",
        return_tensors="pt",
    )["input_ids"].type(torch.long)
    text_enc[text_enc == tokenizer.pad_token_id] = pad_token_id
    return text_enc


def read_jsonl(filepath, parse_idx_set=None, max_lines=None, progress_bar=False):
    data = []
    with open(filepath) as f:
        line_idx = 0
        for line in tqdm(f, total=max_lines, disable=not progress_bar, mininterval=5):
            line = line.strip()
            if len(line) == 0:
                continue
            if parse_idx_set is not None and line_idx not in parse_idx_set:
                data.append(None)
                line_idx += 1
                continue
            m = json.loads(line)
            data.append(m)
            line_idx += 1
            if max_lines is not None and line_idx >= max_lines:
                break
    return data


def write_jsonl(data, filepath):
    with open(filepath, "w") as f:
        for d in data:
            f.write(json.dumps(d, ensure_ascii=False) + "\n")


def load_dataset(
    load_kwargs: dict,
    data_dir: str,
    filename: str,
    info_filename: str,
    metas_filename: str,
    weights_multiplier_map: dict,
    is_finetune: bool,
    use_raw_audio: bool = False,
) -> tuple:
    dataset_names = []
    data_idx_lists = []  # used to randomly sample from the dataset
    data_weights = []
    data = None
    if not use_raw_audio:
        data = np.memmap(os.path.join(data_dir, filename), dtype=np.uint16, mode="r")
        if load_kwargs["semantic_n_codebooks"] + load_kwargs["coarse_n_codebooks"] > 1:
            data = data.reshape(
                -1,
                load_kwargs["semantic_n_codebooks"] + load_kwargs["coarse_n_codebooks"],
            )
            assert (
                data[:100, :, : load_kwargs["semantic_n_codebooks"]].max()
                <= load_kwargs["semantic_vocab_size"]
            )
            assert (
                data[:100, :, load_kwargs["semantic_n_codebooks"] :].max()
                <= load_kwargs["coarse_vocab_size"]
            )
        elif load_kwargs["semantic_n_codebooks"] == 1 and load_kwargs["coarse_n_codebooks"] == 0:
            assert data[:100].max() <= load_kwargs["semantic_vocab_size"]
    with open(os.path.join(data_dir, info_filename)) as f:
        infos = json.load(f)
    # make sure we turn into int since keys in json get auto turned into strings
    for k in infos.keys():
        if "idx_map" in infos[k]:
            infos[k]["idx_map"] = {int(k): v for k, v in infos[k]["idx_map"].items()}
    for use_condition, task_name in zip(
        ("allow_cover", "allow_overpaint", "allow_underpaint"),
        ("covers", "overpaint", "underpaint"),
    ):
        if not load_kwargs[use_condition]:
            # we can shard the data and cache data with cover even if we don't want to use covers
            # so we remove the covers from the infos
            task_keys = [
                dset_name
                for dset_name in infos.keys()
                if infos[dset_name].get("task", "default") == task_name
            ]
            for task_key in task_keys:
                infos.pop(task_key)

    metas = read_jsonl(os.path.join(data_dir, metas_filename))
    if use_raw_audio:
        assert all("s3_filepath" in m for m in metas)

    artist_to_songs = defaultdict(list)
    for i, m in enumerate(metas):
        if "artists" in m:
            dset_prefix = m["dataset"].split("_")[0]
            artist_to_songs[f"{dset_prefix}__{'__'.join(sorted(m['artists']))}"].append(i)
    artist_to_songs = {k: v for k, v in artist_to_songs.items() if len(v) > 1}
    if load_kwargs["allow_artist"]:
        assert len(artist_to_songs) > 0, "no artist data found"
        print_with_time_master(f"found {len(artist_to_songs):,} samples with matching artists found")

    playlist_to_songs = defaultdict(list)
    for i, m in enumerate(metas):
        for playlist_id in m.get("playlist_ids", []):
            playlist_to_songs[playlist_id].append(i)
    playlist_to_songs = {k: v for k, v in playlist_to_songs.items() if len(v) > 1}
    if load_kwargs["allow_playlist"]:
        assert len(playlist_to_songs) > 0, "no playlist data found"
        print_with_time_master(f"found {len(playlist_to_songs):,} samples with matching playlists found")

    idx_set = set()
    has_cover = False
    has_overpaint = False
    has_underpaint = False
    for dset_name, info in infos.items():
        dataset_names.append(dset_name)

        if info.get("task", "default") == "default":
            idx_list = info["idx_list"][:]
            idx_set |= set(idx_list)
        elif info["task"] == "covers":
            if not load_kwargs["allow_cover"]:
                continue
            assert load_kwargs["pack"], "for now pack needs to be active to do covers"
            has_cover = True
            idx_list = []
            n_covers = 0
            for idx, child_idx_l in info["idx_map"].items():
                dset_prefix = dset_name.split("_")[0]
                idx_list.append(idx)
                idx_set.add(idx)
                idx_set |= set(child_idx_l)
                n_covers += len(child_idx_l)
            info["idx_list"] = idx_list
            print_with_time_master(f"found {len(idx_list):,} samples with {n_covers:,} total covers")
        elif info["task"] == "overpaint":
            if not load_kwargs["allow_overpaint"]:
                continue
            assert load_kwargs["pack"], "for now pack needs to be active to do overpaint"
            has_overpaint = True
            idx_list = []
            for idx, child_idx in info["idx_map"].items():
                # full song to instrumental map
                dset_prefix = dset_name.split("_")[0]
                idx_list.append(idx)
                idx_set.add(idx)
                idx_set.add(child_idx)
            info["idx_list"] = idx_list
            print_with_time_master(f"found {len(idx_list):,} total overpaints")
        elif info["task"] == "underpaint":
            if not load_kwargs["allow_underpaint"]:
                continue
            assert load_kwargs["pack"], "for now pack needs to be active to do underpaint"
            has_underpaint = True
            idx_list = []
            for idx, child_idx in info["idx_map"].items():
                # full song to vocals map
                dset_prefix = dset_name.split("_")[0]
                idx_list.append(idx)
                idx_set.add(idx)
                idx_set.add(child_idx)
            info["idx_list"] = idx_list
            print_with_time_master(f"found {len(idx_list):,} total underpaint")
        else:
            raise ValueError(f"unknown task for {dset_name} in info file")
        random.shuffle(idx_list)

        data_idx_lists.append(idx_list)
        data_weights.append(len(idx_list) * weights_multiplier_map.get(dset_name, 1.0))
    if load_kwargs["allow_cover"]:
        assert has_cover, "no cover data found"
    if load_kwargs["allow_overpaint"]:
        assert has_overpaint, "no overpaint data found"
    if load_kwargs["allow_underpaint"]:
        assert has_underpaint, "no underpaint data found"
    weights_norm = np.sum(data_weights)
    data_weights = [v / weights_norm for v in data_weights]

    if not is_finetune:
        print_with_time_master(f"indexed {len(idx_set) / len(metas) * 100:.1f}% of data")
    for k in weights_multiplier_map.keys():
        assert k in dataset_names
        print_with_time_master(f"weight for {k}: {weights_multiplier_map[k]}")

    # some checks on data vs metas
    if not use_raw_audio:
        for m in metas:
            assert m["offset_idx"] <= len(data)
        assert 0.99 <= max([m["offset_idx"] for m in metas]) / len(data) <= 1.0

    assert len(infos) == len(dataset_names) == len(data_weights) == len(data_idx_lists)
    shard_info = "" if load_kwargs["local_data_shard_dir"] is None else " (sharded)"
    print_with_time_master(f"{len(metas):,} lines of {metas_filename} loaded.{shard_info}")
    del idx_set
    return (
        dataset_names,
        data_idx_lists,
        data_weights,
        data,
        metas,
        infos,
        artist_to_songs,
        playlist_to_songs,
    )


CASE_AUGMENT_FUNCS = [
    str.upper,
    str.lower,
    str.capitalize,
    str.title,
]


def _augment_tag(s):
    # case augment
    if random.random() >= 0.8:
        s = random.choice(CASE_AUGMENT_FUNCS)(s)
    # other misc formatting
    if random.random() >= 0.5:
        s = s.replace("-", " ").strip()
    return s


def _get_control_tags(duration_s, sample_vocal_start_s, do_augment=True):
    control_tags = []
    control_tags.append(f"duration:{int(round(duration_s))}")
    min_durations = [n * 60 for n in range(10) if n * 60 <= duration_s]
    max_durations = [n * 60 for n in range(10) if n * 60 >= duration_s]
    if len(min_durations) > 0:
        control_tags.append(f"min_duration:{int(random.choice(min_durations))}")
    if len(max_durations) > 0:
        control_tags.append(f"max_duration:{int(random.choice(max_durations))}")
    if sample_vocal_start_s is not None:
        if sample_vocal_start_s <= 5:
            control_tags.append("vocals:early")
        if sample_vocal_start_s <= 15:
            control_tags.append("vocals:normal")
        if 10 <= sample_vocal_start_s <= 20:
            control_tags.append("vocals:intro")
    if do_augment:
        if random.random() >= 0.5:
            random.shuffle(control_tags)
            control_tags = control_tags[: random.randint(0, len(control_tags))]
    if len(control_tags) == 0:
        return None
    return "{" + ";".join(control_tags) + "}"


def mask_middle_padding(stream: np.ndarray, old_pad_id: int, new_pad_id=-1):
    B, C, T = stream.shape

    shift_left = stream[:, :, :-2] == old_pad_id
    shift_right = stream[:, :, 2:] == old_pad_id
    current = stream[:, :, 1:-1] == old_pad_id

    # The middle mask now checks if the current (center) token
    # and its immediate neighbors (left and right) are all padding tokens.
    # We use logical AND on shifted views of the 'stream' array.
    middle_mask = shift_left & current & shift_right

    # Apply the mask to a slice of the stream to avoid affecting the edges,
    # since the original mask does not include the first and last columns.
    stream[:, :, 1:-1][middle_mask] = new_pad_id

    return stream


def mask_batch_padding(batch: np.ndarray, cfg):
    """batch has shape (B, C, T) where C is (optional text) + n_semantic + n_coarse streams"""
    B, C, T = batch.shape
    if C == 1 + cfg.semantic_n_codebooks + cfg.coarse_n_codebooks:
        text_offs = 1
        mask_middle_padding(batch[:, :text_offs], cfg.text_pad_token)
    elif C == cfg.semantic_n_codebooks + cfg.coarse_n_codebooks:
        text_offs = 0
    else:
        raise ValueError()
    mask_middle_padding(
        batch[:, text_offs : text_offs + cfg.semantic_n_codebooks], cfg.semantic_pad_token
    )
    mask_middle_padding(batch[:, text_offs + cfg.semantic_n_codebooks :], cfg.coarse_pad_token)

    # mask semantic mask
    batch[:, text_offs : text_offs + cfg.semantic_n_codebooks][
        batch[:, text_offs : text_offs + cfg.semantic_n_codebooks] == cfg.semantic_mask_token
    ] = -1
    return batch


def _augment_tags(tags):
    random.shuffle(tags)
    if random.random() <= 0.5:
        tags = tags[: random.randint(0, len(tags))]
        tags = [_augment_tag(tag) for tag in tags]
    return tags


def _clean_tags(tags):
    return [
        clean_tag[:MAX_TAG_LEN]
        for tag in tags
        if len(clean_tag := _simplify_whitespace(tag, retain_newlines=False)) > 0
    ]


def _clean_inline_tags(m):
    tags = m.group(2).split(";")
    tags = _clean_tags(tags)
    ts = ";".join(tags)[:MAX_TOT_TAGS_LEN]
    if len(ts) > 0:
        return f"[{m.group(1)}: {ts}]"
    return f"[{m.group(1)}]"


def _augment_inline_tags(m):
    tags = m.group(2).split(";")
    tags = _augment_tags(tags)
    merge_char = random.choice([", ", " ", "; ", ",", ";", ". "])
    ts = merge_char.join(tags)
    if len(ts) > 0:
        return f"[{m.group(1)}: {ts}]"
    return f"[{m.group(1)}]"


def build_text(
    tags,
    text,
    duration_s,
    sample_vocal_start_s,
    inference,
    suppress_text,
    enable_control_tags=True,
    passin_control_tags=None,
):
    if suppress_text:
        return ""
    text_elements = []
    # collect tags
    tags = _clean_tags(tags)
    if not inference:
        tags = _augment_tags(tags)
        merge_char = random.choice([", ", " ", "; ", ",", ";", ". "])
        tags_str = f"{merge_char.join(tags[:MAX_N_TAGS])}"[:MAX_TOT_TAGS_LEN]
    else:
        # keep the non-inference behavoir
        tags_str = f"{','.join([tag[:MAX_TAG_LEN] for tag in tags])[:MAX_TOT_TAGS_LEN]}"
    if len(tags_str) > 0:
        text_elements.append(f"[{tags_str}]")
    # get lyrics
    if len(text) > 0 and not inference:
        text = re.sub(r"\[(.*?)\:(.*?)\]", _clean_inline_tags, text)
        # augment tags inside text:
        text = re.sub(r"\[(.*?)\:(.*?)\]", _augment_inline_tags, text)
        if random.random() >= 0.95:
            text = text.lower()
        if random.random() >= 0.95:
            text = re.sub(r"\n+", " ", text)
    if len(text) > 0:
        text_elements.append(text.strip())
    # get control tags
    for n in range(len(text_elements)):
        text_elements[n] = text_elements[n].replace("{", "").replace("}", "")
    if (inference or random.random() >= 0.1) and enable_control_tags:
        control_tags = _get_control_tags(duration_s, sample_vocal_start_s, do_augment=not inference)
        if control_tags is not None:
            text_elements = [control_tags] + text_elements
    # this is hard coded control tags that are passed in by the user
    if passin_control_tags:
        text_elements = [passin_control_tags] + text_elements
    if inference or random.random() >= 0.5:
        text = "\n\n".join(text_elements)
    else:
        text = ""
        for t in text_elements:
            text += random.choice([" ", "\n", "\n\n"]) + t
    text = text.strip()
    return text


def pad_audio_arr(arr, cfg):
    """Pads an arr of size (C, T) to (C, cfg.t_audio)"""
    C, T = arr.shape
    assert T <= cfg.t_audio
    assert C == cfg.semantic_n_codebooks + cfg.coarse_n_codebooks
    pad = np.empty((C, cfg.t_audio - T), dtype=np.int64)
    pad[: cfg.semantic_n_codebooks] = cfg.semantic_pad_token
    pad[cfg.semantic_n_codebooks :] = cfg.coarse_pad_token
    return np.concatenate([arr, pad], axis=-1)


def pad_x_arr(x_arr, cfg, sz: int = None):
    """Pads an x_arr of size (C, T) to (C, cfg.block_size)"""
    C, T = x_arr.shape
    if sz is None:
        sz = cfg.block_size
    assert C == cfg.semantic_n_codebooks + cfg.coarse_n_codebooks + 1
    pad = np.empty((C, sz - T), dtype=np.int64)
    pad[0] = cfg.text_pad_token
    pad[1 : 1 + cfg.semantic_n_codebooks] = cfg.semantic_pad_token
    pad[1 + cfg.semantic_n_codebooks :] = cfg.coarse_pad_token
    return np.concatenate([x_arr, pad], axis=-1)


def build_audio_arr(data_row, cfg, semantic_infer_token=None, include_eos=True, skip_factor=1):
    """Builds an audio arr from a data row.
    Prepends the semantic and coarse streams with the infer token."""
    data_row = data_row[:, ::skip_factor]

    audio_len = (
        cfg.semantic_n_codebooks * cfg.semantic_shift_factor * min(1, cfg.coarse_n_codebooks)
        + max((cfg.coarse_n_codebooks - 1), 0) * cfg.coarse_shift_factor
        + data_row[-1].shape[-1]
    )
    if cfg.coarse_n_codebooks == 0 and include_eos:
        audio_len += 1  # used as eos token
    # build semantic
    y_semantic_arr = np.full(
        (cfg.semantic_n_codebooks, audio_len), cfg.semantic_pad_token, dtype=np.int64
    )
    for n in range(cfg.semantic_n_codebooks):
        offs = n * cfg.semantic_shift_factor
        y_semantic_arr[n, offs : offs + data_row[n].shape[-1]] = data_row[n]
    # build coarse
    if cfg.coarse_n_codebooks > 0:
        y_coarse_arr = np.full((cfg.coarse_n_codebooks, audio_len), cfg.coarse_pad_token, dtype=np.int64)
        for n in range(cfg.coarse_n_codebooks):
            offs = cfg.semantic_n_codebooks * cfg.semantic_shift_factor + n * cfg.coarse_shift_factor
            n2 = cfg.semantic_n_codebooks + n
            y_coarse_arr[n, offs : offs + data_row[n2].shape[-1]] = data_row[n2]

    # combine audio and add x with infer token
    if semantic_infer_token is None:
        semantic_infer_token = cfg.semantic_infer_token
    y_audio_arr = y_semantic_arr
    if cfg.coarse_n_codebooks > 0:
        y_audio_arr = np.concatenate([y_audio_arr, y_coarse_arr], axis=0)
    x_audio_arr = y_audio_arr.copy()
    x_audio_arr = np.concatenate(
        [
            np.array(
                [[semantic_infer_token]] * cfg.semantic_n_codebooks
                + [[cfg.coarse_infer_token]] * cfg.coarse_n_codebooks
            ).repeat(skip_factor, 1),
            x_audio_arr,
        ],
        axis=-1,
    )
    return x_audio_arr


def _get_rand_int(min_val, max_val):
    if max_val <= min_val:
        return min_val
    return random.randint(min_val, max_val)


@dataclass
class SampleData:
    data_row: np.ndarray
    data_meta: dict
    sampling_params: SamplingParams
    is_full_track: bool = False
    data_row_cover: np.ndarray | None = None
    artist_tracks: list[np.ndarray] | None = None
    playlist_tracks: list[np.ndarray] | None = None
    overpaint_track: np.ndarray | None = None
    underpaint_track: np.ndarray | None = None
    idx: int | None = None


def get_sample_codes_or_audio(
    data_sampling_info,
    split,
    sampling_params: SamplingParams,
    dataset_idx=None,
    row_idx=None,  # absolute, overrides dataset_idx
):
    data = data_sampling_info[split]["data"]
    metas = data_sampling_info[split]["metas"]
    idx_lists = data_sampling_info[split]["idx_lists"]
    infos = data_sampling_info[split]["infos"]
    weights = data_sampling_info[split]["weights"]
    dset_names = data_sampling_info[split]["names"]
    artist_to_songs = data_sampling_info[split]["artist_to_songs"]
    playlist_to_songs = data_sampling_info[split]["playlist_to_songs"]
    cfg = data_sampling_info["cfg"]

    if row_idx is not None:
        pass
    elif dataset_idx is not None:
        row_idx = random.choice(idx_lists[dataset_idx])
    else:
        weights = weights[:]
        for allow_task, task_name in zip(
            (
                sampling_params.allow_cover,
                sampling_params.allow_overpaint,
                sampling_params.allow_underpaint,
            ),
            ("covers", "overpaint", "underpaint"),
        ):
            if not allow_task:
                # set weight for cover dataset to 0
                for n, dset_name in enumerate(dset_names):
                    if infos.get(dset_name, {}).get("task", "default") == task_name:
                        weights[n] = 0
            assert sum(weights) > 0
        dataset_idx = random.choices(list(range(len(weights))), weights=weights, k=1)[0]
        row_idx = random.choice(idx_lists[dataset_idx])

    skip_factor = 1
    if (
        sampling_params.allow_skip
        and not sampling_params.inference
        and random.random() <= sampling_params.prob_skip
    ):
        skip_factor = 4 if random.random() <= 0.5 else 2

    def load_data_row(row_idx, dummy_data=False):
        data_meta = metas[row_idx]
        offset_idx = None
        if data is None:
            # TODO: fix this once we encode and know after
            use_n_tokens = int(math.floor(data_meta["duration_s"] * cfg.semantic_rate_hz))
        else:
            offset_idx = data_meta["offset_idx"]
            use_n_tokens = data_meta["n_tokens"]
        phase_offset = random.choice(list(range(0, skip_factor)))
        use_n_tokens = min(math.floor((use_n_tokens - phase_offset) / skip_factor), cfg.t_audio)
        is_full_track = use_n_tokens < cfg.t_audio

        if dummy_data:
            data_row = np.zeros(
                (cfg.semantic_n_codebooks + cfg.coarse_n_codebooks + 1, use_n_tokens),
                dtype=np.int64,
            )

        else:
            if data is None:
                audio = get_sample_oracle_file_segment(data_meta["s3_filepath"])
                audio = audio.convert(24000, 2, 1)
                data_row = audio.array_float
            else:
                data_row = (
                    data[
                        offset_idx + phase_offset : offset_idx
                        + phase_offset
                        + use_n_tokens * skip_factor : skip_factor
                    ]
                    .astype(np.int64)
                    .reshape(cfg.semantic_n_codebooks + cfg.coarse_n_codebooks, -1)
                )

        return data_row, data_meta, is_full_track

    data_row, data_meta, is_full_track = load_data_row(row_idx, dummy_data=sampling_params.dummy_data)

    if (
        sampling_params.allow_cover
        and not sampling_params.inference
        and dataset_idx is not None
        and (dataset_name := data_sampling_info[split]["names"][dataset_idx])
        and infos[dataset_name].get("task", "default") == "covers"
    ):
        # sample a random cover
        cover_list = infos[dataset_name]["idx_map"][row_idx]
        cover_idx = random.choice(cover_list)
        data_row_cover, _, _ = load_data_row(cover_idx, dummy_data=sampling_params.dummy_data)
    else:
        data_row_cover = None
    used_cover = data_row_cover is not None

    # Artist audio to audio
    dset_prefix = data_meta["dataset"].split("_")[0]
    p_use_artist = (
        sampling_params.prob_artist * 5 if used_cover else sampling_params.prob_artist
    )  # increase chance for cover cause voice beautifier
    artist_tracks = []
    if (
        sampling_params.allow_artist
        and not sampling_params.inference
        and "artists" in data_meta
        and random.random() < p_use_artist
        and len(artist_to_songs.get(f"{dset_prefix}__{'__'.join(sorted(data_meta['artists']))}", []))
        >= 2
    ):
        # sample a random song from the artist
        track_idx_list = [
            _idx for _idx in artist_to_songs[f"{dset_prefix}__{'__'.join(sorted(data_meta['artists']))}"]
        ]
        track_idx_list.remove(row_idx)
        if len(track_idx_list) > 0:
            # sample multiple artist segments
            n_max_tracks = 5
            max_track_duration_s = 60
            # assert max_track_duration_s * n_max_tracks * cfg.semantic_rate_hz <= cfg.t_audio
            for _ in range(random.randint(1, n_max_tracks)):
                track_idx = random.choice(list(track_idx_list))
                data_row_track, _, _ = load_data_row(track_idx, dummy_data=sampling_params.dummy_data)
                # use only part of artist
                left_idx = random.randint(0, data_row_track.shape[-1])
                use_track_duration_s = random.randint(
                    0, min(data_row_track.shape[-1] - left_idx, max_track_duration_s)
                )
                right_idx = random.randint(left_idx, left_idx + use_track_duration_s)
                assert right_idx - left_idx <= max_track_duration_s, (right_idx, left_idx)
                data_row_track = data_row_track[:, left_idx:right_idx]
                if not 0 <= left_idx and left_idx < right_idx and right_idx <= data_row_track.shape[-1]:
                    continue
                artist_tracks.append(data_row_track)

    # Playlist audio to audio
    playlist_tracks = []
    if (
        sampling_params.allow_playlist
        and not sampling_params.inference
        and "playlist_ids" in data_meta
        and random.random() < sampling_params.prob_playlist
        and len(
            [
                True
                for playlist_id in data_meta["playlist_ids"]
                if len(playlist_to_songs.get(playlist_id, [])) >= 2
            ]
        )
        >= 1
    ):
        # sample a random song from the artist
        playlist_ids = [
            playlist_id
            for playlist_id in data_meta["playlist_ids"]
            if len(playlist_to_songs.get(playlist_id, [])) >= 2
        ]
        playlist_id = random.choice(playlist_ids)
        track_idx_list = [_idx for _idx in playlist_to_songs[playlist_id]]
        track_idx_list.remove(row_idx)
        if len(track_idx_list) > 0:
            # sample multiple artist segments
            n_max_tracks = 5
            max_track_duration_s = 60
            # assert max_track_duration_s * n_max_tracks * cfg.semantic_rate_hz <= cfg.t_audio
            for _ in range(random.randint(1, n_max_tracks)):
                track_idx = random.choice(list(track_idx_list))
                data_row_track, _, _ = load_data_row(track_idx, dummy_data=sampling_params.dummy_data)
                # use only part of artist
                left_idx = random.randint(0, data_row_track.shape[-1])
                use_track_duration_s = random.randint(
                    0, min(data_row_track.shape[-1] - left_idx, max_track_duration_s)
                )
                right_idx = random.randint(left_idx, left_idx + use_track_duration_s)
                assert right_idx - left_idx <= max_track_duration_s, (right_idx, left_idx)
                data_row_track = data_row_track[:, left_idx:right_idx]
                if not 0 <= left_idx and left_idx < right_idx and right_idx <= data_row_track.shape[-1]:
                    continue
                playlist_tracks.append(data_row_track)

    # Overpaint
    overpaint_track = None
    if (
        sampling_params.allow_overpaint
        and not sampling_params.inference
        and dataset_idx is not None
        and (dataset_name := data_sampling_info[split]["names"][dataset_idx])
        and infos[dataset_name].get("task", "default") == "overpaint"
    ):
        # map is instrumental to full
        row_idx_child = infos[dataset_name]["idx_map"][row_idx]
        overpaint_track, _, _ = load_data_row(row_idx_child, dummy_data=sampling_params.dummy_data)

    # Underpaint
    underpaint_track = None
    if (
        sampling_params.allow_underpaint
        and not sampling_params.inference
        and dataset_idx is not None
        and (dataset_name := data_sampling_info[split]["names"][dataset_idx])
        and infos[dataset_name].get("task", "default") == "underpaint"
    ):
        # map is vocals to full song
        row_idx_child = infos[dataset_name]["idx_map"][row_idx]
        underpaint_track, _, _ = load_data_row(row_idx_child, dummy_data=sampling_params.dummy_data)

    sample_data = SampleData(
        data_row=data_row,
        data_meta=data_meta,
        sampling_params=sampling_params,
        skip_factor=skip_factor,
        is_full_track=is_full_track,
        data_row_cover=data_row_cover,
        artist_tracks=artist_tracks,
        playlist_tracks=playlist_tracks,
        overpaint_track=overpaint_track,
        underpaint_track=underpaint_track,
        idx=row_idx,
    )

    return sample_data


def get_sample_from_row(model_cfg: GPTConfig, tokenizer_fp: str, sample_data: SampleData):
    cfg = model_cfg
    sampling_params = sample_data.sampling_params
    data_row = sample_data.data_row
    data_meta = sample_data.data_meta

    def sample_skip_factor():
        if sampling_params.allow_skip and random.random() <= sampling_params.prob_skip:
            return random.randint(2, 6)
        return 1

    is_full_track = sample_data.is_full_track
    data_row_cover = sample_data.data_row_cover
    artist_tracks = sample_data.artist_tracks
    playlist_tracks = sample_data.playlist_tracks
    overpaint_track = sample_data.overpaint_track
    underpaint_track = sample_data.underpaint_track

    if data_row.shape[-1] > cfg.block_size - cfg.t_text:
        print(
            f"WARNING: data row is long {data_meta}: {data_row.shape[-1]} > {cfg.block_size - cfg.t_text}"
        )
        data_row = data_row[:, : cfg.block_size - cfg.t_text]

    # build main audio array
    sample_tags = data_meta.get("tags", [])
    sample_vocal_start_s = None
    audio_blocks = []
    if (
        sampling_params.allow_infill
        and not sampling_params.inference
        and data_row.shape[-1] >= 1
        and random.random() <= sampling_params.prob_infill
    ):
        # a_b_c make b, where a and c can both be empty to allow arbitrary handling
        t = data_row.shape[-1]
        b_duration = _get_rand_int(1, t)
        a_right_idx = _get_rand_int(0, t - b_duration)
        b_right_idx = a_right_idx + b_duration

        if random.random() <= 0.05:
            a_left_idx = 0
        elif random.random() <= 0.05:
            a_left_idx = a_right_idx
        else:
            a_left_idx = _get_rand_int(0, a_right_idx)
        if random.random() <= 0.05:
            c_right_idx = t
        elif random.random() <= 0.05:
            c_right_idx = b_right_idx
        else:
            c_right_idx = _get_rand_int(b_right_idx, t)
        sample_text = ""
        if len(data_meta.get("text_lines", [])) > 0:
            text_lines = data_meta["text_lines"]
            text_left_idx = 0
            text_right_idx = len(text_lines)
            for idx, m_line in enumerate(text_lines):
                if m_line["start_s"] <= a_right_idx / cfg.semantic_rate_hz:
                    text_left_idx = idx
                if m_line["end_s"] >= b_right_idx / cfg.semantic_rate_hz:
                    text_right_idx = idx + 1
                    break
            text_left_idx = _get_rand_int(0, text_left_idx)
            text_right_idx = _get_rand_int(text_right_idx, len(text_lines))
            sample_text = "\n".join([m["text"] for m in text_lines[text_left_idx:text_right_idx]])
        elif "text" in data_meta:
            sample_text = data_meta["text"]
        a_audio_arr = build_audio_arr(
            data_row[:, a_left_idx:a_right_idx],
            cfg,
            semantic_infer_token=cfg.semantic_history_token,
            include_eos=False,
            skip_factor=sample_skip_factor(),
        )
        c_audio_arr = build_audio_arr(
            data_row[:, b_right_idx:c_right_idx],
            cfg,
            semantic_infer_token=cfg.semantic_future_token,
            include_eos=False,
            skip_factor=sample_skip_factor(),
        )
        b_audio_arr = build_audio_arr(
            data_row[:, a_right_idx:b_right_idx],
            cfg,
            include_eos=True,
            skip_factor=sample_skip_factor(),
        )
        x_audio_arr = np.concatenate([c_audio_arr, a_audio_arr, b_audio_arr], axis=-1)
        a_audio_arr = torch.from_numpy(a_audio_arr)
        b_audio_arr = torch.from_numpy(b_audio_arr)
        c_audio_arr = torch.from_numpy(c_audio_arr)
        audio_blocks.extend(
            [
                Block(
                    spec=SuffixBlockType,
                    inputs={"semantic_input": c_audio_arr},
                    targets={"semantic_output": Block.shift_left(c_audio_arr, cfg.semantic_pad_token)},
                ),
                Block(
                    spec=PrefixBlockType,
                    inputs={"semantic_input": a_audio_arr},
                    targets={"semantic_output": Block.shift_left(a_audio_arr, cfg.semantic_pad_token)},
                ),
                Block(
                    spec=CausalSemanticBlockType,
                    inputs={"semantic_input": b_audio_arr},
                    targets={"semantic_output": Block.shift_left(b_audio_arr, cfg.semantic_pad_token)},
                ),
            ]
        )
        sample_duration_s = (c_right_idx - a_left_idx) / cfg.semantic_rate_hz
        text = build_text(
            sample_tags,
            sample_text,
            sample_duration_s,
            sample_vocal_start_s,
            sampling_params.inference,
            sampling_params.suppress_text,
        )
    else:
        if "text_lines" in data_meta and len(data_meta["text_lines"]) > 0:
            sample_vocal_start_s = data_meta["text_lines"][0]["start_s"]
        x_audio_arr = build_audio_arr(
            data_row, cfg, include_eos=is_full_track, skip_factor=sample_skip_factor()
        )
        x_audio_arr = torch.from_numpy(x_audio_arr)
        audio_blocks.append(
            Block(
                spec=CausalSemanticBlockType,
                inputs={"semantic_input": x_audio_arr},
                targets={"semantic_output": Block.shift_left(x_audio_arr, cfg.semantic_pad_token)},
            )
        )
        sample_text = data_meta.get("text", "")
        sample_duration_s = data_row.shape[-1] / cfg.semantic_rate_hz
        text = build_text(
            sample_tags,
            sample_text,
            sample_duration_s,
            sample_vocal_start_s,
            sampling_params.inference,
            sampling_params.suppress_text,
        )

    # build text arr
    if not sampling_params.inference and random.random() >= 0.98:  # sometimes do unconditional
        text = ""
    x_text_arr = tokenize_batch(
        [text],
        max_tokens=cfg.t_text,
        pad_token_id=cfg.text_pad_token,
        tokenizer_fp=tokenizer_fp,
    )[0]  # (N_TEXT_TOKENS)
    if (
        sampling_params.min_text_offs is not None
        and sampling_params.min_text_offs > x_text_arr.shape[-1]
    ):
        x_text_arr = F.pad(
            x_text_arr,
            (0, sampling_params.min_text_offs - x_text_arr.shape[-1]),
            "constant",
            cfg.text_pad_token,
        )

    if sampling_params.text_loss:
        x_text_arr = F.pad(
            x_text_arr,
            (1, 0),
            "constant",
            cfg.text_infer_token,
        )
        if x_text_arr[-1] != cfg.text_pad_token:
            x_text_arr = F.pad(
                x_text_arr,
                (0, 1),
                "constant",
                cfg.text_pad_token,
            )

    # Covers
    if data_row_cover is not None:
        if random.random() <= 0.5:
            # randomly drop out some of the cover
            data_row_cover = data_row_cover[:, : random.randint(0, data_row_cover.shape[-1])]
        # prepend cover to x_audio_arr
        cover_audio_arr = build_audio_arr(
            data_row_cover,
            cfg,
            semantic_infer_token=cfg.semantic_cover_token,
            include_eos=False,
            skip_factor=sample_skip_factor(),
        )
        cover_audio_arr = torch.from_numpy(cover_audio_arr)
        audio_blocks.insert(
            0,
            Block(
                spec=CoverBlockType,
                inputs={"semantic_input": cover_audio_arr},
                targets={"semantic_output": Block.shift_left(cover_audio_arr, cfg.semantic_pad_token)},
            ),
        )

    # Artist audio to audio
    if artist_tracks is not None:
        for data_row_track in artist_tracks:
            artist_audio_arr = build_audio_arr(
                data_row_track,
                cfg,
                semantic_infer_token=cfg.semantic_artist_token,
                include_eos=False,
                skip_factor=sample_skip_factor(),
            )
            artist_audio_arr = torch.from_numpy(artist_audio_arr)
            audio_blocks.insert(
                0,
                Block(
                    spec=ArtistBlockType,
                    inputs={"semantic_input": artist_audio_arr},
                    targets={
                        "semantic_output": Block.shift_left(artist_audio_arr, cfg.semantic_pad_token)
                    },
                ),
            )

    # Playlist audio to audio
    if playlist_tracks is not None:
        for data_row_track in playlist_tracks:
            playlist_audio_arr = build_audio_arr(
                data_row_track,
                cfg,
                semantic_infer_token=cfg.semantic_playlist_token,
                include_eos=False,
                skip_factor=sample_skip_factor(),
            )
            playlist_audio_arr = torch.from_numpy(playlist_audio_arr)
            audio_blocks.insert(
                0,
                Block(
                    spec=PlaylistBlockType,
                    inputs={"semantic_input": playlist_audio_arr},
                    targets={
                        "semantic_output": Block.shift_left(playlist_audio_arr, cfg.semantic_pad_token)
                    },
                ),
            )
    # Overpaint
    if overpaint_track is not None:
        # prepend instrumental to x_audio_arr
        overpaint_audio_arr = build_audio_arr(
            overpaint_track,
            cfg,
            semantic_infer_token=cfg.semantic_overpaint_token,
            include_eos=False,
            skip_factor=sample_skip_factor(),
        )
        overpaint_audio_arr = torch.from_numpy(overpaint_audio_arr)
        audio_blocks.insert(
            0,
            Block(
                spec=OverpaintBlockType,
                inputs={"semantic_input": overpaint_audio_arr},
                targets={
                    "semantic_output": Block.shift_left(overpaint_audio_arr, cfg.semantic_pad_token)
                },
            ),
        )

    # Underpaint
    if underpaint_track is not None:
        # prepend vocals to x_audio_arr
        underpaint_audio_arr = build_audio_arr(
            underpaint_track,
            cfg,
            semantic_infer_token=cfg.semantic_underpaint_token,
            include_eos=False,
            skip_factor=sample_skip_factor(),
        )
        underpaint_audio_arr = torch.from_numpy(underpaint_audio_arr)
        audio_blocks.insert(
            0,
            Block(
                spec=UnderpaintBlockType,
                inputs={"semantic_input": underpaint_audio_arr},
                targets={
                    "semantic_output": Block.shift_left(underpaint_audio_arr, cfg.semantic_pad_token)
                },
            ),
        )

    # build blocks

    # add semantic pad to text
    x_text_arr = x_text_arr.unsqueeze(0)
    text_block = Block(
        spec=TextBlockType,
        inputs={
            "text_input": x_text_arr,
        },
    )

    text_is_post = (
        True
        if sampling_params.text_loss and not sampling_params.inference and random.random() < 0.1
        else False
    )
    if text_is_post:
        blocks = audio_blocks + [text_block]
    else:
        blocks = [text_block] + audio_blocks

    block_sequence = BlockSequence(blocks)
    # print(block_sequence)

    sample_info = {
        "text": text,
        "n_tokens_text": len(x_text_arr),
        "text_is_post": text_is_post,
        "blocks": block_sequence,
    }
    return sample_info


def get_batch(
    batch_size_tokens: int,
    cfg: GPTConfig,
    sample_generator,
    sampling_params: SamplingParams,
    return_idx=False,
):
    if return_idx:
        raise NotImplementedError("return_idx not implemented")

    packed_sequences = PackedBlockSequence([])

    if sampling_params.mask_padding:
        raise NotImplementedError("mask_padding not implemented")

    if sampling_params.pack:
        # fill up block_size with samples
        cur_len = 0
        n_loop = 0
        for sample in tqdm(sample_generator, desc="Packing batch", disable=True):
            if cur_len >= batch_size_tokens:
                break
            packed_sequences.append(sample["blocks"])
            cur_len += sample["blocks"].n_tokens
            n_loop += 1

        # crop so total length is batch_size_tokens
        packed_sequences = packed_sequences.crop_to_max_tokens(batch_size_tokens)
        assert packed_sequences.n_tokens == batch_size_tokens

        extra_factor = 1 if not sampling_params.allow_skip else 4
        if (
            n_loop > int(round(batch_size_tokens / cfg.block_size)) * 16 * extra_factor
        ):  # should almost never trigger on 8k blocks
            print(f"warning, {n_loop} loops in dataloading")

    return packed_sequences


class CustomAudioDataset(IterableDataset):
    def __init__(
        self,
        data_sampling_info,
        split,
        dataset_idx=None,
        sampling_params: SamplingParams | None = None,
    ):
        self.data_sampling_info = data_sampling_info
        self.dataset_idx = dataset_idx
        self.split = split
        self.sampling_params = sampling_params

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

    def sample_generator(self):
        while True:
            try:
                sample_data = get_sample_codes_or_audio(
                    self.data_sampling_info,
                    self.split,
                    self.sampling_params,
                    dataset_idx=self.dataset_idx,
                )
                packed_sequences = get_sample_from_row(
                    self.data_sampling_info["cfg"],
                    self.data_sampling_info["tokenizer_fp"],
                    sample_data,
                )
            except Exception as e:
                print(f"Error in sample_generator: {e}")
                continue
            yield packed_sequences

    def __iter__(self):
        return self

    def __next__(self):
        # Randomly sample from the dataset
        # ++ Pass the stored SamplingParams object to get_batch ++
        batch = get_batch(
            self.data_sampling_info["batch_size_tokens"],
            self.data_sampling_info["cfg"],
            self.sample_generator(),
            self.sampling_params,
        )
        return batch


from suno_utils.tasks.mert_25 import (
    preload_models as preload_semantic_models_,
    encode as encode_semantic,
)


def preload_semantic_models(device):
    preload_semantic_models_(
        checkpoint_filepath="s3://suno-data/georg/models/semantic/mert_25.pt",
        centroids_filepath="s3://suno-data/georg/models/semantic/mert_25_2x4k.npy",
        device=device,
    )


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]


class OracleWavAudioDataset(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",
        dataset_idx=None,
        sampling_params: SamplingParams | None = None,
    ):
        self.model_cfg = model_cfg
        self.train_cfg = train_cfg
        self.metas_path = metas_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

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

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

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

        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)
                self.weights.append(meta.get("weight", 0))
                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"])

        self.tokenizer = _load_tokenizer(self.tokenizer_fp)
        self.random_cache = []

    def __iter__(self):
        return self

    def _load_audio_for_mert(self, local_filepath, s3_filepath, start_s=0, max_duration_s=60 * 30):
        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)
        mock = False
        try:
            audio = get_sample_oracle_file_segment(
                local_filepath=local_filepath,
                s3_filepath=s3_filepath,
                start_s=start_s,
                max_duration_s=max_duration_s,
                mock=mock,
            )
            audio = audio.convert(24000, 2, 1)  # MERT sample rate
        except Exception as e:
            print(f"Error loading audio for {s3_filepath}: {e}")
            print(f"start_s: {start_s}, max_duration_s: {max_duration_s}")
            print(f"audio: {audio.array_float.shape}")
            raise e
        return audio.array_float

    def _load_cover_audio(self, main_meta):
        covers = main_meta.get("cover_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_mert(
                cover_meta["local_filepath"], cover_meta["s3_filepath"]
            )
        else:
            data_row_cover = None
        return data_row_cover

    def _load_artist_audio(self, main_meta):
        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, 5)
        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_mert(
                track_meta["local_filepath"],
                track_meta["s3_filepath"],
                start_s=start_s,
                max_duration_s=dur_s,
            )
            artist_audio_arrs.append(data_row_artist)

        return artist_audio_arrs

    def _load_playlist_audio(self, main_meta):
        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, 5)
        min_segment_len = 5
        max_segment_len = 60
        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_mert(
                track_meta["local_filepath"],
                track_meta["s3_filepath"],
                start_s=start_s,
                max_duration_s=dur_s,
            )
            playlist_audio_arrs.append(data_row_playlist)
        return playlist_audio_arrs

    def _load_overpaint_audio(self, main_meta):
        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_mert(
            overpaint_meta["local_filepath"], overpaint_meta["s3_filepath"]
        )
        return data_row_overpaint

    def _load_underpaint_audio(self, main_meta):
        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_mert(
            underpaint_meta["local_filepath"], underpaint_meta["s3_filepath"]
        )
        return data_row_underpaint

    def __next__(self):
        try:
            return self._next()
        except Exception as e:
            print(f"Error in __next__: {e}")
            # raise e
            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()

        # load wav
        data_row = self._load_audio_for_mert(main_meta["local_filepath"], main_meta["s3_filepath"])

        # TODO: add cover, artist, playlist, overpaint, underpaint
        if self.sampling_params.allow_cover and random.random() < self.sampling_params.prob_cover:
            data_row_cover = self._load_cover_audio(main_meta)
        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 self.sampling_params.allow_artist and random.random() < p_use_artist:
            data_row_artist = self._load_artist_audio(main_meta)
        else:
            data_row_artist = None

        if self.sampling_params.allow_playlist and random.random() < self.sampling_params.prob_playlist:
            data_row_playlist = self._load_playlist_audio(main_meta)
        else:
            data_row_playlist = None

        if (
            self.sampling_params.allow_overpaint
            and random.random() < self.sampling_params.prob_overpaint
        ):
            data_row_overpaint = self._load_overpaint_audio(main_meta)
        else:
            data_row_overpaint = None

        if (
            self.sampling_params.allow_underpaint
            and random.random() < self.sampling_params.prob_underpaint
        ):
            data_row_underpaint = self._load_underpaint_audio(main_meta)
        else:
            data_row_underpaint = None

        sample_data = SampleData(
            data_row=data_row,
            data_meta=main_meta,
            sampling_params=self.sampling_params,
            is_full_track=True,
            data_row_cover=data_row_cover,
            artist_tracks=data_row_artist,
            playlist_tracks=data_row_playlist,
            overpaint_track=data_row_overpaint,
            underpaint_track=data_row_underpaint,
        )
        return sample_data


class RawAudioDataset(IterableDataset):
    def __init__(
        self,
        sample_data_dl,
        batch_size_tokens: int,
        cfg: GPTConfig,
        sampling_params: SamplingParams | None = None,
        device: str | None = None,
        tokenizer_fp: str | None = None,
    ):
        self.sample_data_dl = sample_data_dl
        self.batch_size_tokens = batch_size_tokens
        self.model_cfg = cfg
        self.sampling_params = sampling_params
        self.device = device
        self.tokenizer_fp = tokenizer_fp

    def sample_generator_fn(self, sample_data_dl):
        while True:
            sample_data = next(sample_data_dl)

            def encode_semantic_track(track, fuzz_edges=True):
                assert isinstance(track, np.ndarray)
                track = torch.from_numpy(track).unsqueeze(0).contiguous()
                codes = encode_semantic(
                    track, pad_to_chunksize=True, device=self.device, batch_size=48
                ).T[[0]]  # (1, T)

                if fuzz_edges and random.random() < 0.2:
                    # crop up to 5s from both sides so we aren't always on the chunk boundary
                    l_cut = random.randint(0, 5 * self.model_cfg.semantic_rate_hz)
                    r_cut = random.randint(1, 5 * self.model_cfg.semantic_rate_hz)  # avoid -0 edge case
                    fuzzed_codes = codes[:, l_cut:-r_cut]
                    if fuzzed_codes.shape[1] > 0:  # if we crop too much, dont fuzz
                        codes = fuzzed_codes

                if codes.shape[1] > self.model_cfg.block_size - self.model_cfg.t_text:
                    # TODO: why does this happen?
                    codes = codes[:, : self.model_cfg.block_size - self.model_cfg.t_text]
                return codes

            # encode all data
            if sample_data.data_row is not None:
                sample_data.data_row = encode_semantic_track(sample_data.data_row)  # (1, T)
            if sample_data.data_row_cover is not None:
                sample_data.data_row_cover = encode_semantic_track(sample_data.data_row_cover)  # (1, T)
            if sample_data.artist_tracks is not None:
                sample_data.artist_tracks = [
                    encode_semantic_track(track) for track in sample_data.artist_tracks
                ]  # (N, T)
            if sample_data.playlist_tracks is not None:
                sample_data.playlist_tracks = [
                    encode_semantic_track(track) for track in sample_data.playlist_tracks
                ]  # (N, T)
            if sample_data.overpaint_track is not None:
                sample_data.overpaint_track = encode_semantic_track(
                    sample_data.overpaint_track
                )  # (1, T)
            if sample_data.underpaint_track is not None:
                sample_data.underpaint_track = encode_semantic_track(
                    sample_data.underpaint_track
                )  # (1, T)

            sample = get_sample_from_row(self.model_cfg, self.tokenizer_fp, sample_data)
            yield sample

    def __iter__(self):
        self.sample_generator = self.sample_generator_fn(self.sample_data_dl)
        return self

    def __next__(self):
        batch = get_batch(
            self.batch_size_tokens,
            self.model_cfg,
            self.sample_generator,
            self.sampling_params,
        )
        return batch
