MCPcopy Create free account
hub / github.com/espnet/espnet / _forward

Method _forward

espnet2/tts/prodiff/prodiff.py:589–701  ·  view source on GitHub ↗

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,
    )

Source from the content-addressed store, hash-verified

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))

Callers 2

forwardMethod · 0.95
inferenceMethod · 0.95

Calls 6

_source_maskMethod · 0.95
make_pad_maskFunction · 0.90
toMethod · 0.80
sizeMethod · 0.80
inferenceMethod · 0.45

Tested by

no test coverage detected