Frontend + Encoder. Args: speech: (Batch, Length, ...) speech_lengths: (Batch, )
(
self,
speech: torch.Tensor,
speech_lengths: torch.Tensor,
)
| 115 | return {"feats": feats, "feats_lengths": feats_lengths} |
| 116 | |
| 117 | def encode( |
| 118 | self, |
| 119 | speech: torch.Tensor, |
| 120 | speech_lengths: torch.Tensor, |
| 121 | ): |
| 122 | """Frontend + Encoder. |
| 123 | Args: |
| 124 | speech: (Batch, Length, ...) |
| 125 | speech_lengths: (Batch, ) |
| 126 | """ |
| 127 | with autocast(False): |
| 128 | # 1. Extract feats |
| 129 | feats, feats_lengths = self._extract_feats(speech, speech_lengths) |
| 130 | |
| 131 | # 2. Data augmentation |
| 132 | if self.specaug is not None and self.training: |
| 133 | feats, feats_lengths = self.specaug(feats, feats_lengths) |
| 134 | |
| 135 | # 3. Normalization for feature: e.g. Global-CMVN, Utterance-CMVN |
| 136 | if self.normalize is not None: |
| 137 | feats, feats_lengths = self.normalize(feats, feats_lengths) |
| 138 | |
| 139 | # Pre-encoder, e.g. used for raw input data |
| 140 | if self.preencoder is not None: |
| 141 | feats, feats_lengths = self.preencoder(feats, feats_lengths) |
| 142 | |
| 143 | # 4. Forward encoder |
| 144 | if min(speech_lengths) == max(speech_lengths): # for clipping, set speech_lengths as None |
| 145 | speech_lengths = None |
| 146 | encoder_out = self.encoder(feats, speech_lengths, mask=True, features_only=False) |
| 147 | |
| 148 | return encoder_out |
| 149 | |
| 150 | def _extract_feats( |
| 151 | self, speech: torch.Tensor, speech_lengths: torch.Tensor |
no test coverage detected