shift_tolerance

Forgiving the beat head a few frames of timing error.

Beat annotations are not precise down to a frame — annotators disagree, players are not synchronous, and human perception has limits. The evaluation already accepts this: beat_f_measure() scores a prediction as correct anywhere within ±70 ms. The training loss does not, so it punishes predictions the metric would have accepted. [BT24] fixes that by comparing the max-pooled prediction to the label instead of the prediction itself.

The side effect is the point. Under a plain BCE the cheapest way to cover an imprecise annotation is a wide, blurred peak; here only the largest prediction in the window is read, so a single sharp peak is optimal — and sharp peaks are what musicality.postprocess.pick_peaks() needs. Our alternative so far has been to blur the target instead (sigma_frames in BeatDataset), which [BT24] names and rejects as mitigating slow convergence without helping with the blur.

The two are alternatives, not additions: ±3 frames of pooling on top of the ±4 frames sigma_frames: 1.5 smears over is ±162 ms of combined tolerance against a ±70 ms metric. Pair a non-zero tolerance_frames with sigma_frames: 0. See plans/09_lessons_from_literature.md §2.1.

[BT24] (1,2,3,4,5)

Foscarin, Schlüter and Widmer, “Beat this! Accurate beat tracking without DBN postprocessing”, ISMIR 2024, §3.3.

Functions

resolve_tolerance(loss, tolerance_frames)

Validate a configured beat loss and return the window radius it implies.

shift_tolerant_bce(logits, beat_y[, ...])

Weighted BCE on the beat head, optionally forgiving small timing errors.

sliding_windowed_max(x, radius)

Each frame takes the largest value within radius frames of it.

resolve_tolerance(loss, tolerance_frames)[source]

Validate a configured beat loss and return the window radius it implies.

The named mode is what a config and a checkpoint carry; the radius is what the maths needs. Keeping the translation here means the two beat losses share one definition of what loss: may say, and neither has to encode “classical” as a magic zero.

Parameters:
  • loss (str) – "bce" for the plain weighted cross-entropy every existing checkpoint was trained with, or "shift_tolerant" for the max-pooled variant in shift_tolerant_bce().

  • tolerance_frames (int) – The configured radius. Read only under "shift_tolerant".

Returns:

0 under "bce", otherwise tolerance_frames.

Raises:

ValueError – If loss is not a known mode, or if it asks for shift tolerance with a radius of zero — a config that names the objective and then disables it is a mistake, not a preference.

Return type:

int

shift_tolerant_bce(logits, beat_y, pos_weight=6.0, pos_weight_alpha=1.11, tolerance_frames=0, ignore_frames=None)[source]

Weighted BCE on the beat head, optionally forgiving small timing errors.

With tolerance_frames=0 this is exactly binary_cross_entropy_with_logits() under beat_pos_weight() — bit for bit, so it is a drop-in for the beat term of any loss here. Above zero it becomes [BT24]’s shift-tolerant weighted BCE:

\[\mathcal{L}_{st} = -\frac{1}{BT} \sum_{i,t} w_i \, b_{i,t} \log m_r(\hat{b})_{i,t} + \big(1 - m_{\rho}(b)_{i,t}\big) \log\big(1 - m_r(\hat{b})_{i,t}\big)\]

with \(m_k\) the max over a ±k-frame window, \(r\) = tolerance_frames and \(\rho\) = ignore_frames.

Two halves that would otherwise contradict each other. The first says “the highest prediction within ±r of the annotation should be high”, which accepts a peak r frames off the annotation. The second says “every other frame should be low” — including that same frame. So the negative term is switched off near each annotation: a peak at +r, pooled over ±r, reaches +2r, which is where [BT24]’s default ignore_frames = 2 * tolerance_frames comes from.

Warning

That default is too wide for our fastest corpus. It ignores 4r + 1 frames per beat, and at 43.07 fps jtd’s 193 BPM leaves only 13.4 frames between beats — so r = 3 retains 3.8% of frames as negatives, on 63.8% of the merged split’s tracks. The negative term all but vanishes, and a model that fires everywhere scores well. Hence ignore_frames being its own knob rather than derived: r = 3 (to keep ±70 ms) with ignore_frames = 4 (to keep a third of jtd’s frames) accepts a mild contradiction at the edge of the window, which biases peaks towards its centre — no bad thing. Frames surviving the band, at ρ = 6 against ρ = 4: jtd 3.8% against 33.4%, ballroom 37.7% against 56.9%, rwc_classical’s 10th percentile 71.7% against 80.4%.

Parameters:
  • logits (Tensor) – Raw per-frame model output, shape (B, T) — unactivated.

  • beat_y (Tensor) – Beat target channel, shape (B, T), values in [0, 1]. Shift tolerance assumes this is sharp (sigma_frames: 0); see the module docstring on why the two do not compose.

  • pos_weight (Tensor | float | str) – Positive-class weight, or "auto" to derive one per sample — see beat_pos_weight(). Derived here rather than by the caller because "auto" has to see which negatives survive the ignore band: at ballroom’s 125 BPM with sharp targets and ρ = 4 the honest ratio is 13.2, against 22.1 from the raw target alone.

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

  • tolerance_frames (int) – Half-width r of the prediction window, in frames. 0 disables shift tolerance entirely.

  • ignore_frames (int | None) – Half-width ρ of the band around each annotation where the negative term is switched off. None uses [BT24]’s 2 * tolerance_frames. Read only when tolerance_frames > 0.

Returns:

Scalar mean loss, shape ().

Raises:

ValueError – If either radius is negative.

Return type:

Tensor

sliding_windowed_max(x, radius)[source]

Each frame takes the largest value within radius frames of it.

Stride 1, so the time axis is unchanged. The padding is -inf rather than a repeat or a wrap, which is what we want at the clip edges — a crop boundary is not a beat, and align_time() has already trimmed the clip by the time a loss sees it.

Parameters:
  • x (Tensor) – Tensor whose last axis is time, any leading shape — (B, T) for a single channel, (B, C, T) for a block of them.

  • radius (int) – Half-width of the window in frames. 0 is the identity, which is how every caller’s default reduces to its previous behaviour.

Returns:

Tensor of the same shape as x.

Return type:

Tensor