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
|
Task tag declared explicitly by the run that trained the checkpoint. |
|
Load a beat-only or beat-phase checkpoint, auto-detecting which. |
|
Load a full track (no cropping/padding), mono, at |
|
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_parametersdict, as saved by Lightning’ssave_hyperparameters().- Returns:
"beat_only"or"beat_phase".- Raises:
KeyError – If
taskis missing — the checkpoint predates this field and must be retrained.ValueError – If
taskisn’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 fromframe_head.{weight,bias}toframe_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
.ckptfile.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 (seedetect_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"(seedetect_task()).wav (Tensor) – Mono waveform, shape
(1, N)(e.g. fromload_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 whentask="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()’slist[dict]for beat-phase, orreadout_beat_only()’snp.ndarrayfor beat-only.- Raises:
ValueError – Unknown task.
- Return type:
list[dict] | ndarray