beat_position

Bar position as a softmax over the G positions in a bar.

The current beat-phase objective, paired with BeatPhaseModule and a backbone emitting 1 + G frame-level channels.

Functions

beat_position_loss(logits, target[, ...])

Beat BCE plus a softmax cross-entropy over bar position.

beat_position_loss(logits, target, pos_weight=5.0, phase_conditioning='beat', pos_weight_alpha=1.11, position_norm='global', loss='bce', tolerance_frames=3, ignore_frames=None, return_terms=False)[source]

Beat BCE plus a softmax cross-entropy over bar position.

The successor to beat_phase_loss(). That loss models bar position as two independent binary detectors (one and last), which leaves the positions in between with identical supervision — negative on both heads — so the model is never asked the discriminative question the metric measures, “is this beat a 1 or a 3?”. Here every position gets its own logit and they compete inside a single softmax, so raising the score for position 1 necessarily lowers position 3.

beat stays an independent sigmoid: “is there a beat here?” is a genuine binary question over time, not a pick-one-of-G, and folding it into the softmax would couple a head that already works to one that doesn’t.

\[\mathcal{L} = \underbrace{\frac{1}{BT} \sum_{i,t} \ell(\hat{b}_{i,t}, b_{i,t})}_{\text{beat}} - \underbrace{\frac{\sum_{i,t} w_{i,t} \sum_{p} q_{i,t,p} \log \hat{q}_{i,t,p}}{\sum_{i,t} w_{i,t}}}_{\text{position}}\]

where \(\ell\) is weighted binary cross-entropy — shift-tolerant once tolerance_frames is set — \(q\) is the target’s normalized position block, \(\hat{q}\) the softmax over the model’s position logits, and \(w\) the per-frame phase weight (see phase_conditioning).

No pos_weight is needed on the position term: bar positions occur equally often, so the softmax is already balanced. pos_weight here is a scalar for the beat head alone.

Parameters:
  • logits (Tensor) – Raw per-frame model output, shape (B, 1 + G, T) — beat first, then one logit per bar position (see musicality.models.tcn.TCNTempoNet with frame_level=True and n_outputs=1 + G).

  • target (Tensor) – Ground-truth target, shape (B, 2 + G, T) — beat, the normalized position block, then mask. Built by BeatDataset with target_layout="positions".

  • pos_weight (Tensor | float | str) – Positive-class weight for the beat BCE term only. A number, or "auto" to derive one per sample from the target — see beat_pos_weight().

  • phase_conditioning (str) – "beat" (default) weights the position term by mask * beat, so it is optimized only where a beat actually is — which is where musicality.postprocess reads it. "mask" weights by mask alone, supervising every frame. See phase_weight() for why the former matters.

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

  • position_norm (str) –

    How the position term is averaged.

    • "global" (default): one weighted mean over the whole batch. That makes it a micro-average over beats, so a clip’s influence is proportional to how many beats it happens to contain — and a 16 s crop holds 51 beats at 193 BPM but 15 at 56. Measured on the merged split, jtd takes 73.7% of the position gradient against 63.8% of the tracks, while ballroom gets 18.0% against 24.1%.

    • "per_item": normalize each clip by its own weight first, then average over clips. Every annotated clip carries exactly 1/n_valid whatever its tempo, which is a macro-average over tracks — the same shape as the per-genre metric this is graded by.

    See plans/04_beat_phase_generalization_and_data_prep.md §2.6a.

  • loss (str) – Which objective the beat term uses — "bce" (the default, the pre-existing loss bit for bit) or "shift_tolerant". It also switches the position term’s widening on, since the two are the same decision seen from either head; see tolerance_frames. "shift_tolerant" wants sigma_frames: 0 alongside it.

  • tolerance_frames (int) –

    Half-width, in frames, of the timing error both heads are forgiven. Read only under loss="shift_tolerant"; the default is TOLERANCE_FRAMES. It means two different things to the two terms, because the decoder reads them two different ways:

    • The beat term becomes shift_tolerant_bce(). That head is scanned over time by musicality.postprocess.pick_peaks(), so where its peak sits is the answer, and tolerance means forgiving a peak that is a frame or two off.

    • The position term is widened instead — both the gate and the target. That head is never scanned: the decoder rounds a beat time to a frame and reads that one column. Its answer is a label, not a time, so there is no peak to forgive; what it needs is to be right across the whole window the lookup might land in, which shift tolerance on the beat head has just made r frames wide.

    See musicality.losses.shift_tolerance on why smearing and tolerance do not compose.

  • ignore_frames (int | None) – Half-width of the band around each beat where the beat term’s negative half is switched off. None derives it as 2 * tolerance_frames. Read only under loss="shift_tolerant", and it does not touch the position term, which has no negative class.

  • return_terms (bool) – Return the two terms separately instead of their sum, as (beat_term, position_term). The sum is what optimisation needs; the split is what tells a rising loss apart from a rising error — position CE can climb on overconfidence alone while beat BCE and position accuracy both improve (plans/07 §1.3, measured between epochs 79 and 107 of v6). Logged per epoch by BeatPhaseModule.

Returns:

Scalar mean loss, shape () — or, with return_terms, the pair of scalars that sums to it.

Return type:

Tensor | tuple[Tensor, Tensor]