Calculate forward propagation. Args: xs_pad (torch.Tensor): Input tensor (#batch, L, input_size). ilens (torch.Tensor): Input length (#batch). prev_states (torch.Tensor): Not to be used now. Returns: torch.Tensor: Output tensor (#batch,
(
self,
xs_pad: torch.Tensor,
ilens: torch.Tensor,
prev_states: torch.Tensor = None,
)
| 407 | return xs_pad, olens, None |
| 408 | |
| 409 | def forward_hidden( |
| 410 | self, |
| 411 | xs_pad: torch.Tensor, |
| 412 | ilens: torch.Tensor, |
| 413 | prev_states: torch.Tensor = None, |
| 414 | ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: |
| 415 | """Calculate forward propagation. |
| 416 | Args: |
| 417 | xs_pad (torch.Tensor): Input tensor (#batch, L, input_size). |
| 418 | ilens (torch.Tensor): Input length (#batch). |
| 419 | prev_states (torch.Tensor): Not to be used now. |
| 420 | Returns: |
| 421 | torch.Tensor: Output tensor (#batch, L, output_size). |
| 422 | torch.Tensor: Output length (#batch). |
| 423 | torch.Tensor: Not to be used now. |
| 424 | """ |
| 425 | masks = (~make_pad_mask(ilens)[:, None, :]).to(xs_pad.device) |
| 426 | if ( |
| 427 | isinstance(self.embed, Conv2dSubsampling) |
| 428 | or isinstance(self.embed, Conv2dSubsampling6) |
| 429 | or isinstance(self.embed, Conv2dSubsampling8) |
| 430 | ): |
| 431 | short_status, limit_size = check_short_utt(self.embed, xs_pad.size(1)) |
| 432 | if short_status: |
| 433 | raise TooShortUttError( |
| 434 | f"has {xs_pad.size(1)} frames and is too short for subsampling " |
| 435 | + f"(it needs more than {limit_size} frames), return empty results", |
| 436 | xs_pad.size(1), |
| 437 | limit_size, |
| 438 | ) |
| 439 | xs_pad, masks = self.embed(xs_pad, masks) |
| 440 | else: |
| 441 | xs_pad = self.embed(xs_pad) |
| 442 | num_layer = len(self.encoders) |
| 443 | for idx, encoder in enumerate(self.encoders): |
| 444 | xs_pad, masks = encoder(xs_pad, masks) |
| 445 | if idx == num_layer // 2 - 1: |
| 446 | hidden_feature = xs_pad |
| 447 | if isinstance(xs_pad, tuple): |
| 448 | xs_pad = xs_pad[0] |
| 449 | hidden_feature = hidden_feature[0] |
| 450 | if self.normalize_before: |
| 451 | xs_pad = self.after_norm(xs_pad) |
| 452 | self.hidden_feature = self.after_norm(hidden_feature) |
| 453 | |
| 454 | olens = masks.squeeze(1).sum(1) |
| 455 | return xs_pad, olens, None |
nothing calls this directly
no test coverage detected