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

Method forward

funasr/models/scama/encoder.py:393–478  ·  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,
        ind: int = 0,
    )

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 9

output_sizeMethod · 0.95
make_pad_maskFunction · 0.90
check_short_uttFunction · 0.90
TooShortUttErrorClass · 0.90
gen_chunk_maskMethod · 0.80
split_chunkMethod · 0.80
get_mask_shfit_chunkMethod · 0.80
softmaxMethod · 0.45

Tested by

no test coverage detected