Calculate forward propagation without loss. Args: xs (Tensor): Batch of padded target features (B, T_feats, odim). ilens (LongTensor): Batch of the lengths of each target (B,). Returns: Tensor: Weight value if not joint training else model output
(
self,
xs: torch.Tensor,
ilens: torch.Tensor,
ys: Optional[torch.Tensor] = None,
olens: Optional[torch.Tensor] = None,
ds: Optional[torch.Tensor] = None,
ps: Optional[torch.Tensor] = None,
es: Optional[torch.Tensor] = None,
spembs: Optional[torch.Tensor] = None,
sids: Optional[torch.Tensor] = None,
lids: Optional[torch.Tensor] = None,
is_inference: bool = False,
alpha: float = 1.0,
use_teacher_forcing: bool = False,
)
| 587 | return loss, stats, after_outs if after_outs is not None else before_outs |
| 588 | |
| 589 | def _forward( |
| 590 | self, |
| 591 | xs: torch.Tensor, |
| 592 | ilens: torch.Tensor, |
| 593 | ys: Optional[torch.Tensor] = None, |
| 594 | olens: Optional[torch.Tensor] = None, |
| 595 | ds: Optional[torch.Tensor] = None, |
| 596 | ps: Optional[torch.Tensor] = None, |
| 597 | es: Optional[torch.Tensor] = None, |
| 598 | spembs: Optional[torch.Tensor] = None, |
| 599 | sids: Optional[torch.Tensor] = None, |
| 600 | lids: Optional[torch.Tensor] = None, |
| 601 | is_inference: bool = False, |
| 602 | alpha: float = 1.0, |
| 603 | use_teacher_forcing: bool = False, |
| 604 | ) -> Sequence[torch.Tensor]: |
| 605 | """Calculate forward propagation without loss. |
| 606 | |
| 607 | Args: |
| 608 | xs (Tensor): Batch of padded target features (B, T_feats, odim). |
| 609 | ilens (LongTensor): Batch of the lengths of each target (B,). |
| 610 | |
| 611 | Returns: |
| 612 | Tensor: Weight value if not joint training else model outputs. |
| 613 | |
| 614 | """ |
| 615 | # forward encoder |
| 616 | x_masks = self._source_mask(ilens) |
| 617 | hs, _ = self.encoder(xs, x_masks) # (B, T_text, adim) |
| 618 | |
| 619 | # integrate with GST |
| 620 | if self.use_gst: |
| 621 | style_embs = self.gst(ys) |
| 622 | hs = hs + style_embs.unsqueeze(1) |
| 623 | |
| 624 | # integrate with SID and LID embeddings |
| 625 | if self.spks is not None: |
| 626 | sid_embs = self.sid_emb(sids.view(-1)) |
| 627 | hs = hs + sid_embs.unsqueeze(1) |
| 628 | if self.langs is not None: |
| 629 | lid_embs = self.lid_emb(lids.view(-1)) |
| 630 | hs = hs + lid_embs.unsqueeze(1) |
| 631 | |
| 632 | # integrate speaker embedding |
| 633 | if self.spk_embed_dim is not None: |
| 634 | hs = self._integrate_with_spk_embed(hs, spembs) |
| 635 | |
| 636 | # forward duration predictor and variance predictors |
| 637 | d_masks = make_pad_mask(ilens).to(xs.device) |
| 638 | |
| 639 | if self.stop_gradient_from_pitch_predictor: |
| 640 | p_outs = self.pitch_predictor(hs.detach(), d_masks.unsqueeze(-1)) |
| 641 | else: |
| 642 | p_outs = self.pitch_predictor(hs, d_masks.unsqueeze(-1)) |
| 643 | if self.stop_gradient_from_energy_predictor: |
| 644 | e_outs = self.energy_predictor(hs.detach(), d_masks.unsqueeze(-1)) |
| 645 | else: |
| 646 | e_outs = self.energy_predictor(hs, d_masks.unsqueeze(-1)) |
no test coverage detected