import time
import random
import torch
import lightning as L
import torch.nn.functional as F
import numpy as np
from torch import nn
from sklearn import metrics as skm
from einops import rearrange


def print_model_params(model):
    total_params = 0
    trainable_params = 0
    for _, parameter in model.named_parameters():
        params = parameter.numel()
        total_params += params
        if parameter.requires_grad:
            trainable_params += params
    print(f"Total Params: {total_params:,}, Trainable_params: {trainable_params:,}")


class DittoLitModule(L.LightningModule):
    def __init__(
        self,
        model,
        learning_rate=1e-4,
        gradient_clip_val=1.0,
    ):
        super().__init__()
        self.lr = learning_rate
        self.model = model
        print_model_params(model)
        self.loss_function = nn.CrossEntropyLoss()
        self.gradient_clip_val = gradient_clip_val
        self.save_hyperparameters(ignore=["model"])
        (
            self.inference_music_embeddings,
            self.inference_text_embeddings,
            self.inference_ids,
            self.start_ss,
            self.lyrics,
        ) = [], [], [], [], []
        self.counter = 0

    def get_metrics(self, z_a, z_b, logit_scale):
        metrics = {}

        # scaled dot product similarity
        logits_per_a = logit_scale * z_a @ z_b.t()
        logits_per_b = logits_per_a.t()
        labels = torch.arange(z_a.shape[0]).long().to(self.device)

        # get cross entropy loss
        loss = (
            self.loss_function(logits_per_a, labels)
            + self.loss_function(logits_per_b, labels)
        ) / 2
        metrics["loss"] = loss

        # get accuracy
        metrics["acc_a"] = skm.accuracy_score(
            labels.detach().cpu().numpy(),
            logits_per_a.argmax(dim=1).detach().cpu().numpy(),
        )
        metrics["acc_b"] = skm.accuracy_score(
            labels.detach().cpu().numpy(),
            logits_per_b.argmax(dim=1).detach().cpu().numpy(),
        )

        # get ranking metrics
        logits = {
            "a_to_b": logits_per_a.detach().cpu(),
            "b_to_a": logits_per_b.detach().cpu(),
        }
        ground_truth = torch.arange(len(z_b)).view(-1, 1)
        for name, logit in logits.items():
            ranking = torch.argsort(logit, descending=True)
            preds = torch.where(ranking == ground_truth)[1]
            preds = preds.detach().cpu().numpy()
            metrics[f"{name}_mean_rank"] = preds.mean() + 1
            metrics[f"{name}_mdeidan_rank"] = np.floor(np.median(preds)) + 1
            for k in [1, 5, 10]:
                metrics[f"{name}_R@{k}"] = np.mean(preds < k)
            metrics[f"{name}_mAP@10"] = np.mean(
                np.where(preds < 10, 1 / (preds + 1), 0.0)
            )
        return metrics

    def random_masking(self, x, mask_prob=0.125, mask_hop_s=0.5):
        """random masking of 500ms with given probability"""
        b, t = x.shape
        len_masking_raw = int(24000 * mask_hop_s)

        # get random mask indices
        start_indices = torch.rand(b, t // len_masking_raw) < mask_prob
        time_domain_masked_indices = torch.nonzero(
            start_indices.repeat_interleave(len_masking_raw, dim=1)
        )

        # mask with random values
        masking_noise = (
            torch.randn(time_domain_masked_indices.shape[0], dtype=x.dtype) * 0.1
        )  # 0 mean 0.1 std
        x[tuple(time_domain_masked_indices.t())] = masking_noise.to(self.device)

        return x

    def sequence_masking(self, x, max_ratio=0.6):
        b, t = x.shape
        mask_len = random.randint(1, int(t * max_ratio))
        masking_noise = torch.randn(b, mask_len, dtype=x.dtype) * 0.1
        x[:, -mask_len:] = masking_noise.to(self.device)
        return x

    def step(self, batch, stage):
        # get batch
        inps, task = batch
        task = task[0]
        inp1, inp2 = inps

        # forward based on task
        if task in [
            "self_sim",
            "self_vox_sim",
            "artist_sim",
            "artist_vox_sim",
            "album_sim",
        ]:
            if stage == "train":
                random_length1 = random.randint(24000 * 5, 24000 * 15)
                random_length2 = random.randint(24000 * 5, 24000 * 15)
                inp1 = inp1[:, :random_length1]
                inp2 = inp2[:, :random_length2]
                inp1 = inp1.to(self.device)
                inp2 = inp2.to(self.device)
            outputs = self.model(inp1, inp2, task)
            modality_ab = "m2m"
            modality_ba = "m2m"
        elif task in ["self_lyric_sim"]:
            if stage == "train":
                random_length1 = random.randint(50, 313)
                random_length2 = random.randint(50, 313)
                inp1 = [line[:random_length1] for line in inp1]
                inp2 = [line[:random_length2] for line in inp2]
            outputs = self.model(inp1, inp2, task)
            modality_ab = "t2t"
            modality_ba = "t2t"
        elif task in ["genre_sim", "lyric_sim"]:
            if stage == "train":
                random_length1 = random.randint(24000 * 5, 24000 * 15)
                random_length2 = random.randint(50, 313)
                inp1 = inp1[:, :random_length1]
                inp1 = inp1.to(self.device)
                inp2 = [line[:random_length2] for line in inp2]
            outputs = self.model(inp1, inp2, task)
            modality_ab = "m2t"
            modality_ba = "t2m"

        # gather multi-gpu outputs
        if self.trainer.world_size > 1:
            gathered_outputs = self.all_gather(outputs, sync_grads=(stage == "train"))
            inp1_emb = rearrange(gathered_outputs[0], "n b c -> (n b) c")
            inp2_emb = rearrange(gathered_outputs[1], "n b c -> (n b) c")
        else:
            inp1_emb, inp2_emb = outputs[0], outputs[1]
        logit_scale = outputs[2]

        # get metrics
        metrics = self.get_metrics(inp1_emb, inp2_emb, logit_scale)

        # log metrics
        self.log(
            "loss_%s_%s" % (stage, task),
            metrics["loss"],
            prog_bar=True,
            sync_dist=True,
            batch_size=len(inp1_emb),
        )
        self.log(
            "loss_%s" % stage,
            metrics["loss"],
            prog_bar=True,
            sync_dist=True,
            batch_size=len(inp1_emb),
        )
        self.log(
            "mAP10-%s_%s_%s" % (modality_ab, stage, task),
            metrics["a_to_b_mAP@10"],
            prog_bar=True,
            sync_dist=True,
            batch_size=len(inp1_emb),
        )
        self.log(
            "mAP10-%s_%s_%s" % (modality_ba, stage, task),
            metrics["b_to_a_mAP@10"],
            prog_bar=True,
            sync_dist=True,
            batch_size=len(inp1_emb),
        )
        return metrics

    def training_step(self, batch, batch_idx):
        try:
            metrics = self.step(batch, "train")
        except Exception as e:
            print(f"Error in training: {str(e)}")
        return metrics["loss"]

    def validation_step(self, batch, batch_idx):
        metrics = self.step(batch, "validation")
        return metrics["loss"]

    def configure_optimizers(self):
        optimizer = torch.optim.AdamW(
            [
                {"params": self.model.music_projection.parameters(), "lr": self.lr},
                {"params": self.model.text_projection.parameters(), "lr": self.lr},
                {"params": self.model.music_encoder.parameters(), "lr": self.lr / 10},
                {"params": self.model.text_encoder.parameters(), "lr": self.lr / 10},
            ],
            lr=self.lr,
        )
        return {"optimizer": optimizer, "gradient_clip_val": self.gradient_clip_val}


# import sys
# import warnings
# warnings.filterwarnings("ignore", category=FutureWarning)
# import omegaconf
# import lightning as L
# from torch.utils import data
# from lightning.pytorch.callbacks import ModelCheckpoint
# from lightning.pytorch.loggers import WandbLogger

# from ditto_v2.data_loaders.multi import MultiTaskDataset, MultiTaskBatchSampler
# from ditto_v2.models.ditto import Ditto
# from ditto_v2.modules.lightning_module import DittoLitModule


# def main(cfg):
#     model = Ditto(
#         latent_dim=cfg.model.latent_dim,
#         model_path=cfg.model.model_path,
#         is_flash=True
#     )

#     # lightning module
#     lit_module = DittoLitModule(
#         model=model,
#         learning_rate=cfg.optim.learning_rate,
#     )

#     # data loaders
#     train_dataset = MultiTaskDataset(split="train")
#     valid_dataset = MultiTaskDataset(split="valid")
#     train_dataloader = data.DataLoader(
#         dataset=train_dataset,
#         batch_sampler=MultiTaskBatchSampler(train_dataset, batch_size=cfg.data.batch_size),
#         num_workers=cfg.data.num_workers,
#     )
#     validation_dataloader = data.DataLoader(
#         dataset=valid_dataset,
#         batch_sampler=MultiTaskBatchSampler(valid_dataset, batch_size=cfg.data.batch_size),
#         num_workers=cfg.data.num_workers,
#     )

#     # callbacks
#     callbacks = [
#         ModelCheckpoint(
#             save_last=True,
#             save_top_k=cfg.core.save_top_k,
#             monitor="loss_validation",
#             mode="min",
#             dirpath="/home/minz/logs/%s" % cfg.core.version,
#         )
#     ]

#     # logger
#     logger = WandbLogger(
#         name=cfg.core.version, save_dir="/app/suno/minz/wandb_logs", log_model="all"
#     )

#     # trainer
#     trainer = L.Trainer(
#         accelerator="gpu",
#         devices=cfg.core.devices,
#         num_nodes=cfg.core.num_nodes,
#         strategy="deepspeed",
#         precision=cfg.core.precision,
#         # limit_train_batches=cfg.data.limit_train,
#         profiler="simple",  # "simple" or "advanced"
#         callbacks=callbacks,
#         max_epochs=cfg.core.max_epochs,
#         logger=logger,
#         use_distributed_sampler=False,
#     )
#     trainer.fit(
#         lit_module,
#         train_dataloader,
#         validation_dataloader,
#         ckpt_path=cfg.core.ckpt_path,
#     )


# if __name__ == "__main__":
#     cfg = omegaconf.OmegaConf.load(sys.argv[1])
#     main(cfg)
