from copy import deepcopy
import json
import math
import os
import contextlib
import random
import traceback
from tqdm import tqdm

import numpy as np
import torch
import torch.nn.functional as F
from torch.utils.data import IterableDataset

from modules.gpt import GPTConfig
from utils.bct import Block, BlockSequence, PackedBlockSequence
from data_types import AudioType
from text_utils import get_chorus_section_offset
from block_types import BLOCK_TYPE_NAME_TO_ID
from data_types import SamplingParams, SampleData, DataBundle
from spectral_features import (
    calculate_loudness_seq,
    calculate_spectral_centroid_seq,
    calculate_spectral_complexity_seq,
)
from ordering_utils import create_interleaved_order_delay


def extract_stems_captions_keywords(data_meta, stem_type, p_dropout_tags=0.1):
    """Extract and parse captions from stems_captions for task='add' scenarios.

    Args:
        data_meta: Metadata dictionary containing stems_captions
        stem_type: String like "add Bass, Drums" indicating target stems

    Returns:
        List of keyword strings parsed from captions, or None if not applicable
    """
    if not stem_type or not stem_type.startswith("add "):
        return None

    stems_captions = data_meta.get("stems_captions", {})
    if not stems_captions:
        return None

    # Extract target stem names from "add Bass, Drums"
    target_stems = [name.strip() for name in stem_type[4:].split(",")]

    # Caption types to include (excluding musical_long_description)
    caption_types = ["musical_role", "musical_keywords", "sonic_qualities"]
    # stem_isolation hallucinates a lot, isnt accurate
    # caption_types = ["musical_role", "musical_keywords", "sonic_qualities", "stem_isolation"]

    if random.random() < p_dropout_tags:
        return target_stems

    all_keywords = [] + target_stems
    for stem_name in target_stems:
        if stem_name in stems_captions:
            stem_captions = stems_captions[stem_name]
            for caption_entry in stem_captions:
                if caption_entry.get("prompt_type") in caption_types:
                    caption_text = caption_entry.get("caption", "")
                    if caption_text:
                        # Parse comma-separated keywords from caption
                        keywords = [kw.strip() for kw in caption_text.split(",") if kw.strip()]
                        all_keywords.extend(keywords)
                        # print(caption_entry.get("prompt_type"), keywords)

    # Deduplicate keywords while preserving order
    if all_keywords:
        seen = set()
        deduped_keywords = []
        for keyword in all_keywords:
            keyword_lower = keyword.lower()
            if keyword_lower not in seen:
                seen.add(keyword_lower)
                deduped_keywords.append(keyword)
        return deduped_keywords

    return None


def extract_vocal_captions(data_meta):
    """Extract vocal captions from stems_captions for vocal conditioning.

    Args:
        data_meta: Metadata dictionary containing stems_captions

    Returns:
        List of vocal keyword strings from voice_description_keywords captions, or None if not found
    """
    stems_captions = data_meta.get("stems_captions", {})
    if not stems_captions:
        return None

    # Vocal stem keywords to look for
    # vocal_keywords = ["Vocals", "Backing_Vocals", "Vox"]

    all_vocal_keywords = []
    for stem_name, stem_captions in stems_captions.items():
        # Method 1: Check if this stem contains vocal keywords (case insensitive)
        # is_vocal_stem = any(vocal_kw.lower() in stem_name.lower() for vocal_kw in vocal_keywords)

        # Method 2: Check if this stem is exactly "Vocals"
        is_vocal_stem = stem_name == "Vocals"

        if is_vocal_stem:
            for caption_entry in stem_captions:
                if caption_entry.get("prompt_type") == "voice_description_keywords":
                    caption_text = caption_entry.get("caption", "")
                    if caption_text:
                        # Parse comma-separated keywords from caption
                        keywords = [kw.strip() for kw in caption_text.split(",") if kw.strip()]
                        all_vocal_keywords.extend(keywords)

    # Deduplicate keywords while preserving order
    if all_vocal_keywords:
        seen = set()
        deduped_keywords = []
        for keyword in all_vocal_keywords:
            keyword_lower = keyword.lower()
            if keyword_lower not in seen:
                seen.add(keyword_lower)
                deduped_keywords.append(keyword)
        return deduped_keywords

    return None


def process_text_lines(text_lines, sampling_params):
    """Process text_lines with optional timestamp augmentation."""
    if not text_lines:
        return ""

    # Apply timestamp augmentation 40% of the time if we have timestamps
    if not sampling_params.inference and random.random() < sampling_params.prob_timestamp_augment:
        from text_utils import format_timestamped_lyrics

        return format_timestamped_lyrics(
            text_lines,
            augment_probability=1.0,  # We already passed the probability check
        )
    else:
        return "\n".join([line[2] for line in text_lines])


from block_types import (
    TextBlockType,
    MMBertTextBlockType,
    HootTextBlockType,
    DittoBlockType,
    CausalSemanticBlockType,
    ContinuousSemanticBlockType,
    ArtistBlockType,
    PlaylistBlockType,
    UnderpaintBlockType,
    OverpaintBlockType,
    VoxBlockType,
    RemixBlockType,
    SampleSourceBlockType,
    MashupBlockType,
    StemBlockType,
    SampleBlockType,
    CoverBlockType,
    PrefixBlockType,
    SuffixBlockType,
    NonCausalSemanticBlockType,
    DiffusionBlockType,
    InterleavedSemanticBlockType,
    VAEArtistBlockType,
    VAEPlaylistBlockType,
    VAEUnderpaintBlockType,
    VAEOverpaintBlockType,
    VAEVoxBlockType,
    VAEStemBlockType,
    VAESampleBlockType,
    VAECoverBlockType,
    VAEPrefixBlockType,
    VAESuffixBlockType,
    CondAudioTextBlockType,
    CondAudioBlockType,
    VAECondAudioBlockType,
)
from text_utils import build_text, tokenize_batch, randomize_lyrics, load_tokenizer
from utils.helpers import print_with_time
from models import (
    load_mmbert_tokenizer_and_encoder,
    load_hoot_tokenizer_and_encoder,
    load_midi_model,
    load_ditto_model,
    load_vae_model,
    preload_semantic_models,
)

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

RELOAD_MEMMAP = False
MOCK_SUBSAMPLE_RATE = 100


def get_alphas_sigmas(t):
    """Returns the scaling factors for the clean image (alpha) and for the
    noise (sigma), given a timestep."""
    return torch.cos(t * math.pi / 2), torch.sin(t * math.pi / 2)


def make_diffusion_inputs_targets(x: torch.Tensor, t: torch.Tensor):
    alphas, sigmas = get_alphas_sigmas(t)
    alphas = alphas.unsqueeze(1)
    sigmas = sigmas.unsqueeze(1)

    noise = torch.randn_like(x)
    noised_inputs = x * alphas + noise * sigmas
    targets = noise * alphas - x * sigmas
    return noised_inputs, targets


def interleave_audio_array(audio_array, K=4, DELAY=2, PAD: int = -1):
    interleaved_order = create_interleaved_order_delay(len(audio_array), K=K, DELAY=DELAY, PAD=-1)
    interleaved_audio_array = [audio_array[i] if i != -1 else PAD for i in interleaved_order]
    return interleaved_audio_array


def make_diffusion_block(data: torch.Tensor, sampling_params: "SamplingParams"):
    """Helper function to create a diffusion block from continuous data (VAE or continuous semantic).

    Args:
        data: Continuous data tensor (VAE latents or continuous semantic embeddings)
        sampling_params: Sampling parameters containing output_distribution

    Returns:
        Block with diffusion inputs/outputs using appropriate keys based on output_distribution
    """
    t_single = sampling_params.noise_rng.draw(1)[:, 0].to(torch.bfloat16)
    t_single = torch.where(torch.rand_like(t_single) < 0.01, torch.ones_like(t_single), t_single)
    t = t_single.expand(data.shape[0])
    noised_inputs, targets = make_diffusion_inputs_targets(data, t)

    # Use appropriate input/output keys based on output distribution
    if sampling_params.output_distribution == "continuous_semantic":
        input_key = "continuous_semantic_input"
        output_key = "continuous_semantic_output"
    elif sampling_params.output_distribution == "vae":
        input_key = "vae_input"
        output_key = "vae_output"
    else:
        raise ValueError(
            f"Unsupported output_distribution for diffusion: {sampling_params.output_distribution}"
        )

    return Block(
        spec=DiffusionBlockType,
        inputs={
            input_key: noised_inputs,
            "timestep_input": t.unsqueeze(-1),
        },
        targets={output_key: targets},
    )


def build_output_block(
    output_data: DataBundle,
    cfg: GPTConfig,
    sampling_params: SamplingParams,
    skip_factor: int = 1,
) -> Block:
    """Build output block based on the output paradigm (GPT or diffusion).

    Args:
        output_data: DataBundle containing the output data
        cfg: GPT configuration
        sampling_params: Sampling parameters containing output_paradigm and output_distribution
        skip_factor: Skip factor for sampling

    Returns:
        Block: CausalSemanticBlockType, ContinuousSemanticBlockType, or DiffusionBlockType block
    """
    output_paradigm = sampling_params.output_paradigm
    output_distribution = sampling_params.output_distribution
    is_continuous_input = sampling_params.use_continuous_semantic_input or sampling_params.use_vae_input

    if output_paradigm == "diffusion":
        # Build diffusion block with continuous data (VAE or continuous semantic)
        if output_distribution in ["vae", "continuous_semantic"]:
            audio_arr = build_audio_arr(
                output_data, cfg, include_eos=True, skip_factor=skip_factor, is_discrete=False
            )
            return make_diffusion_block(audio_arr, sampling_params)
        else:
            raise ValueError(
                f"output_distribution '{output_distribution}' not supported for diffusion. Use 'vae' or 'continuous_semantic'."
            )
    elif output_paradigm == "gpt":
        # Build GPT block based on output distribution
        if output_distribution == "semantic":
            # Discrete semantic tokens
            # Decide whether to use interleaved blocks
            use_interleaved = random.random() < sampling_params.interleave_probability
            block_spec = InterleavedSemanticBlockType if use_interleaved else CausalSemanticBlockType
            shift_n = InterleavedSemanticBlockType.chunk_size if use_interleaved else 1
            cs = InterleavedSemanticBlockType.chunk_size if use_interleaved else None

            sem_audio_arr = build_audio_arr(
                output_data,
                cfg,
                include_eos=True,
                skip_factor=skip_factor,
                shift_factor=cfg.semantic_shift_factor,
                is_discrete=True,
                chunk_size=cs,
            )

            # Create targets from the original (unmasked) tokens
            targets = {
                "semantic_output": Block.shift_left(sem_audio_arr, cfg.semantic_pad_token, n=shift_n)
            }

            # Apply token dropout augmentation to inputs only (not targets)
            sem_audio_input = sem_audio_arr.clone()
            if not sampling_params.inference and random.random() < sampling_params.prob_token_dropout:
                # Determine number of tokens to drop based on token_dropout_pct
                seq_len = sem_audio_input.shape[0]
                num_tokens_to_drop = int(seq_len * sampling_params.token_dropout_pct)

                if num_tokens_to_drop > 0:
                    # Randomly select token indices to mask (without replacement)
                    mask_indices = np.random.choice(seq_len, size=num_tokens_to_drop, replace=False)
                    # Replace selected tokens with mask token in the input only
                    sem_audio_input[mask_indices] = cfg.semantic_mask_token

            # Apply codebook dropout: mask out higher-level codebooks with probability
            n_codebooks = cfg.semantic_n_codebooks
            if (
                not sampling_params.inference
                and cfg.semantic_n_codebooks > 1
                and random.random() < sampling_params.dropout_codebook_pct
            ):
                # pick a random number of codebooks to drop out
                n_codebooks = random.randint(1, cfg.semantic_n_codebooks - 1)

            # Multi-codebook: split into separate inputs per codebook
            # Use block_spec from interleave logic (can be InterleavedSemanticBlockType or CausalSemanticBlockType)
            inputs = {f"semantic_input_{n}": sem_audio_input[:, n] for n in range(n_codebooks)}
            targets = {
                f"semantic_output_{n}": targets["semantic_output"][:, n] for n in range(n_codebooks)
            }

            return Block(spec=block_spec, inputs=inputs, targets=targets)
        elif output_distribution == "continuous_semantic":
            # Continuous semantic embeddings
            continuous_audio_arr = build_audio_arr(
                output_data,
                cfg,
                include_eos=True,
                skip_factor=skip_factor,
                is_discrete=False,
            )
            return Block(
                spec=ContinuousSemanticBlockType,
                inputs={"continuous_semantic_input": continuous_audio_arr},
                targets={"continuous_semantic_output": continuous_audio_arr},
            )
        elif output_distribution == "vae":
            # VAE latents (autoregressive)
            vae_audio_arr = build_audio_arr(
                output_data,
                cfg,
                include_eos=True,
                skip_factor=skip_factor,
                is_discrete=False,
            )
            return Block(
                spec=CausalSemanticBlockType,
                inputs={"vae_input": vae_audio_arr},
                targets={"vae_output": vae_audio_arr},
            )
        else:
            raise ValueError(f"Unknown output_distribution for GPT: {output_distribution}")
    else:
        raise ValueError(f"Unknown output_paradigm: {output_paradigm}. Expected 'diffusion' or 'gpt'.")


def make_conditioning_block(
    data: torch.Tensor,
    block_type: str,
    use_vae_input: bool,
    is_continuous_semantic: bool = False,
    is_noncausal: bool = False,
    debug_text: str = None,
):
    """Helper function to create conditioning blocks.

    Args:
        data: Audio data tensor
        block_type: Type of conditioning block ('artist', 'playlist', etc.)
        use_vae_input: Whether to create VAE input block (True) or semantic block (False)
        is_continuous_semantic: Whether to use continuous semantic input (for semantic block)
        is_noncausal: Whether to use noncausal input (for semantic block)
        debug_text: Optional debug text for the block

    Returns:
        Block: Conditioning block with no targets
    """
    # Map block types to their corresponding BlockType classes
    semantic_block_types = {
        "artist": ArtistBlockType,
        "playlist": PlaylistBlockType,
        "underpaint": UnderpaintBlockType,
        "overpaint": OverpaintBlockType,
        "vox": VoxBlockType,
        "remix": RemixBlockType,
        "sample_source": SampleSourceBlockType,
        "mashup": MashupBlockType,
        "stem": StemBlockType,
        "sample": SampleBlockType,
        "cover": CoverBlockType,
        "prefix": PrefixBlockType,
        "suffix": SuffixBlockType,
        "cond_audio": CondAudioBlockType,
    }

    vae_block_types = {
        "artist": VAEArtistBlockType,
        "playlist": VAEPlaylistBlockType,
        "underpaint": VAEUnderpaintBlockType,
        "overpaint": VAEOverpaintBlockType,
        "vox": VAEVoxBlockType,
        "stem": VAEStemBlockType,
        "sample": VAESampleBlockType,
        "cover": VAECoverBlockType,
        "prefix": VAEPrefixBlockType,
        "suffix": VAESuffixBlockType,
        "cond_audio": VAECondAudioBlockType,
    }

    # Determine block type based on data characteristics
    # VAE data has last dimension 128 (or other large VAE dim), discrete semantic data has smaller last dim
    inputs = dict()
    if use_vae_input:
        block_spec = vae_block_types[block_type]
        input_key = "vae_input"
        assert data.shape[-1] == 128, data.shape

        # add noise
        if random.random() <= 0.8:
            noise = torch.randn_like(data) * random.random() * 2.0
            data += noise

        inputs = {input_key: data}

    else:
        block_spec = semantic_block_types[block_type]
        block_spec.is_causal = not is_noncausal
        if is_continuous_semantic:
            input_key = "continuous_semantic_input"
            inputs["block_type_input"] = torch.full(
                (data.shape[0],), BLOCK_TYPE_NAME_TO_ID[block_type], dtype=torch.long
            )
            inputs[input_key] = data
        else:
            # For multi-codebook data, create separate inputs for each codebook
            inputs = {}
            for i in range(data.shape[-1]):
                codebook_data = data[:, i]
                inputs[f"semantic_input_{i}"] = codebook_data

    return Block(
        spec=block_spec,
        inputs=inputs,
        debug_text=debug_text,
    )


# Instruction templates for different operation types
# Each operation type has multiple natural language variations
INSTRUCTION_TEMPLATES = {
    "stem_add": [
        "add {}",
        "generate {} based on this",
        "accompany this using {}",
        "create {} based on the audio",
        "make {}",
        "make {} according to this",
        "produce {} according to this",
        "build {} on the given music",
    ],
    "stem_extract": [
        "extract {}",
        "isolate {}",
        "separate {}",
        "pull out {}",
        "get {} only",
        "focus on {}",
        "keep only {}",
        "solo {}",
        "single out {}",
        "filter to {}",
    ],
    "stem_remove": [
        "remove {}",
        "take out {}",
        "delete {}",
        "subtract {}",
        "eliminate {}",
        "drop {}",
        "exclude {}",
        "strip out {}",
        "cut {}",
        "filter out {}",
    ],
    "cover": [
        "cover this song",
        "remake this track",
        "create a cover version",
        "do a cover of this",
        "make a version of this",
        "recreate this song",
        "perform this track",
        "interpret this music",
        "reimagine this piece",
        "reinterpret this song",
        "use this as reference",
        "inspired by this clip",
        "inspired by this sound",
    ],
    "artist": [
        "in the style of this artist",
        "following this artist's approach",
        "using this artist as reference",
        "inspired by this artist",
        "matching this artist's style",
        "emulating this artist",
        "channeling this artist",
        "following this artist's sound",
        "based on this artist",
        "in the manner of this artist",
    ],
    "playlist": [
        "similar to these songs",
        "following this playlist style",
        "matching this playlist vibe",
        "in the style of this playlist",
        "inspired by these tracks",
        "following this musical direction",
        "based on these references",
        "matching this collection",
        "similar to this selection",
        "following these examples",
    ],
    "vox": [
        "add vocals to this",
        "put vocals on this track",
        "include singing",
        "add voice to this",
        "overlay vocals",
        "include vocal parts",
        "add vocal elements",
        "put voice on this",
        "include vocal performance",
        "add vocal track",
    ],
    "underpaint": [
        "add instrumental to this vocal, output full song",
        "create music to this voice",
        "generate music using this vocal",
        "create backing track",
        "add instrumental backing into a song",
        "overlay music on vocals",
        "add musical accompaniment on the vocal",
        "create instrumental for vocals",
        "add backing music",
        "put music behind vocals",
        "create accompaniment",
    ],
    "overpaint": [
        "add vocals to this instrumental",
        "put vocals on this music",
        "overlay vocals on track",
        "add singing to this",
        "include vocals with music",
        "add vocal melody",
        "put voice on this music",
        "overlay vocal performance",
        "add vocal elements to track",
        "include singing with instrumental",
    ],
    "remix": [
        "remix this song",
        "create a remix of this track",
        "make a new version of this song",
        "flip this song into a remix",
        "remix this into a different style",
        "transform this into a remix",
        "remix this with a different beat",
        "remix this with new production",
    ],
    "sample_source": [
        "sample this song",
        "use this track as a sample source",
        "use this as the main sample",
        "sample this audio",
        "chop up this track for samples",
        "extract samples from this",
        "pull samples out of this song",
        "use this as sampling material",
        "cut samples from this song",
        "harvest samples from this track",
        "use this for sample digging",
        "slice this track for sampling",
        "create from samples of this",
    ],
    "mashup": [
        "mash up these songs",
        "create a mashup of these tracks",
        "blend these songs together",
        "mix these tracks into one song",
        "combine these songs into a mashup",
        "weave these tracks together",
        "layer these songs into a mashup",
        "fuse these songs together",
        "blend these into a single track",
        "mash these tracks together",
        "overlay these tracks together",
        "stitch these songs into a mashup",
    ],
    "sample": [
        "based on this sample",
        "use this sound in the song",
        "referencing this audio",
        "based on this snippet",
        "use this as sample",
        "directly use this sound",
    ],
    "prefix": [
        "continue from this",
        "build on this start",
        "extend this beginning",
        "continue this opening",
        "continue this",
        "follow from this intro",
        "build upon this start",
        "extend from this point",
        "continue this sequence",
        "follow this beginning",
        "build from this intro",
    ],
    "suffix": [
        "lead up to this",
        "build towards this ending",
        "create intro for this",
        "create previous section for this",
        "lead into this outro",
        "lead into it",
        "build up to this finale",
        "create buildup to this",
        "prepare for this ending",
        "lead towards this conclusion",
        "build anticipation for this",
        "create approach to this",
    ],
}


def make_text_conditioning_pair(
    data: torch.Tensor | list[torch.Tensor],
    block_type: str,
    use_vae_input: bool,
    tokenizer_fp: str = None,
    stem_type: str = None,
    is_continuous_semantic: bool = False,
    is_noncausal: bool = False,
    use_mmbert: bool = False,
):
    """Create a text/conditioning block pair that replaces indicator tokens.

    This function is parallel to make_conditioning_block() but creates structured
    block pairs where text description and conditioning content are separate
    blocks that stay together as a tuple.

    Args:
        data: Audio data tensor (semantic or VAE), or list of tensors for multiple conditioning blocks
        block_type: Type of conditioning block ('artist', 'playlist', 'stem', etc.)
        use_vae_input: Whether to create VAE input block (True) or semantic block (False)
        tokenizer_fp: Path to tokenizer for text processing
        stem_type: Optional stem type information (e.g., "add Bass, Drums")
        is_continuous_semantic: Whether to use continuous semantic input (for semantic block)
        is_noncausal: Whether to use noncausal input (for semantic block)
        use_mmbert: Whether to add mmBERT-encoded text block after text description

    Returns:
        tuple: (text_block, [mmbert_block,] content_block1, content_block2, ...)
               - text block, optional mmBERT block, followed by one or more content blocks
    """
    # Create the text description block using instruction templates

    def get_cond_audio_text(block_type: str, stem_type: str = None) -> str:
        """Generate text description for conditioning audio based on block type.

        Args:
            block_type: Type of conditioning (e.g., "artist", "stem", "vox")
            stem_type: Optional stem type information (e.g., "add Bass, Drums", "extract Vocals", "remove Drums")

        Returns:
            Text description without brackets
        """
        if block_type == "stem" and stem_type is not None:
            # For stem blocks, parse operation and instruments
            # Format: "operation instruments" (e.g., "add Bass, Drums", "extract Vocals", "remove Bass")
            parts = stem_type.split(" ", 1)
            if len(parts) == 2:
                operation, instruments = parts
                template_key = f"stem_{operation}"
                if template_key in INSTRUCTION_TEMPLATES:
                    template = random.choice(INSTRUCTION_TEMPLATES[template_key])
                    return template.format(instruments)
            # Fallback to original stem_type if format is unexpected
            return stem_type
        elif block_type in INSTRUCTION_TEMPLATES:
            # Use random template for this block type
            return random.choice(INSTRUCTION_TEMPLATES[block_type])
        else:
            # Fallback to simple block type for unknown types
            return block_type

    # Get text description and wrap in brackets
    cond_audio_text = get_cond_audio_text(block_type, stem_type)
    cond_audio_text = f"[{cond_audio_text}]"

    # Tokenize the text description
    text_tokens = tokenize_batch([cond_audio_text], tokenizer_fp=tokenizer_fp)[0]

    text_block = Block(
        spec=CondAudioTextBlockType,
        inputs={"text_input": text_tokens},
        debug_text=cond_audio_text,  # Store the original text for debugging
    )

    # Optionally add mmBERT-encoded text block
    mmbert_block = None
    if use_mmbert:
        mmbert_tokenizer, mmbert_encoder = load_mmbert_tokenizer_and_encoder()

        # Check if encoder is already on GPU to avoid redundant transfers
        current_device = next(mmbert_encoder.parameters()).device
        if current_device.type != "cuda":
            mmbert_encoder = mmbert_encoder.cuda()

        # Encode text with mmBERT
        mmbert_text_arr = mmbert_tokenizer(cond_audio_text, return_tensors="pt").input_ids.cuda()
        mmbert_text_arr = mmbert_encoder(mmbert_text_arr).last_hidden_state[0].bfloat16()

        # Truncate if exceeds max tokens
        if mmbert_text_arr.shape[0] > 1024:
            mmbert_text_arr = mmbert_text_arr[:1024]

        mmbert_block = Block(
            spec=MMBertTextBlockType,
            inputs={"mmbert_text_input": mmbert_text_arr},
            debug_text=cond_audio_text,
        )
        mmbert_block.spec.is_causal = not is_noncausal

    # Normalize data to list if single tensor
    data_list = data if isinstance(data, list) else [data]

    # Create conditioning content blocks for each data tensor
    content_blocks = []
    for data_tensor in data_list:
        content_block = make_conditioning_block(
            data_tensor,
            "cond_audio",
            use_vae_input,
            is_continuous_semantic,
            is_noncausal,
            debug_text=block_type,
        )
        content_blocks.append(content_block)

    # Return tuple with text block, optional mmBERT block, followed by all content blocks
    if mmbert_block is not None:
        return (text_block, mmbert_block, *content_blocks)
    else:
        return (text_block, *content_blocks)


def build_audio_arr_discrete(
    data_row_DT,
    cfg: GPTConfig,
    semantic_infer_token=None,
    include_eos=True,
    skip_factor=1,
    shift_factor=0,
    chunk_size=None,
    semantic_delay=2,
):
    """Builds an audio arr from discrete data row with multi-codebook support and optional interleaving.

    Args:
        data_row_DT: [D, T] tensor for D codebooks with T tokens each
        cfg: GPT configuration
        semantic_infer_token: Token to prepend to sequence
        include_eos: Whether to add EOS padding
        skip_factor: Downsample factor
        shift_factor: Temporal shift between codebooks (RVQ)
        chunk_size: If provided, applies interleaving within each codebook
        semantic_delay: DELAY parameter for interleaving

    Returns:
        [T', D] tensor where T' includes both interleaving padding and codebook delays
    """
    # Handle both numpy arrays and tensors
    if not isinstance(data_row_DT, torch.Tensor):
        data_row_DT = torch.from_numpy(data_row_DT)
    data_row_DT = data_row_DT.to(torch.int64)

    assert cfg.coarse_n_codebooks == 0
    data_row_DT = data_row_DT[:, ::skip_factor]

    d, t = data_row_DT.shape
    assert cfg.semantic_n_codebooks == d

    # INTERLEAVING: Apply to each codebook independently
    if chunk_size is not None:
        interleaved_codebooks = []
        for i in range(d):
            codebook_data = data_row_DT[i].tolist()
            interleaved = interleave_audio_array(
                codebook_data, K=chunk_size, DELAY=semantic_delay, PAD=cfg.semantic_mask_token
            )
            # Add initial padding chunk
            pad_chunk = [cfg.semantic_mask_token] * chunk_size
            interleaved = pad_chunk + interleaved
            interleaved_codebooks.append(interleaved)

        # Update data_row_DT with interleaved data
        max_len = max(len(cb) for cb in interleaved_codebooks)
        data_row_DT = torch.full((d, max_len), cfg.semantic_mask_token, dtype=torch.int64)
        for i, cb in enumerate(interleaved_codebooks):
            data_row_DT[i, : len(cb)] = torch.tensor(cb, dtype=torch.int64)
        t = max_len

        # MODIFIED SHIFT: Use semantic_delay * chunk_size for inter-codebook delays
        effective_shift = shift_factor * chunk_size
    else:
        # STANDARD RVQ SHIFT: Use shift_factor as before
        effective_shift = shift_factor

    # Apply hierarchical delay pattern across codebooks
    audio_len = (d - 1) * effective_shift + t
    x_audio_arr = torch.full((d, audio_len), cfg.semantic_pad_token, dtype=torch.int64)

    for i in range(d):
        offs = i * effective_shift
        x_audio_arr[i, offs : offs + t] = data_row_DT[i]

    # Add prefix/suffix
    if semantic_infer_token is None:
        semantic_infer_token = cfg.semantic_infer_token
    prefix = torch.tensor([semantic_infer_token] * skip_factor, dtype=torch.int64)
    suffix = torch.tensor([cfg.semantic_pad_token] * int(include_eos), dtype=torch.int64)
    prefix = prefix.unsqueeze(0).repeat(d, 1)
    suffix = suffix.unsqueeze(0).repeat(d, 1)
    x_audio_arr_TD = torch.cat([prefix, x_audio_arr, suffix], dim=-1).T

    # Pad to multiple of chunk_size if interleaving
    if chunk_size is not None:
        total_len = x_audio_arr_TD.shape[0]
        pad_amount = chunk_size - (total_len % chunk_size)
        if pad_amount < chunk_size:
            padding = torch.full((pad_amount, d), cfg.semantic_pad_token, dtype=torch.int64)
            x_audio_arr_TD = torch.cat([x_audio_arr_TD, padding], dim=0)
        assert x_audio_arr_TD.shape[0] % chunk_size == 0

    assert x_audio_arr_TD.shape[1] == d
    return x_audio_arr_TD


def build_audio_arr_continuous(data_row_DT, skip_factor=1):
    """Builds an audio arr from a continuous data row (ndim=768).

    Note: block_type embeddings, don't need semantic_infer_token
    """
    if isinstance(data_row_DT, np.ndarray):
        data_row_DT = torch.from_numpy(data_row_DT)

    if skip_factor > 1:
        data_row_DT = data_row_DT[..., ::skip_factor]

    return data_row_DT.T.bfloat16()


def build_audio_arr(
    data_row_DT,
    cfg: GPTConfig,
    semantic_infer_token=None,
    include_eos=True,
    skip_factor=1,
    shift_factor=1,
    is_discrete: bool = True,
    chunk_size=None,
    semantic_delay=2,
):
    """Builds an audio arr from a data row."""

    # Handle DataBundle from main branch
    if hasattr(data_row_DT, "quantized_data_DT"):
        if is_discrete:
            return build_audio_arr_discrete(
                data_row_DT.quantized_data_DT,
                cfg,
                semantic_infer_token,
                include_eos,
                skip_factor,
                shift_factor,
                chunk_size,
                semantic_delay,
            )
        else:
            return build_audio_arr_continuous(data_row_DT.unquantized_data_DT, skip_factor)

    # Handle direct tensor/array input
    if hasattr(data_row_DT, "shape"):
        d = data_row_DT.shape[0] if len(data_row_DT.shape) > 1 else 1
        is_vae = d == 128

        if is_vae:
            # VAE data: use continuous build function
            return build_audio_arr_continuous(data_row_DT, skip_factor)
        else:
            return build_audio_arr_discrete(
                data_row_DT,
                cfg,
                semantic_infer_token,
                include_eos,
                skip_factor,
                shift_factor,
                chunk_size,
                semantic_delay,
            )

    raise ValueError(f"Unexpected data_row_DT type: {type(data_row_DT)}")


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


def resample_sequence_to_target_length(sequence: torch.Tensor, target_length: int) -> torch.Tensor:
    """Resample a sequence to a target length using linear interpolation.

    Args:
        sequence: Input tensor of shape [seq_len, feature_dim]
        target_length: Desired output sequence length

    Returns:
        Resampled tensor of shape [target_length, feature_dim]
    """
    if sequence.shape[0] == target_length:
        return sequence

    # F.interpolate expects [batch, channels, length] format
    # Transpose to [1, feature_dim, seq_len]
    sequence_transposed = sequence.unsqueeze(0).transpose(1, 2)

    # Interpolate to target length
    resampled = F.interpolate(
        sequence_transposed, size=target_length, mode="linear", align_corners=False
    )

    # Transpose back to [target_length, feature_dim]
    return resampled.squeeze(0).transpose(0, 1)


def _get_audio_type_tag(audio_type: AudioType) -> str:
    return f"audio_type: {audio_type}"


def get_sample_from_row(
    model_cfg: GPTConfig,
    tokenizer_fp: str,
    sample_data: SampleData,
    save_debug_build_text=False,
    debug_output_dir="/tmp/debug_output",
):
    sample_data = deepcopy(sample_data)  # make sure we dont modify the original

    cfg = model_cfg
    sampling_params = sample_data.sampling_params
    use_vae_input = sampling_params.use_vae_input
    use_vae_output = sampling_params.output_distribution == "vae"
    is_diffusion = sampling_params.output_paradigm == "diffusion"
    data_meta = sample_data.data_meta

    # Randomly decide whether to use text conditioning pairs for this sample
    use_text_cond_pairs = random.random() < sampling_params.prob_text_conditioning_pairs

    ########### Figure out the output start and end tokens of the output ###########
    # Use quantized data for shape calculations
    t = len(sample_data.data_latents)
    output_start_tok = 0
    output_end_tok = t
    hook_only = (
        random.random() < sampling_params.prob_hook
        and not sampling_params.inference
        and data_meta.get("hook_offset_s") is not None
        and data_meta.get("hook_offset_s") + 30 < t / cfg.semantic_rate_hz
    )
    chorus_offset_s = get_chorus_section_offset(data_meta)
    start_from_chorus = (
        random.random() < sampling_params.prob_start_from_chorus
        and not sampling_params.inference
        and not hook_only  # Don't apply start_from_chorus if hook_only is already active
        and chorus_offset_s is not None
        and chorus_offset_s + 30 < t / cfg.semantic_rate_hz  # Ensure enough content after chorus
    )
    start_from_active = (
        random.random() < sampling_params.prob_start_from
        and not sampling_params.inference
        and not hook_only  # Don't apply start_from if hook_only is already active
        and not start_from_chorus  # Don't apply start_from if start_from_chorus is already active
        and t / cfg.semantic_rate_hz > 10  # Minimum 10s duration
    )
    # first check if we are using diffusion, this does a random crop of the song
    if is_diffusion:
        # sample a segment of the song
        duration_tok = int(
            cfg.semantic_rate_hz * max(random.random() * sampling_params.max_audio_duration_s, 3)
        )
        output_start_tok = max(int(random.random() * (t - duration_tok)), 0)
        output_end_tok = output_start_tok + duration_tok
    elif hook_only:  # remove the intro
        output_start_tok = int(data_meta["hook_offset_s"] * cfg.semantic_rate_hz)
        output_end_tok = t
    elif start_from_chorus:  # start from chorus section
        output_start_tok = int(chorus_offset_s * cfg.semantic_rate_hz)
        output_end_tok = t
    elif start_from_active:
        # Choose random start point between 0 and 95% of the song
        max_start_toks = int(t * 0.95)
        output_start_tok = random.randint(0, max_start_toks)
        output_end_tok = t
    else:
        output_start_tok = 0

    sample_data.data_latents = sample_data.data_latents[:, output_start_tok:output_end_tok]
    if use_vae_input or use_vae_output:
        sample_data.vae = sample_data.vae[:, output_start_tok:output_end_tok]
        assert len(sample_data.vae) == len(sample_data.data_latents)
    # Note: stem_track and stem_output_mix are now in the data_latents bundle and get sliced automatically

    # Slice all spectral features to match the audio segment (convert tokens to feature rate indices)
    # Handle None values when features weren't calculated due to probability settings
    if sample_data.loudness_25hz is not None:
        # Loudness at semantic rate (25Hz) - tight alignment with tokens
        loudness_rate = cfg.semantic_rate_hz
        loudness_start_idx = int(output_start_tok / cfg.semantic_rate_hz * loudness_rate)
        loudness_end_idx = int(output_end_tok / cfg.semantic_rate_hz * loudness_rate)
        sample_data.loudness_25hz = sample_data.loudness_25hz[loudness_start_idx:loudness_end_idx]

    # Loudness contour and spectral features at variable rate (0.2-1Hz) - use stored rate
    contour_rate = sample_data.contour_rate_hz if sample_data.contour_rate_hz is not None else 1.0
    contour_start_idx = int(output_start_tok / cfg.semantic_rate_hz * contour_rate)
    contour_end_idx = int(output_end_tok / cfg.semantic_rate_hz * contour_rate)

    if sample_data.loudness_seq is not None:
        sample_data.loudness_seq = sample_data.loudness_seq[contour_start_idx:contour_end_idx]
    if sample_data.spectral_centroid_seq is not None:
        sample_data.spectral_centroid_seq = sample_data.spectral_centroid_seq[
            contour_start_idx:contour_end_idx
        ]
    if sample_data.spectral_complexity_seq is not None:
        sample_data.spectral_complexity_seq = sample_data.spectral_complexity_seq[
            contour_start_idx:contour_end_idx
        ]

    if sample_data.audio_type == AudioType.SFX:
        assert len(sample_data.data_latents) >= 1  # allow sound effects to be very short
    else:
        assert len(sample_data.data_latents) > 5  # minimum 5 tokens

    data_row = sample_data.data_latents if not use_vae_input else sample_data.vae

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

    # For conditioning tracks, always use discrete versions even in diffusion mode
    # since they go through semantic processing pipeline
    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
    vox_track = sample_data.vox_track
    remix_track = sample_data.remix_track
    sample_source_track = sample_data.sample_source_track
    mashup_tracks = sample_data.mashup_tracks
    stem_track = sample_data.stem_track

    # build main audio array
    sample_tags = data_meta.get("tags", [])

    # Replace sample_tags with stems_captions keywords when task="add"
    stems_caption_keywords = extract_stems_captions_keywords(data_meta, sample_data.stem_type)
    if stems_caption_keywords is not None:
        sample_tags = stems_caption_keywords

    # Extract vocal captions for vocal conditioning
    vocal_tags = extract_vocal_captions(data_meta)

    # Extract vocal pitch range with IQR_1.5x strategy
    vocal_pitch_hz_min = None
    vocal_pitch_hz_max = None
    vocal_pitch_range = data_meta.get("vocal_pitch_range", [])
    for entry in vocal_pitch_range:
        if entry.get("strategy") == "IQR_1.5x":
            vocal_pitch_hz_min = entry.get("min_f")
            vocal_pitch_hz_max = entry.get("max_f")
            break

    sample_vocal_start_s = None
    audio_blocks = []
    sem_output_data = None
    vae_output_data = None

    is_continuous_semantic = sampling_params.use_continuous_semantic_input
    is_continuous_input = sampling_params.use_continuous_semantic_input or use_vae_input
    is_noncausal = sampling_params.use_noncausal_input

    # Calculate mmBERT usage once for entire sample for consistency
    use_mmbert_for_sample = (
        sampling_params.use_mmbert and random.random() < sampling_params.prob_use_mmbert
    )

    if (
        sampling_params.allow_infill
        and not sampling_params.inference
        and len(data_row) >= 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 = len(data_row)
        b_duration = _get_rand_int(1, t)

        if random.random() <= 0.05:
            # prepaint
            a_right_idx = 0
            a_left_idx = 0
        else:
            # infill or extend
            a_right_idx = _get_rand_int(0, t - b_duration)
            a_left_idx = _get_rand_int(0, a_right_idx)
        b_right_idx = a_right_idx + b_duration

        if random.random() <= 0.05:
            # extend
            c_right_idx = b_right_idx
        else:
            # infill or prepaint
            c_right_idx = _get_rand_int(b_right_idx, t)

        sample_text = ""
        if len(data_meta.get("text_aligned", [])) > 0:
            text_lines = data_meta["text_aligned"]
            text_left_idx = 0  # line that starts before the section
            text_right_idx = len(text_lines)  # line that starts after the section
            for idx, m_line in enumerate(text_lines):
                line_start_s, line_end_s = m_line[0], m_line[1]
                section_start_s = (output_start_tok + a_right_idx) / cfg.semantic_rate_hz
                section_end_s = (output_start_tok + b_right_idx) / cfg.semantic_rate_hz
                if line_start_s <= section_start_s:
                    text_left_idx = idx
                if line_end_s >= section_end_s:
                    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))
            selected_text_lines = text_lines[text_left_idx:text_right_idx]
            sample_text = process_text_lines(selected_text_lines, sampling_params)
        elif "text" in data_meta:
            sample_text = data_meta["text"]

        # a is left, b is middle, c is right
        # a is history, c is future, b is normal infer token
        a_audio_arr = None
        if a_left_idx < a_right_idx:
            if use_text_cond_pairs:
                # NEW MODE: Use shared text_desc_token for prefix
                a_audio_arr = build_audio_arr(
                    data_row[:, a_left_idx:a_right_idx],
                    cfg,
                    semantic_infer_token=cfg.semantic_cond_audio_token,
                    include_eos=False,
                    skip_factor=sample_skip_factor(),
                    shift_factor=cfg.semantic_shift_factor,
                    is_discrete=not is_continuous_input,
                )
            else:
                # ORIGINAL MODE: Use history token for prefix
                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(),
                    shift_factor=cfg.semantic_shift_factor,
                    is_discrete=not is_continuous_input,
                )
        c_audio_arr = None
        if b_right_idx < c_right_idx:
            if use_text_cond_pairs:
                # NEW MODE: Use shared text_desc_token for suffix
                c_audio_arr = build_audio_arr(
                    data_row[:, b_right_idx:c_right_idx],
                    cfg,
                    semantic_infer_token=cfg.semantic_cond_audio_token,
                    include_eos=False,
                    skip_factor=sample_skip_factor(),
                    shift_factor=cfg.semantic_shift_factor,
                    is_discrete=not is_continuous_input,
                )
            else:
                # ORIGINAL MODE: Use future token for suffix
                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(),
                    shift_factor=cfg.semantic_shift_factor,
                    is_discrete=not is_continuous_input,
                )
        if c_audio_arr is not None:
            if not (is_diffusion and random.random() > 0.2):  # 20% chance to add vae suffix
                if use_text_cond_pairs:
                    # NEW MODE: Use text/conditioning pairs for suffix blocks
                    suffix_pair = make_text_conditioning_pair(
                        data=c_audio_arr,
                        block_type="suffix",
                        use_vae_input=use_vae_input,
                        tokenizer_fp=tokenizer_fp,
                        is_continuous_semantic=is_continuous_semantic,
                        is_noncausal=is_noncausal,
                        use_mmbert=use_mmbert_for_sample,
                    )
                    audio_blocks.append(suffix_pair)  # Add tuple to blocks list
                else:
                    # ORIGINAL MODE: Use simple conditioning block for suffix
                    suffix_block = make_conditioning_block(
                        c_audio_arr, "suffix", use_vae_input, is_continuous_semantic, is_noncausal
                    )
                    audio_blocks.append(suffix_block)
        if a_audio_arr is not None:
            if use_text_cond_pairs:
                # NEW MODE: Use text/conditioning pairs for prefix blocks
                prefix_pair = make_text_conditioning_pair(
                    data=a_audio_arr,
                    block_type="prefix",
                    use_vae_input=use_vae_input,
                    tokenizer_fp=tokenizer_fp,
                    is_continuous_semantic=is_continuous_semantic,
                    is_noncausal=is_noncausal,
                    use_mmbert=use_mmbert_for_sample,
                )
                audio_blocks.append(prefix_pair)  # Add tuple to blocks list
            else:
                # ORIGINAL MODE: Use simple conditioning block for prefix
                prefix_block = make_conditioning_block(
                    a_audio_arr, "prefix", use_vae_input, is_continuous_semantic, is_noncausal
                )
                audio_blocks.append(prefix_block)

        sem_output_data = data_row[:, a_right_idx:b_right_idx]

        # Crop loudness_25hz for infill mode if it was calculated
        if sample_data.loudness_25hz is not None:
            sample_data.loudness_25hz = sample_data.loudness_25hz[a_right_idx:b_right_idx]

        # Crop all spectral features for infill mode (use different rates) if they were calculated
        contour_rate = sample_data.contour_rate_hz if sample_data.contour_rate_hz is not None else 1.0
        a_contour_right_idx = int(a_right_idx / cfg.semantic_rate_hz * contour_rate)
        b_contour_right_idx = int(b_right_idx / cfg.semantic_rate_hz * contour_rate)

        if sample_data.loudness_seq is not None:
            sample_data.loudness_seq = sample_data.loudness_seq[a_contour_right_idx:b_contour_right_idx]
        if sample_data.spectral_centroid_seq is not None:
            sample_data.spectral_centroid_seq = sample_data.spectral_centroid_seq[
                a_contour_right_idx:b_contour_right_idx
            ]
        if sample_data.spectral_complexity_seq is not None:
            sample_data.spectral_complexity_seq = sample_data.spectral_complexity_seq[
                a_contour_right_idx:b_contour_right_idx
            ]
        # Note: stem_track and stem_output_mix are now in the data_latents bundle and get sliced automatically

        if use_vae_output:
            vae_output_data = data_row[:, a_right_idx:c_right_idx]
        sample_duration_toks = b_right_idx - a_right_idx
        sample_duration_s = sample_duration_toks / cfg.semantic_rate_hz
    else:
        if "text_aligned" in data_meta and len(data_meta["text_aligned"]) > 0:
            sample_vocal_start_s = data_meta["text_aligned"][0][0]
        sem_output_data = data_row
        if use_vae_output:
            vae_output_data = sample_data.vae

        # Process text_lines with timestamp augmentation if available
        if "text_aligned" in data_meta and len(data_meta["text_aligned"]) > 0:
            sample_text = process_text_lines(data_meta["text_aligned"], sampling_params)
        else:
            sample_text = data_meta.get("text", "")
        sample_duration_toks = len(data_row)
        sample_duration_s = sample_duration_toks / cfg.semantic_rate_hz

    # Build output block(s) based on paradigm (GPT or diffusion)
    output_data = vae_output_data if use_vae_output else sem_output_data

    output_block = build_output_block(output_data, cfg, sampling_params)
    # REPA TARGETS
    if sample_data.stem_type is not None and sampling_params.repa_mixed_semantic:
        # Add continuous semantic encoding of the full mix as an additional target (repa)
        assert not sampling_params.allow_skip
        # pad on both sides with zeros
        mix_audio_arr = torch.tensor(sem_output_data.stem_output_mix_data_DT.T)
        mix_audio_arr = F.pad(mix_audio_arr, (0, 0, 1, 1), "constant", 0)
        # print(mix_audio_arr.shape, len(output_block))
        output_block.targets["repa_mixed_semantic_output"] = mix_audio_arr.bfloat16()

    # Add continuous semantic embeddings as an additional target (repa) - applies to all samples
    if sampling_params.repa_semantic and sem_output_data.unquantized_data_DT is not None:
        assert not sampling_params.allow_skip
        # Transpose from [D, T] to [T, D] format and pad on both sides with zeros
        continuous_audio_arr = torch.tensor(sem_output_data.unquantized_data_DT.T)  # [T, 768]
        continuous_audio_arr = F.pad(continuous_audio_arr, (0, 0, 1, 1), "constant", 0)  # [T+2, 768]
        output_block.targets["repa_semantic_output"] = continuous_audio_arr.bfloat16()

    # Add hoot encoder embeddings as an additional target (repa) - applies to all samples
    if sampling_params.repa_hoot and sem_output_data.hoot_embeddings_data_DT is not None:
        assert not sampling_params.allow_skip
        # Transpose from [D, T] to [T, D] format and pad on both sides with zeros
        hoot_arr = torch.tensor(sem_output_data.hoot_embeddings_data_DT.T)  # [T, 512]
        hoot_arr = F.pad(hoot_arr, (0, 0, 1, 1), "constant", 0)  # [T+2, 512]
        output_block.targets["repa_hoot_output"] = hoot_arr.bfloat16()

    # Add midi encoder embeddings as an additional target (repa) - applies to all samples
    if sampling_params.repa_midi and sem_output_data.midi_embeddings_data_DT is not None:
        assert not sampling_params.allow_skip
        # Transpose from [D, T] to [T, D] format and pad on both sides with zeros
        midi_arr = sem_output_data.midi_embeddings_data_DT.T  # [T, 128]
        midi_arr = F.pad(midi_arr, (0, 0, 1, 1), "constant", 0)  # [T+2, 128]
        output_block.targets["repa_midi_output"] = midi_arr.bfloat16()

    audio_blocks.append(output_block)

    if (
        sampling_params.allow_lyrics_randomize
        and not sampling_params.inference
        and random.random() < sampling_params.prob_lyrics_randomize
    ):
        sample_text = randomize_lyrics(sample_text)
        sample_tags = sample_tags + ["shuffle mode"]

    if (
        sampling_params.allow_mumble
        and not sampling_params.inference
        and len(sample_text.strip()) > 50  # ensure theres enough lyrics to mumble
        and random.random() < sampling_params.prob_mumble
    ):
        # sample a random span of text to mumble
        start_idx = random.randint(0, len(sample_text) - 1)
        end_idx = random.randint(start_idx + 1, len(sample_text))
        # half the time mumble everything
        if random.random() < 0.5:
            start_idx, end_idx = 0, len(sample_text)
        mumble_text = sample_text[start_idx:end_idx]
        from better_profanity import profanity

        is_safe = not profanity.contains_profanity(mumble_text)
        if random.random() < 0.5:  # so we dont bias unsafe, always mark half as unsafe
            is_safe = False

        if is_safe:
            sample_text = sample_text[:start_idx] + "[mumble]" + sample_text[end_idx:]
        else:
            sample_text = sample_text[:start_idx] + "[unsafe mumble]" + sample_text[end_idx:]

    if (
        sampling_params.allow_mumble
        and not sampling_params.inference
        and len(sample_text.strip()) > 50  # ensure theres enough lyrics to mumble
        and random.random() < sampling_params.prob_mumble
    ):
        sample_text = "[unsafe mumble mode]"

    if (
        sampling_params.allow_sfx
        and not sampling_params.inference
        and sample_data.audio_type == AudioType.SFX
        and random.random() < sampling_params.prob_sfx
    ):
        sample_tags = [_get_audio_type_tag(sample_data.audio_type)] + sample_tags

    critical_control_tags = []
    if sample_data.audio_type == AudioType.MUSIC:
        if random.random() < 0.05:  # dont always add music tag, we default to music
            critical_control_tags = [_get_audio_type_tag(sample_data.audio_type)]
        # stem_type is now handled in text description blocks, not control tags
    else:
        # don't want to accidentally output speech or sfx
        critical_control_tags = [_get_audio_type_tag(sample_data.audio_type)]
    if sample_data.stem_type is not None:
        if random.random() < 0.50:
            critical_control_tags.append("add stem")  # generic add stem tag
        else:
            critical_control_tags.append(sample_data.stem_type)  # "add Piano"

    text = build_text(
        sample_tags,
        sample_text,
        sample_duration_s,
        sample_duration_toks,
        sample_vocal_start_s=sample_vocal_start_s,
        hook_only=hook_only,
        start_from_chorus=start_from_chorus,
        sample_offset_toks=output_start_tok,
        inference=sampling_params.inference,
        suppress_text=sampling_params.suppress_text,
        max_possible_duration_s=cfg.block_size // cfg.semantic_rate_hz,
        critical_control_tags=critical_control_tags,
        audio_sample_start_times_s=sample_data.audio_sample_start_times_s,
        audio_sample_sources=sample_data.audio_sample_sources,
        semantic_rate_hz=cfg.semantic_rate_hz,
        loudness_25hz=sample_data.loudness_25hz,  # For activity tags at 25Hz
        loudness_seq=sample_data.loudness_seq,  # Loudness contour with randomized rate
        spectral_centroid_seq=sample_data.spectral_centroid_seq,
        spectral_complexity_seq=sample_data.spectral_complexity_seq,
        contour_rate_hz=sample_data.contour_rate_hz,  # Rate for loudness_seq and spectral features
        contour_is_warped=sample_data.contour_is_warped,  # Whether warping was applied
        vocal_tags=vocal_tags,
        vocal_pitch_hz_min=vocal_pitch_hz_min,
        vocal_pitch_hz_max=vocal_pitch_hz_max,
    )

    # Debug: save build_text output if enabled
    if save_debug_build_text:
        os.makedirs(debug_output_dir, exist_ok=True)

        # Create debug entry
        debug_entry = {
            "sample_id": data_meta.get("id", "unknown"),
            "stem_type": getattr(sample_data, "stem_type", None),
            "original_tags": data_meta.get("tags", []),
            "used_sample_tags": sample_tags,
            "stems_captions_applied": stems_caption_keywords is not None,
            "final_text": text,
            "critical_control_tags": critical_control_tags,
            "audio_sample_sources": sample_data.audio_sample_sources,
        }

        # Save to file
        debug_file = os.path.join(debug_output_dir, "build_text_debug.jsonl")
        with open(debug_file, "a", encoding="utf-8") as f:
            f.write(json.dumps(debug_entry, ensure_ascii=False) + "\n")

    # build text arr
    if not sampling_params.inference and random.random() >= 0.98:  # sometimes do unconditional
        text = ""
    x_text_arr = tokenize_batch(
        [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,
        )

    # Note: We'll add infer token and padding later, only if text_is_post is True
    # 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, len(data_row_cover))]

        if use_text_cond_pairs:
            # NEW MODE: Use shared text_desc_token + text description
            cover_audio_arr = build_audio_arr(
                data_row_cover,
                cfg,
                semantic_infer_token=cfg.semantic_cond_audio_token,
                include_eos=False,
                skip_factor=sample_skip_factor(),
                is_discrete=not is_continuous_input,
            )
            cover_pair = make_text_conditioning_pair(
                data=cover_audio_arr,
                block_type="cover",
                use_vae_input=use_vae_input,
                tokenizer_fp=tokenizer_fp,
                is_continuous_semantic=is_continuous_semantic,
                is_noncausal=is_noncausal,
                use_mmbert=use_mmbert_for_sample,
            )
            audio_blocks.insert(0, cover_pair)
        else:
            # ORIGINAL MODE: Use type-specific token, no text description
            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(),
                is_discrete=not is_continuous_input,
            )
            cover_block = make_conditioning_block(
                cover_audio_arr, "cover", use_vae_input, is_continuous_semantic, is_noncausal
            )
            audio_blocks.insert(0, cover_block)

    # Artist audio to audio
    if artist_tracks is not None:
        if use_text_cond_pairs:
            # NEW MODE: Use shared text_desc_token + text description
            # Process all artist tracks together into a single tuple
            artist_audio_arrs = [
                build_audio_arr(
                    data_row_track,
                    cfg,
                    semantic_infer_token=cfg.semantic_cond_audio_token,
                    include_eos=False,
                    skip_factor=sample_skip_factor(),
                    is_discrete=not is_continuous_input,
                )
                for data_row_track in artist_tracks
            ]
            artist_pair = make_text_conditioning_pair(
                data=artist_audio_arrs,
                block_type="artist",
                use_vae_input=use_vae_input,
                tokenizer_fp=tokenizer_fp,
                is_continuous_semantic=is_continuous_semantic,
                is_noncausal=is_noncausal,
                use_mmbert=use_mmbert_for_sample,
            )
            audio_blocks.insert(0, artist_pair)
        else:
            # ORIGINAL MODE: Use type-specific token, no text description
            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(),
                    is_discrete=not is_continuous_input,
                )
                artist_block = make_conditioning_block(
                    artist_audio_arr, "artist", use_vae_input, is_continuous_semantic, is_noncausal
                )
                audio_blocks.insert(0, artist_block)

    # Playlist audio to audio
    if playlist_tracks is not None:
        if use_text_cond_pairs:
            # NEW MODE: Process all playlist tracks together into a single tuple
            playlist_audio_arrs = [
                build_audio_arr(
                    data_row_track,
                    cfg,
                    semantic_infer_token=cfg.semantic_cond_audio_token,
                    include_eos=False,
                    skip_factor=sample_skip_factor(),
                    is_discrete=not is_continuous_input,
                )
                for data_row_track in playlist_tracks
            ]
            playlist_pair = make_text_conditioning_pair(
                data=playlist_audio_arrs,
                block_type="playlist",
                use_vae_input=use_vae_input,
                tokenizer_fp=tokenizer_fp,
                is_continuous_semantic=is_continuous_semantic,
                is_noncausal=is_noncausal,
                use_mmbert=use_mmbert_for_sample,
            )
            audio_blocks.insert(0, playlist_pair)
        else:
            # ORIGINAL MODE: Use type-specific token, no text description
            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(),
                    is_discrete=not is_continuous_input,
                )
                playlist_block = make_conditioning_block(
                    playlist_audio_arr, "playlist", use_vae_input, is_continuous_semantic, is_noncausal
                )
                audio_blocks.insert(0, playlist_block)

    # Mashup audio to audio (multiple source tracks)
    if mashup_tracks is not None:
        if use_text_cond_pairs:
            # NEW MODE: Process all mashup tracks together into a single tuple
            mashup_audio_arrs = [
                build_audio_arr(
                    data_row_track,
                    cfg,
                    semantic_infer_token=cfg.semantic_cond_audio_token,
                    include_eos=False,
                    skip_factor=sample_skip_factor(),
                    is_discrete=not is_continuous_input,
                )
                for data_row_track in mashup_tracks
            ]
            mashup_pair = make_text_conditioning_pair(
                data=mashup_audio_arrs,
                block_type="mashup",
                use_vae_input=use_vae_input,
                tokenizer_fp=tokenizer_fp,
                is_continuous_semantic=is_continuous_semantic,
                is_noncausal=is_noncausal,
                use_mmbert=use_mmbert_for_sample,
            )
            audio_blocks.insert(0, mashup_pair)
        else:
            # ORIGINAL MODE: Use shared cond audio token, no text description
            # Process all mashup tracks and insert in reverse order to maintain original order
            for data_row_track in reversed(mashup_tracks):
                mashup_audio_arr = build_audio_arr(
                    data_row_track,
                    cfg,
                    semantic_infer_token=cfg.semantic_cond_audio_token,
                    include_eos=False,
                    skip_factor=sample_skip_factor(),
                    is_discrete=not is_continuous_input,
                )
                mashup_block = make_conditioning_block(
                    mashup_audio_arr,
                    "mashup",
                    use_vae_input,
                    is_continuous_semantic,
                    is_noncausal,
                )
                audio_blocks.insert(0, mashup_block)
    # Overpaint
    if overpaint_track is not None:
        if use_text_cond_pairs:
            # NEW MODE: Use shared text_desc_token + text description
            overpaint_audio_arr = build_audio_arr(
                overpaint_track,
                cfg,
                semantic_infer_token=cfg.semantic_cond_audio_token,
                include_eos=False,
                skip_factor=sample_skip_factor(),
                is_discrete=not is_continuous_input,
            )
            overpaint_pair = make_text_conditioning_pair(
                data=overpaint_audio_arr,
                block_type="overpaint",
                use_vae_input=use_vae_input,
                tokenizer_fp=tokenizer_fp,
                is_continuous_semantic=is_continuous_semantic,
                is_noncausal=is_noncausal,
                use_mmbert=use_mmbert_for_sample,
            )
            audio_blocks.insert(0, overpaint_pair)
        else:
            # ORIGINAL MODE: Use type-specific token, no text description
            overpaint_audio_arr = build_audio_arr(
                overpaint_track,
                cfg,
                semantic_infer_token=cfg.semantic_overpaint_token,
                include_eos=False,
                skip_factor=sample_skip_factor(),
                is_discrete=not is_continuous_input,
            )
            overpaint_block = make_conditioning_block(
                overpaint_audio_arr, "overpaint", use_vae_input, is_continuous_semantic, is_noncausal
            )
            audio_blocks.insert(0, overpaint_block)

    # Underpaint
    if underpaint_track is not None:
        if use_text_cond_pairs:
            # NEW MODE: Use shared text_desc_token + text description
            underpaint_audio_arr = build_audio_arr(
                underpaint_track,
                cfg,
                semantic_infer_token=cfg.semantic_cond_audio_token,
                include_eos=False,
                skip_factor=sample_skip_factor(),
                is_discrete=not is_continuous_input,
            )
            underpaint_pair = make_text_conditioning_pair(
                data=underpaint_audio_arr,
                block_type="underpaint",
                use_vae_input=use_vae_input,
                tokenizer_fp=tokenizer_fp,
                is_continuous_semantic=is_continuous_semantic,
                is_noncausal=is_noncausal,
                use_mmbert=use_mmbert_for_sample,
            )
            audio_blocks.insert(0, underpaint_pair)
        else:
            # ORIGINAL MODE: Use type-specific token, no text description
            underpaint_audio_arr = build_audio_arr(
                underpaint_track,
                cfg,
                semantic_infer_token=cfg.semantic_underpaint_token,
                include_eos=False,
                skip_factor=sample_skip_factor(),
                is_discrete=not is_continuous_input,
            )
            underpaint_block = make_conditioning_block(
                underpaint_audio_arr, "underpaint", use_vae_input, is_continuous_semantic, is_noncausal
            )
            audio_blocks.insert(0, underpaint_block)

    # Remix conditioning
    if remix_track is not None:
        if use_text_cond_pairs:
            # NEW MODE: Use shared text_desc_token + text description
            remix_audio_arr = build_audio_arr(
                remix_track,
                cfg,
                semantic_infer_token=cfg.semantic_cond_audio_token,
                include_eos=False,
                skip_factor=sample_skip_factor(),
                is_discrete=not is_continuous_input,
            )
            remix_pair = make_text_conditioning_pair(
                data=remix_audio_arr,
                block_type="remix",
                use_vae_input=use_vae_input,
                tokenizer_fp=tokenizer_fp,
                is_continuous_semantic=is_continuous_semantic,
                is_noncausal=is_noncausal,
                use_mmbert=use_mmbert_for_sample,
            )
            audio_blocks.insert(0, remix_pair)
        else:
            # ORIGINAL MODE: Use shared cond audio token for remix
            remix_audio_arr = build_audio_arr(
                remix_track,
                cfg,
                semantic_infer_token=cfg.semantic_cond_audio_token,
                include_eos=False,
                skip_factor=sample_skip_factor(),
                is_discrete=not is_continuous_input,
            )
            remix_block = make_conditioning_block(
                remix_audio_arr, "remix", use_vae_input, is_continuous_semantic, is_noncausal
            )
            audio_blocks.insert(0, remix_block)

    # Sample source conditioning
    if sample_source_track is not None:
        if use_text_cond_pairs:
            # NEW MODE: Use shared text_desc_token + text description
            sample_source_audio_arr = build_audio_arr(
                sample_source_track,
                cfg,
                semantic_infer_token=cfg.semantic_cond_audio_token,
                include_eos=False,
                skip_factor=sample_skip_factor(),
                is_discrete=not is_continuous_input,
            )
            sample_source_pair = make_text_conditioning_pair(
                data=sample_source_audio_arr,
                block_type="sample_source",
                use_vae_input=use_vae_input,
                tokenizer_fp=tokenizer_fp,
                is_continuous_semantic=is_continuous_semantic,
                is_noncausal=is_noncausal,
                use_mmbert=use_mmbert_for_sample,
            )
            audio_blocks.insert(0, sample_source_pair)
        else:
            # ORIGINAL MODE: Use shared cond audio token for sample source
            sample_source_audio_arr = build_audio_arr(
                sample_source_track,
                cfg,
                semantic_infer_token=cfg.semantic_cond_audio_token,
                include_eos=False,
                skip_factor=sample_skip_factor(),
                is_discrete=not is_continuous_input,
            )
            sample_source_block = make_conditioning_block(
                sample_source_audio_arr,
                "sample_source",
                use_vae_input,
                is_continuous_semantic,
                is_noncausal,
            )
            audio_blocks.insert(0, sample_source_block)

    # Vox
    if vox_track is not None:
        if use_text_cond_pairs:
            # NEW MODE: Use shared text_desc_token + text description
            vox_audio_arr = build_audio_arr(
                vox_track,
                cfg,
                semantic_infer_token=cfg.semantic_cond_audio_token,
                include_eos=False,
                skip_factor=sample_skip_factor(),
                is_discrete=not is_continuous_input,
            )
            vox_pair = make_text_conditioning_pair(
                data=vox_audio_arr,
                block_type="vox",
                use_vae_input=use_vae_input,
                tokenizer_fp=tokenizer_fp,
                is_continuous_semantic=is_continuous_semantic,
                is_noncausal=is_noncausal,
                use_mmbert=use_mmbert_for_sample,
            )
            audio_blocks.insert(0, vox_pair)
        else:
            # ORIGINAL MODE: Use type-specific token, no text description
            vox_audio_arr = build_audio_arr(
                vox_track,
                cfg,
                semantic_infer_token=cfg.semantic_vox_token,
                include_eos=False,
                skip_factor=sample_skip_factor(),
                is_discrete=not is_continuous_input,
            )
            vox_block = make_conditioning_block(
                vox_audio_arr, "vox", use_vae_input, is_continuous_semantic, is_noncausal
            )
            audio_blocks.insert(0, vox_block)

    # Stem
    if stem_track is not None:
        if use_text_cond_pairs:
            # NEW MODE: Use shared text_desc_token + text description
            stem_audio_arr = build_audio_arr(
                stem_track,
                cfg,
                semantic_infer_token=cfg.semantic_cond_audio_token,
                include_eos=False,
                skip_factor=sample_skip_factor(),
                is_discrete=not is_continuous_input,
            )
            # Create text/conditioning pair for stem block with stem_type
            stem_pair = make_text_conditioning_pair(
                data=stem_audio_arr,
                block_type="stem",
                use_vae_input=use_vae_input,
                tokenizer_fp=tokenizer_fp,
                stem_type=sample_data.stem_type,
                is_continuous_semantic=is_continuous_semantic,
                is_noncausal=is_noncausal,
                use_mmbert=use_mmbert_for_sample,
            )
            audio_blocks.insert(0, stem_pair)  # Add tuple to blocks list
        else:
            # ORIGINAL MODE: Use type-specific token, no text description
            stem_audio_arr = build_audio_arr(
                stem_track,
                cfg,
                semantic_infer_token=cfg.semantic_stem_token,
                include_eos=False,
                skip_factor=sample_skip_factor(),
                is_discrete=not is_continuous_input,
            )
            stem_block = make_conditioning_block(
                stem_audio_arr, "stem", use_vae_input, is_continuous_semantic, is_noncausal
            )
            audio_blocks.insert(0, stem_block)

    # Audio samples - process all samples uniformly
    if (
        sample_data.audio_sample_tracks is not None
        and len(sample_data.audio_sample_tracks) > 0
        and not sampling_params.inference
    ):
        if use_text_cond_pairs:
            # NEW MODE: Process all audio samples together into a single tuple
            sample_audio_arrs = [
                build_audio_arr(
                    audio_sample,
                    cfg,
                    semantic_infer_token=cfg.semantic_cond_audio_token,
                    include_eos=False,
                    skip_factor=sample_skip_factor(),
                    is_discrete=not is_continuous_input,
                )
                for audio_sample in sample_data.audio_sample_tracks
            ]
            # Create text/conditioning pair for all audio samples
            sample_pair = make_text_conditioning_pair(
                data=sample_audio_arrs,
                block_type="sample",
                use_vae_input=use_vae_input,
                tokenizer_fp=tokenizer_fp,
                is_continuous_semantic=is_continuous_semantic,
                is_noncausal=is_noncausal,
                use_mmbert=use_mmbert_for_sample,
            )
            audio_blocks.insert(0, sample_pair)
        else:
            # ORIGINAL MODE: Use type-specific token, no text description
            # Process all audio samples and insert in reverse order to maintain original order
            for audio_sample in reversed(sample_data.audio_sample_tracks):
                sample_audio_arr = build_audio_arr(
                    audio_sample,
                    cfg,
                    semantic_infer_token=cfg.semantic_sample_token,
                    include_eos=False,
                    skip_factor=sample_skip_factor(),
                    is_discrete=not is_continuous_input,
                )
                sample_block = make_conditioning_block(
                    sample_audio_arr, "sample", use_vae_input, is_continuous_semantic, is_noncausal
                )
                audio_blocks.insert(0, sample_block)

    # build blocks

    # Determine if text will be generated (post) or used as conditioning (pre)
    text_is_post = (
        sampling_params.text_loss
        and not sampling_params.inference
        and random.random() < sampling_params.prob_text_loss
    )

    # Create text block with targets if text_is_post
    if text_is_post:
        # Add infer token when text is being generated
        x_text_arr = F.pad(
            x_text_arr,
            (1, 0),
            "constant",
            cfg.text_infer_token,
        )
        # shift_left will add padding at the end, so we don't need to add it manually
        text_block = Block(
            spec=TextBlockType,
            inputs={
                "text_input": x_text_arr,
            },
            targets={
                "text_output": Block.shift_left(x_text_arr, cfg.text_pad_token),
            },
        )
    else:
        text_block = Block(
            spec=TextBlockType,
            inputs={
                "text_input": x_text_arr,
            },
        )
        text_block.spec.is_causal = not is_noncausal
    text_blocks = [text_block]

    if sampling_params.use_hoot and sample_data.hoot_audio_arr is not None:
        hoot_text_block = Block(
            spec=HootTextBlockType,
            inputs={"hoot_input": sample_data.hoot_audio_arr},
        )
        hoot_text_block.spec.is_causal = not is_noncausal
        text_blocks = [hoot_text_block] + text_blocks

    if sample_data.ditto_embeddings is not None:
        # Create a block for each ditto embedding, similar to playlist tracks
        for ditto_embedding in sample_data.ditto_embeddings:
            ditto_block = Block(
                spec=DittoBlockType,
                inputs={"ditto_input": ditto_embedding},
            )
            ditto_block.spec.is_causal = not is_noncausal
            text_blocks = [ditto_block] + text_blocks

    if sampling_params.use_mmbert and random.random() < sampling_params.prob_use_mmbert:
        mmbert_tokenizer, mmbert_encoder = load_mmbert_tokenizer_and_encoder()
        mmbert_encoder = mmbert_encoder.cuda()
        mmbert_text_arr = mmbert_tokenizer(text, return_tensors="pt").input_ids.cuda()
        mmbert_text_arr = mmbert_encoder(mmbert_text_arr).last_hidden_state[0].bfloat16()
        mmbert_text_block = Block(
            spec=MMBertTextBlockType,
            inputs={"mmbert_text_input": mmbert_text_arr},
        )
        mmbert_text_block.noncausal = is_noncausal
        text_blocks = [mmbert_text_block] + text_blocks

    # text_is_post already determined above
    if text_is_post:
        blocks = audio_blocks + [text_block]
    else:
        blocks = text_blocks + audio_blocks

    # augment conditioning blocks
    if not sampling_params.inference:
        n_out = 1
        cond_blocks = [
            b for b in blocks[:-n_out] if random.random() > sampling_params.prob_dropout_blocks
        ]
        if random.random() < sampling_params.prob_shuffle_cond_blocks:
            random.shuffle(cond_blocks)
        blocks = cond_blocks + blocks[-n_out:]

    # Flatten block tuples back into a single flat list
    # Text/conditioning pairs (tuples) get expanded to separate consecutive blocks
    flattened_blocks = []
    for block in blocks:
        if isinstance(block, tuple):
            # This is a text/conditioning pair - add text block followed by all content blocks
            flattened_blocks.extend(block)
        else:
            # Regular single block
            flattened_blocks.append(block)

    block_sequence = BlockSequence(flattened_blocks)
    # print(block_sequence)
    # block_sequence.save("test_sequence.pt")

    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

        # Store average sequence length before cropping
        if len(packed_sequences) > 0:
            packed_sequences.avg_seq_length_before_crop = packed_sequences.n_tokens / len(
                packed_sequences
            )

        # 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


from suno_utils.tasks.mert_25 import encode_both as encode_mert
from suno_utils.tasks.musicfm_v3 import encode as encode_musicfm
from suno_utils.tasks.dac_vae_fixed_25hz import encode as encode_vae
from suno_utils.tasks.ditto_v2 import encode as encode_ditto


def song_to_samples(
    song_data, song_meta: dict, cfg: GPTConfig, n_samples=1, sample_duration_range=(5, 15)
):
    """
    Standalone function to convert a song into multiple samples.

    Args:
        song_data: Audio data tensor/array for the song
        song_meta: Metadata dictionary for the song
        cfg: GPT configuration object
        n_samples: Number of samples to generate from this song
        sample_duration_range: Tuple of (min_duration, max_duration) in seconds

    Returns:
        List of sample dictionaries, each containing 'data', 'meta', and 'sample_info'
    """
    samples = []

    if song_data is None or song_data.shape[-1] < cfg.semantic_rate_hz * sample_duration_range[0]:
        # Song too short to sample from
        return samples

    song_duration_s = song_data.shape[-1] / cfg.semantic_rate_hz

    for i in range(n_samples):
        # Determine sample duration
        min_dur, max_dur = sample_duration_range
        sample_dur_s = random.uniform(min_dur, min(max_dur, song_duration_s))
        sample_dur_tokens = int(sample_dur_s * cfg.semantic_rate_hz)

        # Determine sample start position
        max_start_s = song_duration_s - sample_dur_s
        if max_start_s <= 0:
            start_s = 0
        else:
            start_s = random.uniform(0, max_start_s)
        start_token = int(start_s * cfg.semantic_rate_hz)
        end_token = start_token + sample_dur_tokens

        # Extract sample data
        sample_data = song_data[:, start_token:end_token]

        # Create sample metadata
        sample_meta = song_meta.copy()
        sample_meta.update(
            {
                "sample_start_s": start_s,
                "sample_duration_s": sample_dur_s,
                "sample_end_s": start_s + sample_dur_s,
                "original_song_duration_s": song_duration_s,
                "is_sample": True,
                "parent_song_id": song_meta.get("song_id", "unknown"),
            }
        )

        # Add control tags for the sample
        if "text_aligned" in song_meta and len(song_meta["text_aligned"]) > 0:
            # Find which lyrics correspond to this sample
            sample_text_lines = []
            for line in song_meta["text_aligned"]:
                line_start = line[0]
                line_end = line[1]
                # Include lines that overlap with the sample
                if line_start < start_s + sample_dur_s and line_end > start_s:
                    sample_text_lines.append(line)
            sample_meta["text_aligned"] = sample_text_lines
            sample_meta["text"] = "\n".join([line[2] for line in sample_text_lines])

        sample_info = {
            "sample_id": f"{song_meta.get('song_id', 'unknown')}_sample_{i}",
            "sample_type": "excerpt",
            "extraction_method": "random_temporal",
        }

        samples.append({"data": sample_data, "meta": sample_meta, "sample_info": sample_info})

    return samples


def get_samples_for_song(song_idx, data, metas, cfg, sample_mapping=None):
    """
    Helper function to get all samples associated with a given song.

    Args:
        song_idx: Index of the song in the dataset
        data: Dataset audio data
        metas: Dataset metadata
        cfg: GPT configuration
        sample_mapping: Optional dict mapping song indices to their sample indices

    Returns:
        List of sample data for the song
    """
    if sample_mapping is None:
        # Generate samples on the fly
        song_data = data[song_idx] if song_idx < len(data) else None
        song_meta = metas[song_idx] if song_idx < len(metas) else {}
        return song_to_samples(song_data, song_meta, cfg)
    else:
        # Use pre-computed sample mapping
        sample_indices = sample_mapping.get(song_idx, [])
        samples = []
        for sample_idx in sample_indices:
            if sample_idx < len(data):
                sample_data = data[sample_idx]
                sample_meta = metas[sample_idx].copy()
                sample_meta["is_sample"] = True
                sample_meta["parent_song_id"] = song_idx
                samples.append(
                    {
                        "data": sample_data,
                        "meta": sample_meta,
                        "sample_info": {"sample_id": sample_idx, "sample_type": "dataset_sample"},
                    }
                )
        return samples


class BCTDataset(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,
        save_debug_build_text: bool = False,
        debug_output_dir: str = "/tmp/debug_output",
    ):
        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
        self.save_debug_build_text = save_debug_build_text
        self.debug_output_dir = debug_output_dir

    def sample_generator_fn(self, sample_data_dl):
        while True:
            try:
                sample_data = next(sample_data_dl)
                yield from self.make_sample(sample_data)
            except Exception as e:
                print(f"Error in sample_generator_fn: {e}")
                print(traceback.format_exc())

    @torch.no_grad()
    def make_sample(self, sample_data):
        def encode_semantic_track(track) -> DataBundle:
            assert isinstance(track, np.ndarray)
            preload_semantic_models(self.model_cfg.semantic_type)
            if self.model_cfg.semantic_type == "mert":
                track = torch.from_numpy(track).unsqueeze(0).contiguous()
                unquantized_data, quantized_data = encode_mert(
                    track, pad_to_chunksize=True, batch_size=48
                )
                output_data = DataBundle(
                    quantized_data_DT=quantized_data.T[: self.model_cfg.semantic_n_codebooks],
                    unquantized_data_DT=unquantized_data.T,
                )
            elif self.model_cfg.semantic_type.startswith("musicfm"):
                track = torch.from_numpy(track)
                with open(os.devnull, "w") as devnull:
                    with contextlib.redirect_stdout(devnull):
                        quantized_data = encode_musicfm(track, batch_size=48)
                output_data = DataBundle(
                    quantized_data_DT=quantized_data.T[: self.model_cfg.semantic_n_codebooks],
                )
            else:
                raise ValueError(f"Unknown semantic type: {self.model_cfg.semantic_type}")

            return output_data

        def encode_vae_track(track) -> DataBundle:
            load_vae_model()
            wav = torch.from_numpy(track.array_float).cuda()
            vae = encode_vae(wav)
            vae = torch.from_numpy(vae) * self.sampling_params.vae_scale_factor
            return DataBundle(unquantized_data_DT=vae.T)

        def encode_track(track, semantic=True) -> DataBundle:
            if semantic:
                return encode_semantic_track(track)
            else:
                return encode_vae_track(track)

        # encode all data
        def encode_ditto_track(track):
            load_ditto_model()
            audio_arr = torch.from_numpy(track).unsqueeze(0)
            # Sample multiple ditto embeddings, similar to playlist conditioning
            sr = 24_000
            duration_s = audio_arr.shape[-1] / sr
            if duration_s >= 5.0:
                # Randomly sample how many ditto embeddings to use
                n_embeddings = random.randint(
                    self.sampling_params.min_num_ditto_embeddings,
                    self.sampling_params.max_num_ditto_embeddings,
                )
                ditto_embeddings = []
                for _ in range(n_embeddings):
                    # Sample a random section from 5-240 seconds for each embedding
                    dur_s = random.uniform(5, 240)
                    dur_s = min(dur_s, duration_s)
                    start_s = random.uniform(0, duration_s - dur_s)
                    end_s = start_s + dur_s
                    audio_arr_segment = audio_arr[..., int(start_s * sr) : int(end_s * sr)]

                    ditto_embedding_np = encode_ditto([audio_arr_segment], task="self_sim")[0]
                    ditto_embedding = torch.from_numpy(ditto_embedding_np)
                    # reshape to (1, 128) and move to device with bfloat16 dtype
                    ditto_embeddings.append(
                        ditto_embedding.unsqueeze(0).to(
                            device=self.device, dtype=torch.bfloat16, non_blocking=True
                        )
                    )
                return ditto_embeddings
            else:
                return None

        audio_array = sample_data.data_row
        if audio_array.ndim > 1:
            # Handle multi-channel audio
            audio_array = audio_array.mean(axis=0)

        # Calculate spectral features with randomization and optional time-warping
        # Randomize frame rate (0.2-1Hz) and smoothing (6-9) for augmentation
        contour_rate = random.uniform(0.2, 1.0)
        contour_smoothing = random.randint(6, 9)

        # Decide whether to apply time-warping (with random warp ratio)
        apply_warp = random.random() < self.sampling_params.prob_warp_contours
        warp_ratio = random.uniform(0.0, 0.4) if apply_warp else 0.0

        # Store the contour rate and warping flag for later use
        sample_data.contour_rate_hz = contour_rate
        sample_data.contour_is_warped = apply_warp

        # Loudness at semantic rate (25Hz) for tight alignment with tokens (no warping)
        if random.random() < self.sampling_params.prob_loudness_25hz:
            sample_data.loudness_25hz = calculate_loudness_seq(
                audio_array, sample_rate=24000, target_rate=self.model_cfg.semantic_rate_hz
            )

        # Loudness contour with randomized rate, smoothing, and warping (for conditioning)
        if random.random() < self.sampling_params.prob_contour_loudness_seq:
            sample_data.loudness_seq = calculate_loudness_seq(
                audio_array,
                sample_rate=24000,
                target_rate=contour_rate,
                smoothing_kernel_size=contour_smoothing,
                normalize=True,
                apply_warp=apply_warp,
                warp_ratio=warp_ratio,
            )

        # Spectral centroid with randomized rate, smoothing, and warping
        if random.random() < self.sampling_params.prob_contour_spectral_centroid_seq:
            sample_data.spectral_centroid_seq = calculate_spectral_centroid_seq(
                audio_array,
                sample_rate=24000,
                target_rate=contour_rate,
                smoothing_kernel_size=contour_smoothing,
                normalize=True,
                loudness_seq=sample_data.loudness_seq if sample_data.loudness_seq is not None else None,
                apply_warp=apply_warp,
                warp_ratio=warp_ratio,
            )

        # Spectral complexity with randomized rate, smoothing, and warping
        if random.random() < self.sampling_params.prob_contour_spectral_complexity_seq:
            sample_data.spectral_complexity_seq = calculate_spectral_complexity_seq(
                audio_array,
                sample_rate=24000,
                target_rate=contour_rate,
                smoothing_kernel_size=contour_smoothing,
                normalize=True,
                loudness_seq=sample_data.loudness_seq if sample_data.loudness_seq is not None else None,
                apply_warp=apply_warp,
                warp_ratio=warp_ratio,
            )

        if sample_data.data_row is not None:
            sample_data.data_latents = encode_track(sample_data.data_row)
        if sample_data.data_row_cover is not None:
            sample_data.data_row_cover = encode_track(sample_data.data_row_cover)
        if sample_data.artist_tracks is not None:
            sample_data.artist_tracks = [encode_track(track) for track in sample_data.artist_tracks]
        if sample_data.ditto_track is not None:
            sample_data.ditto_embeddings = encode_ditto_track(sample_data.ditto_track)
        if sample_data.playlist_tracks is not None:
            sample_data.playlist_tracks = [encode_track(track) for track in sample_data.playlist_tracks]
        if sample_data.mashup_tracks is not None:
            sample_data.mashup_tracks = [encode_track(track) for track in sample_data.mashup_tracks]
        if sample_data.overpaint_track is not None:
            sample_data.overpaint_track = encode_track(sample_data.overpaint_track)  # (1, T)
        if sample_data.underpaint_track is not None:
            sample_data.underpaint_track = encode_track(sample_data.underpaint_track)
        if sample_data.remix_track is not None:
            sample_data.remix_track = encode_track(sample_data.remix_track)
        if sample_data.sample_source_track is not None:
            sample_data.sample_source_track = encode_track(sample_data.sample_source_track)
        if sample_data.stem_track is not None:
            sample_data.stem_track = encode_track(sample_data.stem_track)
            # Add stem_track to data_latents bundle for conditioning
            sample_data.data_latents.stem_data_DT = sample_data.stem_track.quantized_data_DT
        if sample_data.stem_output_mix is not None:
            sample_data.stem_output_mix = encode_track(sample_data.stem_output_mix)
            assert (
                sample_data.stem_output_mix.quantized_data_DT.shape
                == sample_data.data_latents.quantized_data_DT.shape
            ), f"{sample_data.stem_output_mix.quantized_data_DT.shape} != {sample_data.data_latents.quantized_data_DT.shape}"
            # Add stem_output_mix to data_latents bundle for target
            sample_data.data_latents.stem_output_mix_data_DT = (
                sample_data.stem_output_mix.unquantized_data_DT
            )
        if sample_data.audio_sample_tracks is not None:
            sample_data.audio_sample_tracks = [
                encode_track(track) if track is not None else None
                for track in sample_data.audio_sample_tracks
            ]
        if sample_data.vox_track is not None:
            if self.model_cfg.semantic_type == "mert":
                sample_data.vox_track = sample_data.vox_track[0]
            sample_data.vox_track = encode_track(sample_data.vox_track)

        # use hoot conditioning if we have raw audio. For now only hoot uses raw audio.
        if sample_data.use_hoot and sample_data.raw_audio is not None:
            from suno_utils.tasks.hoot import encode as hoot_encode

            hoot_tokenizer, hoot_encoder = load_hoot_tokenizer_and_encoder()
            logits = torch.tensor(hoot_encode([sample_data.raw_audio], return_logits=True, n_gpus=1)[0])
            logits[:, -1] = -float("inf")  # dont include pad token
            greedy_decoded = logits.argmax(dim=-1)
            # print(logits.shape, hoot_tokenizer.decode(greedy_decoded))
            sample_data.hoot_audio_arr = greedy_decoded
        else:
            sample_data.hoot_audio_arr = None

        # Extract hoot encoder embeddings for repa_hoot auxiliary loss (independent of use_hoot)
        if sample_data.use_repa_hoot:
            assert sample_data.raw_audio is not None
            from suno_utils.tasks.hoot import encode_embeddings

            hoot_tokenizer, hoot_encoder = load_hoot_tokenizer_and_encoder()

            # Extract encoder embeddings using hoot module function
            # Returns numpy array with shape [T_hoot, 512]
            encoder_embeddings = encode_embeddings(sample_data.raw_audio, n_gpus=1)

            # Convert to torch tensor for resampling
            encoder_embeddings = torch.from_numpy(encoder_embeddings).float()

            # Resample to match semantic token rate (25Hz)
            # Semantic tokens: duration_s * 25
            semantic_length = sample_data.data_latents.quantized_data_DT.shape[1]
            encoder_embeddings = resample_sequence_to_target_length(encoder_embeddings, semantic_length)
            assert encoder_embeddings.shape[0] == semantic_length

            # normalize each embedding
            encoder_embeddings = (
                encoder_embeddings - encoder_embeddings.mean(dim=0)
            ) / encoder_embeddings.std(dim=0)

            # Store in data_latents bundle as numpy array [D, T] to match other features
            sample_data.data_latents.hoot_embeddings_data_DT = (
                encoder_embeddings.transpose(0, 1).numpy()  # [512, T]
            )

        # Extract midi encoder embeddings for repa_midi auxiliary loss
        if sample_data.use_repa_midi:
            from suno_utils.tasks.midi_transcription import encode_embeddings as midi_encode_embeddings

            midi_model = load_midi_model()

            # Extract encoder embeddings from middle layer (layer 12)
            # Returns tensor with shape [T_midi, 1024]
            midi_embeddings = midi_encode_embeddings(
                midi_model,
                mert_latents=torch.from_numpy(sample_data.data_latents.unquantized_data_DT.T),
                target_layer=12,
            ).bfloat16()

            # Resample to match semantic token rate (25Hz)
            # Semantic tokens: duration_s * 25
            semantic_length = sample_data.data_latents.quantized_data_DT.shape[1]
            midi_embeddings_TD = resample_sequence_to_target_length(midi_embeddings, semantic_length)
            assert midi_embeddings_TD.shape[0] == semantic_length
            # print(midi_embeddings_TD.shape, midi_embeddings_TD.std(), midi_embeddings_TD.mean())
            # normalize each embedding
            midi_embeddings_TD = (
                midi_embeddings_TD - midi_embeddings_TD.mean(dim=0)
            ) / midi_embeddings_TD.std(dim=0)

            sample_data.data_latents.midi_embeddings_data_DT = midi_embeddings_TD.T.cpu()

        # encode vae if needed
        if (
            self.sampling_params.use_vae_input or self.sampling_params.output_distribution == "vae"
        ) and sample_data.raw_audio is not None:
            sample_data.vae = encode_vae_track(sample_data.raw_audio)

        # free raw audio memory after encoding
        if sample_data.raw_audio is not None:
            del sample_data.raw_audio

        for _ in range(self.sampling_params.n_samples_per_meta):
            sample = get_sample_from_row(
                self.model_cfg,
                self.tokenizer_fp,
                sample_data,
                self.save_debug_build_text,
                self.debug_output_dir,
            )
            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
