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 (#batc
(
self,
xs_pad: torch.Tensor,
ilens: torch.Tensor,
prev_states: torch.Tensor = None,
)
| 515 | return self._output_size |
| 516 | |
| 517 | def forward( |
| 518 | self, |
| 519 | xs_pad: torch.Tensor, |
| 520 | ilens: torch.Tensor, |
| 521 | prev_states: torch.Tensor = None, |
| 522 | ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: |
| 523 | """Calculate forward propagation. |
| 524 | |
| 525 | Args: |
| 526 | xs_pad (torch.Tensor): Input tensor (#batch, L, input_size). |
| 527 | ilens (torch.Tensor): Input length (#batch). |
| 528 | prev_states (torch.Tensor): Not to be used now. |
| 529 | |
| 530 | Returns: |
| 531 | torch.Tensor: Output tensor (#batch, L, output_size). |
| 532 | torch.Tensor: Output length (#batch). |
| 533 | torch.Tensor: Not to be used now. |
| 534 | |
| 535 | """ |
| 536 | |
| 537 | masks = (~make_pad_mask(ilens)[:, None, :]).to(xs_pad.device) |
| 538 | |
| 539 | if ( |
| 540 | isinstance(self.embed, Conv2dSubsampling) |
| 541 | or isinstance(self.embed, Conv2dSubsampling2) |
| 542 | or isinstance(self.embed, Conv2dSubsampling6) |
| 543 | or isinstance(self.embed, Conv2dSubsampling8) |
| 544 | ): |
| 545 | short_status, limit_size = check_short_utt(self.embed, xs_pad.size(1)) |
| 546 | if short_status: |
| 547 | raise TooShortUttError( |
| 548 | f"has {xs_pad.size(1)} frames and is too short for subsampling " |
| 549 | + f"(it needs more than {limit_size} frames), return empty results", |
| 550 | xs_pad.size(1), |
| 551 | limit_size, |
| 552 | ) |
| 553 | xs_pad, masks = self.embed(xs_pad, masks) |
| 554 | elif self.embed is not None: |
| 555 | xs_pad = self.embed(xs_pad) |
| 556 | |
| 557 | xs_pad, masks = self.encoders(xs_pad, masks) |
| 558 | |
| 559 | if isinstance(xs_pad, tuple): |
| 560 | xs_pad = xs_pad[0] |
| 561 | |
| 562 | xs_pad = self.after_norm(xs_pad) |
| 563 | olens = masks.squeeze(1).sum(1) |
| 564 | return xs_pad, olens, None |
nothing calls this directly
no test coverage detected