beat_phase

Bar position as two independent binary detectors (one and last).

The original beat-phase objective, superseded by beat_position_loss(). Kept because checkpoints trained under it are still evaluated and compared; new runs should use the position-softmax loss instead.

Functions

beat_phase_loss(logits, target[, ...])

Masked multi-head frame-wise BCE loss for beat-phase detection.

beat_phase_loss(logits, target, pos_weight=8.0, phase_conditioning='mask')[source]

Masked multi-head frame-wise BCE loss for beat-phase detection.

Sums three per-frame binary-cross-entropy terms (beat, one, last). beat is supervised on every frame. one/last are supervised on a weighted subset, w, controlled by phase_conditioning — see phase_weight().

Both phase terms are normalized by the sum of those weights rather than the total frame count. That keeps them on the same scale as the beat term regardless of how much weight there is, so neither a batch light on position-annotated tracks nor a slow track with few beats has its phase loss silently shrink toward zero.

\[\mathcal{L} = \underbrace{\frac{1}{BT} \sum_{i,t} \ell(\hat{b}_{i,t}, b_{i,t})}_{\text{beat}} + \underbrace{\frac{\sum_{i,t} w_{i,t} \, \ell(\hat{o}_{i,t}, o_{i,t})}{\sum_{i,t} w_{i,t}}}_{\text{one}} + \underbrace{\frac{\sum_{i,t} w_{i,t} \, \ell(\hat{l}_{i,t}, l_{i,t})}{\sum_{i,t} w_{i,t}}}_{\text{last}}\]

where \(\ell\) is per-frame weighted binary cross-entropy (with pos_weight) and \(w\) is the phase weight.

Parameters:
  • logits (Tensor) – Raw per-frame model output, shape (B, 3, T) — beat/one/last, unactivated (see musicality.models.tcn.TCNTempoNet with frame_level=True).

  • target (Tensor) – Ground-truth target, shape (B, 4, T) — beat/one/last/mask channels.

  • pos_weight (Tensor | float) – Positive-class weight applied to every head’s BCE term, compensating for beat/one/last frames being a small fraction of all frames. Scalar (shared across heads) or shape (3,) for a per-head weight. Default 8.0 is a rough starting point, not tuned per dataset.

  • phase_conditioning (str) – Which frames the one/last terms are averaged over — "mask" (default) or "beat". See phase_weight(), including why pos_weight must be retuned alongside it.

Returns:

Scalar mean loss, shape ().

Return type:

Tensor