mirror of
https://github.com/wassname/TTS.git
synced 2026-09-09 11:16:00 +08:00
best model ever changes
This commit is contained in:
+3
-5
@@ -9,11 +9,11 @@ from TTS.layers.tacotron import Prenet, Encoder, Decoder, CBHG
|
||||
class Tacotron(nn.Module):
|
||||
def __init__(self, embedding_dim=256, linear_dim=1025, mel_dim=80,
|
||||
freq_dim=1025, r=5, padding_idx=None,
|
||||
use_memory_mask=False):
|
||||
use_atten_mask=False):
|
||||
super(Tacotron, self).__init__()
|
||||
self.mel_dim = mel_dim
|
||||
self.linear_dim = linear_dim
|
||||
self.use_memory_mask = use_memory_mask
|
||||
self.use_atten_mask = use_atten_mask
|
||||
self.embedding = nn.Embedding(len(symbols), embedding_dim,
|
||||
padding_idx=padding_idx)
|
||||
print(" | > Embedding dim : {}".format(len(symbols)))
|
||||
@@ -33,9 +33,7 @@ class Tacotron(nn.Module):
|
||||
# (B, T', in_dim)
|
||||
encoder_outputs = self.encoder(inputs)
|
||||
|
||||
if self.use_memory_mask:
|
||||
input_lengths = input_lengths
|
||||
else:
|
||||
if not self.use_atten_mask:
|
||||
input_lengths = None
|
||||
|
||||
# (B, T', mel_dim*r)
|
||||
|
||||
Reference in New Issue
Block a user