beat_dataset¶
PyTorch Dataset for this project’s own beat/bar-position-annotated datasets (see docs/source/data.rst’s “Data format” section).
Functions
|
Fold an annotated bar-position cycle onto |
|
Smear a 0/1 spike train into a soft target with a Gaussian bump per event. |
|
Return |
|
Channel names for the |
Classes
|
Dataset returning waveforms and frame-level beat-phase targets, read entirely from this project's own tracks/+annotations/ format. |
- class BeatDataset(name=None, data_home=None, *, refs=None, sample_rate=22050, duration=10.0, hop_length=512, sigma_frames=1.5, group_size=4, binary_only=False, random_crop=False, target_layout='one_last')[source]¶
Bases:
DatasetDataset returning waveforms and frame-level beat-phase targets, read entirely from this project’s own tracks/+annotations/ format.
Loads every track with a migrated
.beatsannotation and a resolvabletracks/<id>.wav.The target is a 4-channel tensor of shape
(4, n_frames)wheren_frames = n_samples // hop_length:beat— any beat.one— position 1 of the group (the downbeat, for the defaultgroup_size=4bar-position case).last— the last beat of the group (bar position 4 forgroup_size=4; e.g. phrase position 8 for a phrase-annotated dataset withgroup_size=8).mask— constant 1.0 across all frames if this track carries usable position annotations, else constant 0.0. Two cases give 0.0: no position annotation at all (e.g.rwc_popular), and a bar length that cannot be folded ontogroup_size(seefold_positions()). Both still contribute theirbeatchannel; the position channels should be excluded from the loss for those tracks via this mask.
Each channel is Gaussian-smeared (see
gaussian_smear()) rather than a hard 0/1 spike.- Parameters:
name (str | None) – Dataset name (e.g.
"ballroom","swing") — must already be migrated to this project’s own format (seetools/migrate_mirdata_dataset.py/tools/migrate_rwc_genre.py). Mutually exclusive with refs.data_home (Path | None) – Dataset directory. Defaults to
DATA_DIR/<name>(DATA_DIRfrommusicality.dataformats). Ignored if refs is given.refs (list[TrackRef] | None) – Explicit list of tracks to load, bypassing name/data_home resolution entirely — e.g. tracks pulled from several source datasets (see
load_refs()). Mutually exclusive with name.sample_rate (int) – Target sample rate. Audio is resampled if needed.
duration (float) – Clip duration in seconds. Longer clips are truncated, shorter clips are zero-padded.
hop_length (int) – Frame hop size in samples used to build the frame targets.
sigma_frames (float) – Gaussian smearing width, in frames, applied to each target channel.
group_size (int) – Number of beats per group that the target counts across —
4(default) for bar-position (1-4) datasets,8for a phrase-position (1-8) dataset. Annotations counting a longer bar are folded onto1..group_sizeat load time (seefold_positions()), so a track annotated 1-8 trains agroup_size=4head as two bars of four rather than falling off the end of the channel list. Tracks whose bar length is not a multiple ofgroup_sizekeep their beats but have their position supervision masked off.binary_only (bool) – If
True, drop tracks whose beats-per-bar (the annotated position cycle length) isn’t a multiple of 2 — e.g. ballroom’s waltz/Viennese waltz tracks, which cycle1, 2, 3in triple meter — as well as tracks with no position annotation at all, since their meter can’t be confirmed. Independent ofgroup_size: a track only needs an even beats-per-bar count, not one equal togroup_size.random_crop (bool) – If
True, draw theduration-second window at a random offset into the track on every access. Use for training (so the model doesn’t just memorize one fixed window per track across epochs). IfFalse, always take a fixed window from the middle of the track — deterministic/reproducible for validation, and more representative than the start, which is often a sparse intro.target_layout (str) –
Which phase target to build.
"one_last"(default): the 4-channel target described above — two independent binary detectors for positions1andgroup_size, with positions in between carrying no supervision at all."positions": a(2 + group_size, n_frames)target —beat, then one channel per bar position, thenmask(seeposition_target_channels()). The position block is normalized to a per-frame probability distribution, so it pairs with a softmax head and a cross-entropy loss rather than per-channel sigmoids. Every position gets its own supervised channel, which makes “is this a 1 or a 3?” a question the model is actually asked. See docs/beat_phase_improvement_review.md section 3.
- fold_positions(positions, group_size)[source]¶
Fold an annotated bar-position cycle onto
1..group_size.Annotations count positions across whatever the annotator treated as one bar, and that is not always
group_sizebeats. A track annotated1..8againstgroup_size=4is two bars of four, so its beats 5-8 are the 1-4 of the next bar and belong in the same four channels.Without folding, positions above
group_sizematch no channel at all: the position block is all-zero on those frames, so they hit the uniform-target fallback inBeatDataset.__getitem__()and the loss teaches “every bar position is equally likely” at full beat weight — on beats that are perfectly well annotated.- Parameters:
positions (ndarray) – 1-indexed bar positions, shape
(n_beats,).group_size (int) – Beats per group the position head predicts over.
- Returns:
Folded positions, or
Nonewhen the annotated cycle is not a multiple of group_size (e.g. a 6-beat bar against a 4-way head). No consistent folding exists there — 1,2,3,4,1,2 puts the downbeat 4 beats after one bar and 2 after the next — so the caller must drop that track’s position supervision rather than fold it wrongly.- Return type:
ndarray | None
- gaussian_smear(spike, sigma)[source]¶
Smear a 0/1 spike train into a soft target with a Gaussian bump per event.
The kernel is left unnormalized (peak value 1.0 at its center), so an isolated spike keeps peak 1.0 after convolution. Overlapping bumps are clipped back to 1.0 rather than allowed to sum above it.
- Parameters:
spike (ndarray) – Binary spike train, shape
(n_frames,).sigma (float) – Gaussian standard deviation, in frames.
0returnsspikeunchanged.
- Returns:
Smeared target, shape
(n_frames,), values in[0, 1].- Return type:
ndarray
- indices_for_split(dataset, name, split, val_split, binary_only=False)[source]¶
Return
datasetindices forsplit, reusing a training run’s cached train/val split (seeSplitter) so"val"means genuinely held-out tracks.- Parameters:
dataset (BeatDataset) – The
BeatDatasetto select indices from.name (str) – Dataset name (e.g.
"ballroom") used to look up the split.split (str) –
"train","val", or"all"(every index — no split file is read in this case).val_split (float) – Fraction of the dataset held out for validation. Must match how the split was created.
binary_only (bool) – Must match how the split was created — see
split_name().
- Return type:
list[int]
- position_target_channels(group_size)[source]¶
Channel names for the
target_layout="positions"target.("beat", "pos_1", ..., "pos_<group_size>", "mask")—beatfirst andmasklast, matchingTARGET_CHANNELS, sotarget[0]andtarget[-1]mean the same thing under both layouts.- Parameters:
group_size (int)
- Return type:
tuple[str, …]