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

beat_pos_weight(beat_y, pos_weight[, alpha, ...])

Positive-class weight for a beat BCE term — passed through, or derived per sample from the target when pos_weight is "auto".

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_weight is "auto".

A fixed pos_weight is 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.0 is exact inverse-frequency weighting.

  • neg_weight (Tensor | None) – Per-frame weight the negative term actually carries, shape (B, T), replacing \(\bar{n}_i\) above. Only shift_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. None uses 1 - 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) inside binary_cross_entropy_with_logits().

Return type:

Tensor