MCPcopy Create free account
hub / github.com/FireRedTeam/FireRedTTS2 / forward

Method forward

fireredtts2/codec/decoder.py:492–521  ·  view source on GitHub ↗

Forward pass of the ISTFTHead module. Args: x (Tensor): Input tensor of shape (B, L, H), where B is the batch size, L is the sequence length, and H denotes the model dimension. Returns: Tensor: Reconstructed time-domain audio

(self, x: torch.Tensor, x_len: torch.Tensor)

Source from the content-addressed store, hash-verified

490 )
491
492 def forward(self, x: torch.Tensor, x_len: torch.Tensor) -> torch.Tensor:
493 """
494 Forward pass of the ISTFTHead module.
495
496 Args:
497 x (Tensor): Input tensor of shape (B, L, H), where B is the batch size,
498 L is the sequence length, and H denotes the model dimension.
499
500 Returns:
501 Tensor: Reconstructed time-domain audio signal of shape (B, T), where T is the length of the output signal.
502 """
503 x_pred = self.out(x)
504 x_pred = x_pred.transpose(1, 2)
505 mag, p = x_pred.chunk(2, dim=1)
506 mag = torch.exp(mag)
507 mag = torch.clip(
508 mag, max=1e2
509 ) # safeguard to prevent excessively large magnitudes
510 # wrapping happens here. These two lines produce real and imaginary value
511 x = torch.cos(p)
512 y = torch.sin(p)
513 # recalculating phase here does not produce anything new
514 # only costs time
515 # phase = torch.atan2(y, x)
516 # S = mag * torch.exp(phase * 1j)
517 # better directly produce the complex value
518 S = mag * (x + 1j * y)
519 audio = self.istft(S)
520 audio_length = x_len * self.hop_length
521 return audio, audio_length
522
523 def forward_chunk(
524 self, x: torch.Tensor, cache: torch.Tensor = None, last_chunk: bool = False

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected