import torch import lightning as L import numpy as np from torch import nn from sklearn import metrics class TaggingLitModule(L.LightningModule): def __init__( self, model, learning_rate=1e-4, ): super().__init__() self.lr = learning_rate self.model = model self.tags = np.load( "/app/suno/data/audio_mono_24khz/ytm_tagged/metadata/tags.npy" ) self.loss_function = nn.BCEWithLogitsLoss() self.eval_logits, self.eval_binaries = [], [] self.save_hyperparameters(ignore=["model"]) def training_step(self, batch, batch_idx): wav, binaries = batch out = self.model(wav) loss = self.loss_function(out, binaries) self.log("train_loss", loss, prog_bar=True, sync_dist=True) return loss def validation_step(self, batch, batch_idx): wav, binaries = batch logits = self.model(wav) loss = self.loss_function(logits, binaries) self.eval_logits.append(logits.float().detach().cpu()) self.eval_binaries.append(binaries.float().detach().cpu()) return loss def test_step(self, batch, batch_idx): wav, binaries = batch logits = self.model(wav) loss = self.loss_function(logits, binaries) self.log("test_loss", loss, prog_bar=True, sync_dist=True) self.eval_logits.append(logits.float().detach().cpu()) self.eval_binaries.append(binaries.float().detach().cpu()) return loss def on_validation_epoch_end(self): logits = torch.cat(self.eval_logits, dim=0) binaries = torch.cat(self.eval_binaries, dim=0) loss = self.loss_function(logits, binaries) roc_auc, pr_auc = self.get_auc_scores(logits, binaries) self.log("valid_loss", loss.cuda(), sync_dist=True) self.log("valid_roc_auc", roc_auc.cuda(), sync_dist=True) self.log("valid_pr_auc", pr_auc.cuda(), sync_dist=True) self.eval_logits, self.eval_binaries = [], [] def on_test_epoch_end(self): logits = torch.cat(self.eval_logits, dim=0) binaries = torch.cat(self.eval_binaries, dim=0) loss = self.loss_function(logits, binaries) roc_auc, pr_auc = self.get_auc_scores(logits, binaries) self.log("test_loss", loss.cuda(), sync_dist=True) self.log("test_roc_auc", roc_auc.cuda(), sync_dist=True) self.log("test_pr_auc", pr_auc.cuda(), sync_dist=True) self.eval_logits, self.eval_binaries = [], [] def configure_optimizers(self): optimizer = torch.optim.AdamW( [ {"params": self.model.frontend.parameters(), "lr": self.lr / 10}, {"params": self.model.projection.parameters(), "lr": self.lr}, ], lr=self.lr, ) return [optimizer] def get_auc_scores(self, logits, targets): try: roc_auc = metrics.roc_auc_score( targets, nn.Sigmoid()(logits), average="macro" ) pr_auc = metrics.average_precision_score( targets, nn.Sigmoid()(logits), average="macro" ) print("roc_auc: %.4f" % roc_auc) print("pr_auc: %.4f" % pr_auc) # tag-wise score for debugging roc_aucs = metrics.roc_auc_score( targets, nn.Sigmoid()(logits), average=None ) pr_aucs = metrics.average_precision_score( targets, nn.Sigmoid()(logits), average=None ) for i in range(len(self.tags)): print("%s: %.4f, %.4f" % (self.tags[i], roc_aucs[i], pr_aucs[i])) except ValueError as e: roc_auc, pr_auc = 0, 0 print("auc not available yet") return torch.tensor(roc_auc), torch.tensor(pr_auc)