MCPcopy Create free account
hub / github.com/modelscope/FunASR / forward

Method forward

funasr/models/branchformer/encoder.py:517–564  ·  view source on GitHub ↗

Calculate forward propagation. Args: xs_pad (torch.Tensor): Input tensor (#batch, L, input_size). ilens (torch.Tensor): Input length (#batch). prev_states (torch.Tensor): Not to be used now. Returns: torch.Tensor: Output tensor (#batc

(
        self,
        xs_pad: torch.Tensor,
        ilens: torch.Tensor,
        prev_states: torch.Tensor = None,
    )

Source from the content-addressed store, hash-verified

515 return self._output_size
516
517 def forward(
518 self,
519 xs_pad: torch.Tensor,
520 ilens: torch.Tensor,
521 prev_states: torch.Tensor = None,
522 ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
523 """Calculate forward propagation.
524
525 Args:
526 xs_pad (torch.Tensor): Input tensor (#batch, L, input_size).
527 ilens (torch.Tensor): Input length (#batch).
528 prev_states (torch.Tensor): Not to be used now.
529
530 Returns:
531 torch.Tensor: Output tensor (#batch, L, output_size).
532 torch.Tensor: Output length (#batch).
533 torch.Tensor: Not to be used now.
534
535 """
536
537 masks = (~make_pad_mask(ilens)[:, None, :]).to(xs_pad.device)
538
539 if (
540 isinstance(self.embed, Conv2dSubsampling)
541 or isinstance(self.embed, Conv2dSubsampling2)
542 or isinstance(self.embed, Conv2dSubsampling6)
543 or isinstance(self.embed, Conv2dSubsampling8)
544 ):
545 short_status, limit_size = check_short_utt(self.embed, xs_pad.size(1))
546 if short_status:
547 raise TooShortUttError(
548 f"has {xs_pad.size(1)} frames and is too short for subsampling "
549 + f"(it needs more than {limit_size} frames), return empty results",
550 xs_pad.size(1),
551 limit_size,
552 )
553 xs_pad, masks = self.embed(xs_pad, masks)
554 elif self.embed is not None:
555 xs_pad = self.embed(xs_pad)
556
557 xs_pad, masks = self.encoders(xs_pad, masks)
558
559 if isinstance(xs_pad, tuple):
560 xs_pad = xs_pad[0]
561
562 xs_pad = self.after_norm(xs_pad)
563 olens = masks.squeeze(1).sum(1)
564 return xs_pad, olens, None

Callers

nothing calls this directly

Calls 3

make_pad_maskFunction · 0.90
check_short_uttFunction · 0.90
TooShortUttErrorClass · 0.90

Tested by

no test coverage detected