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