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

Method forward_hidden

funasr/models/mfcca/mfcca_encoder.py:409–455  ·  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 (#batch,

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

Source from the content-addressed store, hash-verified

407 return xs_pad, olens, None
408
409 def forward_hidden(
410 self,
411 xs_pad: torch.Tensor,
412 ilens: torch.Tensor,
413 prev_states: torch.Tensor = None,
414 ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
415 """Calculate forward propagation.
416 Args:
417 xs_pad (torch.Tensor): Input tensor (#batch, L, input_size).
418 ilens (torch.Tensor): Input length (#batch).
419 prev_states (torch.Tensor): Not to be used now.
420 Returns:
421 torch.Tensor: Output tensor (#batch, L, output_size).
422 torch.Tensor: Output length (#batch).
423 torch.Tensor: Not to be used now.
424 """
425 masks = (~make_pad_mask(ilens)[:, None, :]).to(xs_pad.device)
426 if (
427 isinstance(self.embed, Conv2dSubsampling)
428 or isinstance(self.embed, Conv2dSubsampling6)
429 or isinstance(self.embed, Conv2dSubsampling8)
430 ):
431 short_status, limit_size = check_short_utt(self.embed, xs_pad.size(1))
432 if short_status:
433 raise TooShortUttError(
434 f"has {xs_pad.size(1)} frames and is too short for subsampling "
435 + f"(it needs more than {limit_size} frames), return empty results",
436 xs_pad.size(1),
437 limit_size,
438 )
439 xs_pad, masks = self.embed(xs_pad, masks)
440 else:
441 xs_pad = self.embed(xs_pad)
442 num_layer = len(self.encoders)
443 for idx, encoder in enumerate(self.encoders):
444 xs_pad, masks = encoder(xs_pad, masks)
445 if idx == num_layer // 2 - 1:
446 hidden_feature = xs_pad
447 if isinstance(xs_pad, tuple):
448 xs_pad = xs_pad[0]
449 hidden_feature = hidden_feature[0]
450 if self.normalize_before:
451 xs_pad = self.after_norm(xs_pad)
452 self.hidden_feature = self.after_norm(hidden_feature)
453
454 olens = masks.squeeze(1).sum(1)
455 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