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)