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:
tsukumi
2024-03-06 23:43:25 +00:00
parent a52fda7a88
commit 89825e68d8
10 changed files with 58 additions and 61 deletions

2
app.py
View File

@@ -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

View File

@@ -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

View File

@@ -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"

View File

@@ -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):

View File

@@ -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

View File

@@ -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),
) )
), ),
] ]

View File

@@ -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:

View File

@@ -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 = (

View File

@@ -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 = (

View File

@@ -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