import sys

sys.path.append("/home/minz/glockenspiel/musicfm-training")

from torch import nn


class MusicEncoder(nn.Module):
    """
    Music audio encoder
    """

    def __init__(
        self,
        model_name="musicfm_mertlong",
        layer_ix=12,
        is_flash=True,
    ):
        super(MusicEncoder, self).__init__()

        self.model_name = model_name
        self.layer_ix = layer_ix
        self.is_flash = is_flash

        self.music_encoder = self.get_encoder()

    def get_encoder(self):
        if self.model_name == "musicfm_mertlong":
            from musicfm.models.musicfm_mertlong import MusicFM_MERTLong

            encoder = MusicFM_MERTLong(
                is_flash=self.is_flash, is_cls=True, is_strict=False
            )
            self.hidden_dim = 1024
        elif self.model_name == "musicfm_concat":
            from musicfm.models.musicfm_mertlong import MusicFM_MERTLong

            encoder = MusicFM_MERTLong(
                is_flash=self.is_flash,
                is_cls=True,
                model_path="/home/minz/logs/musicfm_concat/musicfm_concat_epoch=51.pt",
                is_strict=False,
            )
            self.hidden_dim = 1024
        else:
            raise ValueError("%s is not supported yet." % self.model_name)

        return encoder

    def get_embeddings(self, wav):
        if self.model_name == "musicfm_mertlong":
            emb = self.music_encoder.get_latent(wav, self.layer_ix, is_cls=True)
        elif self.model_name == "musicfm_concat":
            emb = self.music_encoder.get_latent(wav, self.layer_ix, is_cls=True)

        return emb

    def forward(self, wav):
        self.music_encoder.eval()
        emb = self.get_embeddings(wav)
        return emb[
            :, 0, :
        ]  # we take the first token to represent the sequence by appending [CLS] token
