from copy import deepcopy import json import math import os import soxr import random import traceback import contextlib from tqdm import tqdm import numpy as np import torch import torch.nn.functional as F from torch.utils.data import IterableDataset from pretrained_models import load_vae_model, preload_mert_models, preload_musicfm_models from audioloader import AudioConfig from suno_utils.tasks.dac_vae_fixed_25hz import encode as encode_vae from suno_utils.tasks.mert_25 import encode as encode_mert from suno_utils.tasks.musicfm_v3 import encode as encode_musicfm def get_batch( batch_size: int, audio_cfg: AudioConfig, sample_generator, ): cur_len = 0 stacked_batch = [] for sample in tqdm(sample_generator, desc="stacking batch", disable=True): if cur_len >= batch_size: break stacked_batch.append(sample) cur_len += 1 return stacked_batch def _resample_for_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[np.newaxis, :] def _resample_for_musicfm(arr): assert arr.ndim == 2 assert arr.shape[0] == 2 out_arr = soxr.resample(arr.T, 48_000, 16_000).astype(np.float32) return out_arr.T class PreprocessDataset(IterableDataset): def __init__( self, sample_data_dl, batch_size: int, audio_cfg: AudioConfig, ): self.sample_data_dl = sample_data_dl self.batch_size = batch_size self.audio_cfg = audio_cfg 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()) def make_sample(self, sample_data): def get_vae(track, vae_scale_factor=0.4): load_vae_model() wav = torch.from_numpy(track).cuda() vae = encode_vae(wav) vae = torch.from_numpy(vae) * vae_scale_factor return vae.T # (C, T) def get_mert(track): preload_mert_models() track = _resample_for_mert(track) track = torch.from_numpy(track) embs = encode_mert( track, pad_to_chunksize=False, batch_size=48, do_clustering=False, ).T # (C, T) return embs def get_musicfm(track): preload_musicfm_models() track = _resample_for_musicfm(track) track = torch.from_numpy(track) with open(os.devnull, "w") as devnull: with contextlib.redirect_stdout(devnull): embs = encode_musicfm( track, pad_to_chunksize=False, batch_size=48, token_type="emb_pre" ).T # (C, T) return embs # process data if self.audio_cfg.is_vae: sample_data.data_vae = get_vae(sample_data.data_wav) if self.audio_cfg.is_mert: sample_data.data_mert = get_mert(sample_data.data_wav) if self.audio_cfg.is_musicfm: sample_data.data_musicfm = get_musicfm(sample_data.data_wav) yield sample_data 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, self.audio_cfg, self.sample_generator, ) return batch