MCPcopy Create free account
hub / github.com/espnet/espnet / encode

Method encode

espnet2/st/espnet_model.py:478–536  ·  view source on GitHub ↗

Frontend + Encoder. Note that this method is used by st_inference.py Args: speech: (Batch, Length, ...) speech_lengths: (Batch, )

(
        self,
        speech: torch.Tensor,
        speech_lengths: torch.Tensor,
        return_int_enc: bool = False,
    )

Source from the content-addressed store, hash-verified

476 return {"feats": feats, "feats_lengths": feats_lengths}
477
478 def encode(
479 self,
480 speech: torch.Tensor,
481 speech_lengths: torch.Tensor,
482 return_int_enc: bool = False,
483 ) -> Tuple[torch.Tensor, torch.Tensor]:
484 """Frontend + Encoder. Note that this method is used by st_inference.py
485
486 Args:
487 speech: (Batch, Length, ...)
488 speech_lengths: (Batch, )
489 """
490 with autocast("cuda", enabled=False):
491 # 1. Extract feats
492 feats, feats_lengths = self._extract_feats(speech, speech_lengths)
493
494 # 2. Data augmentation
495 if self.specaug is not None and self.training:
496 feats, feats_lengths = self.specaug(feats, feats_lengths)
497
498 # 3. Normalization for feature: e.g. Global-CMVN, Utterance-CMVN
499 if self.normalize is not None:
500 feats, feats_lengths = self.normalize(feats, feats_lengths)
501
502 # Pre-encoder, e.g. used for raw input data
503 if self.preencoder is not None:
504 feats, feats_lengths = self.preencoder(feats, feats_lengths)
505
506 # 4. Forward encoder
507 # feats: (Batch, Length, Dim)
508 # -> encoder_out: (Batch, Length2, Dim2)
509 encoder_out, encoder_out_lens, _ = self.encoder(feats, feats_lengths)
510
511 if return_int_enc:
512 int_encoder_out, int_encoder_out_lens = encoder_out, encoder_out_lens
513
514 if self.hier_encoder is not None:
515 encoder_out, encoder_out_lens, _ = self.hier_encoder(
516 encoder_out, encoder_out_lens
517 )
518
519 # Post-encoder, e.g. NLU
520 if self.postencoder is not None:
521 encoder_out, encoder_out_lens = self.postencoder(
522 encoder_out, encoder_out_lens
523 )
524
525 assert encoder_out.size(0) == speech.size(0), (
526 encoder_out.size(),
527 speech.size(0),
528 )
529 assert encoder_out.size(1) <= encoder_out_lens.max(), (
530 encoder_out.size(),
531 encoder_out_lens.max(),
532 )
533
534 if return_int_enc:
535 return encoder_out, encoder_out_lens, int_encoder_out, int_encoder_out_lens

Callers 7

forwardMethod · 0.95
__init__Method · 0.45
preprocessMethod · 0.45
find_lengthMethod · 0.45
_codec_encode_batchMethod · 0.45
inference_workerFunction · 0.45
forwardMethod · 0.45

Calls 2

_extract_featsMethod · 0.95
sizeMethod · 0.80

Tested by

no test coverage detected