import json import torch 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 Conv2dSubsampling from musicfm.modules.discrete_conformer import Wav2Vec2ConformerEncoderDiscrete, Wav2Vec2ConformerConfig class MusicFM7q(nn.Module): """ MusicFM pretrained with MERT-long data (~23k hours) Quantized in the seventh layer. Input: 128-band mel spectrogram Frontend: 2-layer Residual convolution Backend: 12-layer Conformer Quantizer: 1 codebooks for mel spectrogram or spectrogram Internal quantization """ def __init__( self, num_codebooks=1, codebook_dim=16, codebook_size=4096, num_latent_codebook=4096, latent_codebook_dim=16, 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, stat_path="/app/suno/minz/models/mertlong_stats.json", model_path="/app/suno/minz/models/musicfm_7q_step=151524.pt", ): super(MusicFM7q, 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 * 4, codebook_dim, codebook_size, seed=seed+i)) # two residual convolution layers + one projection layer self.conv = Conv2dSubsampling(1, conv_dim, encoder_dim, strides=[2, 2], n_bands=n_mels) # Conformer config = Wav2Vec2ConformerConfig.from_pretrained( "facebook/wav2vec2-conformer-rope-large-960h-ft" ) config.num_hidden_layers = encoder_depth config.hidden_size = encoder_dim config.num_latent_codebook = num_latent_codebook config.latent_codebook_dim = latent_codebook_dim config.is_flash = is_flash self.conformer = Wav2Vec2ConformerEncoderDiscrete(config) # projection self.linear = nn.Linear(encoder_dim, num_codebooks * codebook_size * len(features)) # loss function self.loss = nn.CrossEntropyLoss() # 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=True) 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/2/2 * 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): """ 2-layer conv + w2v-conformer """ x = self.conv(x) out, gumbel_states, gumbel_tokens = 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, gumbel_states, gumbel_tokens @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 4 steps """ for key in x.keys(): if key == "chromagram": x[key] = rearrange(x[key], "b f t -> b t f") else: x[key] = rearrange(x[key], "b f (t s) -> b t (s f)", s=4) 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): # preprocessing x = self.preprocessing(x, features=["melspec_2048"]) x = self.normalize(x) # encoding logits, hidden_emb, gumbel_emb, gumbel_tokens = self.encoder(x["melspec_2048"]) return logits, hidden_emb, gumbel_emb, gumbel_tokens def get_latent(self, x, layer_ix=12): _, hidden_states, gumbel_emb, gumbel_tokens = self.get_predictions(x) emb = hidden_states[layer_ix] return emb, gumbel_emb, gumbel_tokens 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