espnet2.rst.rst_model.W2VBert2Encoder
espnet2.rst.rst_model.W2VBert2Encoder
class espnet2.rst.rst_model.W2VBert2Encoder(model_tag: str = 'facebook/w2v-bert-2.0', target_layer: int = 8, lora_rank: int = 64, lora_alpha: int = 16, lora_dropout: float = 0.1, input_sr: int = 16000, freeze_base: bool = True)
Bases: Module
w2v-BERT 2.0 encoder pair for the restoration feature predictor.
The teacher is a frozen copy that extracts the target features from clean speech; the student is a LoRA-adapted copy trained to produce the same features from degraded speech. Both are truncated at target_layer.
Initialize internal Module state, shared by both nn.Module and ScriptModule.
encode(ssl_inputs: Dict[str, Tensor], teacher: bool = False) → Tuple[Tensor, Tensor]
Features and a 0/1 frame mask from the student or the frozen teacher.
extract_clean_features(ssl_inputs: Dict[str, Tensor]) → Tensor
forward(ssl_inputs: Dict[str, Tensor]) → Tuple[Tensor, OrderedDict]
Define the computation performed at every call.
Should be overridden by all subclasses.
NOTE
Although the recipe for forward pass needs to be defined within this function, one should call the Module instance afterwards instead of this since the former takes care of running the registered hooks while the latter silently ignores them.
property ssl_dim : int
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
