Refactor: moved model, attentions definitions and inference code to style_bert_vits2/models/
The code has not yet been cleaned up, just moved.
This commit is contained in:
2
app.py
2
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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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"
|
||||
@@ -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):
|
||||
@@ -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
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
),
|
||||
]
|
||||
@@ -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:
|
||||
13
train_ms.py
13
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 = (
|
||||
|
||||
@@ -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 = (
|
||||
|
||||
15
webui.py
15
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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user