beat_only

Beat detection alone, with no bar-position term.

The objective behind the beat-only task and BeatModule: a single frame-wise sigmoid over “is there a beat here?”. It is the beat term of beat_position_loss() on its own, and shares that loss’s positive-class weighting and shift tolerance.

Functions

beat_only_loss(logits, beat_y[, pos_weight, ...])

Frame-wise weighted BCE against the beat target.

beat_only_loss(logits, beat_y, pos_weight=6.0, pos_weight_alpha=1.11, loss='bce', tolerance_frames=3, ignore_frames=None)[source]

Frame-wise weighted BCE against the beat target.

\[\mathcal{L} = \frac{1}{BT} \sum_{i,t} \ell(\hat{b}_{i,t}, b_{i,t})\]

where \(\ell\) is binary cross-entropy weighted by pos_weight on the positive class.

Parameters:
  • logits (Tensor) – Raw per-frame model output, shape (B, T) — unactivated (see musicality.models.tcn.TCNTempoNet with frame_level=True and n_outputs=1).

  • beat_y (Tensor) – Beat target channel, shape (B, T), values in [0, 1].

  • pos_weight (Tensor | float | str) – Positive-class weight, compensating for beat frames being a small fraction of all frames. A number, or "auto" to derive one per sample from the target — see beat_pos_weight().

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

  • loss (str) –

    Which objective to compare against the target.

    "bce" (the default)

    The plain weighted cross-entropy above — what every existing checkpoint was trained with, bit for bit.

    "shift_tolerant"

    The max-pooled variant from Beat This! (ISMIR 2024), which stops punishing a peak that is a frame or two off the annotation. See musicality.losses.shift_tolerance, and pair it with sigma_frames: 0 — it replaces target smearing rather than adding to it.

  • tolerance_frames (int) – Half-width, in frames, of the window the model’s peak may sit anywhere in without penalty. Read only under loss="shift_tolerant"; the default is TOLERANCE_FRAMES.

  • ignore_frames (int | None) – Half-width of the band around each beat where the negative term is switched off. None derives it as 2 * tolerance_frames. Read only under loss="shift_tolerant".

Returns:

Scalar mean loss, shape ().

Raises:

ValueError – If loss is not a known mode, or names shift tolerance with a zero radius.

Return type:

Tensor