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]