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

Method forward

funasr/models/sanm/encoder.py:392–461  ·  view source on GitHub ↗

Embed positions in tensor. Args: xs_pad: input tensor (B, L, D) ilens: input length (B) prev_states: Not to be used now. Returns: position embedded tensor and mask

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

Source from the content-addressed store, hash-verified

390 return self._output_size
391
392 def forward(
393 self,
394 xs_pad: torch.Tensor,
395 ilens: torch.Tensor,
396 prev_states: torch.Tensor = None,
397 ctc: CTC = None,
398 ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
399 """Embed positions in tensor.
400
401 Args:
402 xs_pad: input tensor (B, L, D)
403 ilens: input length (B)
404 prev_states: Not to be used now.
405 Returns:
406 position embedded tensor and mask
407 """
408 masks = (~make_pad_mask(ilens)[:, None, :]).to(xs_pad.device)
409 xs_pad = xs_pad * self.output_size() ** 0.5
410 if self.embed is None:
411 xs_pad = xs_pad
412 elif (
413 isinstance(self.embed, Conv2dSubsampling)
414 or isinstance(self.embed, Conv2dSubsampling2)
415 or isinstance(self.embed, Conv2dSubsampling6)
416 or isinstance(self.embed, Conv2dSubsampling8)
417 ):
418 short_status, limit_size = check_short_utt(self.embed, xs_pad.size(1))
419 if short_status:
420 raise TooShortUttError(
421 f"has {xs_pad.size(1)} frames and is too short for subsampling "
422 + f"(it needs more than {limit_size} frames), return empty results",
423 xs_pad.size(1),
424 limit_size,
425 )
426 xs_pad, masks = self.embed(xs_pad, masks)
427 else:
428 xs_pad = self.embed(xs_pad)
429
430 # xs_pad = self.dropout(xs_pad)
431 encoder_outs = self.encoders0(xs_pad, masks)
432 xs_pad, masks = encoder_outs[0], encoder_outs[1]
433 intermediate_outs = []
434 if len(self.interctc_layer_idx) == 0:
435 encoder_outs = self.encoders(xs_pad, masks)
436 xs_pad, masks = encoder_outs[0], encoder_outs[1]
437 else:
438 for layer_idx, encoder_layer in enumerate(self.encoders):
439 encoder_outs = encoder_layer(xs_pad, masks)
440 xs_pad, masks = encoder_outs[0], encoder_outs[1]
441
442 if layer_idx + 1 in self.interctc_layer_idx:
443 encoder_out = xs_pad
444
445 # intermediate outputs are also normalized
446 if self.normalize_before:
447 encoder_out = self.after_norm(encoder_out)
448
449 intermediate_outs.append((layer_idx + 1, encoder_out))

Callers

nothing calls this directly

Calls 5

output_sizeMethod · 0.95
make_pad_maskFunction · 0.90
check_short_uttFunction · 0.90
TooShortUttErrorClass · 0.90
softmaxMethod · 0.45

Tested by

no test coverage detected