pos_weight¶
Positive-class weighting for the frame-wise beat BCE term.
Beat frames are a small fraction of all frames, so the beat head needs a
pos_weight to stop it collapsing to “never a beat”. How small a fraction
depends on the tempo, which a single configured number cannot track — hence
pos_weight="auto", which derives the weight per sample from the target
already in hand. Shared by every loss here that carries a beat head:
beat_only_loss() and
beat_position_loss().
See docs/beat_phase_pos_weight_notes.md for the measurements.
Functions
|
Positive-class weight for a beat BCE term — passed through, or derived per sample from the target when |
- beat_pos_weight(beat_y, pos_weight, alpha=1.11, neg_weight=None)[source]¶
Positive-class weight for a beat BCE term — passed through, or derived per sample from the target when
pos_weightis"auto".A fixed
pos_weightis only correct at one tempo. The beat target is a Gaussian smeared to peak 1.0, so its mass per beat is a constant ~3.75 frames regardless of tempo, while the beat period is not: 20.7 frames at ballroom’s 125 BPM, 13.4 at jtd’s 193, 46.1 at classical’s 10th percentile of 56. The true neg:pos ratio therefore spans 2.6–11.3 across the corpora we train on, against a single configured 5.Since the ratio is just a function of the target already in hand, derive it rather than tune it:
\[w_i = \alpha \, \frac{\bar{n}_i}{\bar{b}_i}, \qquad \bar{b}_i = \frac{1}{T} \sum_t b_{i,t}, \qquad \bar{n}_i = 1 - \bar{b}_i\]Deriving it per sample also means it tracks time-stretch augmentation, which silently invalidates a hand-tuned value on every augmented clip.
- Parameters:
beat_y (Tensor) – Beat target channel, shape
(B, T), values in[0, 1].pos_weight (Tensor | float | str) – A number (or tensor) to pass through unchanged, or the string
"auto"to derive one per sample.alpha (float) – Scale on the derived ratio. Defaults to
AUTO_POS_WEIGHT_ALPHA;1.0is exact inverse-frequency weighting.neg_weight (Tensor | None) – Per-frame weight the negative term actually carries, shape
(B, T), replacing \(\bar{n}_i\) above. Onlyshift_tolerant_bce()passes it, because that loss ignores negatives near each annotation and the ratio has to count the frames that survive rather than every non-beat frame. On ballroom with sharp targets that is 13.2 against the 22.1 the raw target implies — a 1.7x difference the beat head would otherwise absorb as over-weighted positives.Noneuses1 - mean(beat_y), which is the plain BCE’s(1 - y).
- Returns:
Scalar tensor when passed through, shape
(B, 1)when derived — which broadcasts against(B, T)insidebinary_cross_entropy_with_logits().- Return type:
Tensor