improve robustness of defining wavernn in config file

This commit is contained in:
Branislav Gerazov
2021-03-08 02:54:21 +01:00
committed by Eren Gölge
parent 5e2bc8c99f
commit ed56944c4a
+3 -3
View File
@@ -71,10 +71,10 @@ def setup_generator(c):
MyModel = importlib.import_module('TTS.vocoder.models.' +
c.generator_model.lower())
# this is to preserve the WaveRNN class name (instead of Wavernn)
if c.generator_model != 'WaveRNN':
MyModel = getattr(MyModel, to_camel(c.generator_model))
if c.generator_model.lower() == 'wavernn':
MyModel = getattr(MyModel, 'WaveRNN')
else:
MyModel = getattr(MyModel, c.generator_model)
MyModel = getattr(MyModel, to_camel(c.generator_model))
if c.generator_model.lower() in 'wavernn':
model = MyModel(
rnn_dims=c.wavernn_model_params['rnn_dims'],