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

Method forward

funasr/models/conformer/encoder.py:559–636  ·  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,
        ctc: CTC = None,
    )

Source from the content-addressed store, hash-verified

557 return self._output_size
558
559 def forward(
560 self,
561 xs_pad: torch.Tensor,
562 ilens: torch.Tensor,
563 prev_states: torch.Tensor = None,
564 ctc: CTC = None,
565 ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
566 """Calculate forward propagation.
567
568 Args:
569 xs_pad (torch.Tensor): Input tensor (#batch, L, input_size).
570 ilens (torch.Tensor): Input length (#batch).
571 prev_states (torch.Tensor): Not to be used now.
572
573 Returns:
574 torch.Tensor: Output tensor (#batch, L, output_size).
575 torch.Tensor: Output length (#batch).
576 torch.Tensor: Not to be used now.
577
578 """
579 masks = (~make_pad_mask(ilens)[:, None, :]).to(xs_pad.device)
580
581 if (
582 isinstance(self.embed, Conv2dSubsampling)
583 or isinstance(self.embed, Conv2dSubsampling2)
584 or isinstance(self.embed, Conv2dSubsampling6)
585 or isinstance(self.embed, Conv2dSubsampling8)
586 or isinstance(self.embed, Conv2dSubsamplingPad)
587 ):
588 short_status, limit_size = check_short_utt(self.embed, xs_pad.size(1))
589 if short_status:
590 raise TooShortUttError(
591 f"has {xs_pad.size(1)} frames and is too short for subsampling "
592 + f"(it needs more than {limit_size} frames), return empty results",
593 xs_pad.size(1),
594 limit_size,
595 )
596 xs_pad, masks = self.embed(xs_pad, masks)
597 else:
598 xs_pad = self.embed(xs_pad)
599
600 intermediate_outs = []
601 if len(self.interctc_layer_idx) == 0:
602 xs_pad, masks = self.encoders(xs_pad, masks)
603 else:
604 for layer_idx, encoder_layer in enumerate(self.encoders):
605 xs_pad, masks = encoder_layer(xs_pad, masks)
606
607 if layer_idx + 1 in self.interctc_layer_idx:
608 encoder_out = xs_pad
609 if isinstance(encoder_out, tuple):
610 encoder_out = encoder_out[0]
611
612 # intermediate outputs are also normalized
613 if self.normalize_before:
614 encoder_out = self.after_norm(encoder_out)
615
616 intermediate_outs.append((layer_idx + 1, encoder_out))

Callers

nothing calls this directly

Calls 4

make_pad_maskFunction · 0.90
check_short_uttFunction · 0.90
TooShortUttErrorClass · 0.90
softmaxMethod · 0.45

Tested by

no test coverage detected