import torch
import lightning as L

from musicfm.modules.lr_scheduler import TriStageLRScheduler


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 MusicFMLitModule(L.LightningModule):
    def __init__(
            self,
            model,
            learning_rate=1e-4,
            warmup_steps=30_000,
            hold_steps=270_000,
            decay_steps=50_000,
        ):
        super().__init__()
        self.lr = learning_rate
        self.warmup_steps = warmup_steps
        self.hold_steps = hold_steps
        self.decay_steps = decay_steps

        self.model = model
        print_model_params(model)
        self.save_hyperparameters(ignore=["model"])

    def step(self, batch, stage):
        _, _, losses, accuracies = self.model(batch)
        losses["overall"] = torch.mean(torch.stack([losses[key] for key in losses.keys()]))
        accuracies["overall"] = torch.mean(torch.stack([accuracies[key] for key in accuracies.keys()]))
        for key in losses.keys():
            prog_bar = True if key == "overall" else False
            self.log("loss_%s/%s" % (key, stage), losses[key], prog_bar=prog_bar, sync_dist=True)
            self.log("acc_%s/%s" % (key, stage), accuracies[key], prog_bar=prog_bar, sync_dist=True)
        return losses, accuracies

    def training_step(self, batch, batch_idx):
        loss, acc = self.step(batch, "train")
        return loss["overall"]

    def validation_step(self, batch, batch_idx):
        is_train = False
        loss, acc = self.step(batch, "validation")
        return loss["overall"]

    def configure_optimizers(self):
        optimizer = torch.optim.AdamW(self.model.parameters(), lr=self.lr)
        # optimizer = torch.optim.Adam(self.model.parameters(), lr=self.lr)
        scheduler = {
            "scheduler": TriStageLRScheduler(optimizer, self.lr, self.warmup_steps, self.hold_steps, self.decay_steps),
            "interval": "step",
            "name": "learning_rate",
        }
        return [optimizer], [scheduler]