espnet2.rst.rst_vocoder_model.ESPnetRestorationVocoderModel
espnet2.rst.rst_vocoder_model.ESPnetRestorationVocoderModel
class espnet2.rst.rst_vocoder_model.ESPnetRestorationVocoderModel(ssl_encoder: Module, vocoder: Module, discriminator: Module, use_predicted_feat: bool = False, input_sr: int = 16000, output_sr: int = 48000, segment_duration: float = 1.0, mel_loss_weight: float = 15.0, adv_loss_weight: float = 2.0, fm_loss_weight: float = 1.0, mel_loss_conf: Dict | None = None)
Bases: RestorationVocoderFeatures, AbsGANESPnetModel
Vocoder generator + discriminator with a frozen feature encoder.
- Parameters:
- ssl_encoder – stage-1 encoder (W2VBert2Encoder or XeusEncoder). Always frozen here; the teacher branch provides ground-truth features (pretrain) and the LoRA student branch provides predicted ones (finetune).
- vocoder – generator mapping (B, T, D) features to (B, T * 960) audio (DACVocoder or HiFiGANVocoder; anything with
generateandupsample_factor). - discriminator – returns a list, one entry per sub-discriminator, of [feature maps …, logits].
- use_predicted_feat – False for stage 2 (teacher on clean speech), True for stage 3 (student on degraded speech).
- input_sr – encoder sample rate (16 kHz).
- output_sr – vocoder sample rate (48 kHz).
- segment_duration – length in seconds of the aligned feature/waveform excerpt the generator and discriminator actually train on. The encoder runs on the longer context the collate function provides, so the excerpt’s features are computed with the surrounding speech in view, as they are at inference.
- mel_loss_weight – generator loss weights (Sidon: 15 / 2 / 1).
- adv_loss_weight – generator loss weights (Sidon: 15 / 2 / 1).
- fm_loss_weight – generator loss weights (Sidon: 15 / 2 / 1).
- mel_loss_conf – overrides for MelSpectrogramLoss.
collect_feats(**batch)
forward(speech_ref1: Tensor, speech_ref1_lengths: Tensor, noisy_speech: Tensor | None = None, noisy_speech_lengths: Tensor | None = None, vocoder_crop_start: Tensor | None = None, forward_generator: bool = True, **kwargs)
Return the generator loss or the discrimiantor loss.
This method must have an argument “forward_generator” to switch the generator loss calculation and the discrimiantor loss calculation. If forward_generator is true, return the generator loss with optim_idx 0. If forward_generator is false, return the discrimiantor loss with optim_idx 1.
- Parameters:forward_generator (bool) – Whether to return the generator loss or the discrimiantor loss. This must have the default value.
- Returns:
- loss (Tensor): Loss scalar tensor.
- stats (Dict[str, float]): Statistics to be monitored.
- weight (Tensor): Weight tensor to summarize losses.
- optim_idx (int): Optimizer index (0 for G and 1 for D).
- Return type: Dict[str, Any]
train(mode: bool = True)
Set the module in training mode.
This has an effect only on certain modules. See the documentation of particular modules for details of their behaviors in training/evaluation mode, i.e., whether they are affected, e.g. Dropout, BatchNorm, etc.
- Parameters:mode (bool) – whether to set training mode (
True) or evaluation mode (False). Default:True. - Returns: self
- Return type: Module
