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 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 (oneandlast), 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.beatstays 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_framesis 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 (seephase_conditioning).No
pos_weightis needed on the position term: bar positions occur equally often, so the softmax is already balanced.pos_weighthere is a scalar for thebeathead alone.- Parameters:
logits (Tensor) – Raw per-frame model output, shape
(B, 1 + G, T)— beat first, then one logit per bar position (seemusicality.models.tcn.TCNTempoNetwithframe_level=Trueandn_outputs=1 + G).target (Tensor) – Ground-truth target, shape
(B, 2 + G, T)— beat, the normalized position block, then mask. Built byBeatDatasetwithtarget_layout="positions".pos_weight (Tensor | float | str) – Positive-class weight for the
beatBCE term only. A number, or"auto"to derive one per sample from the target — seebeat_pos_weight().phase_conditioning (str) –
"beat"(default) weights the position term bymask * beat, so it is optimized only where a beat actually is — which is wheremusicality.postprocessreads it."mask"weights bymaskalone, supervising every frame. Seephase_weight()for why the former matters.pos_weight_alpha (float) – Scale on the derived
pos_weight. Read only whenpos_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 exactly1/n_validwhatever 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
beatterm uses —"bce"(the default, the pre-existing loss bit for bit) or"shift_tolerant". It also switches thepositionterm’s widening on, since the two are the same decision seen from either head; seetolerance_frames."shift_tolerant"wantssigma_frames: 0alongside it.tolerance_frames (int) –
Half-width, in frames, of the timing error both heads are forgiven. Read only under
loss="shift_tolerant"; the default isTOLERANCE_FRAMES. It means two different things to the two terms, because the decoder reads them two different ways:The
beatterm becomesshift_tolerant_bce(). That head is scanned over time bymusicality.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
positionterm 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 maderframes wide.
See
musicality.losses.shift_toleranceon why smearing and tolerance do not compose.ignore_frames (int | None) – Half-width of the band around each beat where the
beatterm’s negative half is switched off.Nonederives it as2 * tolerance_frames. Read only underloss="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 byBeatPhaseModule.
- Returns:
Scalar mean loss, shape
()— or, withreturn_terms, the pair of scalars that sums to it.- Return type:
Tensor | tuple[Tensor, Tensor]