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
|
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).
beatis supervised on every frame.one/lastare supervised on a weighted subset,w, controlled byphase_conditioning— seephase_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 (seemusicality.models.tcn.TCNTempoNetwithframe_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. Default8.0is a rough starting point, not tuned per dataset.phase_conditioning (str) – Which frames the
one/lastterms are averaged over —"mask"(default) or"beat". Seephase_weight(), including whypos_weightmust be retuned alongside it.
- Returns:
Scalar mean loss, shape
().- Return type:
Tensor