import json import os import random import traceback import numpy as np from torch.utils.data import IterableDataset from tqdm import tqdm from dataclasses import dataclass from enum import Enum from oracle_dataset import get_sample_oracle_file_segment 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] @dataclass class AudioConfig: sample_rate: int = 48000 n_channels: int = 2 duration_s: float = 1.0 is_vae: bool = False is_mert: bool = False is_musicfm: bool = False @dataclass class SampleData: data_wav: np.ndarray # Raw audio in 48kHz (channel, length) filepath: str start_s: float data_mert: np.ndarray | None = None # MERT embedding data_musicfm: np.ndarray | None = None # MusicFM3 embedding class AudioLoaderDataset(IterableDataset): """Gets SampleData objects from the oracle dataset, which contain raw wav data""" def __init__( self, audio_cfg: AudioConfig, metas_path: str, split="train", ): self.audio_cfg = audio_cfg self.metas_path = metas_path self.split = split # load metas as a memmap to save memory self.metas = JSONLMemmap(self.metas_path, verbose=False) self.random_cache = [] def __iter__(self): return self def pad_audio(self, track): assert len(track.shape) == 2 if track.shape[1] < int(48000 * self.audio_cfg.duration_s): pad_length = int(48000 * self.audio_cfg.duration_s) - track.shape[1] track = np.pad(track, ((0, 0), (0, pad_length)), mode="constant", constant_values=0) return track def loudness(self, audio_array): import pyloudnorm as pyln m = pyln.Meter(48000) # create BS.1770 meter assert len(audio_array.shape) == 2 lufs_db = m.integrated_loudness(audio_array.T) return lufs_db def normalize_volume(self, array_float, target_db=-16): loudness = self.loudness(array_float) # Handle edge cases where loudness measurement fails if not np.isfinite(loudness) or loudness < -80: # For very quiet/silent audio, return the original array # or apply minimal normalization max_val = np.abs(array_float).max() if max_val > 0: return array_float / max_val * 0.1 # Scale to 10% of max else: return array_float # Return silent audio as-is gain_factor = np.log(10) / 20 gain = target_db - loudness # Limit gain to reasonable bounds to prevent overflow gain = np.clip(gain, -40, 40) # Limit to ±40dB adjustment gain = np.exp(gain * gain_factor) # Additional check to ensure gain is finite if not np.isfinite(gain): gain = 1.0 norm_arr = array_float * gain # Prevent clipping if np.abs(norm_arr).max() > 1: norm_arr = norm_arr / np.abs(norm_arr).max() return norm_arr def clip_audio(self, array_float): if np.abs(array_float).max() > 1: print("Clipping audio by ", np.abs(array_float).max()) return array_float / np.abs(array_float).max() else: return array_float def _load_audio( self, local_filepath, s3_filepath=None, expected_duration_s=180, audio_stats=None, max_duration_s=1, target_loudness_db=-16, ): try: start_s = ( max(0.0, random.random() * int(expected_duration_s - self.audio_cfg.duration_s)) if self.split == "train" else expected_duration_s / 2 ) 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 # normalize gain gain_db = 0.0 if audio_stats is not None: loudness_db = audio_stats.get("loudness", None) if loudness_db is not None: gain_db = target_loudness_db - float(loudness_db) gain_db = np.clip(gain_db, -12, 12) audio, _ = audio.apply_gain(gain_db) target_sample_length = int(48000 * self.audio_cfg.duration_s) array_float = audio.array_float[:, :target_sample_length] array_float = self.pad_audio(array_float) array_float = self.clip_audio(array_float) return array_float, local_filepath, start_s 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}" ) 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)), k=10000) self.random_cache.extend(choices) idx = self.random_cache.pop() return self.metas[idx] def _next(self): # randomly sample a meta main_meta = self._sample_meta() # load wav data_wav, filepath, start_s = self._load_audio( main_meta["local_filepath"], main_meta.get("s3_filepath", None), main_meta["duration_s"], main_meta.get("audio_stats", None), ) sample_data = SampleData(data_wav=data_wav, filepath=filepath, start_s=start_s) return sample_data