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

Method forward

funasr/models/ct_transformer_streaming/encoder.py:355–429  ·  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,
        vad_indexes: torch.Tensor,
        prev_states: torch.Tensor = None,
        ctc: CTC = None,
    )

Source from the content-addressed store, hash-verified

353 return self._output_size
354
355 def forward(
356 self,
357 xs_pad: torch.Tensor,
358 ilens: torch.Tensor,
359 vad_indexes: torch.Tensor,
360 prev_states: torch.Tensor = None,
361 ctc: CTC = None,
362 ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
363 """Embed positions in tensor.
364
365 Args:
366 xs_pad: input tensor (B, L, D)
367 ilens: input length (B)
368 prev_states: Not to be used now.
369 Returns:
370 position embedded tensor and mask
371 """
372 masks = (~make_pad_mask(ilens)[:, None, :]).to(xs_pad.device)
373 sub_masks = subsequent_mask(masks.size(-1), device=xs_pad.device).unsqueeze(0)
374 no_future_masks = masks & sub_masks
375 xs_pad *= self.output_size() ** 0.5
376 if self.embed is None:
377 xs_pad = xs_pad
378 elif (
379 isinstance(self.embed, Conv2dSubsampling)
380 or isinstance(self.embed, Conv2dSubsampling2)
381 or isinstance(self.embed, Conv2dSubsampling6)
382 or isinstance(self.embed, Conv2dSubsampling8)
383 ):
384 short_status, limit_size = check_short_utt(self.embed, xs_pad.size(1))
385 if short_status:
386 raise TooShortUttError(
387 f"has {xs_pad.size(1)} frames and is too short for subsampling "
388 + f"(it needs more than {limit_size} frames), return empty results",
389 xs_pad.size(1),
390 limit_size,
391 )
392 xs_pad, masks = self.embed(xs_pad, masks)
393 else:
394 xs_pad = self.embed(xs_pad)
395
396 # xs_pad = self.dropout(xs_pad)
397 mask_tup0 = [masks, no_future_masks]
398 encoder_outs = self.encoders0(xs_pad, mask_tup0)
399 xs_pad, _ = encoder_outs[0], encoder_outs[1]
400 intermediate_outs = []
401
402 for layer_idx, encoder_layer in enumerate(self.encoders):
403 if layer_idx + 1 == len(self.encoders):
404 # This is last layer.
405 coner_mask = torch.ones(
406 masks.size(0),
407 masks.size(-1),
408 masks.size(-1),
409 device=xs_pad.device,
410 dtype=torch.bool,
411 )
412 for word_index, length in enumerate(ilens):

Callers

nothing calls this directly

Calls 6

output_sizeMethod · 0.95
make_pad_maskFunction · 0.90
subsequent_maskFunction · 0.90
check_short_uttFunction · 0.90
TooShortUttErrorClass · 0.90
vad_maskFunction · 0.90

Tested by

no test coverage detected