import json import tqdm import random import torch import numpy as np from torch import nn from einops import rearrange from musicfm.modules.random_quantizer import RandomProjectionQuantizer from musicfm.modules.features import MelSTFT from musicfm.modules.conv import DilatedConv2d from suno_utils.tasks.mert_25 import _array_chunk class MusicFM100Hz(nn.Module): """ MusicFM pretrained with MERT-long data (~23k hours) Input: 128-band mel spectrogram Frontend: 2-layer Residual convolution Backend: 12-layer Conformer Quantizer: 1 codebooks for mel spectrogram or spectrogram """ def __init__( self, num_codebooks=1, codebook_dim=16, codebook_size=4096, features=["melspec_2048"], hop_length=240, n_mels=128, conv_dim=512, encoder_dim=1024, encoder_depth=12, mask_hop=0.4, mask_prob=0.6, is_flash=False, is_cls=False, is_strict=True, stat_path="/app/suno/minz/models/mertlong_stats.json", model_path=None, ): super(MusicFM100Hz, self).__init__() # global variables self.hop_length = hop_length self.mask_hop = mask_hop self.mask_prob = mask_prob self.num_codebooks = num_codebooks self.codebook_size = codebook_size self.features = features # load feature mean / std stats with open(stat_path, "r") as f: self.stat = json.load(f) # feature extractor self.preprocessor_melspec_2048 = MelSTFT( n_fft=2048, hop_length=hop_length, is_db=True ) # random quantizer seed = 142 for feature in self.features: for i in range(num_codebooks): setattr( self, "quantizer_%s_%d" % (feature, i), RandomProjectionQuantizer( n_mels, codebook_dim, codebook_size, seed=seed + i ), ) # two residual convolution layers + one projection layer self.conv = DilatedConv2d( 1, conv_dim, encoder_dim, dilations=[2, 2], n_bands=n_mels ) # Conformer if is_flash: from musicfm.modules.flash_conformer import ( Wav2Vec2ConformerEncoder, Wav2Vec2ConformerConfig, ) else: from musicfm.modules.non_flash_conformer import ( Wav2Vec2ConformerEncoder, Wav2Vec2ConformerConfig, ) config = Wav2Vec2ConformerConfig.from_pretrained( "facebook/wav2vec2-conformer-rope-large-960h-ft" ) config.num_hidden_layers = encoder_depth config.hidden_size = encoder_dim self.conformer = Wav2Vec2ConformerEncoder(config) # projection self.linear = nn.Linear( encoder_dim, num_codebooks * codebook_size * len(features) ) # loss function self.loss = nn.CrossEntropyLoss() # cls token if is_cls: random.seed(seed) self.cls_token = nn.Parameter(torch.randn(encoder_dim)) # load model if model_path: S = torch.load(model_path)["state_dict"] SS = {k[6:]: v for k, v in S.items()} self.load_state_dict(SS, strict=is_strict) def masking(self, x): """random masking of 400ms with given probability""" mx = x.clone() b, t = mx.shape len_masking_raw = int(24000 * self.mask_hop) len_masking_token = int(24000 / self.hop_length * self.mask_hop) # get random mask indices start_indices = torch.rand(b, t // len_masking_raw) < self.mask_prob time_domain_masked_indices = torch.nonzero( start_indices.repeat_interleave(len_masking_raw, dim=1) ) token_domain_masked_indices = torch.nonzero( start_indices.repeat_interleave(len_masking_token, dim=1) ) # mask with random values masking_noise = ( torch.randn(time_domain_masked_indices.shape[0], dtype=x.dtype) * 0.1 ) # 0 mean 0.1 std mx[tuple(time_domain_masked_indices.t())] = masking_noise.to(x.device) return mx, token_domain_masked_indices @torch.no_grad() def preprocessing(self, x, features): """extract classic audio features""" # check precision if x.dtype == torch.float16: precision = 16 elif x.dtype == torch.bfloat16: precision = "bf16" else: precision = 32 out = {} for key in features: layer = getattr(self, "preprocessor_%s" % key) out[key] = layer.float()(x.float())[..., :-1] if precision == 16: out[key] = out[key].half() elif precision == "bf16": out[key] = out[key].bfloat16() return out def encoder(self, x, is_cls=False): """2-layer conv + w2v-conformer""" x = self.conv(x) if is_cls: cls_token = self.cls_token.repeat(x.shape[0], 1, 1) x = torch.cat((cls_token, x), dim=1) out = self.conformer(x, output_hidden_states=True) hidden_emb = out["hidden_states"] last_emb = out["last_hidden_state"] logits = self.linear(last_emb) logit_dict = {} ix = 0 for key in self.features: for i in range(self.num_codebooks): logit_dict["%s_%d" % (key, i)] = logits[ :, :, ix * self.codebook_size : (ix + 1) * self.codebook_size ] return logit_dict, hidden_emb @torch.no_grad() def normalize(self, x): """normalize the input audio to have zero mean unit variance""" for key in x.keys(): x[key] = (x[key] - self.stat["%s_mean" % key]) / self.stat["%s_std" % key] return x @torch.no_grad() def rearrange(self, x): """rearrange the batch to flatten every 1 steps""" for key in x.keys(): x[key] = rearrange(x[key], "b f t -> b t f") return x @torch.no_grad() def tokenize(self, x): out = {} for key in x.keys(): for i in range(self.num_codebooks): layer = getattr(self, "quantizer_%s_%d" % (key, i)) out["%s_%d" % (key, i)] = layer(x[key]) return out def get_targets(self, x): x = self.preprocessing(x, features=self.features) x = self.normalize(x) x = self.rearrange(x) target_tokens = self.tokenize(x) return target_tokens def get_predictions(self, x, is_cls=False): # preprocessing x = self.preprocessing(x, features=["melspec_2048"]) x = self.normalize(x) # encoding logits, hidden_emb = self.encoder(x["melspec_2048"], is_cls) return logits, hidden_emb def get_latent(self, x, layer_ix=12, is_cls=False): # preprocessing x = self.preprocessing(x, features=["melspec_2048"]) x = self.normalize(x) # conv x = self.conv(x["melspec_2048"]) if is_cls: cls_token = self.cls_token.repeat(x.shape[0], 1, 1) x = torch.cat((cls_token, x), dim=1) # conformer out = self.conformer.partial_encode( x, layer_ix=layer_ix + 1, output_hidden_states=True ) return out["hidden_states"][layer_ix] @torch.no_grad() def encode_arrays(self, x, batch_size=16, layer_ix=7): # make batch subsplit_arrays = [] for n_array, arr in enumerate(x): for sub_arr in _array_chunk(24000 * 30, arr, step_size=24000 * 20, dim=1): if sub_arr.size(-1) == 24000 * 30: subsplit_arrays.append(sub_arr) # encode them encoded_arrays = [] num_iter = len(subsplit_arrays) // batch_size if len(subsplit_arrays) % batch_size > 0: num_iter += 1 for i in tqdm.tqdm(range(num_iter)): inp = torch.cat( subsplit_arrays[i * batch_size : (i + 1) * batch_size] ).cuda() emb = self.get_latent(inp, layer_ix) emb = rearrange(emb, "b t c -> (b t) c") encoded_arrays.append(emb.cpu().detach().numpy()) return np.concatenate(encoded_arrays) def get_loss(self, logits, target_tokens, masked_indices): losses = {} accuracies = {} for key in logits.keys(): masked_logits = logits[key][tuple(masked_indices.t())] masked_tokens = target_tokens[key][tuple(masked_indices.t())] losses[key] = self.loss(masked_logits, masked_tokens) accuracies[key] = ( torch.sum(masked_logits.argmax(-1) == masked_tokens) / masked_tokens.numel() ) return losses, accuracies def forward(self, x): # get target feature tokens target_tokens = self.get_targets(x) # masking x, masked_indices = self.masking(x) # forward logits, hidden_emb = self.get_predictions(x) # get loss losses, accuracies = self.get_loss(logits, target_tokens, masked_indices) return logits, hidden_emb, losses, accuracies