espnet2.beats.espnet_model.BeatsPretrainModel
Less than 1 minute
espnet2.beats.espnet_model.BeatsPretrainModel
class espnet2.beats.espnet_model.BeatsPretrainModel(encoder: AbsEncoder, decoder: Module, ignore_id: int = -2, label_smoothing: float = 0.1, waveform_input: bool = False, contrastive_loss_weight: float = 0.0)
Bases: AbsESPnetModel
Beats Pretraining model
Initialize internal Module state, shared by both nn.Module and ScriptModule.
collect_feats(speech: Tensor, speech_lengths: Tensor, target: Tensor, target_lengths: Tensor, **kwargs) → Dict[str, Tensor]
forward(speech: Tensor, speech_lengths: Tensor, target: Tensor, target_lengths: Tensor, **kwargs) → Tuple[Tensor, Dict[str, Tensor], Tensor]
Encoder + Predictor + Calc loss
- Parameters:
- speech – (Batch, Length, Dim). Either raw speech or features. If raw speech, then should be single channel ie Dim=1.
- speech_lengths – (Batch, )
- target – (Batch, Length)
- target_lengths – (Batch,)
patch_contrastive_loss(unmasked_patch_emb, temperature=0.1)
