import json import os import random from collections import defaultdict import traceback import copy import numpy as np import orjson import re import struct import webrtcvad from scipy.ndimage.morphology import binary_dilation from scipy.signal import butter, lfilter import soxr import torch from suno_utils.audio import Audio from torch.utils.data import IterableDataset from tqdm import tqdm from data_types import SampleData, SamplingParams, AudioType from modules.gpt import GPTConfig, GPTTrainConfig from oracle_dataset import get_sample_oracle_file_segment from sample_extraction import create_sample_from_stems, create_sample_from_full_song, apply_audio_effects from text_utils import load_tokenizer MOCK_SUBSAMPLE_RATE = 100 class JSONLMemmap: def __init__(self, path, verbose=False): self.path = path self._index = self.get_line_index_map(verbose=verbose) def get_line_index_map(self, verbose=False): line_positions = [0] with open(self.path, "rb") as f: for line in tqdm(f, disable=not verbose, desc="Making memmap index"): line_positions.append(len(line) + line_positions[-1]) return np.array(line_positions[:-1]) def get_line_from_index(self, index: int): with open(self.path, "rb") as f: f.seek(self._index[index]) return json.loads(f.readline()) def __getitem__(self, index: int): return self.get_line_from_index(index) def __len__(self): return len(self._index) def __iter__(self): for i in range(len(self)): yield self[i] def _resample_to_mert(arr): assert arr.ndim == 2 assert arr.shape[0] == 2 out_arr = soxr.resample(arr.T, 48_000, 24_000).mean(axis=1).astype(np.float32) return out_arr def _resample_to_musicfm(arr): assert arr.ndim == 2 assert arr.shape[0] == 2 out_arr = soxr.resample(arr.T, 48_000, 16_000).astype(np.float32).T # (2, T) return out_arr def trim_long_silences(wav_stereo, sampling_rate): """ Ensures that segments without voice in the waveform remain no longer than a threshold determined by the VAD parameters in params.py. :param wav: the raw waveform as a numpy array of floats :return: the same waveform with silences trimmed away (length <= original wav length) """ vad_moving_average_width = 8 vad_max_silence_length = 6 vad_window_length = 30 int16_max = (2**15) - 1 # stereo to mono wav_mono = wav_stereo.mean(axis=0) # Compute the voice detection window size samples_per_window = (vad_window_length * sampling_rate) // 1000 # Trim the end of the audio to have a multiple of the window size wav_mono = wav_mono[: len(wav_mono) - (len(wav_mono) % samples_per_window)] wav_stereo = wav_stereo[:, : len(wav_mono) - (len(wav_mono) % samples_per_window)] # Convert the float waveform to 16-bit mono PCM pcm_wave = struct.pack("%dh" % len(wav_mono), *(np.round(wav_mono * int16_max)).astype(np.int16)) # Perform voice activation detection voice_flags = [] vad = webrtcvad.Vad(mode=3) for window_start in range(0, len(wav_mono), samples_per_window): window_end = window_start + samples_per_window voice_flags.append( vad.is_speech(pcm_wave[window_start * 2 : window_end * 2], sample_rate=sampling_rate) ) voice_flags = np.array(voice_flags) # Smooth the voice detection with a moving average def moving_average(array, width): array_padded = np.concatenate((np.zeros((width - 1) // 2), array, np.zeros(width // 2))) ret = np.cumsum(array_padded, dtype=float) ret[width:] = ret[width:] - ret[:-width] return ret[width - 1 :] / width audio_mask = moving_average(voice_flags, vad_moving_average_width) audio_mask = np.round(audio_mask).astype(bool) # Dilate the voiced regions audio_mask = binary_dilation(audio_mask, np.ones(vad_max_silence_length + 1)) audio_mask = np.repeat(audio_mask, samples_per_window) return wav_stereo[:, audio_mask == True] def bandpass_filter(data, sr, lowcut=300.0, highcut=3400.0, order=5): nyq = 0.5 * sr low = lowcut / nyq high = highcut / nyq b, a = butter(order, [low, high], btype="band") # Handle stereo input if len(data.shape) == 2: # stereo return np.array([lfilter(b, a, channel) for channel in data]) else: # mono return lfilter(b, a, data) def add_room_noise(wave, sr, snr_db=None): """Overlay white noise at a given SNR (dB) per *mix* (not per channel).""" if snr_db is None: snr_db = random.uniform(10.0, 30.0) noise = np.random.normal(0.0, 1.0, size=wave.shape) noise *= _snr_scale(wave, noise, snr_db) return np.clip(wave + noise, -1.0, 1.0) def _soft_clip(x, alpha=2.0): return np.tanh(alpha * x) / np.tanh(alpha) def _snr_scale(clean, noise, snr_db): power_clean = np.mean(clean**2) power_noise = np.mean(noise**2) + 1e-12 target_noise_power = power_clean / (10.0 ** (snr_db / 10.0)) return np.sqrt(target_noise_power / power_noise) def soft_clipping(wave, alpha=None): if alpha is None: alpha = random.uniform(1.5, 3.0) return _soft_clip(wave, alpha) def heavy_compression(wave, threshold_db=-18.0, ratio=None): if ratio is None: ratio = random.uniform(4.0, 10.0) eps = 1e-12 mag = np.abs(wave) + eps db = 20.0 * np.log10(mag) over = db - threshold_db gain_db = np.where(over > 0.0, -over * (1.0 - 1.0 / ratio), 0.0) gain_lin = 10.0 ** (gain_db / 20.0) return wave * gain_lin def loudness_wobble(wave, sr, depth_db=3.0, lfo_hz=None): if lfo_hz is None: lfo_hz = random.uniform(0.5, 3.0) t = np.arange(wave.shape[0]) / sr lfo = (1.0 + (10.0 ** (depth_db / 20.0) - 1.0) * np.sin(2 * np.pi * lfo_hz * t)) / ( 10.0 ** (depth_db / 20.0) ) return np.clip(wave * lfo[:, None] if wave.ndim == 2 else wave * lfo, -1.0, 1.0) def random_amp(audio, min_amp=0.3, max_amp=1.0): return audio * np.random.uniform(min_amp, max_amp) def interleave_pitch_shift(y_stereo, sr, semitone_range=(-0.6, 0.6)): duration = int(y_stereo.shape[1] / sr) hopsize = random.sample([0.5, 1], 1)[0] indices = [int(sr * hopsize * i) for i in range(int((sr * duration) // (sr * hopsize) + 1))] y_stereo = torch.tensor(y_stereo.astype(np.float32)) for i in range(len(indices) - 1): start = indices[i] end = min(indices[i + 1], y_stereo.shape[1]) part = y_stereo[:, start : end + 1] if random.random() < 0.1: n_semitone = random.uniform(*semitone_range) shifted = apply_audio_effects(part, sr, pitch_semitones=n_semitone, rate_factor=1.0) shifted = shifted[:, : end - start] if (end - start) != shifted.shape[1]: shifted = torch.nn.functional.pad(shifted, (0, int(end - start - shifted.shape[1]))) y_stereo[:, start:end] = shifted return y_stereo.numpy() class AudioLoaderDataset(IterableDataset): """Gets SampleData objects from the oracle dataset, which contain raw wav data""" def __init__( self, model_cfg: GPTConfig, train_cfg: GPTTrainConfig, metas_path: str, batch_size_tokens: int, tokenizer_fp: str, device: str, split="train", info_path: str | None = None, dataset_idx=None, sampling_params: SamplingParams | None = None, stem_active_sections_weight: float = 4, ): self.model_cfg = model_cfg self.train_cfg = train_cfg self.metas_path = metas_path self.info_path = info_path self.batch_size_tokens = batch_size_tokens self.tokenizer_fp = tokenizer_fp self.device = device self.dataset_idx = dataset_idx self.split = split self.sampling_params = sampling_params self.stem_active_sections_weight = stem_active_sections_weight assert model_cfg.semantic_type.startswith("mert") or model_cfg.semantic_type.startswith( "musicfm" ) if model_cfg.semantic_type.startswith("mert"): self.resample_fn = _resample_to_mert self.semantic_sample_rate = 24_000 self.semantic_n_channels = 1 elif model_cfg.semantic_type.startswith("musicfm"): self.resample_fn = _resample_to_musicfm self.semantic_sample_rate = 16_000 self.semantic_n_channels = 2 if split == "val": assert self.sampling_params.inference # load metas as a memmap to save memory self.metas = JSONLMemmap(self.metas_path, verbose=False) if self.info_path is not None: print(f"Filtering metas with {self.info_path}") with open(self.info_path, "r") as f: self.info = json.load(f) self.valid_ids = {} for key, val in self.info.items(): if isinstance(val, list): self.valid_ids[key] = set(val) print(f"{key}: {len(val):,}") else: self.valid_ids = None # make essential dicts self.weights = [] self.id_to_index = {} self.artist_to_ids = defaultdict(list) self.playlist_to_ids = defaultdict(list) self.artist_to_vox_ids = defaultdict(list) meta_count = 0 with open(self.metas_path, "r") as f: for i, line in tqdm( enumerate(f), disable=True, total=len(self.metas), desc="Loading metas", mininterval=10, ): meta = orjson.loads(line) if self.valid_ids is not None and meta["id"] not in self.valid_ids["audio_ids"]: weight = 0 # set weight to 0 if not in valid ids else: meta_count += 1 # count valid ids weight = meta.get("weight", 0) if "stem_active_sections" in meta: weight *= self.stem_active_sections_weight self.weights.append(weight) self.id_to_index[meta["id"]] = i for artist_id in meta.get("artist_ids", []): self.artist_to_ids[artist_id].append(meta["id"]) for playlist_id in meta.get("playlist_ids", []): self.playlist_to_ids[playlist_id].append(meta["id"]) if "underpaint_id" in meta and meta.get("artist_ids"): self.artist_to_vox_ids[meta["artist_ids"][0]].append(meta["id"]) if self.valid_ids is not None: prev_len = len(self.metas) filtered_pct = meta_count / prev_len print(f"Filtered to {meta_count:,} metas from {prev_len:,} ({filtered_pct:.2%})") self.tokenizer = load_tokenizer(self.tokenizer_fp) self.random_cache = [] def __iter__(self): return self def _load_audio_for_semantic( self, local_filepath, s3_filepath=None, expected_duration_s=180, start_s=0, max_duration_s=60 * 30, mock=False, load_for_semantic=True, ): model_max_duration_s = self.model_cfg.block_size / self.model_cfg.semantic_rate_hz max_duration_s = min(max_duration_s, model_max_duration_s) try: if mock: use_duration_s = min(max_duration_s, expected_duration_s) # subsample to avoid ipc limit audio_arr = Audio.from_silence( use_duration_s, self.semantic_sample_rate, n_channels=self.semantic_n_channels ).array_float[..., ::MOCK_SUBSAMPLE_RATE] else: audio = get_sample_oracle_file_segment( local_filepath=local_filepath, s3_filepath=s3_filepath, start_s=start_s, max_duration_s=max_duration_s, ) assert audio.sample_rate == 48_000 assert audio.n_channels == 2 if load_for_semantic: audio_arr = self.resample_fn(audio.array_float) else: return audio except Exception as e: print(traceback.format_exc()) print( f"Host {os.environ['HOSTNAME']} Error loading audio for {s3_filepath} " f"(local_fp: {local_filepath}, start_s: {start_s}, max_duration_s: {max_duration_s}): {e}" ) return audio_arr, audio.array_float def _load_cover_audio(self, main_meta, mock=False, load_for_semantic=True): covers = main_meta.get("cover_ids", []) if self.valid_ids is not None: covers = [id for id in covers if id in self.valid_ids] if len(covers) > 0: cover_id = random.choice(covers) cover_meta = self.metas[self.id_to_index[cover_id]] data_row_cover, _ = self._load_audio_for_semantic( cover_meta["local_filepath"], cover_meta["s3_filepath"], cover_meta["duration_s"], mock=mock, load_for_semantic=load_for_semantic, ) else: data_row_cover = None return data_row_cover def _load_artist_audio(self, main_meta, mock=False, load_for_semantic=True): artists = main_meta.get("artist_ids", []) if len(artists) == 0: return None artist_id = random.choice(artists) artist_song_ids = self.artist_to_ids[artist_id] # remove main_meta id from artist_song_ids artist_song_ids = [id for id in artist_song_ids if id != main_meta["id"]] if len(artist_song_ids) == 0: return None # sample multiple segments n_segments = random.randint(1, 10) min_segment_len = 5 max_segment_len = 120 artist_audio_arrs = [] for _ in range(n_segments): track_id = random.choice(artist_song_ids) track_meta = self.metas[self.id_to_index[track_id]] dur_s = random.random() * max_segment_len dur_s = min(max(min_segment_len, dur_s), track_meta["duration_s"]) start_s = random.random() * (track_meta["duration_s"] - dur_s) assert start_s >= 0 data_row_artist, _ = self._load_audio_for_semantic( track_meta["local_filepath"], track_meta["s3_filepath"], track_meta["duration_s"], start_s=start_s, max_duration_s=dur_s, mock=mock, load_for_semantic=load_for_semantic, ) artist_audio_arrs.append(data_row_artist) return artist_audio_arrs def _load_playlist_audio(self, main_meta, mock=False, load_for_semantic=True): playlists = main_meta.get("playlist_ids", []) if len(playlists) == 0: return None playlist_id = random.choice(playlists) playlist_song_ids = self.playlist_to_ids[playlist_id] # remove main_meta id from playlist_song_ids playlist_song_ids = [id for id in playlist_song_ids if id != main_meta["id"]] if len(playlist_song_ids) == 0: return None playlist_audio_arrs = [] n_segments = random.randint(1, 10) min_segment_len = 5 max_segment_len = 120 for _ in range(n_segments): track_id = random.choice(playlist_song_ids) track_meta = self.metas[self.id_to_index[track_id]] dur_s = random.random() * max_segment_len dur_s = min(max(min_segment_len, dur_s), track_meta["duration_s"]) start_s = random.random() * (track_meta["duration_s"] - dur_s) assert start_s >= 0 data_row_playlist, _ = self._load_audio_for_semantic( track_meta["local_filepath"], track_meta["s3_filepath"], track_meta["duration_s"], start_s=start_s, max_duration_s=dur_s, mock=mock, load_for_semantic=load_for_semantic, ) playlist_audio_arrs.append(data_row_playlist) return playlist_audio_arrs def _load_overpaint_audio(self, main_meta, mock=False, load_for_semantic=True): overpaint_id = main_meta.get("overpaint_id", None) if overpaint_id is None: return None overpaint_meta = self.metas[self.id_to_index[overpaint_id]] data_row_overpaint, _ = self._load_audio_for_semantic( overpaint_meta["local_filepath"], overpaint_meta["s3_filepath"], overpaint_meta["duration_s"], mock=mock, load_for_semantic=load_for_semantic, ) return data_row_overpaint def _load_underpaint_audio(self, main_meta, mock=False, load_for_semantic=True): underpaint_id = main_meta.get("underpaint_id", None) if underpaint_id is None: return None underpaint_meta = self.metas[self.id_to_index[underpaint_id]] data_row_underpaint, _ = self._load_audio_for_semantic( underpaint_meta["local_filepath"], underpaint_meta["s3_filepath"], underpaint_meta["duration_s"], mock=mock, load_for_semantic=load_for_semantic, ) return data_row_underpaint def _load_sample_source_audio(self, main_meta, mock=False, load_for_semantic=True): sample_source_id = main_meta.get("sample_source_id", None) if sample_source_id is None: return None sample_source_meta = self.metas[self.id_to_index[sample_source_id]] data_row_sample_source, _ = self._load_audio_for_semantic( sample_source_meta["local_filepath"], sample_source_meta["s3_filepath"], sample_source_meta["duration_s"], mock=mock, load_for_semantic=load_for_semantic, ) return data_row_sample_source def _load_remix_source_audio(self, main_meta, mock=False, load_for_semantic=True): remix_source_id = main_meta.get("remix_source_id", None) if remix_source_id is None: return None remix_source_meta = self.metas[self.id_to_index[remix_source_id]] data_row_remix_source, _ = self._load_audio_for_semantic( remix_source_meta["local_filepath"], remix_source_meta["s3_filepath"], remix_source_meta["duration_s"], mock=mock, load_for_semantic=load_for_semantic, ) return data_row_remix_source def _load_mashup_audio(self, main_meta, mock=False, load_for_semantic=True): mashup_song_ids = main_meta.get("mashup_source_ids", []) if len(mashup_song_ids) == 0: return None # select up to 4 unique mashup sources num_to_select = random.randint(1, min(4, len(mashup_song_ids))) mashup_song_ids = random.sample(mashup_song_ids, num_to_select) mashup_audio_arrs = [] min_segment_len = 60 max_segment_len = 180 for track_id in mashup_song_ids: track_meta = self.metas[self.id_to_index[track_id]] dur_s = random.random() * max_segment_len dur_s = min(max(min_segment_len, dur_s), track_meta["duration_s"]) start_s = random.random() * (track_meta["duration_s"] - dur_s) assert start_s >= 0 data_row_mashup, _ = self._load_audio_for_semantic( track_meta["local_filepath"], track_meta["s3_filepath"], track_meta["duration_s"], start_s=start_s, max_duration_s=dur_s, mock=mock, load_for_semantic=load_for_semantic, ) mashup_audio_arrs.append(data_row_mashup) if len(mashup_audio_arrs) == 0: return None return mashup_audio_arrs def _crop_and_patch( self, ref_audio, sample_rate, n_channels, cond_length_s=30.0, patch_length_s=0.5 ): patch_length = int(sample_rate * patch_length_s) num_patch = int(sample_rate * cond_length_s / patch_length) dtype = ref_audio.dtype # Ensure ref_audio is 2D: (n_channels, T) or (T,) → (T, n_channels) if ref_audio.ndim == 1: ref_audio = ref_audio[:, np.newaxis] # (T,) → (T, 1) else: ref_audio = ref_audio.T # (n_channels, T) → (T, n_channels) out_patch = np.zeros((int(sample_rate * cond_length_s), n_channels), dtype=dtype) # crop input duration to be 5s to 30s if (self.split == "train") and (len(ref_audio) > 5 * sample_rate): patch_duration = min(len(ref_audio), random.randint(5 * sample_rate, len(ref_audio))) start_ix = random.randint(0, len(ref_audio) - patch_duration) ref_audio = ref_audio[start_ix : start_ix + patch_duration] if len(ref_audio) < patch_length: if ref_audio.ndim == 1: ref_audio = np.pad(ref_audio, (0, patch_length - len(ref_audio))) elif ref_audio.ndim == 2: ref_audio = np.pad(ref_audio, ((0, patch_length - len(ref_audio)), (0, 0))) if self.split == "train": for i in range(num_patch): start_ix = random.randint(0, max(0, len(ref_audio) - patch_length)) out_patch[i * patch_length : (i + 1) * patch_length] = ref_audio[ start_ix : start_ix + patch_length ] return out_patch.T else: if len(ref_audio) < int(sample_rate * cond_length_s): ref_audio = np.pad(ref_audio, (0, int(sample_rate * cond_length_s) - len(ref_audio))) return ref_audio[max(0, len(ref_audio) - int(sample_rate * cond_length_s)) :].T def _load_vox_audio(self, main_meta, mock=False, load_for_semantic=True): mix_id = main_meta["id"] # when it has artist_ids if main_meta.get("artist_ids", None) is not None: artist_id = main_meta["artist_ids"][0] vox_ids = self.artist_to_vox_ids[artist_id] vox_ids = [id for id in vox_ids if id != mix_id] if len(vox_ids) == 0: vox_ids = self.artist_to_vox_ids[artist_id] if len(vox_ids) == 0: return None if self.split == "train": vox_id = random.choice(vox_ids) else: vox_id = vox_ids[0] vox_meta = self.metas[self.id_to_index[vox_id]] # load vox data_row_vox = self._load_audio_for_semantic( vox_meta["local_filepath"].replace(".opus", "_vocals.opus"), vox_meta["s3_filepath"].replace(".opus", "_vocals.opus"), vox_meta["duration_s"], mock=mock, load_for_semantic=False, ) # when it is cover data without artist_ids elif ( main_meta.get("stems", None) is not None and main_meta.get("stems").get("Vocals", None) is not None ): vox_meta = main_meta data_row_vox = self._load_audio_for_semantic( vox_meta.get("stems").get("Vocals"), vox_meta["s3_filepath"].replace( ".opus", "_vocals.opus" ), # <- this is a placeholder (fake path) vox_meta["duration_s"], mock=mock, load_for_semantic=False, ) else: return None # remove silence trimmed_data_row_vox = trim_long_silences(data_row_vox.array_float, 48000) if trimmed_data_row_vox.shape[-1] > 0: data_row_vox = trimmed_data_row_vox else: data_row_vox = data_row_vox.array_float # add noise prob = 0.5 if self.split == "train": if random.random() < prob: data_row_vox = bandpass_filter(data_row_vox, 48000) if random.random() < prob: data_row_vox = add_room_noise(data_row_vox, 48000) if random.random() < prob: data_row_vox = soft_clipping(data_row_vox) if random.random() < prob: data_row_vox = heavy_compression(data_row_vox) if random.random() < prob: data_row_vox = loudness_wobble(data_row_vox, 48000) if random.random() < prob: data_row_vox = random_amp(data_row_vox) if random.random() < prob: data_row_vox = interleave_pitch_shift(data_row_vox, 48000) data_row_vox = data_row_vox.astype(np.float32) # prevent signal overflow if np.max(np.abs(data_row_vox)) > 1.0: data_row_vox = data_row_vox / np.max(np.abs(data_row_vox)) # resample data_row_vox = self.resample_fn(data_row_vox) data_row_vox = self._crop_and_patch( data_row_vox, sample_rate=self.semantic_sample_rate, n_channels=self.semantic_n_channels, patch_length_s=3.0, ) return data_row_vox def _load_and_mix_stems( self, stem_dict, stem_keys, mock=False, load_for_semantic=True, start_s=0, end_s=None ): """Load and mix multiple stems together.""" stem_audios = [ self._load_audio_for_semantic( stem_dict[k], mock=mock, load_for_semantic=False, ).get_segment(start_s, end_s) for k in stem_keys ] try: mixed_audio = Audio.sum(stem_audios) except Exception as e: print(f"Error mixing stems: {e}") print(f"Stem keys: {stem_keys}") print(f"Stem dict: {stem_dict}") print(f"Start s: {start_s}") print(f"End s: {end_s}") raise e if load_for_semantic: return self.resample_fn(mixed_audio.array_float) else: return mixed_audio.array_float def _load_stem_audio(self, main_meta, task="add", mock=False, load_for_semantic=True): """Load and process stem audio data for training. Randomly splits available stems into input and output subsets, then loads and mixes the stems in each subset. Returns the stem type description and the mixed audio arrays. """ stem_dict = copy.deepcopy(main_meta.get("stems", {})) # filter out keys that contain " and " or "&". multi instrument isnt good. stem_dict = { k: v for k, v in stem_dict.items() if " and " not in k.lower() and "&" not in k.lower() } if len(stem_dict) < 2: return None, None, None, None keys = list(stem_dict.keys()) # Clean up keys by removing trailing single characters/numbers cleaned_keys = [] for key in keys: # Remove trailing single character or number (e.g., "synths 3" -> "synths") cleaned_key = re.sub(r"\s+[a-zA-Z0-9]$", "", key).strip() if cleaned_key: # Only add non-empty keys cleaned_keys.append(cleaned_key) stem_dict[cleaned_key] = stem_dict.pop(key) # Use cleaned keys and remove duplicates while preserving order seen = set() unique_keys = [] for key in cleaned_keys: if key not in seen: seen.add(key) unique_keys.append(key) keys = unique_keys random.shuffle(keys) # Check if we have enough keys after deduplication if len(keys) < 2: return None, None, None, None # Split keys into input and output subsets split_idx = random.randint(1, len(keys) - 1) input_subset = keys[:split_idx] # Randomly choose subsets of each input_subset = random.sample(input_subset, random.randint(1, len(input_subset))) if task == "add": output_subset = keys[split_idx : split_idx + 1] # only add one stem # Randomly set instrument name to "auto" for some percentage of the data if random.random() < 0.0: stem_type = f"{task} auto" else: stem_type = f"{task} {', '.join(output_subset)}" elif task == "extract": output_subset = input_subset output_subset = random.sample(output_subset, 1) stem_type = f"{task} {', '.join(output_subset)}" elif task == "remove": output_subset = input_subset output_subset = random.sample( output_subset, random.randint(1, max(1, len(output_subset) - 1)) ) # print(f"remove, {len(input_subset)} -> {len(output_subset)}") stems_removed = [k for k in input_subset if k not in output_subset] stem_type = f"{task} {', '.join(stems_removed)}" else: raise ValueError(f"Invalid task: {task}") # 50% chance to crop to an active section start_s = 0 end_s = None if ( random.random() < 1 and "stem_active_sections" in main_meta and len(output_subset) == 1 and len(main_meta["stem_active_sections"].get(output_subset[0], [])) > 0 ): active_sections = main_meta["stem_active_sections"][output_subset[0]] # Try to find a random span between section boundaries where the stem is active at least 30% of the time max_attempts = 100 for _ in range(max_attempts): # choose a span of sections section_start, section_end = sorted(random.choices(active_sections, k=2)) span_start = section_start[0] span_end = section_end[1] span_duration = span_end - span_start # Calculate how much of this span overlaps with active sections active_duration = 0 for section_start, section_end in active_sections: overlap_start = max(span_start, section_start) overlap_end = min(span_end, section_end) if overlap_start < overlap_end: active_duration += overlap_end - overlap_start # Check if at least 30% of the span is active activity_ratio = active_duration / span_duration # print( # f"Found span for {main_meta['id']}, using {span_start} to {span_end}, activity ratio: {activity_ratio}" # ) if activity_ratio >= 0.3: start_s = span_start end_s = span_end break else: # Fallback: use a random active section if no good span found print(f"No good span found for {main_meta['id']}, using random active section") start_s, end_s = random.choice(active_sections) if end_s - start_s < 2: # skip if the active section is too short start_s = 0 end_s = None # Load and mix stems output_stem_audio = self._load_and_mix_stems( stem_dict, output_subset, mock=mock, load_for_semantic=load_for_semantic, start_s=start_s, end_s=end_s, ) input_subset_audio = self._load_and_mix_stems( stem_dict, input_subset, mock=mock, load_for_semantic=load_for_semantic, start_s=start_s, end_s=end_s, ) # pad to same length if self.model_cfg.semantic_type.startswith("mert"): input_subset_audio = np.pad( input_subset_audio, (0, max(0, len(output_stem_audio) - len(input_subset_audio))) ) output_stem_audio = np.pad( output_stem_audio, (0, max(0, len(input_subset_audio) - len(output_stem_audio))) ) assert len(input_subset_audio) == len(output_stem_audio) elif self.model_cfg.semantic_type.startswith("musicfm"): input_subset_audio = np.pad( input_subset_audio, ((0, 0), (0, max(0, output_stem_audio.shape[1] - input_subset_audio.shape[1]))), ) output_stem_audio = np.pad( output_stem_audio, ((0, 0), (0, max(0, input_subset_audio.shape[1] - output_stem_audio.shape[1]))), ) assert input_subset_audio.shape[1] == output_stem_audio.shape[1] # Create full mix (input + output stem) for output target full_mix_audio = input_subset_audio + output_stem_audio * 3 return stem_type, input_subset_audio, output_stem_audio, full_mix_audio def __next__(self): return self._next() def _sample_meta(self): """ Randomly sample a meta from the metas list. Weighted sampling is slow for large weighted datasets so we sample 10k at a time. """ if len(self.random_cache) == 0: choices = random.choices(range(len(self.metas)), weights=self.weights, k=10000) self.random_cache.extend(choices) idx = self.random_cache.pop() return self.metas[idx] def _next(self): # randomly sample a meta main_meta = self._sample_meta() main_meta["s3_filepath"] = main_meta.get("s3_filepath", None) # fine if missing audio_type = main_meta.get("audio_type", AudioType.MUSIC) # fine if missing use_raw_audio = self.sampling_params.output_distribution == "vae" # load wav data_row, data_row_48kHz = self._load_audio_for_semantic( main_meta["local_filepath"], main_meta.get("s3_filepath", None), main_meta["duration_s"], mock=self.sampling_params.mock_data, ) # TODO: add cover, artist, playlist, overpaint, underpaint if ( not self.sampling_params.inference and self.sampling_params.allow_cover and random.random() < self.sampling_params.prob_cover ): data_row_cover = self._load_cover_audio( main_meta, mock=self.sampling_params.mock_data, load_for_semantic=not use_raw_audio ) else: data_row_cover = None p_use_artist = ( 0.5 if data_row_cover is not None else self.sampling_params.prob_artist ) # increase chance for cover cause voice beautifier if ( not self.sampling_params.inference and self.sampling_params.allow_artist and random.random() < p_use_artist ): data_row_artist = self._load_artist_audio( main_meta, mock=self.sampling_params.mock_data, load_for_semantic=not use_raw_audio ) else: data_row_artist = None if ( not self.sampling_params.inference and self.sampling_params.use_ditto and random.random() < self.sampling_params.prob_use_ditto ): data_row_ditto = _resample_to_mert(data_row_48kHz) else: data_row_ditto = None if ( not self.sampling_params.inference and self.sampling_params.allow_playlist and random.random() < self.sampling_params.prob_playlist ): data_row_playlist = self._load_playlist_audio( main_meta, mock=self.sampling_params.mock_data, load_for_semantic=not use_raw_audio ) else: data_row_playlist = None if ( not self.sampling_params.inference and self.sampling_params.allow_overpaint and random.random() < self.sampling_params.prob_overpaint ): data_row_overpaint = self._load_overpaint_audio( main_meta, mock=self.sampling_params.mock_data, load_for_semantic=not use_raw_audio ) # sometimes add noise to hide imperfect source sep at input if data_row_overpaint is not None and random.random() < 0.5: noise = np.random.normal(0, random.random() * 0.25, data_row_overpaint.shape) data_row_overpaint = np.clip( data_row_overpaint + noise.astype(data_row_overpaint.dtype), -1.1, 1.1 ) else: data_row_overpaint = None if ( not self.sampling_params.inference and self.sampling_params.allow_underpaint and random.random() < self.sampling_params.prob_underpaint ): data_row_underpaint = self._load_underpaint_audio( main_meta, mock=self.sampling_params.mock_data, load_for_semantic=not use_raw_audio ) # sometimes add noise to hide imperfect source sep at input if data_row_underpaint is not None and random.random() < 0.5: noise = np.random.normal(0, random.random() * 0.25, data_row_underpaint.shape) data_row_underpaint = np.clip( data_row_underpaint + noise.astype(data_row_underpaint.dtype), -1.1, 1.1 ) else: data_row_underpaint = None if ( not self.sampling_params.inference and self.sampling_params.allow_vox and random.random() < self.sampling_params.prob_vox ): data_row_vox = self._load_vox_audio( main_meta, mock=self.sampling_params.mock_data, load_for_semantic=not use_raw_audio ) else: data_row_vox = None if ( not self.sampling_params.inference and self.sampling_params.allow_remix and random.random() < self.sampling_params.prob_remix ): data_row_remix = self._load_remix_source_audio( main_meta, mock=self.sampling_params.mock_data, load_for_semantic=not use_raw_audio ) else: data_row_remix = None if ( not self.sampling_params.inference and self.sampling_params.allow_sample_source and random.random() < self.sampling_params.prob_sample_source ): data_row_sample_source = self._load_sample_source_audio( main_meta, mock=self.sampling_params.mock_data, load_for_semantic=not use_raw_audio ) else: data_row_sample_source = None if ( not self.sampling_params.inference and self.sampling_params.allow_mashup and random.random() < self.sampling_params.prob_mashup ): mashup_tracks = self._load_mashup_audio( main_meta, mock=self.sampling_params.mock_data, load_for_semantic=not use_raw_audio ) else: mashup_tracks = None # Sample conditioning - moved before stem conditioning to preserve intact data_row audio_sample_tracks = None audio_sample_start_times_s = None audio_sample_sources = None if ( not self.sampling_params.inference and self.sampling_params.allow_sample and random.random() < self.sampling_params.prob_sample ): # Sample conditioning mode: extract random number of audio samples max_samples = getattr(self.sampling_params, "max_num_audio_samples", 1) num_samples = random.randint(1, max_samples) prob_stems = getattr( self.sampling_params, "prob_sample_from_stems", 0.95 ) # Default 95% stems, 5% full song extracted_samples = [] extracted_times = [] extracted_sources = [] for _ in range(num_samples): result = None # Get sample permutation probability from sampling params (only for training) sample_permutation_prob = 0.0 if self.split == "train": sample_permutation_prob = getattr( self.sampling_params, "sample_permutation_prob", 0.3 ) if random.random() < prob_stems: # Try stems first (preferably vocal) result = create_sample_from_stems( self, main_meta, data_row, self.semantic_sample_rate, sample_permutation_prob ) # If stems weren't chosen or failed, use full song approach # TODO: also apply smart cropping for full song, instead of the naive approach if result is None: result = create_sample_from_full_song( data_row, main_meta, self.semantic_sample_rate, sample_permutation_prob ) if result is not None: sample_audio, sample_time, source_type = result extracted_samples.append(sample_audio) extracted_times.append(sample_time) extracted_sources.append(source_type) if extracted_samples: audio_sample_tracks = extracted_samples audio_sample_start_times_s = extracted_times audio_sample_sources = extracted_sources stem_type = None data_row_stem = None stem_output_mix = None if ( not self.sampling_params.inference and self.sampling_params.allow_stem and main_meta.get("stems", {}) and (random.random() < self.sampling_params.prob_stem or "stem_active_sections" in main_meta) ): task = random.choices(["add", "extract", "remove"], weights=[0.8, 0.0, 0.0])[0] stem_type, stem_input_mix, isolated_stem, full_mix = self._load_stem_audio( main_meta, task=task, mock=self.sampling_params.mock_data, load_for_semantic=not use_raw_audio, ) if stem_type is not None: data_row_stem = stem_input_mix data_row = isolated_stem stem_output_mix = full_mix # passing 48khz raw audio through dataloaders reduces throughput 20% use_hoot = self.sampling_params.use_hoot and random.random() < self.sampling_params.prob_use_hoot use_repa_hoot = ( self.sampling_params.repa_hoot and random.random() < self.sampling_params.prob_repa_hoot ) use_repa_midi = ( self.sampling_params.repa_midi and random.random() < self.sampling_params.prob_repa_midi ) if ( self.sampling_params.use_vae_input or self.sampling_params.output_distribution == "vae" or use_hoot or use_repa_hoot ): use_raw_audio = True if use_raw_audio: raw_audio = self._load_audio_for_semantic( main_meta["local_filepath"], main_meta["s3_filepath"], main_meta["duration_s"], mock=self.sampling_params.mock_data, load_for_semantic=False, ) else: raw_audio = None sample_data = SampleData( data_row=data_row, data_meta=main_meta, sampling_params=self.sampling_params, data_row_cover=data_row_cover, artist_tracks=data_row_artist, ditto_track=data_row_ditto, playlist_tracks=data_row_playlist, overpaint_track=data_row_overpaint, underpaint_track=data_row_underpaint, vox_track=data_row_vox, remix_track=data_row_remix, sample_source_track=data_row_sample_source, mashup_tracks=mashup_tracks, stem_track=data_row_stem, stem_output_mix=stem_output_mix, audio_sample_tracks=audio_sample_tracks, audio_sample_start_times_s=audio_sample_start_times_s, audio_sample_sources=audio_sample_sources, raw_audio=raw_audio, stem_type=stem_type, audio_type=audio_type, # signals passed from audioloader to data_utils.py use_repa_hoot=use_repa_hoot, use_repa_midi=use_repa_midi, use_hoot=use_hoot, ) return sample_data