import torch
import lightning as L
import numpy as np
from torch import nn
from sklearn import metrics


class KeyLitModule(L.LightningModule):
    def __init__(
            self,
            model,
            dataset="key_aug_tency",
            learning_rate=1e-4,
        ):
        super().__init__()
        self.lr = learning_rate
        self.model = model
        self.loss_function = nn.CrossEntropyLoss()
        self.eval_logits, self.eval_keys = [], []
        self.save_hyperparameters(ignore=["model"])

    def training_step(self, batch, batch_idx):
        wav, keys = batch
        out = self.model(wav)
        loss = self.loss_function(out, keys)
        self.log("train_loss", loss, prog_bar=True, sync_dist=True)
        return loss
    
    def validation_step(self, batch, batch_idx):
        wav, keys = batch
        # out = self.model(wav[0]).mean(dim=0).unsqueeze(0)
        out = self.model(wav)
        loss = self.loss_function(out, keys)
        self.eval_logits.append(out.float().detach().cpu())
        self.eval_keys.append(keys.long().detach().cpu())
        return loss

    def on_validation_epoch_end(self):
        logits = torch.cat(self.eval_logits, dim=0)
        keys = torch.cat(self.eval_keys, dim=0)

        # get loss
        loss = self.loss_function(logits, keys)

        # get accuracy
        prd = logits.argmax(dim=1)
        accuracy = metrics.accuracy_score(keys, prd)
        print("accuracy: %.4f" % accuracy)
    
        # log
        self.log("valid_loss", loss.cuda(), sync_dist=True)
        self.log("valid_acc", torch.tensor(accuracy).cuda(), sync_dist=True)
        self.eval_logits, self.eval_keys = [], []

    def configure_optimizers(self):
        return torch.optim.AdamW(self.model.parameters(), lr=self.lr)