espnet2.rst.rst_model.XeusEncoder
espnet2.rst.rst_model.XeusEncoder
class espnet2.rst.rst_model.XeusEncoder(model_tag: str = 'espnet/xeus', config: str = 'model/config.yaml', checkpoint: str = 'model/xeus_checkpoint_new.pth', target_layer: int = 10, lora_rank: int = 64, lora_alpha: int = 16, lora_dropout: float = 0.1, lora_target_modules: Sequence[str] = ('w_2',), input_sr: int = 16000, freeze_base: bool = True)
Bases: Module
XEUS encoder pair for the restoration feature predictor.
Same teacher/student roles as W2vBert2Encoder (frozen target extractor and LoRA-adapted predictor), truncated at target_layer.
XEUS is loaded the ESPnet way, SSLTask.build_model_from_file on the model/config.yaml + model/xeus_checkpoint_new.pth pair of the espnet/xeus Hub release (or any directory with the same layout). The E-Branchformer blocks above target_layer are dropped from both copies since only the block-target_layer output is used, and the LoRA adapter goes on the second linear of every feed-forward module (w_2, both the macaron and the main FFN) of the remaining blocks.
The release is CC-BY-NC-SA-4.0; see the recipe README.
Initialize internal Module state, shared by both nn.Module and ScriptModule.
encode(ssl_inputs: Dict[str, Tensor], teacher: bool = False) → Tuple[Tensor, Tensor]
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
