beat_phase_module

PyTorch Lightning module for frame-level beat-phase detection (beat/one/last).

Functions

align_time(logits, target)

Crop the longer of (logits, target) along the time axis to match the shorter.

Classes

BeatPhaseModule(model[, pos_weight, ...])

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: LightningModule

LightningModule 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/last as independent sigmoids, trained with beat_phase_loss(). Pairs with BeatDataset’s target_layout="one_last".

  • group_size=G: 1 + G channels — beat as a sigmoid, then a softmax over the G bar positions, trained with beat_position_loss(). Pairs with target_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_outputs is forced regardless of what the config says, mirroring how TempoModule overrides 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). With group_size set 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_size only.

  • pos_weight_alpha (float) – Scale on a derived pos_weight. Read only when pos_weight == "auto".

  • position_norm (str) – "global" (default) or "per_item" — how beat_position_loss() averages the bar-position term. group_size only.

  • 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 where musicality.postprocess.label_bar_position() actually reads them. Must be retuned together with pos_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_size only — 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 wants sigma_frames: 0 in 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; see beat_position_loss(). group_size only.

  • ignore_frames (int | None) – Half-width of the band around each beat where the beat term’s negative half is switched off. None derives it as 2 * tolerance_frames. group_size only.

  • 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 logged acc_one/acc_last metrics. Defaults to True since 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 by peak_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). The ReduceLROnPlateau scheduler needs this as its frequency — Lightning otherwise tries to step it (and read val/loss) every epoch regardless of how often validation runs, raising MisconfigurationException on any epoch without a fresh value.

  • group_size (int | None) – Number of bar positions. None keeps the original one/last sigmoid head; an integer switches to a group_size-way softmax over positions. Saved to the checkpoint’s hyperparameters, so load_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’s task: setting is the visible, single source of truth for what a checkpoint is.

configure_optimizers()[source]
forward(wav)[source]
Parameters:

wav (Tensor)

Return type:

Tensor

training_step(batch, batch_idx)[source]
validation_step(batch, batch_idx)[source]
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 (from torchaudio.MelSpectrogram) isn’t guaranteed to exactly equal BeatDataset’s n_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]