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
