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