torch_audio

Classes

TorchAudioTempoNet([pipeline, dropout, ...])

TorchAudio wav2vec2 encoder with a tempo regression head.

class TorchAudioTempoNet(pipeline='WAV2VEC2_BASE', dropout=0.1, freeze_backbone=True)[source]

Bases: Module

TorchAudio wav2vec2 encoder with a tempo regression head.

Uses a torchaudio.pipelines bundle as backbone (e.g. WAV2VEC2_BASE). Expects raw waveforms at the pipeline’s native sample rate (typically 16 kHz). The backbone is frozen by default.

Input: (B, 1, T) Output: (B,)

Parameters:
  • pipeline (str) – torchaudio.pipelines attribute name (e.g. "WAV2VEC2_BASE").

  • dropout (float) – Dropout probability in the regression head.

  • freeze_backbone (bool) – Whether to freeze the backbone weights.

forward(wav)[source]
Parameters:

wav (Tensor)

Return type:

Tensor