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,
ctc: CTC = None,
)
| 557 | return self._output_size |
| 558 | |
| 559 | def forward( |
| 560 | self, |
| 561 | xs_pad: torch.Tensor, |
| 562 | ilens: torch.Tensor, |
| 563 | prev_states: torch.Tensor = None, |
| 564 | ctc: CTC = None, |
| 565 | ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: |
| 566 | """Calculate forward propagation. |
| 567 | |
| 568 | Args: |
| 569 | xs_pad (torch.Tensor): Input tensor (#batch, L, input_size). |
| 570 | ilens (torch.Tensor): Input length (#batch). |
| 571 | prev_states (torch.Tensor): Not to be used now. |
| 572 | |
| 573 | Returns: |
| 574 | torch.Tensor: Output tensor (#batch, L, output_size). |
| 575 | torch.Tensor: Output length (#batch). |
| 576 | torch.Tensor: Not to be used now. |
| 577 | |
| 578 | """ |
| 579 | masks = (~make_pad_mask(ilens)[:, None, :]).to(xs_pad.device) |
| 580 | |
| 581 | if ( |
| 582 | isinstance(self.embed, Conv2dSubsampling) |
| 583 | or isinstance(self.embed, Conv2dSubsampling2) |
| 584 | or isinstance(self.embed, Conv2dSubsampling6) |
| 585 | or isinstance(self.embed, Conv2dSubsampling8) |
| 586 | or isinstance(self.embed, Conv2dSubsamplingPad) |
| 587 | ): |
| 588 | short_status, limit_size = check_short_utt(self.embed, xs_pad.size(1)) |
| 589 | if short_status: |
| 590 | raise TooShortUttError( |
| 591 | f"has {xs_pad.size(1)} frames and is too short for subsampling " |
| 592 | + f"(it needs more than {limit_size} frames), return empty results", |
| 593 | xs_pad.size(1), |
| 594 | limit_size, |
| 595 | ) |
| 596 | xs_pad, masks = self.embed(xs_pad, masks) |
| 597 | else: |
| 598 | xs_pad = self.embed(xs_pad) |
| 599 | |
| 600 | intermediate_outs = [] |
| 601 | if len(self.interctc_layer_idx) == 0: |
| 602 | xs_pad, masks = self.encoders(xs_pad, masks) |
| 603 | else: |
| 604 | for layer_idx, encoder_layer in enumerate(self.encoders): |
| 605 | xs_pad, masks = encoder_layer(xs_pad, masks) |
| 606 | |
| 607 | if layer_idx + 1 in self.interctc_layer_idx: |
| 608 | encoder_out = xs_pad |
| 609 | if isinstance(encoder_out, tuple): |
| 610 | encoder_out = encoder_out[0] |
| 611 | |
| 612 | # intermediate outputs are also normalized |
| 613 | if self.normalize_before: |
| 614 | encoder_out = self.after_norm(encoder_out) |
| 615 | |
| 616 | intermediate_outs.append((layer_idx + 1, encoder_out)) |
nothing calls this directly
no test coverage detected