inference

Shared inference plumbing for beat and beat-phase checkpoints: loading a Lightning checkpoint (auto-detecting which task it was trained for), and running a forward pass through musicality.postprocess’s readout functions.

Task type is identified entirely from a checkpoint’s own saved hyper_parameters — specifically the task field, declared explicitly by whatever trained the checkpoint and threaded through BeatPhaseModule/ BeatModule’s save_hyperparameters() call. No companion Hydra config file is read at inference time.

Functions

detect_task(hyper_parameters)

Task tag declared explicitly by the run that trained the checkpoint.

load_module(checkpoint_path[, device])

Load a beat-only or beat-phase checkpoint, auto-detecting which.

load_track_waveform(audio_path, sample_rate)

Load a full track (no cropping/padding), mono, at sample_rate.

run_inference(module, task, wav, fps, device)

Run module on one full-track waveform and decode it via the readout function matching task.

detect_task(hyper_parameters)[source]

Task tag declared explicitly by the run that trained the checkpoint.

Parameters:

hyper_parameters (dict) – A checkpoint’s hyper_parameters dict, as saved by Lightning’s save_hyperparameters().

Returns:

"beat_only" or "beat_phase".

Raises:
  • KeyError – If task is missing — the checkpoint predates this field and must be retrained.

  • ValueError – If task isn’t a recognized value.

Return type:

str

load_module(checkpoint_path, device='cpu')[source]

Load a beat-only or beat-phase checkpoint, auto-detecting which.

Migrates checkpoints saved before TCNTempoNet’s frame head gained a dropout layer (Conv1d -> Sequential(Dropout, Conv1d)), which shifted its state dict keys from frame_head.{weight,bias} to frame_head.1.{weight,bias}. Safe to delete once no pre-dropout checkpoints are still in use.

Parameters:
  • checkpoint_path (str | Path) – Path to a Lightning .ckpt file.

  • device (str | device) – Torch device the returned module is moved to.

Returns:

(module, task) — a loaded, eval-mode module already moved to device, and its task tag (see detect_task()).

Return type:

tuple[BeatModule | BeatPhaseModule, str]

load_track_waveform(audio_path, sample_rate)[source]

Load a full track (no cropping/padding), mono, at sample_rate.

Parameters:
  • audio_path (str) – Path to an audio file.

  • sample_rate (int) – Target sample rate — resampled if the file differs.

Returns:

Waveform, shape (1, N).

Return type:

Tensor

run_inference(module, task, wav, fps, device, beat_threshold=0.3, min_distance_frames=1, gate_tolerance=0.2, anchor_threshold=0.5, group_size=4, decoder='greedy', switch_penalty=None, advance='index')[source]

Run module on one full-track waveform and decode it via the readout function matching task.

Parameters:
  • module (BeatModule | BeatPhaseModule) – A loaded, eval-mode module (see load_module()).

  • task (str) – "beat_phase" or "beat_only" (see detect_task()).

  • wav (Tensor) – Mono waveform, shape (1, N) (e.g. from load_track_waveform()) — batched and moved to device internally.

  • fps (float) – Frames per second (sample_rate / hop_length).

  • switch_penalty (float | None) – Forwarded to readout()’s bar-position stage only — ignored when task="beat_only". decoder="global" uses the whole-track maximum-likelihood decode (label_bar_position_global()) instead of the greedy count-forward one, which measurably lowers phase confusion on the same probabilities — see docs/beat_phase_improvement_review.md.

  • device (str | device)

  • beat_threshold (float)

  • min_distance_frames (int)

  • gate_tolerance (float)

  • anchor_threshold (float)

  • group_size (int)

  • decoder (str)

  • switch_penalty

  • advance (str)

Returns:

readout()’s list[dict] for beat-phase, or readout_beat_only()’s np.ndarray for beat-only.

Raises:

ValueError – Unknown task.

Return type:

list[dict] | ndarray