beat_phase_module¶
PyTorch Lightning module for frame-level beat-phase detection (beat/one/last).
Functions
|
Crop the longer of |
Classes
|
LightningModule wrapping a frame-level beat-phase model (beat/one/last). |
- class BeatPhaseModule(model, pos_weight=8.0, phase_conditioning='mask', group_size=None, pos_weight_alpha=1.11, position_norm='global', loss='bce', tolerance_frames=3, ignore_frames=None, lr=0.001, weight_decay=0.0001, threshold=0.5, balanced=True, check_val_every_n_epoch=1, task='beat_phase')[source]¶
Bases:
LightningModuleLightningModule wrapping a frame-level beat-phase model (beat/one/last).
Two output parameterizations, selected by
group_size:group_size=None(default): the original three-channel head —beat/one/lastas independent sigmoids, trained withbeat_phase_loss(). Pairs withBeatDataset’starget_layout="one_last".group_size=G:1 + Gchannels —beatas a sigmoid, then a softmax over theGbar positions, trained withbeat_position_loss(). Pairs withtarget_layout="positions". Positions 2..G-1 gain their own supervised logits, so “is this a 1 or a 3?” becomes a question the model is actually asked — see docs/beat_phase_improvement_review.md section 3.
Either way
n_outputsis forced regardless of what the config says, mirroring howTempoModuleoverrides it for classification mode.- Parameters:
model (DictConfig) – DictConfig for instantiating the backbone (e.g.
TCNTempoNet).pos_weight (float | list[float] | str) – Positive-class weight passed to
beat_phase_loss(). Scalar (shared across heads) or a 3-element sequence (per-head). Withgroup_sizeset there is only one BCE head, so a sequence is rejected up front rather than left to fail as a broadcast error on the first batch."auto"derives it per sample from the target (beat_pos_weight());group_sizeonly.pos_weight_alpha (float) – Scale on a derived
pos_weight. Read only whenpos_weight == "auto".position_norm (str) –
"global"(default) or"per_item"— howbeat_position_loss()averages the bar-position term.group_sizeonly.phase_conditioning (str) – Passed to
beat_phase_loss()—"mask"supervises the one/last heads on every frame,"beat"only on frames at or near a beat, which is wheremusicality.postprocess.label_bar_position()actually reads them. Must be retuned together withpos_weight: the class imbalance the latter compensates for largely disappears under"beat".loss (str) –
"bce"(the default, the pre-existing objective bit for bit) or"shift_tolerant". Saved to the checkpoint’s hyperparameters, so a checkpoint records which objective trained it rather than leaving it to be inferred from a radius.group_sizeonly — the legacy one/last path ignores it.tolerance_frames (int) – Half-width, in frames, of the timing error both heads are forgiven. Read only under
loss="shift_tolerant", which also wantssigma_frames: 0in the config. It means two different things to the two heads — the beat term is forgiven a shifted peak, the position term is supervised across a wider window; seebeat_position_loss().group_sizeonly.ignore_frames (int | None) – Half-width of the band around each beat where the beat term’s negative half is switched off.
Nonederives it as2 * tolerance_frames.group_sizeonly.lr (float) – Learning rate.
weight_decay (float) – L2 regularisation.
threshold (float) – Sigmoid/target threshold used only for the logged metrics, not for the loss itself. Doubles as the peak-picking threshold for
{stage}/f_beat.balanced (bool) – Passed through to
frame_accuracy()for the loggedacc_one/acc_lastmetrics. Defaults toTruesince one/last frames are a small minority and a pooled mean is dominated by the true-negative rate. No longer affects the beat head, which is scored bypeak_f_measure(); kept as a constructor parameter regardless, since it is saved in every existing checkpoint’s hyperparameters and is passed back on load.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.group_size (int | None) – Number of bar positions.
Nonekeeps the original one/last sigmoid head; an integer switches to agroup_size-way softmax over positions. Saved to the checkpoint’s hyperparameters, soload_module()reconstructs the right head and downstream code can tell the two apart.task (str) – Saved into the checkpoint’s hyperparameters for
detect_task()to read back at eval/inference time. Always"beat_phase"for this class; exists as a parameter (rather than hardcoded) so the training run’stask:setting is the visible, single source of truth for what a checkpoint is.
- align_time(logits, target)[source]¶
Crop the longer of
(logits, target)along the time axis to match the shorter.TCNTempoNet’s frame-level output length (fromtorchaudio.MelSpectrogram) isn’t guaranteed to exactly equalBeatDataset’sn_frames— off-by-one in practice. Cropping (rather than padding) keeps both sides comparing only real, non-fabricated frames.- Parameters:
logits (Tensor) – Model output, shape
(B, 3, T_logits).target (Tensor) – Ground-truth target, shape
(B, 4, T_target).
- Returns:
(logits, target)both cropped to(..., min(T_logits, T_target)).- Return type:
tuple[Tensor, Tensor]