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_positions(positions, group_size)

Fold an annotated bar-position cycle onto 1..group_size.

gaussian_smear(spike, sigma)

Smear a 0/1 spike train into a soft target with a Gaussian bump per event.

indices_for_split(dataset, name, split, ...)

Return dataset indices for split, reusing a training run's cached train/val split (see Splitter) so "val" means genuinely held-out tracks.

position_target_channels(group_size)

Channel names for the target_layout="positions" target.

Classes

BeatDataset([name, data_home, refs, ...])

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: Dataset

Dataset returning waveforms and frame-level beat-phase targets, read entirely from this project’s own tracks/+annotations/ format.

Loads every track with a migrated .beats annotation and a resolvable tracks/<id>.wav.

The target is a 4-channel tensor of shape (4, n_frames) where n_frames = n_samples // hop_length:

  • beat — any beat.

  • one — position 1 of the group (the downbeat, for the default group_size=4 bar-position case).

  • last — the last beat of the group (bar position 4 for group_size=4; e.g. phrase position 8 for a phrase-annotated dataset with group_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 onto group_size (see fold_positions()). Both still contribute their beat channel; 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 (see tools/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_DIR from musicality.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, 8 for a phrase-position (1-8) dataset. Annotations counting a longer bar are folded onto 1..group_size at load time (see fold_positions()), so a track annotated 1-8 trains a group_size=4 head as two bars of four rather than falling off the end of the channel list. Tracks whose bar length is not a multiple of group_size keep 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 cycle 1, 2, 3 in triple meter — as well as tracks with no position annotation at all, since their meter can’t be confirmed. Independent of group_size: a track only needs an even beats-per-bar count, not one equal to group_size.

  • random_crop (bool) – If True, draw the duration-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). If False, 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 positions 1 and group_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, then mask (see position_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_size beats. A track annotated 1..8 against group_size=4 is 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_size match no channel at all: the position block is all-zero on those frames, so they hit the uniform-target fallback in BeatDataset.__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 None when 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. 0 returns spike unchanged.

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 dataset indices for split, reusing a training run’s cached train/val split (see Splitter) so "val" means genuinely held-out tracks.

Parameters:
  • dataset (BeatDataset) – The BeatDataset to 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") — beat first and mask last, matching TARGET_CHANNELS, so target[0] and target[-1] mean the same thing under both layouts.

Parameters:

group_size (int)

Return type:

tuple[str, …]