mirror of
https://github.com/2noise/ChatTTS.git
synced 2026-08-29 02:10:59 +08:00
chore(format): run black on dev (#583)
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
This commit is contained in:
committed by
GitHub
parent
9f0b7a048d
commit
20402c367b
+17
-10
@@ -95,11 +95,14 @@ class GFSQ(nn.Module):
|
||||
feat = self.quantizer.get_output_from_indices(x)
|
||||
return feat.transpose_(1, 2) if self.transpose else feat
|
||||
|
||||
|
||||
def __call__(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
def __call__(
|
||||
self, x: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
return super().__call__(x)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
def forward(
|
||||
self, x: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
if self.transpose:
|
||||
x.transpose_(1, 2)
|
||||
feat, ind = self.quantizer(x)
|
||||
@@ -177,7 +180,7 @@ class MelSpectrogramFeatures(torch.nn.Module):
|
||||
n_fft=1024,
|
||||
hop_length=256,
|
||||
n_mels=100,
|
||||
padding: Literal["center", "same"]="center",
|
||||
padding: Literal["center", "same"] = "center",
|
||||
):
|
||||
super().__init__()
|
||||
if padding not in ["center", "same"]:
|
||||
@@ -221,10 +224,10 @@ class DVAE(nn.Module):
|
||||
|
||||
if encoder_config is not None:
|
||||
self.downsample_conv = nn.Sequential(
|
||||
nn.Conv1d(100, dim, 3, 1, 1),
|
||||
nn.GELU(),
|
||||
nn.Conv1d(dim, dim, 4, 2, 1),
|
||||
nn.GELU()
|
||||
nn.Conv1d(100, dim, 3, 1, 1),
|
||||
nn.GELU(),
|
||||
nn.Conv1d(dim, dim, 4, 2, 1),
|
||||
nn.GELU(),
|
||||
)
|
||||
self.preprocessor_mel = MelSpectrogramFeatures()
|
||||
self.encoder: Optional[DVAEDecoder] = DVAEDecoder(**encoder_config)
|
||||
@@ -241,11 +244,15 @@ class DVAE(nn.Module):
|
||||
self.coef.cpu().numpy().astype(np.float32).tobytes()
|
||||
)
|
||||
|
||||
def __call__(self, inp: torch.Tensor, mode: Literal["encode", "decode"]="decode") -> torch.Tensor:
|
||||
def __call__(
|
||||
self, inp: torch.Tensor, mode: Literal["encode", "decode"] = "decode"
|
||||
) -> torch.Tensor:
|
||||
return super().__call__(inp, mode)
|
||||
|
||||
@torch.inference_mode()
|
||||
def forward(self, inp: torch.Tensor, mode: Literal["encode", "decode"]="decode") -> torch.Tensor:
|
||||
def forward(
|
||||
self, inp: torch.Tensor, mode: Literal["encode", "decode"] = "decode"
|
||||
) -> torch.Tensor:
|
||||
if mode == "encode" and hasattr(self, "encoder") and self.vq_layer is not None:
|
||||
mel = self.preprocessor_mel(inp)
|
||||
x: torch.Tensor = self.downsample_conv(mel / self.coef)
|
||||
|
||||
Reference in New Issue
Block a user