tempo_module¶
PyTorch Lightning module for tempo estimation.
Classes
|
LightningModule wrapping a tempo regression or classification model. |
- class TempoModule(model, loss='absolute', classification=None, lr=0.001, weight_decay=0.0001, check_val_every_n_epoch=1)[source]¶
Bases:
LightningModuleLightningModule wrapping a tempo regression or classification model.
- Three loss modes are supported:
"absolute"— plain MAE between predicted and true BPM."relative"— octave-invariant MAE."classification"— softmax over BPM bins with a Gaussian soft target. Requires theclassificationconfig section.
For classification mode the model’s
n_outputsis overridden toclassification.n_binsautomatically.- Parameters:
model (DictConfig) – DictConfig for instantiating the backbone.
loss (str) – Loss name —
"absolute","relative", or"classification".classification (DictConfig | None) – Required when
loss == "classification". Must havebpm_min,bpm_max,n_bins,sigma.lr (float) – Learning rate.
weight_decay (float) – L2 regularisation.
check_val_every_n_epoch (int) – How often the trainer actually runs validation (
cfg.trainer.check_val_every_n_epoch). TheReduceLROnPlateauscheduler needs this as itsfrequency— Lightning otherwise tries to step it (and readval/loss) every epoch regardless of how often validation runs, raisingMisconfigurationExceptionon any epoch without a fresh value.