Fix: support infer 2.2 models (#244)
* Fix: support infer 2.2 models * Fix imports * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
This commit is contained in:
@@ -62,7 +62,7 @@ class OnnxInferenceSession:
|
||||
tone = np.expand_dims(tone, 0)
|
||||
if language.ndim == 1:
|
||||
language = np.expand_dims(language, 0)
|
||||
assert(seq.ndim == 2,tone.ndim == 2,language.ndim == 2)
|
||||
assert (seq.ndim == 2, tone.ndim == 2, language.ndim == 2)
|
||||
g = self.emb_g.run(
|
||||
None,
|
||||
{
|
||||
|
||||
@@ -63,7 +63,7 @@ class OnnxInferenceSession:
|
||||
tone = np.expand_dims(tone, 0)
|
||||
if language.ndim == 1:
|
||||
language = np.expand_dims(language, 0)
|
||||
assert(seq.ndim == 2,tone.ndim == 2,language.ndim == 2)
|
||||
assert (seq.ndim == 2, tone.ndim == 2, language.ndim == 2)
|
||||
g = self.emb_g.run(
|
||||
None,
|
||||
{
|
||||
@@ -82,7 +82,7 @@ class OnnxInferenceSession:
|
||||
"bert_2": bert_en.astype(np.float32),
|
||||
"g": g.astype(np.float32),
|
||||
"vqidx": vqidx.astype(np.int64),
|
||||
"sid": sid.astype(np.int64)
|
||||
"sid": sid.astype(np.int64),
|
||||
},
|
||||
)
|
||||
x, m_p, logs_p, x_mask = enc_rtn[0], enc_rtn[1], enc_rtn[2], enc_rtn[3]
|
||||
|
||||
@@ -63,7 +63,7 @@ class OnnxInferenceSession:
|
||||
tone = np.expand_dims(tone, 0)
|
||||
if language.ndim == 1:
|
||||
language = np.expand_dims(language, 0)
|
||||
assert(seq.ndim == 2,tone.ndim == 2,language.ndim == 2)
|
||||
assert (seq.ndim == 2, tone.ndim == 2, language.ndim == 2)
|
||||
g = self.emb_g.run(
|
||||
None,
|
||||
{
|
||||
|
||||
@@ -15,8 +15,6 @@ from commons import init_weights, get_padding
|
||||
from .text import symbols, num_tones, num_languages
|
||||
|
||||
|
||||
|
||||
|
||||
class DurationDiscriminator(nn.Module): # vits2
|
||||
def __init__(
|
||||
self, in_channels, filter_channels, kernel_size, p_dropout, gin_channels=0
|
||||
@@ -335,8 +333,12 @@ class TextEncoder(nn.Module):
|
||||
def forward(self, x, x_lengths, tone, language, bert, ja_bert, en_bert, g=None):
|
||||
x_mask = torch.ones_like(x).unsqueeze(0)
|
||||
bert_emb = self.bert_proj(bert.transpose(0, 1).unsqueeze(0)).transpose(1, 2)
|
||||
ja_bert_emb = self.ja_bert_proj(ja_bert.transpose(0, 1).unsqueeze(0)).transpose(1, 2)
|
||||
en_bert_emb = self.en_bert_proj(en_bert.transpose(0, 1).unsqueeze(0)).transpose(1, 2)
|
||||
ja_bert_emb = self.ja_bert_proj(ja_bert.transpose(0, 1).unsqueeze(0)).transpose(
|
||||
1, 2
|
||||
)
|
||||
en_bert_emb = self.en_bert_proj(en_bert.transpose(0, 1).unsqueeze(0)).transpose(
|
||||
1, 2
|
||||
)
|
||||
x = (
|
||||
self.emb(x)
|
||||
+ self.tone_emb(tone)
|
||||
@@ -795,7 +797,7 @@ class SynthesizerTrn(nn.Module):
|
||||
n_layers_trans_flow=4,
|
||||
flow_share_parameter=False,
|
||||
use_transformer_flow=True,
|
||||
**kwargs
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.n_vocab = n_vocab
|
||||
|
||||
@@ -62,7 +62,7 @@ class OnnxInferenceSession:
|
||||
tone = np.expand_dims(tone, 0)
|
||||
if language.ndim == 1:
|
||||
language = np.expand_dims(language, 0)
|
||||
assert(seq.ndim == 2,tone.ndim == 2,language.ndim == 2)
|
||||
assert (seq.ndim == 2, tone.ndim == 2, language.ndim == 2)
|
||||
g = self.emb_g.run(
|
||||
None,
|
||||
{
|
||||
|
||||
Reference in New Issue
Block a user