import os import glob import glob import torch import uuid import IPython import numpy as np import torchaudio import itertools from tqdm import tqdm from dac.model.dac2 import DAC from dac.model.discriminator2 import Discriminator as Discriminator_import from dac.nn import loss as loss_import from dac.utils.accelerator import Accelerator from dac.utils import load_model from suno_utils.models.musicfm.modeling_MusicFM import MusicFM_MERTLong import matplotlib.pyplot as plt from sklearn.preprocessing import StandardScaler from sklearn.cluster import KMeans, MiniBatchKMeans from suno_utils.utils.s3 import read_from_s3 USE_VAL = False NUM_FRAMES = int(10 * 24_000) SAMPLE_RATE = 24_000 BATCH_SIZE = 128 class AudioDataset(torch.utils.data.Dataset): def __init__(self): # get audio files if USE_VAL: audio_subsets = glob.glob( os.path.join(f"/app/suno/data/audio_2ch_24khz_lg/val/**") ) else: audio_subsets = glob.glob( os.path.join(f"/app/suno/data/audio_2ch_24khz_lg/train/**") ) audio_files = [] for audio_subset in audio_subsets: # find the first MAX_FILES_PER_SUBSET files with os.scandir(audio_subset) as filepaths: # Use itertools.islice to limit the iterator to the first N entries first_n_files = list(itertools.islice(filepaths, 40000)) first_n_files = [entry.path for entry in first_n_files if entry.is_file()] audio_files += first_n_files print(len(first_n_files), audio_subset) print("Total", len(audio_files)) print(np.random.choice(audio_files, 5)) self.audio_files = audio_files def __len__(self): return len(self.audio_files) def __getitem__(self, idx): audio_file = self.audio_files[idx] num_frames = torchaudio.info(audio_file).num_frames if num_frames > NUM_FRAMES: frame_offset = np.random.randint( 0, torchaudio.info(audio_file).num_frames - NUM_FRAMES - 1 ) else: frame_offset = 0 audio, sr = torchaudio.load( audio_file, frame_offset=frame_offset, num_frames=NUM_FRAMES ) audio = audio.mean(dim=0) if audio.shape[-1] < NUM_FRAMES: audio = torch.cat( [audio, torch.zeros(NUM_FRAMES - audio.shape[-1])], dim=-1 ) assert sr == SAMPLE_RATE audio = audio / audio.abs().max().clamp(1e-8) return idx, audio if __name__ == "__main__": # load MERT model_filepath = "s3://suno-data/minz/models/musicfm_concat_epoch=51.pt" centroids_filepath = "s3://suno-data/minz/models/musicfm_concat_centroids.npy" mert_model = MusicFM_MERTLong( is_flash=False, stat_path="s3://suno-data/minz/models/mertlong_stats.json", model_path=model_filepath, ) mert_model.cuda() mert_model.eval() # setup data if USE_VAL: subset_name = "val" else: subset_name = "train" root_dir = "/app/suno/christian/data/mert/" out_dir = os.path.join(root_dir, subset_name) os.makedirs(out_dir, exist_ok=True) dataset = AudioDataset() dataloader = torch.utils.data.DataLoader( dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=32 ) # embed for batch in tqdm(dataloader): idx, audios = batch with torch.no_grad(): embeddings = mert_model.get_latent(audios.cuda()) # save embeddings as npz for i, (index, emb) in enumerate(zip(idx, embeddings)): np.savez( f"/app/suno/christian/data/mert/{subset_name}/{index:10d}.npz", emb.cpu().numpy(), )