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

from ditto_v2.modules.music_encoder import MusicEncoder
from ditto_v2.modules.text_encoder import TextEncoder
from ditto_v2.modules.mlp import Projection


class Ditto(nn.Module):
    def __init__(
        self,
        latent_dim=512,
        model_path=None,
        is_flash=True,
    ):
        super().__init__()

        self.latent_dim = latent_dim

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

        # 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)
            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
            # random.seed(142)
            # SS["music_encoder.music_encoder.cls_token"] = nn.Parameter(torch.randn(1024))
            self.load_state_dict(SS, strict=True)
            print("model loaded!")

    def get_projection_layers(self):
        music_dim = 1024
        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, task):
        music_emb = self.music_projection.float()(self.music_encoder(wav, task).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_multi(self, wav, text, task):
        # music encoding
        music_emb = self.music_projection.float()(self.music_encoder(wav, task).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()

    def forward_text(self, text1, text2):
        # text encoding
        text1_emb = self.text_projection.float()(self.text_encoder(text1).float())
        text1_emb = F.normalize(text1_emb, dim=-1)

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

        return text1_emb, text2_emb, self.logit_scale.exp()

    def forward_music(self, wav1, wav2, task):
        # music encoding
        music1_emb = self.music_projection.float()(
            self.music_encoder(wav1, task).float()
        )
        music1_emb = F.normalize(music1_emb, dim=-1)

        # music encoding
        music2_emb = self.music_projection.float()(
            self.music_encoder(wav2, task).float()
        )
        music2_emb = F.normalize(music2_emb, dim=-1)

        return music1_emb, music2_emb, self.logit_scale.exp()

    def forward(self, inp1, inp2, task):
        if task in [
            "self_sim",
            "self_vox_sim",
            "artist_sim",
            "artist_vox_sim",
            "album_sim",
        ]:
            return self.forward_music(inp1, inp2, task)
        elif task in ["self_lyric_sim"]:
            return self.forward_text(inp1, inp2)
        elif task in ["genre_sim", "lyric_sim"]:  # music first
            return self.forward_multi(inp1, inp2, task)
