diff --git a/app.py b/app.py index c70a89d..a0a215c 100644 --- a/app.py +++ b/app.py @@ -26,7 +26,7 @@ from style_bert_vits2.constants import ( ) from style_bert_vits2.logging import logger from common.tts_model import ModelHolder -from infer import InvalidToneError +from style_bert_vits2.models.infer import InvalidToneError from style_bert_vits2.text_processing.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone from style_bert_vits2.text_processing.japanese.normalizer import normalize_text diff --git a/common/tts_model.py b/common/tts_model.py index 3f17d15..c924686 100644 --- a/common/tts_model.py +++ b/common/tts_model.py @@ -8,10 +8,6 @@ import torch from gradio.processing_utils import convert_to_16_bit_wav import utils -from infer import get_net_g, infer -from models import SynthesizerTrn -from models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra - from style_bert_vits2.constants import ( DEFAULT_ASSIST_TEXT_WEIGHT, DEFAULT_LENGTH, @@ -23,6 +19,9 @@ from style_bert_vits2.constants import ( DEFAULT_STYLE, DEFAULT_STYLE_WEIGHT, ) +from style_bert_vits2.models.infer import get_net_g, infer +from style_bert_vits2.models.models import SynthesizerTrn +from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra from style_bert_vits2.logging import logger diff --git a/attentions.py b/style_bert_vits2/models/attentions.py similarity index 99% rename from attentions.py rename to style_bert_vits2/models/attentions.py index 1cca086..6d43e08 100644 --- a/attentions.py +++ b/style_bert_vits2/models/attentions.py @@ -4,7 +4,6 @@ from torch import nn from torch.nn import functional as F from style_bert_vits2.models import commons -from style_bert_vits2.logging import logger as logging class LayerNorm(nn.Module): @@ -67,7 +66,7 @@ class Encoder(nn.Module): self.cond_layer_idx = ( kwargs["cond_layer_idx"] if "cond_layer_idx" in kwargs else 2 ) - # logging.debug(self.gin_channels, self.cond_layer_idx) + # logger.debug(self.gin_channels, self.cond_layer_idx) assert ( self.cond_layer_idx < self.n_layers ), "cond_layer_idx should be less than n_layers" diff --git a/infer.py b/style_bert_vits2/models/infer.py similarity index 95% rename from infer.py rename to style_bert_vits2/models/infer.py index 4afd048..9abd378 100644 --- a/infer.py +++ b/style_bert_vits2/models/infer.py @@ -1,13 +1,13 @@ import torch -from style_bert_vits2.models import commons import utils -from models import SynthesizerTrn -from models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra from text import cleaned_text_to_sequence, get_bert from text.cleaner import clean_text -from style_bert_vits2.text_processing.symbols import SYMBOLS from style_bert_vits2.logging import logger +from style_bert_vits2.models import commons +from style_bert_vits2.models.models import SynthesizerTrn +from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra +from style_bert_vits2.text_processing.symbols import SYMBOLS class InvalidToneError(ValueError): diff --git a/models.py b/style_bert_vits2/models/models.py similarity index 99% rename from models.py rename to style_bert_vits2/models/models.py index ef581f2..a8c6695 100644 --- a/models.py +++ b/style_bert_vits2/models/models.py @@ -1,5 +1,4 @@ import math -import warnings import torch from torch import nn @@ -7,10 +6,10 @@ from torch.nn import Conv1d, Conv2d, ConvTranspose1d from torch.nn import functional as F from torch.nn.utils import remove_weight_norm, spectral_norm, weight_norm -import attentions -from style_bert_vits2.models import commons -import modules import monotonic_align +from style_bert_vits2.models import attentions +from style_bert_vits2.models import commons +from style_bert_vits2.models import modules from style_bert_vits2.models.commons import get_padding, init_weights from style_bert_vits2.text_processing.symbols import NUM_LANGUAGES, NUM_TONES, SYMBOLS diff --git a/models_jp_extra.py b/style_bert_vits2/models/models_jp_extra.py similarity index 98% rename from models_jp_extra.py rename to style_bert_vits2/models/models_jp_extra.py index 16cacc7..8c4a4ec 100644 --- a/models_jp_extra.py +++ b/style_bert_vits2/models/models_jp_extra.py @@ -1,17 +1,15 @@ import math + import torch from torch import nn +from torch.nn import Conv1d, Conv2d, ConvTranspose1d from torch.nn import functional as F +from torch.nn.utils import remove_weight_norm, spectral_norm, weight_norm -from style_bert_vits2.models import commons -import modules -import attentions import monotonic_align - -from torch.nn import Conv1d, ConvTranspose1d, Conv2d -from torch.nn.utils import weight_norm, remove_weight_norm, spectral_norm - -from style_bert_vits2.models.commons import init_weights, get_padding +from style_bert_vits2.models import attentions +from style_bert_vits2.models import commons +from style_bert_vits2.models import modules from style_bert_vits2.text_processing.symbols import SYMBOLS, NUM_TONES, NUM_LANGUAGES @@ -529,7 +527,7 @@ class Generator(torch.nn.Module): self.resblocks.append(resblock(ch, k, d)) self.conv_post = Conv1d(ch, 1, 7, 1, padding=3, bias=False) - self.ups.apply(init_weights) + self.ups.apply(commons.init_weights) if gin_channels != 0: self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1) @@ -577,7 +575,7 @@ class DiscriminatorP(torch.nn.Module): 32, (kernel_size, 1), (stride, 1), - padding=(get_padding(kernel_size, 1), 0), + padding=(commons.get_padding(kernel_size, 1), 0), ) ), norm_f( @@ -586,7 +584,7 @@ class DiscriminatorP(torch.nn.Module): 128, (kernel_size, 1), (stride, 1), - padding=(get_padding(kernel_size, 1), 0), + padding=(commons.get_padding(kernel_size, 1), 0), ) ), norm_f( @@ -595,7 +593,7 @@ class DiscriminatorP(torch.nn.Module): 512, (kernel_size, 1), (stride, 1), - padding=(get_padding(kernel_size, 1), 0), + padding=(commons.get_padding(kernel_size, 1), 0), ) ), norm_f( @@ -604,7 +602,7 @@ class DiscriminatorP(torch.nn.Module): 1024, (kernel_size, 1), (stride, 1), - padding=(get_padding(kernel_size, 1), 0), + padding=(commons.get_padding(kernel_size, 1), 0), ) ), norm_f( @@ -613,7 +611,7 @@ class DiscriminatorP(torch.nn.Module): 1024, (kernel_size, 1), 1, - padding=(get_padding(kernel_size, 1), 0), + padding=(commons.get_padding(kernel_size, 1), 0), ) ), ] diff --git a/modules.py b/style_bert_vits2/models/modules.py similarity index 95% rename from modules.py rename to style_bert_vits2/models/modules.py index 68b0b9a..e0885c4 100644 --- a/modules.py +++ b/style_bert_vits2/models/modules.py @@ -1,16 +1,14 @@ import math -import warnings import torch from torch import nn from torch.nn import Conv1d from torch.nn import functional as F from torch.nn.utils import remove_weight_norm, weight_norm +from transforms import piecewise_rational_quadratic_transform from style_bert_vits2.models import commons -from attentions import Encoder -from style_bert_vits2.models.commons import get_padding, init_weights -from transforms import piecewise_rational_quadratic_transform +from style_bert_vits2.models.attentions import Encoder LRELU_SLOPE = 0.1 @@ -231,7 +229,7 @@ class ResBlock1(torch.nn.Module): kernel_size, 1, dilation=dilation[0], - padding=get_padding(kernel_size, dilation[0]), + padding=commons.get_padding(kernel_size, dilation[0]), ) ), weight_norm( @@ -241,7 +239,7 @@ class ResBlock1(torch.nn.Module): kernel_size, 1, dilation=dilation[1], - padding=get_padding(kernel_size, dilation[1]), + padding=commons.get_padding(kernel_size, dilation[1]), ) ), weight_norm( @@ -251,12 +249,12 @@ class ResBlock1(torch.nn.Module): kernel_size, 1, dilation=dilation[2], - padding=get_padding(kernel_size, dilation[2]), + padding=commons.get_padding(kernel_size, dilation[2]), ) ), ] ) - self.convs1.apply(init_weights) + self.convs1.apply(commons.init_weights) self.convs2 = nn.ModuleList( [ @@ -267,7 +265,7 @@ class ResBlock1(torch.nn.Module): kernel_size, 1, dilation=1, - padding=get_padding(kernel_size, 1), + padding=commons.get_padding(kernel_size, 1), ) ), weight_norm( @@ -277,7 +275,7 @@ class ResBlock1(torch.nn.Module): kernel_size, 1, dilation=1, - padding=get_padding(kernel_size, 1), + padding=commons.get_padding(kernel_size, 1), ) ), weight_norm( @@ -287,12 +285,12 @@ class ResBlock1(torch.nn.Module): kernel_size, 1, dilation=1, - padding=get_padding(kernel_size, 1), + padding=commons.get_padding(kernel_size, 1), ) ), ] ) - self.convs2.apply(init_weights) + self.convs2.apply(commons.init_weights) def forward(self, x, x_mask=None): for c1, c2 in zip(self.convs1, self.convs2): @@ -328,7 +326,7 @@ class ResBlock2(torch.nn.Module): kernel_size, 1, dilation=dilation[0], - padding=get_padding(kernel_size, dilation[0]), + padding=commons.get_padding(kernel_size, dilation[0]), ) ), weight_norm( @@ -338,12 +336,12 @@ class ResBlock2(torch.nn.Module): kernel_size, 1, dilation=dilation[1], - padding=get_padding(kernel_size, dilation[1]), + padding=commons.get_padding(kernel_size, dilation[1]), ) ), ] ) - self.convs.apply(init_weights) + self.convs.apply(commons.init_weights) def forward(self, x, x_mask=None): for c in self.convs: diff --git a/train_ms.py b/train_ms.py index 783b157..05cd4cb 100644 --- a/train_ms.py +++ b/train_ms.py @@ -1,6 +1,5 @@ import argparse import datetime -import gc import os import platform @@ -15,11 +14,8 @@ from torch.utils.tensorboard import SummaryWriter from tqdm import tqdm # logging.getLogger("numba").setLevel(logging.WARNING) -from style_bert_vits2.models import commons import default_style import utils -from style_bert_vits2.logging import logger -from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT from config import config from data_utils import ( DistributedBucketSampler, @@ -28,8 +24,15 @@ from data_utils import ( ) from losses import discriminator_loss, feature_loss, generator_loss, kl_loss from mel_processing import mel_spectrogram_torch, spec_to_mel_torch -from models import DurationDiscriminator, MultiPeriodDiscriminator, SynthesizerTrn +from style_bert_vits2.logging import logger +from style_bert_vits2.models import commons +from style_bert_vits2.models.models import ( + DurationDiscriminator, + MultiPeriodDiscriminator, + SynthesizerTrn, +) from style_bert_vits2.text_processing.symbols import SYMBOLS +from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = ( diff --git a/train_ms_jp_extra.py b/train_ms_jp_extra.py index e4e4e56..1a287d9 100644 --- a/train_ms_jp_extra.py +++ b/train_ms_jp_extra.py @@ -6,20 +6,17 @@ import platform import torch import torch.distributed as dist +from huggingface_hub import HfApi from torch.cuda.amp import GradScaler, autocast from torch.nn import functional as F from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data import DataLoader from torch.utils.tensorboard import SummaryWriter from tqdm import tqdm -from huggingface_hub import HfApi # logging.getLogger("numba").setLevel(logging.WARNING) -from style_bert_vits2.models import commons import default_style import utils -from style_bert_vits2.logging import logger -from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT from config import config from data_utils import ( DistributedBucketSampler, @@ -28,13 +25,16 @@ from data_utils import ( ) from losses import WavLMLoss, discriminator_loss, feature_loss, generator_loss, kl_loss from mel_processing import mel_spectrogram_torch, spec_to_mel_torch -from models_jp_extra import ( +from style_bert_vits2.logging import logger +from style_bert_vits2.models import commons +from style_bert_vits2.models.models_jp_extra import ( DurationDiscriminator, MultiPeriodDiscriminator, SynthesizerTrn, WavLMDiscriminator, ) from style_bert_vits2.text_processing.symbols import SYMBOLS +from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = ( diff --git a/webui.py b/webui.py index 90318a1..1f31e94 100644 --- a/webui.py +++ b/webui.py @@ -19,15 +19,16 @@ logging.basicConfig( logger = logging.getLogger(__name__) -import torch -import utils -from infer import infer, latest_version, get_net_g, infer_multilang import gradio as gr -import webbrowser -import numpy as np -from config import config -from tools.translate import translate import librosa +import numpy as np +import torch +import webbrowser + +import utils +from config import config +from style_bert_vits2.models.infer import infer, latest_version, get_net_g, infer_multilang +from tools.translate import translate net_g = None