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,
)
| 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 |
nothing calls this directly
no test coverage detected