import torch
import torch.nn.functional as F
import numpy as np
from torch import nn

from ditto.modules.music_encoder import MusicEncoder
from ditto.modules.text_encoder import TextEncoder
from ditto.modules.mlp import Projection, MLP


class Ditto(nn.Module):
    def __init__(
        self,
        music_encoder_name="musicfm_mertlong",
        text_encoder_name="xlm-roberta",
        latent_dim=512,
        model_path=None,
        is_flash=True,
    ):
        super().__init__()

        self.music_encoder_name = music_encoder_name
        self.text_encoder_name = text_encoder_name
        self.latent_dim = latent_dim

        # prepare encoders
        self.music_encoder = MusicEncoder(
            music_encoder_name, layer_ix=12, is_flash=is_flash
        )
        self.text_encoder = TextEncoder(text_encoder_name)

        # get projection layers
        self.music_projection, self.text_projection = self.get_projection_layers()

        # logit scale
        self.logit_scale = nn.Parameter(torch.ones([]) * np.log(1 / 0.07))

        # load model
        if model_path:
            S = torch.load(model_path)["state_dict"]
            SS = {k[6:]: v for k, v in S.items()}
            SS["logit_scale"] = nn.Parameter(
                torch.ones([]) * np.log(1 / 0.07)
            )  # TODO: include this scale in back propagation
            self.load_state_dict(SS, strict=True)
            print("model loaded!")

    def get_projection_layers(self):
        if self.music_encoder_name == "musicfm_mertlong":
            music_dim = 1024

        elif self.music_encoder_name == "musicfm_concat":
            music_dim = 1024

        if self.text_encoder_name == "xlm-roberta":
            text_dim = 768

        music_projection = Projection(music_dim, self.latent_dim)
        text_projection = Projection(text_dim, self.latent_dim)
        return music_projection, text_projection

    @torch.no_grad()
    def music_to_latent(self, wav):
        music_emb = self.music_projection.float()(self.music_encoder(wav).float())
        music_emb = F.normalize(music_emb, dim=-1)
        return music_emb

    @torch.no_grad()
    def text_to_latent(self, text):
        text_emb = self.text_projection.float()(self.text_encoder(text).float())
        text_emb = F.normalize(text_emb, dim=-1)
        return text_emb

    def forward(self, wav, text):
        # music encoding
        music_emb = self.music_projection.float()(self.music_encoder(wav).float())
        music_emb = F.normalize(music_emb, dim=-1)

        # text encoding
        text_emb = self.text_projection.float()(self.text_encoder(text).float())
        text_emb = F.normalize(text_emb, dim=-1)

        return music_emb, text_emb, self.logit_scale.exp()
