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.
Functions
|
Validate a configured beat loss and return the window radius it implies. |
|
Weighted BCE on the beat head, optionally forgiving small timing errors. |
|
Each frame takes the largest value within |
- 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 inshift_tolerant_bce().tolerance_frames (int) – The configured radius. Read only under
"shift_tolerant".
- Returns:
0under"bce", otherwisetolerance_frames.- Raises:
ValueError – If
lossis 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=0this is exactlybinary_cross_entropy_with_logits()underbeat_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_framesand \(\rho\) =ignore_frames.Two halves that would otherwise contradict each other. The first says “the highest prediction within
±rof the annotation should be high”, which accepts a peakrframes 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 defaultignore_frames = 2 * tolerance_framescomes from.Warning
That default is too wide for our fastest corpus. It ignores
4r + 1frames per beat, and at 43.07 fps jtd’s 193 BPM leaves only 13.4 frames between beats — sor = 3retains 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. Henceignore_framesbeing its own knob rather than derived:r = 3(to keep ±70 ms) withignore_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ρ = 6againstρ = 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 — seebeat_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ρ = 4the 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
rof the prediction window, in frames.0disables shift tolerance entirely.ignore_frames (int | None) – Half-width
ρof the band around each annotation where the negative term is switched off.Noneuses [BT24]’s2 * tolerance_frames. Read only whentolerance_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
radiusframes of it.Stride 1, so the time axis is unchanged. The padding is
-infrather than a repeat or a wrap, which is what we want at the clip edges — a crop boundary is not a beat, andalign_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.
0is 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