from torch import nn from suno_utils.models.musicfm.modeling_MusicFM import MusicFM_MERTLong class MusicEncoder(nn.Module): """ Music audio encoder """ def __init__( self, layer_ix=12, is_flash=True, model_path="/home/minz/logs/musicfm_concat/musicfm_concat_epoch=51.pt", ): super(MusicEncoder, self).__init__() self.model_path = model_path self.layer_ix = layer_ix self.is_flash = is_flash self.music_encoder = self.get_encoder() def get_encoder(self): encoder = MusicFM_MERTLong( is_flash=self.is_flash, is_cls=True, model_path=self.model_path, is_strict=False, num_tasks=10, ) self.hidden_dim = 1024 return encoder def get_embeddings(self, wav, task): if task == "self_sim": task_ix = 0 elif task == "self_vox_sim": task_ix = 1 elif task == "artist_sim": task_ix = 2 elif task == "artist_vox_sim": task_ix = 3 elif task == "album_sim": task_ix = 4 elif task == "genre_sim": task_ix = 5 elif task == "lyric_sim": task_ix = 6 return self.music_encoder.get_latent( wav, self.layer_ix, is_cls=True, cls_task=task_ix ) def forward(self, wav, task): self.music_encoder.eval() emb = self.get_embeddings(wav, task) return emb[ :, 0, : ] # we take the first token to represent the sequence by appending [CLS] token