from dataclasses import dataclass from enum import Enum import numpy as np import torch from suno_utils.audio import Audio @dataclass class DataBundle: """ Data class to hold quantized and unquantized data. """ quantized_data_DT: np.ndarray = None unquantized_data_DT: np.ndarray = None stem_data_DT: np.ndarray = None # Stem conditioning data (for stem tasks) stem_output_mix_data_DT: np.ndarray = None # Full mix (input + output stem) for stem tasks hoot_embeddings_data_DT: np.ndarray = None # Hoot encoder embeddings (for repa_hoot auxiliary loss) midi_embeddings_data_DT: np.ndarray = None # MIDI encoder embeddings (for repa_midi auxiliary loss) def __len__(self): if self.quantized_data_DT is not None: return self.quantized_data_DT.shape[1] elif self.unquantized_data_DT is not None: return self.unquantized_data_DT.shape[1] else: raise ValueError("No data to return length of") def __getitem__(self, key): """Slices both quantized and unquantized data in sync.""" quantized = self.quantized_data_DT[key] if self.quantized_data_DT is not None else None unquantized = self.unquantized_data_DT[key] if self.unquantized_data_DT is not None else None stem = self.stem_data_DT[key] if self.stem_data_DT is not None else None stem_output_mix = ( self.stem_output_mix_data_DT[key] if self.stem_output_mix_data_DT is not None else None ) hoot_embeddings = ( self.hoot_embeddings_data_DT[key] if self.hoot_embeddings_data_DT is not None else None ) midi_embeddings = ( self.midi_embeddings_data_DT[key] if self.midi_embeddings_data_DT is not None else None ) return DataBundle( quantized_data_DT=quantized, unquantized_data_DT=unquantized, stem_data_DT=stem, stem_output_mix_data_DT=stem_output_mix, hoot_embeddings_data_DT=hoot_embeddings, midi_embeddings_data_DT=midi_embeddings, ) # ++ Define a dataclass to hold sampling parameters ++ @dataclass class SamplingParams: # Flags controlling data augmentation and task types allow_infill: bool = False allow_artist: bool = False allow_cover: bool = False allow_overpaint: bool = False allow_underpaint: bool = False allow_vox: bool = False allow_remix: bool = False allow_sample_source: bool = False allow_mashup: bool = False allow_stem: bool = False allow_sample: bool = False allow_skip: bool = False allow_playlist: bool = False allow_mumble: bool = False allow_sfx: bool = False allow_lyrics_randomize: bool = False text_loss: bool = False use_mmbert: bool = False use_vae_input: bool = False output_paradigm: str = "gpt" # "gpt" or "diffusion" output_distribution: str = "semantic" # "semantic" or "vae" use_hoot: bool = False use_ditto: bool = False use_continuous_semantic_input: bool = False use_noncausal_input: bool = False repa_semantic: bool = False # Add continuous semantic as auxiliary target repa_mixed_semantic: bool = False # Add mixed semantic as auxiliary target for stem conditioning repa_hoot: bool = False # Add hoot encoder embeddings as auxiliary target repa_midi: bool = False # Add midi encoder embeddings as auxiliary target # probs prob_infill: float = 0.5 prob_artist: float = 0.05 prob_cover: float = 0.5 prob_overpaint: float = 0.5 prob_underpaint: float = 0.5 prob_vox: float = 0.5 prob_remix: float = 0.5 prob_sample_source: float = 0.5 prob_mashup: float = 0.5 prob_stem: float = 0.5 prob_stem_output: float = 1.0 prob_sample: float = 0.1 prob_sample_from_stems: float = ( 0.95 # Probability of using stems vs full song for sample conditioning ) max_num_audio_samples: int = 10 # Maximum number of audio samples to extract per data item sample_permutation_prob: float = ( 0.3 # Probability of applying pitch/rate permutation to extracted samples ) prob_skip: float = 0.05 prob_playlist: float = 0.1 prob_hook: float = 0.5 prob_start_from: float = 0.02 prob_start_from_chorus: float = 0.02 prob_timestamp_augment: float = 0.4 prob_text_loss: float = 0.1 prob_dropout_blocks: float = 0.02 prob_mumble: float = 0.01 prob_sfx: float = 1.0 prob_lyrics_randomize: float = 0.01 prob_shuffle_cond_blocks: float = 0.5 prob_use_hoot: float = 0.02 prob_use_mmbert: float = 0.8 prob_use_ditto: float = 0.02 prob_repa_hoot: float = 0.1 # Probability of extracting hoot embeddings for repa auxiliary loss prob_repa_midi: float = 0.1 # Probability of extracting midi embeddings for repa auxiliary loss min_num_ditto_embeddings: int = 1 # Minimum number of ditto embeddings to use max_num_ditto_embeddings: int = 5 # Maximum number of ditto embeddings to use prob_text_conditioning_pairs: float = ( 0.5 # Probability of using text conditioning pairs vs original blocks ) prob_warp_contours: float = 0.0 # Probability of applying time-warping to contour features prob_token_dropout: float = 0.0 # Probability of applying token dropout to discrete output blocks token_dropout_pct: float = ( 0.5 # Percentage of tokens to mask when token dropout is applied (0.0-1.0) ) dropout_codebook_pct: float = ( 0.0 # Percentage probability of masking out higher-level semantic codebooks (0.0-1.0) ) # Contour feature calculation probabilities (to reduce computation when not needed) prob_loudness_25hz: float = 0.5 # Probability of calculating loudness_25hz (for activity tags) prob_contour_loudness_seq: float = 0.1 # Probability of calculating loudness contour sequence prob_contour_spectral_centroid_seq: float = 0.05 # Probability of calculating spectral centroid prob_contour_spectral_complexity_seq: float = 0.05 # Probability of calculating spectral complexity # Flags controlling sampling mode inference: bool = False mock_data: bool = False suppress_text: bool = False dataset_idx: int | None = None # Flags/params controlling batching and padding pack: bool = False mask_padding: bool = False min_text_offs: int | None = None interleave_probability: float = 0.0 max_mmbert_tokens: int = 1024 # cap memory usage noise_rng: torch.quasirandom.SobolEngine | None = None vae_scale_factor: float = 0.4 max_audio_duration_s: float = 60 * 30 n_samples_per_meta: int = 1 # Add any other related parameters that are frequently passed together class AudioType(str, Enum): MUSIC = "music" SPEECH = "speech" SFX = "sfx" @dataclass class SampleData: data_row: np.ndarray # Raw audio data (before encoding) data_meta: dict sampling_params: "SamplingParams" data_latents: DataBundle | None = None # Encoded data (quantized + unquantized) data_row_cover: DataBundle | None = None artist_tracks: list[DataBundle] | None = None ditto_track: DataBundle | None = None playlist_tracks: list[DataBundle] | None = None overpaint_track: DataBundle | None = None underpaint_track: DataBundle | None = None vox_track: DataBundle | None = None remix_track: DataBundle | None = None sample_source_track: DataBundle | None = None mashup_tracks: list[DataBundle] | None = None stem_track: DataBundle | None = None stem_output_mix: np.ndarray | None = None # The full mix (input + new stem) for output audio_sample_tracks: list[DataBundle] | None = None # Multiple audio samples for conditioning audio_sample_start_times_s: list[float] | None = None # Audio sample timings for text conditioning audio_sample_sources: list[str] | None = ( None # Source types for each audio sample ("vocal", "drum", "full_mix", etc.") ) hoot_audio_arr: np.ndarray | None = None ditto_embeddings: list[torch.Tensor] | None = None # Multiple ditto embeddings per sample raw_audio: Audio | None = None vae: DataBundle | None = None idx: int | None = None stem_type: str | None = None loudness_25hz: np.ndarray | None = None # Loudness at 25Hz for token alignment loudness_seq: np.ndarray | None = None # Loudness contour with randomized rate/smoothing/warping spectral_centroid_seq: np.ndarray | None = None spectral_complexity_seq: np.ndarray | None = None contour_rate_hz: float | None = ( None # Rate used for loudness_seq and spectral contour features (randomized 0.2-1Hz) ) contour_is_warped: bool = False # Whether time-warping was applied to contour features audio_type: AudioType = AudioType.MUSIC # signals passed from audioloader to data_utils.py use_repa_hoot: bool = False use_repa_midi: bool = False use_hoot: bool = False