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