Nymbo commited on
Commit
888f467
·
verified ·
1 Parent(s): c2a90ae

Modernize: Gradio 6.29, Python 3.12, models from CorentinJ/SV2TTS, librosa/numpy/torch compat

Browse files
README.md CHANGED
@@ -3,9 +3,9 @@ title: Clone Your Voice
3
  emoji: 📚
4
  colorFrom: blue
5
  colorTo: yellow
6
- python_version: 3.8.4
7
  sdk: gradio
8
- sdk_version: 3.0.4
9
  app_file: app.py
10
  pinned: false
11
  ---
 
3
  emoji: 📚
4
  colorFrom: blue
5
  colorTo: yellow
6
+ python_version: "3.12"
7
  sdk: gradio
8
+ sdk_version: 6.29.0
9
  app_file: app.py
10
  pinned: false
11
  ---
app.py CHANGED
@@ -1,380 +1,374 @@
1
- import gradio as gr
2
- import os
3
- from utils.default_models import ensure_default_models
4
- import sys
5
- import traceback
6
- from pathlib import Path
7
- from time import perf_counter as timer
8
- import numpy as np
9
- import torch
10
- from encoder import inference as encoder
11
- from synthesizer.inference import Synthesizer
12
- #from toolbox.utterance import Utterance
13
- from vocoder import inference as vocoder
14
- import time
15
- import librosa
16
- import numpy as np
17
- #import sounddevice as sd
18
- import soundfile as sf
19
- import argparse
20
- from utils.argutils import print_args
21
-
22
- parser = argparse.ArgumentParser(
23
- formatter_class=argparse.ArgumentDefaultsHelpFormatter
24
- )
25
- parser.add_argument("-e", "--enc_model_fpath", type=Path,
26
- default="saved_models/default/encoder.pt",
27
- help="Path to a saved encoder")
28
- parser.add_argument("-s", "--syn_model_fpath", type=Path,
29
- default="saved_models/default/synthesizer.pt",
30
- help="Path to a saved synthesizer")
31
- parser.add_argument("-v", "--voc_model_fpath", type=Path,
32
- default="saved_models/default/vocoder.pt",
33
- help="Path to a saved vocoder")
34
- parser.add_argument("--cpu", action="store_true", help=\
35
- "If True, processing is done on CPU, even when a GPU is available.")
36
- parser.add_argument("--no_sound", action="store_true", help=\
37
- "If True, audio won't be played.")
38
- parser.add_argument("--seed", type=int, default=None, help=\
39
- "Optional random number seed value to make toolbox deterministic.")
40
- args = parser.parse_args()
41
- arg_dict = vars(args)
42
- print_args(args, parser)
43
-
44
- # Maximum of generated wavs to keep on memory
45
- MAX_WAVS = 15
46
- utterances = set()
47
- current_generated = (None, None, None, None) # speaker_name, spec, breaks, wav
48
- synthesizer = None # type: Synthesizer
49
- current_wav = None
50
- waves_list = []
51
- waves_count = 0
52
- waves_namelist = []
53
-
54
- # Hide GPUs from Pytorch to force CPU processing
55
- if arg_dict.pop("cpu"):
56
- os.environ["CUDA_VISIBLE_DEVICES"] = "-1"
57
-
58
- print("Running a test of your configuration...\n")
59
-
60
- if torch.cuda.is_available():
61
- device_id = torch.cuda.current_device()
62
- gpu_properties = torch.cuda.get_device_properties(device_id)
63
- ## Print some environment information (for debugging purposes)
64
- print("Found %d GPUs available. Using GPU %d (%s) of compute capability %d.%d with "
65
- "%.1fGb total memory.\n" %
66
- (torch.cuda.device_count(),
67
- device_id,
68
- gpu_properties.name,
69
- gpu_properties.major,
70
- gpu_properties.minor,
71
- gpu_properties.total_memory / 1e9))
72
- else:
73
- print("Using CPU for inference.\n")
74
-
75
- ## Load the models one by one.
76
- print("Preparing the encoder, the synthesizer and the vocoder...")
77
- ensure_default_models(Path("saved_models"))
78
- #encoder.load_model(args.enc_model_fpath)
79
- #synthesizer = Synthesizer(args.syn_model_fpath)
80
- #vocoder.load_model(args.voc_model_fpath)
81
-
82
- def compute_embedding(in_fpath):
83
-
84
- if not encoder.is_loaded():
85
- model_fpath = args.enc_model_fpath
86
- print("Loading the encoder %s... " % model_fpath)
87
- start = time.time()
88
- encoder.load_model(model_fpath)
89
- print("Done (%dms)." % int(1000 * (time.time() - start)), "append")
90
-
91
-
92
- ## Computing the embedding
93
- # First, we load the wav using the function that the speaker encoder provides. This is
94
-
95
- # Get the wav from the disk. We take the wav with the vocoder/synthesizer format for
96
- # playback, so as to have a fair comparison with the generated audio
97
- print("Step 1- load_preprocess_wav",in_fpath)
98
- wav = Synthesizer.load_preprocess_wav(in_fpath)
99
-
100
- # important: there is preprocessing that must be applied.
101
-
102
- # The following two methods are equivalent:
103
- # - Directly load from the filepath:
104
- print("Step 2- preprocess_wav")
105
- preprocessed_wav = encoder.preprocess_wav(wav)
106
-
107
- # - If the wav is already loaded:
108
- #original_wav, sampling_rate = librosa.load(str(in_fpath))
109
- #preprocessed_wav = encoder.preprocess_wav(original_wav, sampling_rate)
110
-
111
- # Compute the embedding
112
- print("Step 3- embed_utterance")
113
- embed, partial_embeds, _ = encoder.embed_utterance(preprocessed_wav, return_partials=True)
114
-
115
-
116
- print("Loaded file succesfully")
117
-
118
- # Then we derive the embedding. There are many functions and parameters that the
119
- # speaker encoder interfaces. These are mostly for in-depth research. You will typically
120
- # only use this function (with its default parameters):
121
- #embed = encoder.embed_utterance(preprocessed_wav)
122
-
123
- return embed
124
- def create_spectrogram(text,embed):
125
- # If seed is specified, reset torch seed and force synthesizer reload
126
- if args.seed is not None:
127
- torch.manual_seed(args.seed)
128
- synthesizer = Synthesizer(args.syn_model_fpath)
129
-
130
-
131
- # Synthesize the spectrogram
132
- model_fpath = args.syn_model_fpath
133
- print("Loading the synthesizer %s... " % model_fpath)
134
- start = time.time()
135
- synthesizer = Synthesizer(model_fpath)
136
- print("Done (%dms)." % int(1000 * (time.time()- start)), "append")
137
-
138
-
139
- # The synthesizer works in batch, so you need to put your data in a list or numpy array
140
- texts = [text]
141
- embeds = [embed]
142
- # If you know what the attention layer alignments are, you can retrieve them here by
143
- # passing return_alignments=True
144
- specs = synthesizer.synthesize_spectrograms(texts, embeds)
145
- breaks = [spec.shape[1] for spec in specs]
146
- spec = np.concatenate(specs, axis=1)
147
- sample_rate=synthesizer.sample_rate
148
- return spec, breaks , sample_rate
149
-
150
-
151
- def generate_waveform(current_generated):
152
-
153
- speaker_name, spec, breaks = current_generated
154
- assert spec is not None
155
-
156
- ## Generating the waveform
157
- print("Synthesizing the waveform:")
158
- # If seed is specified, reset torch seed and reload vocoder
159
- if args.seed is not None:
160
- torch.manual_seed(args.seed)
161
- vocoder.load_model(args.voc_model_fpath)
162
-
163
- model_fpath = args.voc_model_fpath
164
- # Synthesize the waveform
165
- if not vocoder.is_loaded():
166
- print("Loading the vocoder %s... " % model_fpath)
167
- start = time.time()
168
- vocoder.load_model(model_fpath)
169
- print("Done (%dms)." % int(1000 * (time.time()- start)), "append")
170
-
171
- current_vocoder_fpath= model_fpath
172
- def vocoder_progress(i, seq_len, b_size, gen_rate):
173
- real_time_factor = (gen_rate / Synthesizer.sample_rate) * 1000
174
- line = "Waveform generation: %d/%d (batch size: %d, rate: %.1fkHz - %.2fx real time)" \
175
- % (i * b_size, seq_len * b_size, b_size, gen_rate, real_time_factor)
176
- print(line, "overwrite")
177
-
178
-
179
- # Synthesizing the waveform is fairly straightforward. Remember that the longer the
180
- # spectrogram, the more time-efficient the vocoder.
181
- if current_vocoder_fpath is not None:
182
- print("")
183
- generated_wav = vocoder.infer_waveform(spec, progress_callback=vocoder_progress)
184
- else:
185
- print("Waveform generation with Griffin-Lim... ")
186
- generated_wav = Synthesizer.griffin_lim(spec)
187
-
188
- print(" Done!", "append")
189
-
190
-
191
- ## Post-generation
192
- # There's a bug with sounddevice that makes the audio cut one second earlier, so we
193
- # pad it.
194
- generated_wav = np.pad(generated_wav, (0, Synthesizer.sample_rate), mode="constant")
195
-
196
- # Add breaks
197
- b_ends = np.cumsum(np.array(breaks) * Synthesizer.hparams.hop_size)
198
- b_starts = np.concatenate(([0], b_ends[:-1]))
199
- wavs = [generated_wav[start:end] for start, end, in zip(b_starts, b_ends)]
200
- breaks = [np.zeros(int(0.15 * Synthesizer.sample_rate))] * len(breaks)
201
- generated_wav = np.concatenate([i for w, b in zip(wavs, breaks) for i in (w, b)])
202
-
203
-
204
- # Trim excess silences to compensate for gaps in spectrograms (issue #53)
205
- generated_wav = encoder.preprocess_wav(generated_wav)
206
-
207
-
208
- return generated_wav
209
-
210
-
211
- def save_on_disk(generated_wav,sample_rate):
212
- # Save it on the disk
213
- filename = "cloned_voice.wav"
214
- print(generated_wav.dtype)
215
- #OUT=os.environ['OUT_PATH']
216
- # Returns `None` if key doesn't exist
217
- #OUT=os.environ.get('OUT_PATH')
218
- #result = os.path.join(OUT, filename)
219
- result = filename
220
- print(" > Saving output to {}".format(result))
221
- sf.write(result, generated_wav.astype(np.float32), sample_rate)
222
- print("\nSaved output as %s\n\n" % result)
223
-
224
- return result
225
- def play_audio(generated_wav,sample_rate):
226
- # Play the audio (non-blocking)
227
- if not args.no_sound:
228
-
229
- try:
230
- sd.stop()
231
- sd.play(generated_wav, sample_rate)
232
- except sd.PortAudioError as e:
233
- print("\nCaught exception: %s" % repr(e))
234
- print("Continuing without audio playback. Suppress this message with the \"--no_sound\" flag.\n")
235
- except:
236
- raise
237
-
238
-
239
- def clean_memory():
240
- import gc
241
- #import GPUtil
242
- # To see memory usage
243
- print('Before clean ')
244
- #GPUtil.showUtilization()
245
- #cleaning memory 1
246
- gc.collect()
247
- torch.cuda.empty_cache()
248
- time.sleep(2)
249
- print('After Clean GPU')
250
- #GPUtil.showUtilization()
251
-
252
- def clone_voice(in_fpath, text):
253
- try:
254
- speaker_name = "output"
255
- # Compute embedding
256
- embed=compute_embedding(in_fpath)
257
- print("Created the embedding")
258
- # Generating the spectrogram
259
- spec, breaks, sample_rate = create_spectrogram(text,embed)
260
- current_generated = (speaker_name, spec, breaks)
261
- print("Created the mel spectrogram")
262
-
263
- # Create waveform
264
- generated_wav=generate_waveform(current_generated)
265
- print("Created the the waveform ")
266
-
267
- # Save it on the disk
268
- save_on_disk(generated_wav,sample_rate)
269
-
270
- #Play the audio
271
- #play_audio(generated_wav,sample_rate)
272
-
273
- return
274
- except Exception as e:
275
- print("Caught exception: %s" % repr(e))
276
- print("Restarting\n")
277
-
278
- # Set environment variables
279
- home_dir = os.getcwd()
280
- OUT_PATH=os.path.join(home_dir, "out/")
281
- os.environ['OUT_PATH'] = OUT_PATH
282
-
283
- # create output path
284
- os.makedirs(OUT_PATH, exist_ok=True)
285
-
286
- USE_CUDA = torch.cuda.is_available()
287
-
288
- os.system('pip install -q pydub ffmpeg-normalize')
289
- CONFIG_SE_PATH = "config_se.json"
290
- CHECKPOINT_SE_PATH = "SE_checkpoint.pth.tar"
291
- def greet(Text,Voicetoclone ,input_mic=None):
292
- text= "%s" % (Text)
293
- #reference_files= "%s" % (Voicetoclone)
294
-
295
- clean_memory()
296
- print(text,len(text),type(text))
297
- print(Voicetoclone,type(Voicetoclone))
298
-
299
- if len(text) == 0 :
300
- print("Please add text to the program")
301
- Text="Please add text to the program, thank you."
302
- is_no_text=True
303
- else:
304
- is_no_text=False
305
-
306
-
307
- if Voicetoclone==None and input_mic==None:
308
- print("There is no input audio")
309
- Text="Please add audio input, to the program, thank you."
310
- Voicetoclone='trump.mp3'
311
- if is_no_text:
312
- Text="Please add text and audio, to the program, thank you."
313
-
314
- if input_mic != "" and input_mic != None :
315
- # Get the wav file from the microphone
316
- print('The value of MIC IS :',input_mic,type(input_mic))
317
- Voicetoclone= input_mic
318
-
319
- text= "%s" % (Text)
320
- reference_files= Voicetoclone
321
- print("path url")
322
- print(Voicetoclone)
323
- sample= str(Voicetoclone)
324
- os.environ['sample'] = sample
325
- size= len(reference_files)*sys.getsizeof(reference_files)
326
- size2= size / 1000000
327
- if (size2 > 0.012) or len(text)>2000:
328
- message="File is greater than 30mb or Text inserted is longer than 2000 characters. Please re-try with smaller sizes."
329
- print(message)
330
- raise SystemExit("File is greater than 30mb. Please re-try or Text inserted is longer than 2000 characters. Please re-try with smaller sizes.")
331
- else:
332
-
333
- env_var = 'sample'
334
- if env_var in os.environ:
335
- print(f'{env_var} value is {os.environ[env_var]}')
336
- else:
337
- print(f'{env_var} does not exist')
338
- #os.system(f'ffmpeg-normalize {os.environ[env_var]} -nt rms -t=-27 -o {os.environ[env_var]} -ar 16000 -f')
339
- in_fpath = Path(Voicetoclone)
340
- #in_fpath= in_fpath.replace("\"", "").replace("\'", "")
341
-
342
- out_path=clone_voice(in_fpath, text)
343
-
344
- print(" > text: {}".format(text))
345
-
346
- print("Generated Audio")
347
- return "cloned_voice.wav"
348
-
349
- demo = gr.Interface(
350
- fn=greet,
351
- inputs=[gr.inputs.Textbox(label='What would you like the voice to say? (max. 2000 characters per request)'),
352
- gr.Audio(
353
- type="filepath",
354
- source="upload",
355
- label='Please upload a voice to clone (max. 30mb)'),
356
- gr.inputs.Audio(
357
- source="microphone",
358
- label='or record',
359
- type="filepath",
360
- optional=True)
361
- ],
362
- outputs="audio",
363
-
364
- title = 'Clone Your Voice',
365
- description = 'A simple application that Clone Your Voice. Wait one minute to process.',
366
- article =
367
- '''<div>
368
- <p style="text-align: center"> All you need to do is record your voice, type what you want be say
369
- ,then wait for compiling. After that click on Play/Pause for listen the audio. The audio is saved in an wav format.
370
- For more information visit <a href="https://ruslanmv.com/">ruslanmv.com</a>
371
- </p>
372
- </div>''',
373
-
374
- examples = [["I am the cloned version of Donald Trump. Well. I think what's happening to this country is unbelievably bad. We're no longer a respected country","trump.mp3","trump.mp3"],
375
- ["I am the cloned version of Elon Musk. Persistence is very important. You should not give up unless you are forced to give up.","musk.mp3","musk.mp3"] #,
376
- # ["I am the cloned version of Elizabeth. It has always been easy to hate and destroy. To build and to cherish is much more difficult." ,"queen.mp3","queen.mp3"]
377
- ]
378
-
379
- )
380
  demo.launch()
 
1
+ import gradio as gr
2
+ import os
3
+ from utils.default_models import ensure_default_models
4
+ import sys
5
+ import traceback
6
+ from pathlib import Path
7
+ from time import perf_counter as timer
8
+ import numpy as np
9
+ import torch
10
+ from encoder import inference as encoder
11
+ from synthesizer.inference import Synthesizer
12
+ #from toolbox.utterance import Utterance
13
+ from vocoder import inference as vocoder
14
+ import time
15
+ import librosa
16
+ import numpy as np
17
+ #import sounddevice as sd
18
+ import soundfile as sf
19
+ import argparse
20
+ from utils.argutils import print_args
21
+
22
+ parser = argparse.ArgumentParser(
23
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter
24
+ )
25
+ parser.add_argument("-e", "--enc_model_fpath", type=Path,
26
+ default="saved_models/default/encoder.pt",
27
+ help="Path to a saved encoder")
28
+ parser.add_argument("-s", "--syn_model_fpath", type=Path,
29
+ default="saved_models/default/synthesizer.pt",
30
+ help="Path to a saved synthesizer")
31
+ parser.add_argument("-v", "--voc_model_fpath", type=Path,
32
+ default="saved_models/default/vocoder.pt",
33
+ help="Path to a saved vocoder")
34
+ parser.add_argument("--cpu", action="store_true", help=\
35
+ "If True, processing is done on CPU, even when a GPU is available.")
36
+ parser.add_argument("--no_sound", action="store_true", help=\
37
+ "If True, audio won't be played.")
38
+ parser.add_argument("--seed", type=int, default=None, help=\
39
+ "Optional random number seed value to make toolbox deterministic.")
40
+ args, _unknown = parser.parse_known_args()
41
+ arg_dict = vars(args)
42
+ print_args(args, parser)
43
+
44
+ # Maximum of generated wavs to keep on memory
45
+ MAX_WAVS = 15
46
+ utterances = set()
47
+ current_generated = (None, None, None, None) # speaker_name, spec, breaks, wav
48
+ synthesizer = None # type: Synthesizer
49
+ current_wav = None
50
+ waves_list = []
51
+ waves_count = 0
52
+ waves_namelist = []
53
+
54
+ # Hide GPUs from Pytorch to force CPU processing
55
+ if arg_dict.pop("cpu"):
56
+ os.environ["CUDA_VISIBLE_DEVICES"] = "-1"
57
+
58
+ print("Running a test of your configuration...\n")
59
+
60
+ if torch.cuda.is_available():
61
+ device_id = torch.cuda.current_device()
62
+ gpu_properties = torch.cuda.get_device_properties(device_id)
63
+ ## Print some environment information (for debugging purposes)
64
+ print("Found %d GPUs available. Using GPU %d (%s) of compute capability %d.%d with "
65
+ "%.1fGb total memory.\n" %
66
+ (torch.cuda.device_count(),
67
+ device_id,
68
+ gpu_properties.name,
69
+ gpu_properties.major,
70
+ gpu_properties.minor,
71
+ gpu_properties.total_memory / 1e9))
72
+ else:
73
+ print("Using CPU for inference.\n")
74
+
75
+ ## Load the models one by one.
76
+ print("Preparing the encoder, the synthesizer and the vocoder...")
77
+ ensure_default_models(Path("saved_models"))
78
+ #encoder.load_model(args.enc_model_fpath)
79
+ #synthesizer = Synthesizer(args.syn_model_fpath)
80
+ #vocoder.load_model(args.voc_model_fpath)
81
+
82
+ def compute_embedding(in_fpath):
83
+
84
+ if not encoder.is_loaded():
85
+ model_fpath = args.enc_model_fpath
86
+ print("Loading the encoder %s... " % model_fpath)
87
+ start = time.time()
88
+ encoder.load_model(model_fpath)
89
+ print("Done (%dms)." % int(1000 * (time.time() - start)), "append")
90
+
91
+
92
+ ## Computing the embedding
93
+ # First, we load the wav using the function that the speaker encoder provides. This is
94
+
95
+ # Get the wav from the disk. We take the wav with the vocoder/synthesizer format for
96
+ # playback, so as to have a fair comparison with the generated audio
97
+ print("Step 1- load_preprocess_wav",in_fpath)
98
+ wav = Synthesizer.load_preprocess_wav(in_fpath)
99
+
100
+ # important: there is preprocessing that must be applied.
101
+
102
+ # The following two methods are equivalent:
103
+ # - Directly load from the filepath:
104
+ print("Step 2- preprocess_wav")
105
+ preprocessed_wav = encoder.preprocess_wav(wav)
106
+
107
+ # - If the wav is already loaded:
108
+ #original_wav, sampling_rate = librosa.load(str(in_fpath))
109
+ #preprocessed_wav = encoder.preprocess_wav(original_wav, sampling_rate)
110
+
111
+ # Compute the embedding
112
+ print("Step 3- embed_utterance")
113
+ embed, partial_embeds, _ = encoder.embed_utterance(preprocessed_wav, return_partials=True)
114
+
115
+
116
+ print("Loaded file succesfully")
117
+
118
+ # Then we derive the embedding. There are many functions and parameters that the
119
+ # speaker encoder interfaces. These are mostly for in-depth research. You will typically
120
+ # only use this function (with its default parameters):
121
+ #embed = encoder.embed_utterance(preprocessed_wav)
122
+
123
+ return embed
124
+ def create_spectrogram(text,embed):
125
+ # If seed is specified, reset torch seed and force synthesizer reload
126
+ if args.seed is not None:
127
+ torch.manual_seed(args.seed)
128
+ synthesizer = Synthesizer(args.syn_model_fpath)
129
+
130
+
131
+ # Synthesize the spectrogram
132
+ model_fpath = args.syn_model_fpath
133
+ print("Loading the synthesizer %s... " % model_fpath)
134
+ start = time.time()
135
+ synthesizer = Synthesizer(model_fpath)
136
+ print("Done (%dms)." % int(1000 * (time.time()- start)), "append")
137
+
138
+
139
+ # The synthesizer works in batch, so you need to put your data in a list or numpy array
140
+ texts = [text]
141
+ embeds = [embed]
142
+ # If you know what the attention layer alignments are, you can retrieve them here by
143
+ # passing return_alignments=True
144
+ specs = synthesizer.synthesize_spectrograms(texts, embeds)
145
+ breaks = [spec.shape[1] for spec in specs]
146
+ spec = np.concatenate(specs, axis=1)
147
+ sample_rate=synthesizer.sample_rate
148
+ return spec, breaks , sample_rate
149
+
150
+
151
+ def generate_waveform(current_generated):
152
+
153
+ speaker_name, spec, breaks = current_generated
154
+ assert spec is not None
155
+
156
+ ## Generating the waveform
157
+ print("Synthesizing the waveform:")
158
+ # If seed is specified, reset torch seed and reload vocoder
159
+ if args.seed is not None:
160
+ torch.manual_seed(args.seed)
161
+ vocoder.load_model(args.voc_model_fpath)
162
+
163
+ model_fpath = args.voc_model_fpath
164
+ # Synthesize the waveform
165
+ if not vocoder.is_loaded():
166
+ print("Loading the vocoder %s... " % model_fpath)
167
+ start = time.time()
168
+ vocoder.load_model(model_fpath)
169
+ print("Done (%dms)." % int(1000 * (time.time()- start)), "append")
170
+
171
+ current_vocoder_fpath= model_fpath
172
+ def vocoder_progress(i, seq_len, b_size, gen_rate):
173
+ real_time_factor = (gen_rate / Synthesizer.sample_rate) * 1000
174
+ line = "Waveform generation: %d/%d (batch size: %d, rate: %.1fkHz - %.2fx real time)" \
175
+ % (i * b_size, seq_len * b_size, b_size, gen_rate, real_time_factor)
176
+ print(line, "overwrite")
177
+
178
+
179
+ # Synthesizing the waveform is fairly straightforward. Remember that the longer the
180
+ # spectrogram, the more time-efficient the vocoder.
181
+ if current_vocoder_fpath is not None:
182
+ print("")
183
+ generated_wav = vocoder.infer_waveform(spec, progress_callback=vocoder_progress)
184
+ else:
185
+ print("Waveform generation with Griffin-Lim... ")
186
+ generated_wav = Synthesizer.griffin_lim(spec)
187
+
188
+ print(" Done!", "append")
189
+
190
+
191
+ ## Post-generation
192
+ # There's a bug with sounddevice that makes the audio cut one second earlier, so we
193
+ # pad it.
194
+ generated_wav = np.pad(generated_wav, (0, Synthesizer.sample_rate), mode="constant")
195
+
196
+ # Add breaks
197
+ b_ends = np.cumsum(np.array(breaks) * Synthesizer.hparams.hop_size)
198
+ b_starts = np.concatenate(([0], b_ends[:-1]))
199
+ wavs = [generated_wav[start:end] for start, end, in zip(b_starts, b_ends)]
200
+ breaks = [np.zeros(int(0.15 * Synthesizer.sample_rate))] * len(breaks)
201
+ generated_wav = np.concatenate([i for w, b in zip(wavs, breaks) for i in (w, b)])
202
+
203
+
204
+ # Trim excess silences to compensate for gaps in spectrograms (issue #53)
205
+ generated_wav = encoder.preprocess_wav(generated_wav)
206
+
207
+
208
+ return generated_wav
209
+
210
+
211
+ def save_on_disk(generated_wav,sample_rate):
212
+ # Save it on the disk
213
+ import tempfile
214
+ filename = tempfile.NamedTemporaryFile(suffix="_cloned_voice.wav", delete=False).name
215
+ print(generated_wav.dtype)
216
+ #OUT=os.environ['OUT_PATH']
217
+ # Returns `None` if key doesn't exist
218
+ #OUT=os.environ.get('OUT_PATH')
219
+ #result = os.path.join(OUT, filename)
220
+ result = filename
221
+ print(" > Saving output to {}".format(result))
222
+ sf.write(result, generated_wav.astype(np.float32), sample_rate)
223
+ print("\nSaved output as %s\n\n" % result)
224
+
225
+ return result
226
+ def play_audio(generated_wav,sample_rate):
227
+ # Play the audio (non-blocking)
228
+ if not args.no_sound:
229
+
230
+ try:
231
+ sd.stop()
232
+ sd.play(generated_wav, sample_rate)
233
+ except sd.PortAudioError as e:
234
+ print("\nCaught exception: %s" % repr(e))
235
+ print("Continuing without audio playback. Suppress this message with the \"--no_sound\" flag.\n")
236
+ except:
237
+ raise
238
+
239
+
240
+ def clean_memory():
241
+ import gc
242
+ #import GPUtil
243
+ # To see memory usage
244
+ print('Before clean ')
245
+ #GPUtil.showUtilization()
246
+ #cleaning memory 1
247
+ gc.collect()
248
+ torch.cuda.empty_cache()
249
+ time.sleep(2)
250
+ print('After Clean GPU')
251
+ #GPUtil.showUtilization()
252
+
253
+ def clone_voice(in_fpath, text):
254
+ try:
255
+ speaker_name = "output"
256
+ # Compute embedding
257
+ embed=compute_embedding(in_fpath)
258
+ print("Created the embedding")
259
+ # Generating the spectrogram
260
+ spec, breaks, sample_rate = create_spectrogram(text,embed)
261
+ current_generated = (speaker_name, spec, breaks)
262
+ print("Created the mel spectrogram")
263
+
264
+ # Create waveform
265
+ generated_wav=generate_waveform(current_generated)
266
+ print("Created the the waveform ")
267
+
268
+ # Save it on the disk
269
+ return save_on_disk(generated_wav,sample_rate)
270
+ except Exception as e:
271
+ print("Caught exception: %s" % repr(e))
272
+ raise gr.Error("Voice cloning failed: %s" % e)
273
+
274
+ # Set environment variables
275
+ home_dir = os.getcwd()
276
+ OUT_PATH=os.path.join(home_dir, "out/")
277
+ os.environ['OUT_PATH'] = OUT_PATH
278
+
279
+ # create output path
280
+ os.makedirs(OUT_PATH, exist_ok=True)
281
+
282
+ USE_CUDA = torch.cuda.is_available()
283
+
284
+ CONFIG_SE_PATH = "config_se.json"
285
+ CHECKPOINT_SE_PATH = "SE_checkpoint.pth.tar"
286
+ def greet(Text,Voicetoclone ,input_mic=None):
287
+ text= "%s" % (Text)
288
+ #reference_files= "%s" % (Voicetoclone)
289
+
290
+ clean_memory()
291
+ print(text,len(text),type(text))
292
+ print(Voicetoclone,type(Voicetoclone))
293
+
294
+ if len(text) == 0 :
295
+ print("Please add text to the program")
296
+ Text="Please add text to the program, thank you."
297
+ is_no_text=True
298
+ else:
299
+ is_no_text=False
300
+
301
+
302
+ if Voicetoclone==None and input_mic==None:
303
+ print("There is no input audio")
304
+ Text="Please add audio input, to the program, thank you."
305
+ Voicetoclone='trump.mp3'
306
+ if is_no_text:
307
+ Text="Please add text and audio, to the program, thank you."
308
+
309
+ if input_mic != "" and input_mic != None :
310
+ # Get the wav file from the microphone
311
+ print('The value of MIC IS :',input_mic,type(input_mic))
312
+ Voicetoclone= input_mic
313
+
314
+ text= "%s" % (Text)
315
+ reference_files= Voicetoclone
316
+ print("path url")
317
+ print(Voicetoclone)
318
+ sample= str(Voicetoclone)
319
+ os.environ['sample'] = sample
320
+ size2= os.path.getsize(str(reference_files)) / 1000000 if os.path.exists(str(reference_files)) else 0
321
+ if (size2 > 30) or len(text)>2000:
322
+ message="File is greater than 30mb or Text inserted is longer than 2000 characters. Please re-try with smaller sizes."
323
+ print(message)
324
+ raise gr.Error(message)
325
+ else:
326
+
327
+ env_var = 'sample'
328
+ if env_var in os.environ:
329
+ print(f'{env_var} value is {os.environ[env_var]}')
330
+ else:
331
+ print(f'{env_var} does not exist')
332
+ #os.system(f'ffmpeg-normalize {os.environ[env_var]} -nt rms -t=-27 -o {os.environ[env_var]} -ar 16000 -f')
333
+ in_fpath = Path(Voicetoclone)
334
+ #in_fpath= in_fpath.replace("\"", "").replace("\'", "")
335
+
336
+ out_path=clone_voice(in_fpath, text)
337
+
338
+ print(" > text: {}".format(text))
339
+
340
+ print("Generated Audio")
341
+ return out_path
342
+
343
+ demo = gr.Interface(
344
+ fn=greet,
345
+ inputs=[gr.Textbox(label='What would you like the voice to say? (max. 2000 characters per request)'),
346
+ gr.Audio(
347
+ type="filepath",
348
+ sources=["upload"],
349
+ label='Please upload a voice to clone (max. 30mb)'),
350
+ gr.Audio(
351
+ sources=["microphone"],
352
+ label='or record',
353
+ type="filepath")
354
+ ],
355
+ outputs=gr.Audio(type="filepath"),
356
+ cache_examples=False,
357
+
358
+ title = 'Clone Your Voice',
359
+ description = 'A simple application that Clone Your Voice. Wait one minute to process.',
360
+ article =
361
+ '''<div>
362
+ <p style="text-align: center"> All you need to do is record your voice, type what you want be say
363
+ ,then wait for compiling. After that click on Play/Pause for listen the audio. The audio is saved in an wav format.
364
+ For more information visit <a href="https://ruslanmv.com/">ruslanmv.com</a>
365
+ </p>
366
+ </div>''',
367
+
368
+ examples = [["I am the cloned version of Donald Trump. Well. I think what's happening to this country is unbelievably bad. We're no longer a respected country","trump.mp3","trump.mp3"],
369
+ ["I am the cloned version of Elon Musk. Persistence is very important. You should not give up unless you are forced to give up.","musk.mp3","musk.mp3"] #,
370
+ # ["I am the cloned version of Elizabeth. It has always been easy to hate and destroy. To build and to cherish is much more difficult." ,"queen.mp3","queen.mp3"]
371
+ ]
372
+
373
+ )
 
 
 
 
 
 
374
  demo.launch()
encoder/audio.py CHANGED
@@ -1,117 +1,117 @@
1
- from scipy.ndimage.morphology import binary_dilation
2
- from encoder.params_data import *
3
- from pathlib import Path
4
- from typing import Optional, Union
5
- from warnings import warn
6
- import numpy as np
7
- import librosa
8
- import struct
9
-
10
- try:
11
- import webrtcvad
12
- except:
13
- warn("Unable to import 'webrtcvad'. This package enables noise removal and is recommended.")
14
- webrtcvad=None
15
-
16
- int16_max = (2 ** 15) - 1
17
-
18
-
19
- def preprocess_wav(fpath_or_wav: Union[str, Path, np.ndarray],
20
- source_sr: Optional[int] = None,
21
- normalize: Optional[bool] = True,
22
- trim_silence: Optional[bool] = True):
23
- """
24
- Applies the preprocessing operations used in training the Speaker Encoder to a waveform
25
- either on disk or in memory. The waveform will be resampled to match the data hyperparameters.
26
-
27
- :param fpath_or_wav: either a filepath to an audio file (many extensions are supported, not
28
- just .wav), either the waveform as a numpy array of floats.
29
- :param source_sr: if passing an audio waveform, the sampling rate of the waveform before
30
- preprocessing. After preprocessing, the waveform's sampling rate will match the data
31
- hyperparameters. If passing a filepath, the sampling rate will be automatically detected and
32
- this argument will be ignored.
33
- """
34
- # Load the wav from disk if needed
35
- if isinstance(fpath_or_wav, str) or isinstance(fpath_or_wav, Path):
36
- wav, source_sr = librosa.load(str(fpath_or_wav), sr=None)
37
- else:
38
- wav = fpath_or_wav
39
-
40
- # Resample the wav if needed
41
- if source_sr is not None and source_sr != sampling_rate:
42
- wav = librosa.resample(wav, source_sr, sampling_rate)
43
-
44
- # Apply the preprocessing: normalize volume and shorten long silences
45
- if normalize:
46
- wav = normalize_volume(wav, audio_norm_target_dBFS, increase_only=True)
47
- if webrtcvad and trim_silence:
48
- wav = trim_long_silences(wav)
49
-
50
- return wav
51
-
52
-
53
- def wav_to_mel_spectrogram(wav):
54
- """
55
- Derives a mel spectrogram ready to be used by the encoder from a preprocessed audio waveform.
56
- Note: this not a log-mel spectrogram.
57
- """
58
- frames = librosa.feature.melspectrogram(
59
- wav,
60
- sampling_rate,
61
- n_fft=int(sampling_rate * mel_window_length / 1000),
62
- hop_length=int(sampling_rate * mel_window_step / 1000),
63
- n_mels=mel_n_channels
64
- )
65
- return frames.astype(np.float32).T
66
-
67
-
68
- def trim_long_silences(wav):
69
- """
70
- Ensures that segments without voice in the waveform remain no longer than a
71
- threshold determined by the VAD parameters in params.py.
72
-
73
- :param wav: the raw waveform as a numpy array of floats
74
- :return: the same waveform with silences trimmed away (length <= original wav length)
75
- """
76
- # Compute the voice detection window size
77
- samples_per_window = (vad_window_length * sampling_rate) // 1000
78
-
79
- # Trim the end of the audio to have a multiple of the window size
80
- wav = wav[:len(wav) - (len(wav) % samples_per_window)]
81
-
82
- # Convert the float waveform to 16-bit mono PCM
83
- pcm_wave = struct.pack("%dh" % len(wav), *(np.round(wav * int16_max)).astype(np.int16))
84
-
85
- # Perform voice activation detection
86
- voice_flags = []
87
- vad = webrtcvad.Vad(mode=3)
88
- for window_start in range(0, len(wav), samples_per_window):
89
- window_end = window_start + samples_per_window
90
- voice_flags.append(vad.is_speech(pcm_wave[window_start * 2:window_end * 2],
91
- sample_rate=sampling_rate))
92
- voice_flags = np.array(voice_flags)
93
-
94
- # Smooth the voice detection with a moving average
95
- def moving_average(array, width):
96
- array_padded = np.concatenate((np.zeros((width - 1) // 2), array, np.zeros(width // 2)))
97
- ret = np.cumsum(array_padded, dtype=float)
98
- ret[width:] = ret[width:] - ret[:-width]
99
- return ret[width - 1:] / width
100
-
101
- audio_mask = moving_average(voice_flags, vad_moving_average_width)
102
- audio_mask = np.round(audio_mask).astype(np.bool)
103
-
104
- # Dilate the voiced regions
105
- audio_mask = binary_dilation(audio_mask, np.ones(vad_max_silence_length + 1))
106
- audio_mask = np.repeat(audio_mask, samples_per_window)
107
-
108
- return wav[audio_mask == True]
109
-
110
-
111
- def normalize_volume(wav, target_dBFS, increase_only=False, decrease_only=False):
112
- if increase_only and decrease_only:
113
- raise ValueError("Both increase only and decrease only are set")
114
- dBFS_change = target_dBFS - 10 * np.log10(np.mean(wav ** 2))
115
- if (dBFS_change < 0 and increase_only) or (dBFS_change > 0 and decrease_only):
116
- return wav
117
- return wav * (10 ** (dBFS_change / 20))
 
1
+ from scipy.ndimage import binary_dilation
2
+ from encoder.params_data import *
3
+ from pathlib import Path
4
+ from typing import Optional, Union
5
+ from warnings import warn
6
+ import numpy as np
7
+ import librosa
8
+ import struct
9
+
10
+ try:
11
+ import webrtcvad
12
+ except:
13
+ warn("Unable to import 'webrtcvad'. This package enables noise removal and is recommended.")
14
+ webrtcvad=None
15
+
16
+ int16_max = (2 ** 15) - 1
17
+
18
+
19
+ def preprocess_wav(fpath_or_wav: Union[str, Path, np.ndarray],
20
+ source_sr: Optional[int] = None,
21
+ normalize: Optional[bool] = True,
22
+ trim_silence: Optional[bool] = True):
23
+ """
24
+ Applies the preprocessing operations used in training the Speaker Encoder to a waveform
25
+ either on disk or in memory. The waveform will be resampled to match the data hyperparameters.
26
+
27
+ :param fpath_or_wav: either a filepath to an audio file (many extensions are supported, not
28
+ just .wav), either the waveform as a numpy array of floats.
29
+ :param source_sr: if passing an audio waveform, the sampling rate of the waveform before
30
+ preprocessing. After preprocessing, the waveform's sampling rate will match the data
31
+ hyperparameters. If passing a filepath, the sampling rate will be automatically detected and
32
+ this argument will be ignored.
33
+ """
34
+ # Load the wav from disk if needed
35
+ if isinstance(fpath_or_wav, str) or isinstance(fpath_or_wav, Path):
36
+ wav, source_sr = librosa.load(str(fpath_or_wav), sr=None)
37
+ else:
38
+ wav = fpath_or_wav
39
+
40
+ # Resample the wav if needed
41
+ if source_sr is not None and source_sr != sampling_rate:
42
+ wav = librosa.resample(wav, orig_sr=source_sr, target_sr=sampling_rate)
43
+
44
+ # Apply the preprocessing: normalize volume and shorten long silences
45
+ if normalize:
46
+ wav = normalize_volume(wav, audio_norm_target_dBFS, increase_only=True)
47
+ if webrtcvad and trim_silence:
48
+ wav = trim_long_silences(wav)
49
+
50
+ return wav
51
+
52
+
53
+ def wav_to_mel_spectrogram(wav):
54
+ """
55
+ Derives a mel spectrogram ready to be used by the encoder from a preprocessed audio waveform.
56
+ Note: this not a log-mel spectrogram.
57
+ """
58
+ frames = librosa.feature.melspectrogram(
59
+ y=wav,
60
+ sr=sampling_rate,
61
+ n_fft=int(sampling_rate * mel_window_length / 1000),
62
+ hop_length=int(sampling_rate * mel_window_step / 1000),
63
+ n_mels=mel_n_channels
64
+ )
65
+ return frames.astype(np.float32).T
66
+
67
+
68
+ def trim_long_silences(wav):
69
+ """
70
+ Ensures that segments without voice in the waveform remain no longer than a
71
+ threshold determined by the VAD parameters in params.py.
72
+
73
+ :param wav: the raw waveform as a numpy array of floats
74
+ :return: the same waveform with silences trimmed away (length <= original wav length)
75
+ """
76
+ # Compute the voice detection window size
77
+ samples_per_window = (vad_window_length * sampling_rate) // 1000
78
+
79
+ # Trim the end of the audio to have a multiple of the window size
80
+ wav = wav[:len(wav) - (len(wav) % samples_per_window)]
81
+
82
+ # Convert the float waveform to 16-bit mono PCM
83
+ pcm_wave = struct.pack("%dh" % len(wav), *(np.round(wav * int16_max)).astype(np.int16))
84
+
85
+ # Perform voice activation detection
86
+ voice_flags = []
87
+ vad = webrtcvad.Vad(mode=3)
88
+ for window_start in range(0, len(wav), samples_per_window):
89
+ window_end = window_start + samples_per_window
90
+ voice_flags.append(vad.is_speech(pcm_wave[window_start * 2:window_end * 2],
91
+ sample_rate=sampling_rate))
92
+ voice_flags = np.array(voice_flags)
93
+
94
+ # Smooth the voice detection with a moving average
95
+ def moving_average(array, width):
96
+ array_padded = np.concatenate((np.zeros((width - 1) // 2), array, np.zeros(width // 2)))
97
+ ret = np.cumsum(array_padded, dtype=float)
98
+ ret[width:] = ret[width:] - ret[:-width]
99
+ return ret[width - 1:] / width
100
+
101
+ audio_mask = moving_average(voice_flags, vad_moving_average_width)
102
+ audio_mask = np.round(audio_mask).astype(bool)
103
+
104
+ # Dilate the voiced regions
105
+ audio_mask = binary_dilation(audio_mask, np.ones(vad_max_silence_length + 1))
106
+ audio_mask = np.repeat(audio_mask, samples_per_window)
107
+
108
+ return wav[audio_mask == True]
109
+
110
+
111
+ def normalize_volume(wav, target_dBFS, increase_only=False, decrease_only=False):
112
+ if increase_only and decrease_only:
113
+ raise ValueError("Both increase only and decrease only are set")
114
+ dBFS_change = target_dBFS - 10 * np.log10(np.mean(wav ** 2))
115
+ if (dBFS_change < 0 and increase_only) or (dBFS_change > 0 and decrease_only):
116
+ return wav
117
+ return wav * (10 ** (dBFS_change / 20))
encoder/inference.py CHANGED
@@ -1,178 +1,178 @@
1
- from encoder.params_data import *
2
- from encoder.model import SpeakerEncoder
3
- from encoder.audio import preprocess_wav # We want to expose this function from here
4
- from matplotlib import cm
5
- from encoder import audio
6
- from pathlib import Path
7
- import numpy as np
8
- import torch
9
-
10
- _model = None # type: SpeakerEncoder
11
- _device = None # type: torch.device
12
-
13
-
14
- def load_model(weights_fpath: Path, device=None):
15
- """
16
- Loads the model in memory. If this function is not explicitely called, it will be run on the
17
- first call to embed_frames() with the default weights file.
18
-
19
- :param weights_fpath: the path to saved model weights.
20
- :param device: either a torch device or the name of a torch device (e.g. "cpu", "cuda"). The
21
- model will be loaded and will run on this device. Outputs will however always be on the cpu.
22
- If None, will default to your GPU if it"s available, otherwise your CPU.
23
- """
24
- # TODO: I think the slow loading of the encoder might have something to do with the device it
25
- # was saved on. Worth investigating.
26
- global _model, _device
27
- if device is None:
28
- _device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
29
- elif isinstance(device, str):
30
- _device = torch.device(device)
31
- _model = SpeakerEncoder(_device, torch.device("cpu"))
32
- checkpoint = torch.load(weights_fpath, _device)
33
- _model.load_state_dict(checkpoint["model_state"])
34
- _model.eval()
35
- print("Loaded encoder \"%s\" trained to step %d" % (weights_fpath.name, checkpoint["step"]))
36
-
37
-
38
- def is_loaded():
39
- return _model is not None
40
-
41
-
42
- def embed_frames_batch(frames_batch):
43
- """
44
- Computes embeddings for a batch of mel spectrogram.
45
-
46
- :param frames_batch: a batch mel of spectrogram as a numpy array of float32 of shape
47
- (batch_size, n_frames, n_channels)
48
- :return: the embeddings as a numpy array of float32 of shape (batch_size, model_embedding_size)
49
- """
50
- if _model is None:
51
- raise Exception("Model was not loaded. Call load_model() before inference.")
52
-
53
- frames = torch.from_numpy(frames_batch).to(_device)
54
- embed = _model.forward(frames).detach().cpu().numpy()
55
- return embed
56
-
57
-
58
- def compute_partial_slices(n_samples, partial_utterance_n_frames=partials_n_frames,
59
- min_pad_coverage=0.75, overlap=0.5):
60
- """
61
- Computes where to split an utterance waveform and its corresponding mel spectrogram to obtain
62
- partial utterances of <partial_utterance_n_frames> each. Both the waveform and the mel
63
- spectrogram slices are returned, so as to make each partial utterance waveform correspond to
64
- its spectrogram. This function assumes that the mel spectrogram parameters used are those
65
- defined in params_data.py.
66
-
67
- The returned ranges may be indexing further than the length of the waveform. It is
68
- recommended that you pad the waveform with zeros up to wave_slices[-1].stop.
69
-
70
- :param n_samples: the number of samples in the waveform
71
- :param partial_utterance_n_frames: the number of mel spectrogram frames in each partial
72
- utterance
73
- :param min_pad_coverage: when reaching the last partial utterance, it may or may not have
74
- enough frames. If at least <min_pad_coverage> of <partial_utterance_n_frames> are present,
75
- then the last partial utterance will be considered, as if we padded the audio. Otherwise,
76
- it will be discarded, as if we trimmed the audio. If there aren't enough frames for 1 partial
77
- utterance, this parameter is ignored so that the function always returns at least 1 slice.
78
- :param overlap: by how much the partial utterance should overlap. If set to 0, the partial
79
- utterances are entirely disjoint.
80
- :return: the waveform slices and mel spectrogram slices as lists of array slices. Index
81
- respectively the waveform and the mel spectrogram with these slices to obtain the partial
82
- utterances.
83
- """
84
- assert 0 <= overlap < 1
85
- assert 0 < min_pad_coverage <= 1
86
-
87
- samples_per_frame = int((sampling_rate * mel_window_step / 1000))
88
- n_frames = int(np.ceil((n_samples + 1) / samples_per_frame))
89
- frame_step = max(int(np.round(partial_utterance_n_frames * (1 - overlap))), 1)
90
-
91
- # Compute the slices
92
- wav_slices, mel_slices = [], []
93
- steps = max(1, n_frames - partial_utterance_n_frames + frame_step + 1)
94
- for i in range(0, steps, frame_step):
95
- mel_range = np.array([i, i + partial_utterance_n_frames])
96
- wav_range = mel_range * samples_per_frame
97
- mel_slices.append(slice(*mel_range))
98
- wav_slices.append(slice(*wav_range))
99
-
100
- # Evaluate whether extra padding is warranted or not
101
- last_wav_range = wav_slices[-1]
102
- coverage = (n_samples - last_wav_range.start) / (last_wav_range.stop - last_wav_range.start)
103
- if coverage < min_pad_coverage and len(mel_slices) > 1:
104
- mel_slices = mel_slices[:-1]
105
- wav_slices = wav_slices[:-1]
106
-
107
- return wav_slices, mel_slices
108
-
109
-
110
- def embed_utterance(wav, using_partials=True, return_partials=False, **kwargs):
111
- """
112
- Computes an embedding for a single utterance.
113
-
114
- # TODO: handle multiple wavs to benefit from batching on GPU
115
- :param wav: a preprocessed (see audio.py) utterance waveform as a numpy array of float32
116
- :param using_partials: if True, then the utterance is split in partial utterances of
117
- <partial_utterance_n_frames> frames and the utterance embedding is computed from their
118
- normalized average. If False, the utterance is instead computed from feeding the entire
119
- spectogram to the network.
120
- :param return_partials: if True, the partial embeddings will also be returned along with the
121
- wav slices that correspond to the partial embeddings.
122
- :param kwargs: additional arguments to compute_partial_splits()
123
- :return: the embedding as a numpy array of float32 of shape (model_embedding_size,). If
124
- <return_partials> is True, the partial utterances as a numpy array of float32 of shape
125
- (n_partials, model_embedding_size) and the wav partials as a list of slices will also be
126
- returned. If <using_partials> is simultaneously set to False, both these values will be None
127
- instead.
128
- """
129
- # Process the entire utterance if not using partials
130
- if not using_partials:
131
- frames = audio.wav_to_mel_spectrogram(wav)
132
- embed = embed_frames_batch(frames[None, ...])[0]
133
- if return_partials:
134
- return embed, None, None
135
- return embed
136
-
137
- # Compute where to split the utterance into partials and pad if necessary
138
- wave_slices, mel_slices = compute_partial_slices(len(wav), **kwargs)
139
- max_wave_length = wave_slices[-1].stop
140
- if max_wave_length >= len(wav):
141
- wav = np.pad(wav, (0, max_wave_length - len(wav)), "constant")
142
-
143
- # Split the utterance into partials
144
- frames = audio.wav_to_mel_spectrogram(wav)
145
- frames_batch = np.array([frames[s] for s in mel_slices])
146
- partial_embeds = embed_frames_batch(frames_batch)
147
-
148
- # Compute the utterance embedding from the partial embeddings
149
- raw_embed = np.mean(partial_embeds, axis=0)
150
- embed = raw_embed / np.linalg.norm(raw_embed, 2)
151
-
152
- if return_partials:
153
- return embed, partial_embeds, wave_slices
154
- return embed
155
-
156
-
157
- def embed_speaker(wavs, **kwargs):
158
- raise NotImplemented()
159
-
160
-
161
- def plot_embedding_as_heatmap(embed, ax=None, title="", shape=None, color_range=(0, 0.30)):
162
- import matplotlib.pyplot as plt
163
- if ax is None:
164
- ax = plt.gca()
165
-
166
- if shape is None:
167
- height = int(np.sqrt(len(embed)))
168
- shape = (height, -1)
169
- embed = embed.reshape(shape)
170
-
171
- cmap = cm.get_cmap()
172
- mappable = ax.imshow(embed, cmap=cmap)
173
- cbar = plt.colorbar(mappable, ax=ax, fraction=0.046, pad=0.04)
174
- sm = cm.ScalarMappable(cmap=cmap)
175
- sm.set_clim(*color_range)
176
-
177
- ax.set_xticks([]), ax.set_yticks([])
178
- ax.set_title(title)
 
1
+ from encoder.params_data import *
2
+ from encoder.model import SpeakerEncoder
3
+ from encoder.audio import preprocess_wav # We want to expose this function from here
4
+ from matplotlib import cm
5
+ from encoder import audio
6
+ from pathlib import Path
7
+ import numpy as np
8
+ import torch
9
+
10
+ _model = None # type: SpeakerEncoder
11
+ _device = None # type: torch.device
12
+
13
+
14
+ def load_model(weights_fpath: Path, device=None):
15
+ """
16
+ Loads the model in memory. If this function is not explicitely called, it will be run on the
17
+ first call to embed_frames() with the default weights file.
18
+
19
+ :param weights_fpath: the path to saved model weights.
20
+ :param device: either a torch device or the name of a torch device (e.g. "cpu", "cuda"). The
21
+ model will be loaded and will run on this device. Outputs will however always be on the cpu.
22
+ If None, will default to your GPU if it"s available, otherwise your CPU.
23
+ """
24
+ # TODO: I think the slow loading of the encoder might have something to do with the device it
25
+ # was saved on. Worth investigating.
26
+ global _model, _device
27
+ if device is None:
28
+ _device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
29
+ elif isinstance(device, str):
30
+ _device = torch.device(device)
31
+ _model = SpeakerEncoder(_device, torch.device("cpu"))
32
+ checkpoint = torch.load(weights_fpath, map_location=_device, weights_only=False)
33
+ _model.load_state_dict(checkpoint["model_state"])
34
+ _model.eval()
35
+ print("Loaded encoder \"%s\" trained to step %d" % (weights_fpath.name, checkpoint["step"]))
36
+
37
+
38
+ def is_loaded():
39
+ return _model is not None
40
+
41
+
42
+ def embed_frames_batch(frames_batch):
43
+ """
44
+ Computes embeddings for a batch of mel spectrogram.
45
+
46
+ :param frames_batch: a batch mel of spectrogram as a numpy array of float32 of shape
47
+ (batch_size, n_frames, n_channels)
48
+ :return: the embeddings as a numpy array of float32 of shape (batch_size, model_embedding_size)
49
+ """
50
+ if _model is None:
51
+ raise Exception("Model was not loaded. Call load_model() before inference.")
52
+
53
+ frames = torch.from_numpy(frames_batch).to(_device)
54
+ embed = _model.forward(frames).detach().cpu().numpy()
55
+ return embed
56
+
57
+
58
+ def compute_partial_slices(n_samples, partial_utterance_n_frames=partials_n_frames,
59
+ min_pad_coverage=0.75, overlap=0.5):
60
+ """
61
+ Computes where to split an utterance waveform and its corresponding mel spectrogram to obtain
62
+ partial utterances of <partial_utterance_n_frames> each. Both the waveform and the mel
63
+ spectrogram slices are returned, so as to make each partial utterance waveform correspond to
64
+ its spectrogram. This function assumes that the mel spectrogram parameters used are those
65
+ defined in params_data.py.
66
+
67
+ The returned ranges may be indexing further than the length of the waveform. It is
68
+ recommended that you pad the waveform with zeros up to wave_slices[-1].stop.
69
+
70
+ :param n_samples: the number of samples in the waveform
71
+ :param partial_utterance_n_frames: the number of mel spectrogram frames in each partial
72
+ utterance
73
+ :param min_pad_coverage: when reaching the last partial utterance, it may or may not have
74
+ enough frames. If at least <min_pad_coverage> of <partial_utterance_n_frames> are present,
75
+ then the last partial utterance will be considered, as if we padded the audio. Otherwise,
76
+ it will be discarded, as if we trimmed the audio. If there aren't enough frames for 1 partial
77
+ utterance, this parameter is ignored so that the function always returns at least 1 slice.
78
+ :param overlap: by how much the partial utterance should overlap. If set to 0, the partial
79
+ utterances are entirely disjoint.
80
+ :return: the waveform slices and mel spectrogram slices as lists of array slices. Index
81
+ respectively the waveform and the mel spectrogram with these slices to obtain the partial
82
+ utterances.
83
+ """
84
+ assert 0 <= overlap < 1
85
+ assert 0 < min_pad_coverage <= 1
86
+
87
+ samples_per_frame = int((sampling_rate * mel_window_step / 1000))
88
+ n_frames = int(np.ceil((n_samples + 1) / samples_per_frame))
89
+ frame_step = max(int(np.round(partial_utterance_n_frames * (1 - overlap))), 1)
90
+
91
+ # Compute the slices
92
+ wav_slices, mel_slices = [], []
93
+ steps = max(1, n_frames - partial_utterance_n_frames + frame_step + 1)
94
+ for i in range(0, steps, frame_step):
95
+ mel_range = np.array([i, i + partial_utterance_n_frames])
96
+ wav_range = mel_range * samples_per_frame
97
+ mel_slices.append(slice(*mel_range))
98
+ wav_slices.append(slice(*wav_range))
99
+
100
+ # Evaluate whether extra padding is warranted or not
101
+ last_wav_range = wav_slices[-1]
102
+ coverage = (n_samples - last_wav_range.start) / (last_wav_range.stop - last_wav_range.start)
103
+ if coverage < min_pad_coverage and len(mel_slices) > 1:
104
+ mel_slices = mel_slices[:-1]
105
+ wav_slices = wav_slices[:-1]
106
+
107
+ return wav_slices, mel_slices
108
+
109
+
110
+ def embed_utterance(wav, using_partials=True, return_partials=False, **kwargs):
111
+ """
112
+ Computes an embedding for a single utterance.
113
+
114
+ # TODO: handle multiple wavs to benefit from batching on GPU
115
+ :param wav: a preprocessed (see audio.py) utterance waveform as a numpy array of float32
116
+ :param using_partials: if True, then the utterance is split in partial utterances of
117
+ <partial_utterance_n_frames> frames and the utterance embedding is computed from their
118
+ normalized average. If False, the utterance is instead computed from feeding the entire
119
+ spectogram to the network.
120
+ :param return_partials: if True, the partial embeddings will also be returned along with the
121
+ wav slices that correspond to the partial embeddings.
122
+ :param kwargs: additional arguments to compute_partial_splits()
123
+ :return: the embedding as a numpy array of float32 of shape (model_embedding_size,). If
124
+ <return_partials> is True, the partial utterances as a numpy array of float32 of shape
125
+ (n_partials, model_embedding_size) and the wav partials as a list of slices will also be
126
+ returned. If <using_partials> is simultaneously set to False, both these values will be None
127
+ instead.
128
+ """
129
+ # Process the entire utterance if not using partials
130
+ if not using_partials:
131
+ frames = audio.wav_to_mel_spectrogram(wav)
132
+ embed = embed_frames_batch(frames[None, ...])[0]
133
+ if return_partials:
134
+ return embed, None, None
135
+ return embed
136
+
137
+ # Compute where to split the utterance into partials and pad if necessary
138
+ wave_slices, mel_slices = compute_partial_slices(len(wav), **kwargs)
139
+ max_wave_length = wave_slices[-1].stop
140
+ if max_wave_length >= len(wav):
141
+ wav = np.pad(wav, (0, max_wave_length - len(wav)), "constant")
142
+
143
+ # Split the utterance into partials
144
+ frames = audio.wav_to_mel_spectrogram(wav)
145
+ frames_batch = np.array([frames[s] for s in mel_slices])
146
+ partial_embeds = embed_frames_batch(frames_batch)
147
+
148
+ # Compute the utterance embedding from the partial embeddings
149
+ raw_embed = np.mean(partial_embeds, axis=0)
150
+ embed = raw_embed / np.linalg.norm(raw_embed, 2)
151
+
152
+ if return_partials:
153
+ return embed, partial_embeds, wave_slices
154
+ return embed
155
+
156
+
157
+ def embed_speaker(wavs, **kwargs):
158
+ raise NotImplemented()
159
+
160
+
161
+ def plot_embedding_as_heatmap(embed, ax=None, title="", shape=None, color_range=(0, 0.30)):
162
+ import matplotlib.pyplot as plt
163
+ if ax is None:
164
+ ax = plt.gca()
165
+
166
+ if shape is None:
167
+ height = int(np.sqrt(len(embed)))
168
+ shape = (height, -1)
169
+ embed = embed.reshape(shape)
170
+
171
+ cmap = cm.get_cmap()
172
+ mappable = ax.imshow(embed, cmap=cmap)
173
+ cbar = plt.colorbar(mappable, ax=ax, fraction=0.046, pad=0.04)
174
+ sm = cm.ScalarMappable(cmap=cmap)
175
+ sm.set_clim(*color_range)
176
+
177
+ ax.set_xticks([]), ax.set_yticks([])
178
+ ax.set_title(title)
encoder/model.py CHANGED
@@ -1,135 +1,135 @@
1
- from encoder.params_model import *
2
- from encoder.params_data import *
3
- from scipy.interpolate import interp1d
4
- from sklearn.metrics import roc_curve
5
- from torch.nn.utils import clip_grad_norm_
6
- from scipy.optimize import brentq
7
- from torch import nn
8
- import numpy as np
9
- import torch
10
-
11
-
12
- class SpeakerEncoder(nn.Module):
13
- def __init__(self, device, loss_device):
14
- super().__init__()
15
- self.loss_device = loss_device
16
-
17
- # Network defition
18
- self.lstm = nn.LSTM(input_size=mel_n_channels,
19
- hidden_size=model_hidden_size,
20
- num_layers=model_num_layers,
21
- batch_first=True).to(device)
22
- self.linear = nn.Linear(in_features=model_hidden_size,
23
- out_features=model_embedding_size).to(device)
24
- self.relu = torch.nn.ReLU().to(device)
25
-
26
- # Cosine similarity scaling (with fixed initial parameter values)
27
- self.similarity_weight = nn.Parameter(torch.tensor([10.])).to(loss_device)
28
- self.similarity_bias = nn.Parameter(torch.tensor([-5.])).to(loss_device)
29
-
30
- # Loss
31
- self.loss_fn = nn.CrossEntropyLoss().to(loss_device)
32
-
33
- def do_gradient_ops(self):
34
- # Gradient scale
35
- self.similarity_weight.grad *= 0.01
36
- self.similarity_bias.grad *= 0.01
37
-
38
- # Gradient clipping
39
- clip_grad_norm_(self.parameters(), 3, norm_type=2)
40
-
41
- def forward(self, utterances, hidden_init=None):
42
- """
43
- Computes the embeddings of a batch of utterance spectrograms.
44
-
45
- :param utterances: batch of mel-scale filterbanks of same duration as a tensor of shape
46
- (batch_size, n_frames, n_channels)
47
- :param hidden_init: initial hidden state of the LSTM as a tensor of shape (num_layers,
48
- batch_size, hidden_size). Will default to a tensor of zeros if None.
49
- :return: the embeddings as a tensor of shape (batch_size, embedding_size)
50
- """
51
- # Pass the input through the LSTM layers and retrieve all outputs, the final hidden state
52
- # and the final cell state.
53
- out, (hidden, cell) = self.lstm(utterances, hidden_init)
54
-
55
- # We take only the hidden state of the last layer
56
- embeds_raw = self.relu(self.linear(hidden[-1]))
57
-
58
- # L2-normalize it
59
- embeds = embeds_raw / (torch.norm(embeds_raw, dim=1, keepdim=True) + 1e-5)
60
-
61
- return embeds
62
-
63
- def similarity_matrix(self, embeds):
64
- """
65
- Computes the similarity matrix according the section 2.1 of GE2E.
66
-
67
- :param embeds: the embeddings as a tensor of shape (speakers_per_batch,
68
- utterances_per_speaker, embedding_size)
69
- :return: the similarity matrix as a tensor of shape (speakers_per_batch,
70
- utterances_per_speaker, speakers_per_batch)
71
- """
72
- speakers_per_batch, utterances_per_speaker = embeds.shape[:2]
73
-
74
- # Inclusive centroids (1 per speaker). Cloning is needed for reverse differentiation
75
- centroids_incl = torch.mean(embeds, dim=1, keepdim=True)
76
- centroids_incl = centroids_incl.clone() / (torch.norm(centroids_incl, dim=2, keepdim=True) + 1e-5)
77
-
78
- # Exclusive centroids (1 per utterance)
79
- centroids_excl = (torch.sum(embeds, dim=1, keepdim=True) - embeds)
80
- centroids_excl /= (utterances_per_speaker - 1)
81
- centroids_excl = centroids_excl.clone() / (torch.norm(centroids_excl, dim=2, keepdim=True) + 1e-5)
82
-
83
- # Similarity matrix. The cosine similarity of already 2-normed vectors is simply the dot
84
- # product of these vectors (which is just an element-wise multiplication reduced by a sum).
85
- # We vectorize the computation for efficiency.
86
- sim_matrix = torch.zeros(speakers_per_batch, utterances_per_speaker,
87
- speakers_per_batch).to(self.loss_device)
88
- mask_matrix = 1 - np.eye(speakers_per_batch, dtype=np.int)
89
- for j in range(speakers_per_batch):
90
- mask = np.where(mask_matrix[j])[0]
91
- sim_matrix[mask, :, j] = (embeds[mask] * centroids_incl[j]).sum(dim=2)
92
- sim_matrix[j, :, j] = (embeds[j] * centroids_excl[j]).sum(dim=1)
93
-
94
- ## Even more vectorized version (slower maybe because of transpose)
95
- # sim_matrix2 = torch.zeros(speakers_per_batch, speakers_per_batch, utterances_per_speaker
96
- # ).to(self.loss_device)
97
- # eye = np.eye(speakers_per_batch, dtype=np.int)
98
- # mask = np.where(1 - eye)
99
- # sim_matrix2[mask] = (embeds[mask[0]] * centroids_incl[mask[1]]).sum(dim=2)
100
- # mask = np.where(eye)
101
- # sim_matrix2[mask] = (embeds * centroids_excl).sum(dim=2)
102
- # sim_matrix2 = sim_matrix2.transpose(1, 2)
103
-
104
- sim_matrix = sim_matrix * self.similarity_weight + self.similarity_bias
105
- return sim_matrix
106
-
107
- def loss(self, embeds):
108
- """
109
- Computes the softmax loss according the section 2.1 of GE2E.
110
-
111
- :param embeds: the embeddings as a tensor of shape (speakers_per_batch,
112
- utterances_per_speaker, embedding_size)
113
- :return: the loss and the EER for this batch of embeddings.
114
- """
115
- speakers_per_batch, utterances_per_speaker = embeds.shape[:2]
116
-
117
- # Loss
118
- sim_matrix = self.similarity_matrix(embeds)
119
- sim_matrix = sim_matrix.reshape((speakers_per_batch * utterances_per_speaker,
120
- speakers_per_batch))
121
- ground_truth = np.repeat(np.arange(speakers_per_batch), utterances_per_speaker)
122
- target = torch.from_numpy(ground_truth).long().to(self.loss_device)
123
- loss = self.loss_fn(sim_matrix, target)
124
-
125
- # EER (not backpropagated)
126
- with torch.no_grad():
127
- inv_argmax = lambda i: np.eye(1, speakers_per_batch, i, dtype=np.int)[0]
128
- labels = np.array([inv_argmax(i) for i in ground_truth])
129
- preds = sim_matrix.detach().cpu().numpy()
130
-
131
- # Snippet from https://yangcha.github.io/EER-ROC/
132
- fpr, tpr, thresholds = roc_curve(labels.flatten(), preds.flatten())
133
- eer = brentq(lambda x: 1. - x - interp1d(fpr, tpr)(x), 0., 1.)
134
-
135
- return loss, eer
 
1
+ from encoder.params_model import *
2
+ from encoder.params_data import *
3
+ from scipy.interpolate import interp1d
4
+ from sklearn.metrics import roc_curve
5
+ from torch.nn.utils import clip_grad_norm_
6
+ from scipy.optimize import brentq
7
+ from torch import nn
8
+ import numpy as np
9
+ import torch
10
+
11
+
12
+ class SpeakerEncoder(nn.Module):
13
+ def __init__(self, device, loss_device):
14
+ super().__init__()
15
+ self.loss_device = loss_device
16
+
17
+ # Network defition
18
+ self.lstm = nn.LSTM(input_size=mel_n_channels,
19
+ hidden_size=model_hidden_size,
20
+ num_layers=model_num_layers,
21
+ batch_first=True).to(device)
22
+ self.linear = nn.Linear(in_features=model_hidden_size,
23
+ out_features=model_embedding_size).to(device)
24
+ self.relu = torch.nn.ReLU().to(device)
25
+
26
+ # Cosine similarity scaling (with fixed initial parameter values)
27
+ self.similarity_weight = nn.Parameter(torch.tensor([10.])).to(loss_device)
28
+ self.similarity_bias = nn.Parameter(torch.tensor([-5.])).to(loss_device)
29
+
30
+ # Loss
31
+ self.loss_fn = nn.CrossEntropyLoss().to(loss_device)
32
+
33
+ def do_gradient_ops(self):
34
+ # Gradient scale
35
+ self.similarity_weight.grad *= 0.01
36
+ self.similarity_bias.grad *= 0.01
37
+
38
+ # Gradient clipping
39
+ clip_grad_norm_(self.parameters(), 3, norm_type=2)
40
+
41
+ def forward(self, utterances, hidden_init=None):
42
+ """
43
+ Computes the embeddings of a batch of utterance spectrograms.
44
+
45
+ :param utterances: batch of mel-scale filterbanks of same duration as a tensor of shape
46
+ (batch_size, n_frames, n_channels)
47
+ :param hidden_init: initial hidden state of the LSTM as a tensor of shape (num_layers,
48
+ batch_size, hidden_size). Will default to a tensor of zeros if None.
49
+ :return: the embeddings as a tensor of shape (batch_size, embedding_size)
50
+ """
51
+ # Pass the input through the LSTM layers and retrieve all outputs, the final hidden state
52
+ # and the final cell state.
53
+ out, (hidden, cell) = self.lstm(utterances, hidden_init)
54
+
55
+ # We take only the hidden state of the last layer
56
+ embeds_raw = self.relu(self.linear(hidden[-1]))
57
+
58
+ # L2-normalize it
59
+ embeds = embeds_raw / (torch.norm(embeds_raw, dim=1, keepdim=True) + 1e-5)
60
+
61
+ return embeds
62
+
63
+ def similarity_matrix(self, embeds):
64
+ """
65
+ Computes the similarity matrix according the section 2.1 of GE2E.
66
+
67
+ :param embeds: the embeddings as a tensor of shape (speakers_per_batch,
68
+ utterances_per_speaker, embedding_size)
69
+ :return: the similarity matrix as a tensor of shape (speakers_per_batch,
70
+ utterances_per_speaker, speakers_per_batch)
71
+ """
72
+ speakers_per_batch, utterances_per_speaker = embeds.shape[:2]
73
+
74
+ # Inclusive centroids (1 per speaker). Cloning is needed for reverse differentiation
75
+ centroids_incl = torch.mean(embeds, dim=1, keepdim=True)
76
+ centroids_incl = centroids_incl.clone() / (torch.norm(centroids_incl, dim=2, keepdim=True) + 1e-5)
77
+
78
+ # Exclusive centroids (1 per utterance)
79
+ centroids_excl = (torch.sum(embeds, dim=1, keepdim=True) - embeds)
80
+ centroids_excl /= (utterances_per_speaker - 1)
81
+ centroids_excl = centroids_excl.clone() / (torch.norm(centroids_excl, dim=2, keepdim=True) + 1e-5)
82
+
83
+ # Similarity matrix. The cosine similarity of already 2-normed vectors is simply the dot
84
+ # product of these vectors (which is just an element-wise multiplication reduced by a sum).
85
+ # We vectorize the computation for efficiency.
86
+ sim_matrix = torch.zeros(speakers_per_batch, utterances_per_speaker,
87
+ speakers_per_batch).to(self.loss_device)
88
+ mask_matrix = 1 - np.eye(speakers_per_batch, dtype=int)
89
+ for j in range(speakers_per_batch):
90
+ mask = np.where(mask_matrix[j])[0]
91
+ sim_matrix[mask, :, j] = (embeds[mask] * centroids_incl[j]).sum(dim=2)
92
+ sim_matrix[j, :, j] = (embeds[j] * centroids_excl[j]).sum(dim=1)
93
+
94
+ ## Even more vectorized version (slower maybe because of transpose)
95
+ # sim_matrix2 = torch.zeros(speakers_per_batch, speakers_per_batch, utterances_per_speaker
96
+ # ).to(self.loss_device)
97
+ # eye = np.eye(speakers_per_batch, dtype=int)
98
+ # mask = np.where(1 - eye)
99
+ # sim_matrix2[mask] = (embeds[mask[0]] * centroids_incl[mask[1]]).sum(dim=2)
100
+ # mask = np.where(eye)
101
+ # sim_matrix2[mask] = (embeds * centroids_excl).sum(dim=2)
102
+ # sim_matrix2 = sim_matrix2.transpose(1, 2)
103
+
104
+ sim_matrix = sim_matrix * self.similarity_weight + self.similarity_bias
105
+ return sim_matrix
106
+
107
+ def loss(self, embeds):
108
+ """
109
+ Computes the softmax loss according the section 2.1 of GE2E.
110
+
111
+ :param embeds: the embeddings as a tensor of shape (speakers_per_batch,
112
+ utterances_per_speaker, embedding_size)
113
+ :return: the loss and the EER for this batch of embeddings.
114
+ """
115
+ speakers_per_batch, utterances_per_speaker = embeds.shape[:2]
116
+
117
+ # Loss
118
+ sim_matrix = self.similarity_matrix(embeds)
119
+ sim_matrix = sim_matrix.reshape((speakers_per_batch * utterances_per_speaker,
120
+ speakers_per_batch))
121
+ ground_truth = np.repeat(np.arange(speakers_per_batch), utterances_per_speaker)
122
+ target = torch.from_numpy(ground_truth).long().to(self.loss_device)
123
+ loss = self.loss_fn(sim_matrix, target)
124
+
125
+ # EER (not backpropagated)
126
+ with torch.no_grad():
127
+ inv_argmax = lambda i: np.eye(1, speakers_per_batch, i, dtype=int)[0]
128
+ labels = np.array([inv_argmax(i) for i in ground_truth])
129
+ preds = sim_matrix.detach().cpu().numpy()
130
+
131
+ # Snippet from https://yangcha.github.io/EER-ROC/
132
+ fpr, tpr, thresholds = roc_curve(labels.flatten(), preds.flatten())
133
+ eer = brentq(lambda x: 1. - x - interp1d(fpr, tpr)(x), 0., 1.)
134
+
135
+ return loss, eer
packages.txt ADDED
@@ -0,0 +1 @@
 
 
1
+ ffmpeg
requirements.txt CHANGED
@@ -1,14 +1,12 @@
1
- #torch==1.12.1+cu116 --find-links https://download.pytorch.org/whl/torch/
2
- https://download.pytorch.org/whl/cu116/torch-1.12.1%2Bcu116-cp38-cp38-linux_x86_64.whl#sha256=dda312901220895087cc83d3665464a3dc171d04460c61c31af463efbfb54896
3
- gradio==3.0.4
4
- PyQt5==5.15.6
5
- sounddevice==0.4.3
6
- SoundFile==0.10.3.post1
7
- tqdm==4.62.3
8
- Unidecode==1.3.2
9
- webrtcvad==2.0.10
10
- librosa==0.8.1
11
- inflect==5.3.0
12
- umap-learn==0.5.2
13
- numpy<1.24
14
- IPython
 
1
+ --extra-index-url https://download.pytorch.org/whl/cpu
2
+ torch
3
+ numpy
4
+ scipy
5
+ scikit-learn
6
+ matplotlib
7
+ librosa>=0.10
8
+ soundfile
9
+ webrtcvad-wheels
10
+ inflect
11
+ Unidecode
12
+ tqdm
 
 
synthesizer/audio.py CHANGED
@@ -1,206 +1,206 @@
1
- import librosa
2
- import librosa.filters
3
- import numpy as np
4
- from scipy import signal
5
- from scipy.io import wavfile
6
- import soundfile as sf
7
-
8
-
9
- def load_wav(path, sr):
10
- return librosa.core.load(path, sr=sr)[0]
11
-
12
- def save_wav(wav, path, sr):
13
- wav *= 32767 / max(0.01, np.max(np.abs(wav)))
14
- #proposed by @dsmiller
15
- wavfile.write(path, sr, wav.astype(np.int16))
16
-
17
- def save_wavenet_wav(wav, path, sr):
18
- sf.write(path, wav.astype(np.float32), sr)
19
-
20
- def preemphasis(wav, k, preemphasize=True):
21
- if preemphasize:
22
- return signal.lfilter([1, -k], [1], wav)
23
- return wav
24
-
25
- def inv_preemphasis(wav, k, inv_preemphasize=True):
26
- if inv_preemphasize:
27
- return signal.lfilter([1], [1, -k], wav)
28
- return wav
29
-
30
- #From https://github.com/r9y9/wavenet_vocoder/blob/master/audio.py
31
- def start_and_end_indices(quantized, silence_threshold=2):
32
- for start in range(quantized.size):
33
- if abs(quantized[start] - 127) > silence_threshold:
34
- break
35
- for end in range(quantized.size - 1, 1, -1):
36
- if abs(quantized[end] - 127) > silence_threshold:
37
- break
38
-
39
- assert abs(quantized[start] - 127) > silence_threshold
40
- assert abs(quantized[end] - 127) > silence_threshold
41
-
42
- return start, end
43
-
44
- def get_hop_size(hparams):
45
- hop_size = hparams.hop_size
46
- if hop_size is None:
47
- assert hparams.frame_shift_ms is not None
48
- hop_size = int(hparams.frame_shift_ms / 1000 * hparams.sample_rate)
49
- return hop_size
50
-
51
- def linearspectrogram(wav, hparams):
52
- D = _stft(preemphasis(wav, hparams.preemphasis, hparams.preemphasize), hparams)
53
- S = _amp_to_db(np.abs(D), hparams) - hparams.ref_level_db
54
-
55
- if hparams.signal_normalization:
56
- return _normalize(S, hparams)
57
- return S
58
-
59
- def melspectrogram(wav, hparams):
60
- D = _stft(preemphasis(wav, hparams.preemphasis, hparams.preemphasize), hparams)
61
- S = _amp_to_db(_linear_to_mel(np.abs(D), hparams), hparams) - hparams.ref_level_db
62
-
63
- if hparams.signal_normalization:
64
- return _normalize(S, hparams)
65
- return S
66
-
67
- def inv_linear_spectrogram(linear_spectrogram, hparams):
68
- """Converts linear spectrogram to waveform using librosa"""
69
- if hparams.signal_normalization:
70
- D = _denormalize(linear_spectrogram, hparams)
71
- else:
72
- D = linear_spectrogram
73
-
74
- S = _db_to_amp(D + hparams.ref_level_db) #Convert back to linear
75
-
76
- if hparams.use_lws:
77
- processor = _lws_processor(hparams)
78
- D = processor.run_lws(S.astype(np.float64).T ** hparams.power)
79
- y = processor.istft(D).astype(np.float32)
80
- return inv_preemphasis(y, hparams.preemphasis, hparams.preemphasize)
81
- else:
82
- return inv_preemphasis(_griffin_lim(S ** hparams.power, hparams), hparams.preemphasis, hparams.preemphasize)
83
-
84
- def inv_mel_spectrogram(mel_spectrogram, hparams):
85
- """Converts mel spectrogram to waveform using librosa"""
86
- if hparams.signal_normalization:
87
- D = _denormalize(mel_spectrogram, hparams)
88
- else:
89
- D = mel_spectrogram
90
-
91
- S = _mel_to_linear(_db_to_amp(D + hparams.ref_level_db), hparams) # Convert back to linear
92
-
93
- if hparams.use_lws:
94
- processor = _lws_processor(hparams)
95
- D = processor.run_lws(S.astype(np.float64).T ** hparams.power)
96
- y = processor.istft(D).astype(np.float32)
97
- return inv_preemphasis(y, hparams.preemphasis, hparams.preemphasize)
98
- else:
99
- return inv_preemphasis(_griffin_lim(S ** hparams.power, hparams), hparams.preemphasis, hparams.preemphasize)
100
-
101
- def _lws_processor(hparams):
102
- import lws
103
- return lws.lws(hparams.n_fft, get_hop_size(hparams), fftsize=hparams.win_size, mode="speech")
104
-
105
- def _griffin_lim(S, hparams):
106
- """librosa implementation of Griffin-Lim
107
- Based on https://github.com/librosa/librosa/issues/434
108
- """
109
- angles = np.exp(2j * np.pi * np.random.rand(*S.shape))
110
- S_complex = np.abs(S).astype(np.complex)
111
- y = _istft(S_complex * angles, hparams)
112
- for i in range(hparams.griffin_lim_iters):
113
- angles = np.exp(1j * np.angle(_stft(y, hparams)))
114
- y = _istft(S_complex * angles, hparams)
115
- return y
116
-
117
- def _stft(y, hparams):
118
- if hparams.use_lws:
119
- return _lws_processor(hparams).stft(y).T
120
- else:
121
- return librosa.stft(y=y, n_fft=hparams.n_fft, hop_length=get_hop_size(hparams), win_length=hparams.win_size)
122
-
123
- def _istft(y, hparams):
124
- return librosa.istft(y, hop_length=get_hop_size(hparams), win_length=hparams.win_size)
125
-
126
- ##########################################################
127
- #Those are only correct when using lws!!! (This was messing with Wavenet quality for a long time!)
128
- def num_frames(length, fsize, fshift):
129
- """Compute number of time frames of spectrogram
130
- """
131
- pad = (fsize - fshift)
132
- if length % fshift == 0:
133
- M = (length + pad * 2 - fsize) // fshift + 1
134
- else:
135
- M = (length + pad * 2 - fsize) // fshift + 2
136
- return M
137
-
138
-
139
- def pad_lr(x, fsize, fshift):
140
- """Compute left and right padding
141
- """
142
- M = num_frames(len(x), fsize, fshift)
143
- pad = (fsize - fshift)
144
- T = len(x) + 2 * pad
145
- r = (M - 1) * fshift + fsize - T
146
- return pad, pad + r
147
- ##########################################################
148
- #Librosa correct padding
149
- def librosa_pad_lr(x, fsize, fshift):
150
- return 0, (x.shape[0] // fshift + 1) * fshift - x.shape[0]
151
-
152
- # Conversions
153
- _mel_basis = None
154
- _inv_mel_basis = None
155
-
156
- def _linear_to_mel(spectogram, hparams):
157
- global _mel_basis
158
- if _mel_basis is None:
159
- _mel_basis = _build_mel_basis(hparams)
160
- return np.dot(_mel_basis, spectogram)
161
-
162
- def _mel_to_linear(mel_spectrogram, hparams):
163
- global _inv_mel_basis
164
- if _inv_mel_basis is None:
165
- _inv_mel_basis = np.linalg.pinv(_build_mel_basis(hparams))
166
- return np.maximum(1e-10, np.dot(_inv_mel_basis, mel_spectrogram))
167
-
168
- def _build_mel_basis(hparams):
169
- assert hparams.fmax <= hparams.sample_rate // 2
170
- return librosa.filters.mel(hparams.sample_rate, hparams.n_fft, n_mels=hparams.num_mels,
171
- fmin=hparams.fmin, fmax=hparams.fmax)
172
-
173
- def _amp_to_db(x, hparams):
174
- min_level = np.exp(hparams.min_level_db / 20 * np.log(10))
175
- return 20 * np.log10(np.maximum(min_level, x))
176
-
177
- def _db_to_amp(x):
178
- return np.power(10.0, (x) * 0.05)
179
-
180
- def _normalize(S, hparams):
181
- if hparams.allow_clipping_in_normalization:
182
- if hparams.symmetric_mels:
183
- return np.clip((2 * hparams.max_abs_value) * ((S - hparams.min_level_db) / (-hparams.min_level_db)) - hparams.max_abs_value,
184
- -hparams.max_abs_value, hparams.max_abs_value)
185
- else:
186
- return np.clip(hparams.max_abs_value * ((S - hparams.min_level_db) / (-hparams.min_level_db)), 0, hparams.max_abs_value)
187
-
188
- assert S.max() <= 0 and S.min() - hparams.min_level_db >= 0
189
- if hparams.symmetric_mels:
190
- return (2 * hparams.max_abs_value) * ((S - hparams.min_level_db) / (-hparams.min_level_db)) - hparams.max_abs_value
191
- else:
192
- return hparams.max_abs_value * ((S - hparams.min_level_db) / (-hparams.min_level_db))
193
-
194
- def _denormalize(D, hparams):
195
- if hparams.allow_clipping_in_normalization:
196
- if hparams.symmetric_mels:
197
- return (((np.clip(D, -hparams.max_abs_value,
198
- hparams.max_abs_value) + hparams.max_abs_value) * -hparams.min_level_db / (2 * hparams.max_abs_value))
199
- + hparams.min_level_db)
200
- else:
201
- return ((np.clip(D, 0, hparams.max_abs_value) * -hparams.min_level_db / hparams.max_abs_value) + hparams.min_level_db)
202
-
203
- if hparams.symmetric_mels:
204
- return (((D + hparams.max_abs_value) * -hparams.min_level_db / (2 * hparams.max_abs_value)) + hparams.min_level_db)
205
- else:
206
- return ((D * -hparams.min_level_db / hparams.max_abs_value) + hparams.min_level_db)
 
1
+ import librosa
2
+ import librosa.filters
3
+ import numpy as np
4
+ from scipy import signal
5
+ from scipy.io import wavfile
6
+ import soundfile as sf
7
+
8
+
9
+ def load_wav(path, sr):
10
+ return librosa.core.load(path, sr=sr)[0]
11
+
12
+ def save_wav(wav, path, sr):
13
+ wav *= 32767 / max(0.01, np.max(np.abs(wav)))
14
+ #proposed by @dsmiller
15
+ wavfile.write(path, sr, wav.astype(np.int16))
16
+
17
+ def save_wavenet_wav(wav, path, sr):
18
+ sf.write(path, wav.astype(np.float32), sr)
19
+
20
+ def preemphasis(wav, k, preemphasize=True):
21
+ if preemphasize:
22
+ return signal.lfilter([1, -k], [1], wav)
23
+ return wav
24
+
25
+ def inv_preemphasis(wav, k, inv_preemphasize=True):
26
+ if inv_preemphasize:
27
+ return signal.lfilter([1], [1, -k], wav)
28
+ return wav
29
+
30
+ #From https://github.com/r9y9/wavenet_vocoder/blob/master/audio.py
31
+ def start_and_end_indices(quantized, silence_threshold=2):
32
+ for start in range(quantized.size):
33
+ if abs(quantized[start] - 127) > silence_threshold:
34
+ break
35
+ for end in range(quantized.size - 1, 1, -1):
36
+ if abs(quantized[end] - 127) > silence_threshold:
37
+ break
38
+
39
+ assert abs(quantized[start] - 127) > silence_threshold
40
+ assert abs(quantized[end] - 127) > silence_threshold
41
+
42
+ return start, end
43
+
44
+ def get_hop_size(hparams):
45
+ hop_size = hparams.hop_size
46
+ if hop_size is None:
47
+ assert hparams.frame_shift_ms is not None
48
+ hop_size = int(hparams.frame_shift_ms / 1000 * hparams.sample_rate)
49
+ return hop_size
50
+
51
+ def linearspectrogram(wav, hparams):
52
+ D = _stft(preemphasis(wav, hparams.preemphasis, hparams.preemphasize), hparams)
53
+ S = _amp_to_db(np.abs(D), hparams) - hparams.ref_level_db
54
+
55
+ if hparams.signal_normalization:
56
+ return _normalize(S, hparams)
57
+ return S
58
+
59
+ def melspectrogram(wav, hparams):
60
+ D = _stft(preemphasis(wav, hparams.preemphasis, hparams.preemphasize), hparams)
61
+ S = _amp_to_db(_linear_to_mel(np.abs(D), hparams), hparams) - hparams.ref_level_db
62
+
63
+ if hparams.signal_normalization:
64
+ return _normalize(S, hparams)
65
+ return S
66
+
67
+ def inv_linear_spectrogram(linear_spectrogram, hparams):
68
+ """Converts linear spectrogram to waveform using librosa"""
69
+ if hparams.signal_normalization:
70
+ D = _denormalize(linear_spectrogram, hparams)
71
+ else:
72
+ D = linear_spectrogram
73
+
74
+ S = _db_to_amp(D + hparams.ref_level_db) #Convert back to linear
75
+
76
+ if hparams.use_lws:
77
+ processor = _lws_processor(hparams)
78
+ D = processor.run_lws(S.astype(np.float64).T ** hparams.power)
79
+ y = processor.istft(D).astype(np.float32)
80
+ return inv_preemphasis(y, hparams.preemphasis, hparams.preemphasize)
81
+ else:
82
+ return inv_preemphasis(_griffin_lim(S ** hparams.power, hparams), hparams.preemphasis, hparams.preemphasize)
83
+
84
+ def inv_mel_spectrogram(mel_spectrogram, hparams):
85
+ """Converts mel spectrogram to waveform using librosa"""
86
+ if hparams.signal_normalization:
87
+ D = _denormalize(mel_spectrogram, hparams)
88
+ else:
89
+ D = mel_spectrogram
90
+
91
+ S = _mel_to_linear(_db_to_amp(D + hparams.ref_level_db), hparams) # Convert back to linear
92
+
93
+ if hparams.use_lws:
94
+ processor = _lws_processor(hparams)
95
+ D = processor.run_lws(S.astype(np.float64).T ** hparams.power)
96
+ y = processor.istft(D).astype(np.float32)
97
+ return inv_preemphasis(y, hparams.preemphasis, hparams.preemphasize)
98
+ else:
99
+ return inv_preemphasis(_griffin_lim(S ** hparams.power, hparams), hparams.preemphasis, hparams.preemphasize)
100
+
101
+ def _lws_processor(hparams):
102
+ import lws
103
+ return lws.lws(hparams.n_fft, get_hop_size(hparams), fftsize=hparams.win_size, mode="speech")
104
+
105
+ def _griffin_lim(S, hparams):
106
+ """librosa implementation of Griffin-Lim
107
+ Based on https://github.com/librosa/librosa/issues/434
108
+ """
109
+ angles = np.exp(2j * np.pi * np.random.rand(*S.shape))
110
+ S_complex = np.abs(S).astype(complex)
111
+ y = _istft(S_complex * angles, hparams)
112
+ for i in range(hparams.griffin_lim_iters):
113
+ angles = np.exp(1j * np.angle(_stft(y, hparams)))
114
+ y = _istft(S_complex * angles, hparams)
115
+ return y
116
+
117
+ def _stft(y, hparams):
118
+ if hparams.use_lws:
119
+ return _lws_processor(hparams).stft(y).T
120
+ else:
121
+ return librosa.stft(y=y, n_fft=hparams.n_fft, hop_length=get_hop_size(hparams), win_length=hparams.win_size)
122
+
123
+ def _istft(y, hparams):
124
+ return librosa.istft(y, hop_length=get_hop_size(hparams), win_length=hparams.win_size)
125
+
126
+ ##########################################################
127
+ #Those are only correct when using lws!!! (This was messing with Wavenet quality for a long time!)
128
+ def num_frames(length, fsize, fshift):
129
+ """Compute number of time frames of spectrogram
130
+ """
131
+ pad = (fsize - fshift)
132
+ if length % fshift == 0:
133
+ M = (length + pad * 2 - fsize) // fshift + 1
134
+ else:
135
+ M = (length + pad * 2 - fsize) // fshift + 2
136
+ return M
137
+
138
+
139
+ def pad_lr(x, fsize, fshift):
140
+ """Compute left and right padding
141
+ """
142
+ M = num_frames(len(x), fsize, fshift)
143
+ pad = (fsize - fshift)
144
+ T = len(x) + 2 * pad
145
+ r = (M - 1) * fshift + fsize - T
146
+ return pad, pad + r
147
+ ##########################################################
148
+ #Librosa correct padding
149
+ def librosa_pad_lr(x, fsize, fshift):
150
+ return 0, (x.shape[0] // fshift + 1) * fshift - x.shape[0]
151
+
152
+ # Conversions
153
+ _mel_basis = None
154
+ _inv_mel_basis = None
155
+
156
+ def _linear_to_mel(spectogram, hparams):
157
+ global _mel_basis
158
+ if _mel_basis is None:
159
+ _mel_basis = _build_mel_basis(hparams)
160
+ return np.dot(_mel_basis, spectogram)
161
+
162
+ def _mel_to_linear(mel_spectrogram, hparams):
163
+ global _inv_mel_basis
164
+ if _inv_mel_basis is None:
165
+ _inv_mel_basis = np.linalg.pinv(_build_mel_basis(hparams))
166
+ return np.maximum(1e-10, np.dot(_inv_mel_basis, mel_spectrogram))
167
+
168
+ def _build_mel_basis(hparams):
169
+ assert hparams.fmax <= hparams.sample_rate // 2
170
+ return librosa.filters.mel(sr=hparams.sample_rate, n_fft=hparams.n_fft, n_mels=hparams.num_mels,
171
+ fmin=hparams.fmin, fmax=hparams.fmax)
172
+
173
+ def _amp_to_db(x, hparams):
174
+ min_level = np.exp(hparams.min_level_db / 20 * np.log(10))
175
+ return 20 * np.log10(np.maximum(min_level, x))
176
+
177
+ def _db_to_amp(x):
178
+ return np.power(10.0, (x) * 0.05)
179
+
180
+ def _normalize(S, hparams):
181
+ if hparams.allow_clipping_in_normalization:
182
+ if hparams.symmetric_mels:
183
+ return np.clip((2 * hparams.max_abs_value) * ((S - hparams.min_level_db) / (-hparams.min_level_db)) - hparams.max_abs_value,
184
+ -hparams.max_abs_value, hparams.max_abs_value)
185
+ else:
186
+ return np.clip(hparams.max_abs_value * ((S - hparams.min_level_db) / (-hparams.min_level_db)), 0, hparams.max_abs_value)
187
+
188
+ assert S.max() <= 0 and S.min() - hparams.min_level_db >= 0
189
+ if hparams.symmetric_mels:
190
+ return (2 * hparams.max_abs_value) * ((S - hparams.min_level_db) / (-hparams.min_level_db)) - hparams.max_abs_value
191
+ else:
192
+ return hparams.max_abs_value * ((S - hparams.min_level_db) / (-hparams.min_level_db))
193
+
194
+ def _denormalize(D, hparams):
195
+ if hparams.allow_clipping_in_normalization:
196
+ if hparams.symmetric_mels:
197
+ return (((np.clip(D, -hparams.max_abs_value,
198
+ hparams.max_abs_value) + hparams.max_abs_value) * -hparams.min_level_db / (2 * hparams.max_abs_value))
199
+ + hparams.min_level_db)
200
+ else:
201
+ return ((np.clip(D, 0, hparams.max_abs_value) * -hparams.min_level_db / hparams.max_abs_value) + hparams.min_level_db)
202
+
203
+ if hparams.symmetric_mels:
204
+ return (((D + hparams.max_abs_value) * -hparams.min_level_db / (2 * hparams.max_abs_value)) + hparams.min_level_db)
205
+ else:
206
+ return ((D * -hparams.min_level_db / hparams.max_abs_value) + hparams.min_level_db)
synthesizer/inference.py CHANGED
@@ -134,7 +134,7 @@ class Synthesizer:
134
  train the synthesizer.
135
  """
136
  print("Loading fpath and hparams.sample_rate :",str(fpath), hparams.sample_rate)
137
- wav = librosa.load(str(fpath), hparams.sample_rate)[0]
138
  if hparams.rescale:
139
  wav = wav / np.abs(wav).max() * hparams.rescaling_max
140
  return wav
 
134
  train the synthesizer.
135
  """
136
  print("Loading fpath and hparams.sample_rate :",str(fpath), hparams.sample_rate)
137
+ wav = librosa.load(str(fpath), sr=hparams.sample_rate)[0]
138
  if hparams.rescale:
139
  wav = wav / np.abs(wav).max() * hparams.rescaling_max
140
  return wav
synthesizer/models/tacotron.py CHANGED
@@ -1,519 +1,519 @@
1
- import os
2
- import numpy as np
3
- import torch
4
- import torch.nn as nn
5
- import torch.nn.functional as F
6
- from pathlib import Path
7
- from typing import Union
8
-
9
-
10
- class HighwayNetwork(nn.Module):
11
- def __init__(self, size):
12
- super().__init__()
13
- self.W1 = nn.Linear(size, size)
14
- self.W2 = nn.Linear(size, size)
15
- self.W1.bias.data.fill_(0.)
16
-
17
- def forward(self, x):
18
- x1 = self.W1(x)
19
- x2 = self.W2(x)
20
- g = torch.sigmoid(x2)
21
- y = g * F.relu(x1) + (1. - g) * x
22
- return y
23
-
24
-
25
- class Encoder(nn.Module):
26
- def __init__(self, embed_dims, num_chars, encoder_dims, K, num_highways, dropout):
27
- super().__init__()
28
- prenet_dims = (encoder_dims, encoder_dims)
29
- cbhg_channels = encoder_dims
30
- self.embedding = nn.Embedding(num_chars, embed_dims)
31
- self.pre_net = PreNet(embed_dims, fc1_dims=prenet_dims[0], fc2_dims=prenet_dims[1],
32
- dropout=dropout)
33
- self.cbhg = CBHG(K=K, in_channels=cbhg_channels, channels=cbhg_channels,
34
- proj_channels=[cbhg_channels, cbhg_channels],
35
- num_highways=num_highways)
36
-
37
- def forward(self, x, speaker_embedding=None):
38
- x = self.embedding(x)
39
- x = self.pre_net(x)
40
- x.transpose_(1, 2)
41
- x = self.cbhg(x)
42
- if speaker_embedding is not None:
43
- x = self.add_speaker_embedding(x, speaker_embedding)
44
- return x
45
-
46
- def add_speaker_embedding(self, x, speaker_embedding):
47
- # SV2TTS
48
- # The input x is the encoder output and is a 3D tensor with size (batch_size, num_chars, tts_embed_dims)
49
- # When training, speaker_embedding is also a 2D tensor with size (batch_size, speaker_embedding_size)
50
- # (for inference, speaker_embedding is a 1D tensor with size (speaker_embedding_size))
51
- # This concats the speaker embedding for each char in the encoder output
52
-
53
- # Save the dimensions as human-readable names
54
- batch_size = x.size()[0]
55
- num_chars = x.size()[1]
56
-
57
- if speaker_embedding.dim() == 1:
58
- idx = 0
59
- else:
60
- idx = 1
61
-
62
- # Start by making a copy of each speaker embedding to match the input text length
63
- # The output of this has size (batch_size, num_chars * tts_embed_dims)
64
- speaker_embedding_size = speaker_embedding.size()[idx]
65
- e = speaker_embedding.repeat_interleave(num_chars, dim=idx)
66
-
67
- # Reshape it and transpose
68
- e = e.reshape(batch_size, speaker_embedding_size, num_chars)
69
- e = e.transpose(1, 2)
70
-
71
- # Concatenate the tiled speaker embedding with the encoder output
72
- x = torch.cat((x, e), 2)
73
- return x
74
-
75
-
76
- class BatchNormConv(nn.Module):
77
- def __init__(self, in_channels, out_channels, kernel, relu=True):
78
- super().__init__()
79
- self.conv = nn.Conv1d(in_channels, out_channels, kernel, stride=1, padding=kernel // 2, bias=False)
80
- self.bnorm = nn.BatchNorm1d(out_channels)
81
- self.relu = relu
82
-
83
- def forward(self, x):
84
- x = self.conv(x)
85
- x = F.relu(x) if self.relu is True else x
86
- return self.bnorm(x)
87
-
88
-
89
- class CBHG(nn.Module):
90
- def __init__(self, K, in_channels, channels, proj_channels, num_highways):
91
- super().__init__()
92
-
93
- # List of all rnns to call `flatten_parameters()` on
94
- self._to_flatten = []
95
-
96
- self.bank_kernels = [i for i in range(1, K + 1)]
97
- self.conv1d_bank = nn.ModuleList()
98
- for k in self.bank_kernels:
99
- conv = BatchNormConv(in_channels, channels, k)
100
- self.conv1d_bank.append(conv)
101
-
102
- self.maxpool = nn.MaxPool1d(kernel_size=2, stride=1, padding=1)
103
-
104
- self.conv_project1 = BatchNormConv(len(self.bank_kernels) * channels, proj_channels[0], 3)
105
- self.conv_project2 = BatchNormConv(proj_channels[0], proj_channels[1], 3, relu=False)
106
-
107
- # Fix the highway input if necessary
108
- if proj_channels[-1] != channels:
109
- self.highway_mismatch = True
110
- self.pre_highway = nn.Linear(proj_channels[-1], channels, bias=False)
111
- else:
112
- self.highway_mismatch = False
113
-
114
- self.highways = nn.ModuleList()
115
- for i in range(num_highways):
116
- hn = HighwayNetwork(channels)
117
- self.highways.append(hn)
118
-
119
- self.rnn = nn.GRU(channels, channels // 2, batch_first=True, bidirectional=True)
120
- self._to_flatten.append(self.rnn)
121
-
122
- # Avoid fragmentation of RNN parameters and associated warning
123
- self._flatten_parameters()
124
-
125
- def forward(self, x):
126
- # Although we `_flatten_parameters()` on init, when using DataParallel
127
- # the model gets replicated, making it no longer guaranteed that the
128
- # weights are contiguous in GPU memory. Hence, we must call it again
129
- self._flatten_parameters()
130
-
131
- # Save these for later
132
- residual = x
133
- seq_len = x.size(-1)
134
- conv_bank = []
135
-
136
- # Convolution Bank
137
- for conv in self.conv1d_bank:
138
- c = conv(x) # Convolution
139
- conv_bank.append(c[:, :, :seq_len])
140
-
141
- # Stack along the channel axis
142
- conv_bank = torch.cat(conv_bank, dim=1)
143
-
144
- # dump the last padding to fit residual
145
- x = self.maxpool(conv_bank)[:, :, :seq_len]
146
-
147
- # Conv1d projections
148
- x = self.conv_project1(x)
149
- x = self.conv_project2(x)
150
-
151
- # Residual Connect
152
- x = x + residual
153
-
154
- # Through the highways
155
- x = x.transpose(1, 2)
156
- if self.highway_mismatch is True:
157
- x = self.pre_highway(x)
158
- for h in self.highways: x = h(x)
159
-
160
- # And then the RNN
161
- x, _ = self.rnn(x)
162
- return x
163
-
164
- def _flatten_parameters(self):
165
- """Calls `flatten_parameters` on all the rnns used by the WaveRNN. Used
166
- to improve efficiency and avoid PyTorch yelling at us."""
167
- [m.flatten_parameters() for m in self._to_flatten]
168
-
169
- class PreNet(nn.Module):
170
- def __init__(self, in_dims, fc1_dims=256, fc2_dims=128, dropout=0.5):
171
- super().__init__()
172
- self.fc1 = nn.Linear(in_dims, fc1_dims)
173
- self.fc2 = nn.Linear(fc1_dims, fc2_dims)
174
- self.p = dropout
175
-
176
- def forward(self, x):
177
- x = self.fc1(x)
178
- x = F.relu(x)
179
- x = F.dropout(x, self.p, training=True)
180
- x = self.fc2(x)
181
- x = F.relu(x)
182
- x = F.dropout(x, self.p, training=True)
183
- return x
184
-
185
-
186
- class Attention(nn.Module):
187
- def __init__(self, attn_dims):
188
- super().__init__()
189
- self.W = nn.Linear(attn_dims, attn_dims, bias=False)
190
- self.v = nn.Linear(attn_dims, 1, bias=False)
191
-
192
- def forward(self, encoder_seq_proj, query, t):
193
-
194
- # print(encoder_seq_proj.shape)
195
- # Transform the query vector
196
- query_proj = self.W(query).unsqueeze(1)
197
-
198
- # Compute the scores
199
- u = self.v(torch.tanh(encoder_seq_proj + query_proj))
200
- scores = F.softmax(u, dim=1)
201
-
202
- return scores.transpose(1, 2)
203
-
204
-
205
- class LSA(nn.Module):
206
- def __init__(self, attn_dim, kernel_size=31, filters=32):
207
- super().__init__()
208
- self.conv = nn.Conv1d(1, filters, padding=(kernel_size - 1) // 2, kernel_size=kernel_size, bias=True)
209
- self.L = nn.Linear(filters, attn_dim, bias=False)
210
- self.W = nn.Linear(attn_dim, attn_dim, bias=True) # Include the attention bias in this term
211
- self.v = nn.Linear(attn_dim, 1, bias=False)
212
- self.cumulative = None
213
- self.attention = None
214
-
215
- def init_attention(self, encoder_seq_proj):
216
- device = next(self.parameters()).device # use same device as parameters
217
- b, t, c = encoder_seq_proj.size()
218
- self.cumulative = torch.zeros(b, t, device=device)
219
- self.attention = torch.zeros(b, t, device=device)
220
-
221
- def forward(self, encoder_seq_proj, query, t, chars):
222
-
223
- if t == 0: self.init_attention(encoder_seq_proj)
224
-
225
- processed_query = self.W(query).unsqueeze(1)
226
-
227
- location = self.cumulative.unsqueeze(1)
228
- processed_loc = self.L(self.conv(location).transpose(1, 2))
229
-
230
- u = self.v(torch.tanh(processed_query + encoder_seq_proj + processed_loc))
231
- u = u.squeeze(-1)
232
-
233
- # Mask zero padding chars
234
- u = u * (chars != 0).float()
235
-
236
- # Smooth Attention
237
- # scores = torch.sigmoid(u) / torch.sigmoid(u).sum(dim=1, keepdim=True)
238
- scores = F.softmax(u, dim=1)
239
- self.attention = scores
240
- self.cumulative = self.cumulative + self.attention
241
-
242
- return scores.unsqueeze(-1).transpose(1, 2)
243
-
244
-
245
- class Decoder(nn.Module):
246
- # Class variable because its value doesn't change between classes
247
- # yet ought to be scoped by class because its a property of a Decoder
248
- max_r = 20
249
- def __init__(self, n_mels, encoder_dims, decoder_dims, lstm_dims,
250
- dropout, speaker_embedding_size):
251
- super().__init__()
252
- self.register_buffer("r", torch.tensor(1, dtype=torch.int))
253
- self.n_mels = n_mels
254
- prenet_dims = (decoder_dims * 2, decoder_dims * 2)
255
- self.prenet = PreNet(n_mels, fc1_dims=prenet_dims[0], fc2_dims=prenet_dims[1],
256
- dropout=dropout)
257
- self.attn_net = LSA(decoder_dims)
258
- self.attn_rnn = nn.GRUCell(encoder_dims + prenet_dims[1] + speaker_embedding_size, decoder_dims)
259
- self.rnn_input = nn.Linear(encoder_dims + decoder_dims + speaker_embedding_size, lstm_dims)
260
- self.res_rnn1 = nn.LSTMCell(lstm_dims, lstm_dims)
261
- self.res_rnn2 = nn.LSTMCell(lstm_dims, lstm_dims)
262
- self.mel_proj = nn.Linear(lstm_dims, n_mels * self.max_r, bias=False)
263
- self.stop_proj = nn.Linear(encoder_dims + speaker_embedding_size + lstm_dims, 1)
264
-
265
- def zoneout(self, prev, current, p=0.1):
266
- device = next(self.parameters()).device # Use same device as parameters
267
- mask = torch.zeros(prev.size(), device=device).bernoulli_(p)
268
- return prev * mask + current * (1 - mask)
269
-
270
- def forward(self, encoder_seq, encoder_seq_proj, prenet_in,
271
- hidden_states, cell_states, context_vec, t, chars):
272
-
273
- # Need this for reshaping mels
274
- batch_size = encoder_seq.size(0)
275
-
276
- # Unpack the hidden and cell states
277
- attn_hidden, rnn1_hidden, rnn2_hidden = hidden_states
278
- rnn1_cell, rnn2_cell = cell_states
279
-
280
- # PreNet for the Attention RNN
281
- prenet_out = self.prenet(prenet_in)
282
-
283
- # Compute the Attention RNN hidden state
284
- attn_rnn_in = torch.cat([context_vec, prenet_out], dim=-1)
285
- attn_hidden = self.attn_rnn(attn_rnn_in.squeeze(1), attn_hidden)
286
-
287
- # Compute the attention scores
288
- scores = self.attn_net(encoder_seq_proj, attn_hidden, t, chars)
289
-
290
- # Dot product to create the context vector
291
- context_vec = scores @ encoder_seq
292
- context_vec = context_vec.squeeze(1)
293
-
294
- # Concat Attention RNN output w. Context Vector & project
295
- x = torch.cat([context_vec, attn_hidden], dim=1)
296
- x = self.rnn_input(x)
297
-
298
- # Compute first Residual RNN
299
- rnn1_hidden_next, rnn1_cell = self.res_rnn1(x, (rnn1_hidden, rnn1_cell))
300
- if self.training:
301
- rnn1_hidden = self.zoneout(rnn1_hidden, rnn1_hidden_next)
302
- else:
303
- rnn1_hidden = rnn1_hidden_next
304
- x = x + rnn1_hidden
305
-
306
- # Compute second Residual RNN
307
- rnn2_hidden_next, rnn2_cell = self.res_rnn2(x, (rnn2_hidden, rnn2_cell))
308
- if self.training:
309
- rnn2_hidden = self.zoneout(rnn2_hidden, rnn2_hidden_next)
310
- else:
311
- rnn2_hidden = rnn2_hidden_next
312
- x = x + rnn2_hidden
313
-
314
- # Project Mels
315
- mels = self.mel_proj(x)
316
- mels = mels.view(batch_size, self.n_mels, self.max_r)[:, :, :self.r]
317
- hidden_states = (attn_hidden, rnn1_hidden, rnn2_hidden)
318
- cell_states = (rnn1_cell, rnn2_cell)
319
-
320
- # Stop token prediction
321
- s = torch.cat((x, context_vec), dim=1)
322
- s = self.stop_proj(s)
323
- stop_tokens = torch.sigmoid(s)
324
-
325
- return mels, scores, hidden_states, cell_states, context_vec, stop_tokens
326
-
327
-
328
- class Tacotron(nn.Module):
329
- def __init__(self, embed_dims, num_chars, encoder_dims, decoder_dims, n_mels,
330
- fft_bins, postnet_dims, encoder_K, lstm_dims, postnet_K, num_highways,
331
- dropout, stop_threshold, speaker_embedding_size):
332
- super().__init__()
333
- self.n_mels = n_mels
334
- self.lstm_dims = lstm_dims
335
- self.encoder_dims = encoder_dims
336
- self.decoder_dims = decoder_dims
337
- self.speaker_embedding_size = speaker_embedding_size
338
- self.encoder = Encoder(embed_dims, num_chars, encoder_dims,
339
- encoder_K, num_highways, dropout)
340
- self.encoder_proj = nn.Linear(encoder_dims + speaker_embedding_size, decoder_dims, bias=False)
341
- self.decoder = Decoder(n_mels, encoder_dims, decoder_dims, lstm_dims,
342
- dropout, speaker_embedding_size)
343
- self.postnet = CBHG(postnet_K, n_mels, postnet_dims,
344
- [postnet_dims, fft_bins], num_highways)
345
- self.post_proj = nn.Linear(postnet_dims, fft_bins, bias=False)
346
-
347
- self.init_model()
348
- self.num_params()
349
-
350
- self.register_buffer("step", torch.zeros(1, dtype=torch.long))
351
- self.register_buffer("stop_threshold", torch.tensor(stop_threshold, dtype=torch.float32))
352
-
353
- @property
354
- def r(self):
355
- return self.decoder.r.item()
356
-
357
- @r.setter
358
- def r(self, value):
359
- self.decoder.r = self.decoder.r.new_tensor(value, requires_grad=False)
360
-
361
- def forward(self, x, m, speaker_embedding):
362
- device = next(self.parameters()).device # use same device as parameters
363
-
364
- self.step += 1
365
- batch_size, _, steps = m.size()
366
-
367
- # Initialise all hidden states and pack into tuple
368
- attn_hidden = torch.zeros(batch_size, self.decoder_dims, device=device)
369
- rnn1_hidden = torch.zeros(batch_size, self.lstm_dims, device=device)
370
- rnn2_hidden = torch.zeros(batch_size, self.lstm_dims, device=device)
371
- hidden_states = (attn_hidden, rnn1_hidden, rnn2_hidden)
372
-
373
- # Initialise all lstm cell states and pack into tuple
374
- rnn1_cell = torch.zeros(batch_size, self.lstm_dims, device=device)
375
- rnn2_cell = torch.zeros(batch_size, self.lstm_dims, device=device)
376
- cell_states = (rnn1_cell, rnn2_cell)
377
-
378
- # <GO> Frame for start of decoder loop
379
- go_frame = torch.zeros(batch_size, self.n_mels, device=device)
380
-
381
- # Need an initial context vector
382
- context_vec = torch.zeros(batch_size, self.encoder_dims + self.speaker_embedding_size, device=device)
383
-
384
- # SV2TTS: Run the encoder with the speaker embedding
385
- # The projection avoids unnecessary matmuls in the decoder loop
386
- encoder_seq = self.encoder(x, speaker_embedding)
387
- encoder_seq_proj = self.encoder_proj(encoder_seq)
388
-
389
- # Need a couple of lists for outputs
390
- mel_outputs, attn_scores, stop_outputs = [], [], []
391
-
392
- # Run the decoder loop
393
- for t in range(0, steps, self.r):
394
- prenet_in = m[:, :, t - 1] if t > 0 else go_frame
395
- mel_frames, scores, hidden_states, cell_states, context_vec, stop_tokens = \
396
- self.decoder(encoder_seq, encoder_seq_proj, prenet_in,
397
- hidden_states, cell_states, context_vec, t, x)
398
- mel_outputs.append(mel_frames)
399
- attn_scores.append(scores)
400
- stop_outputs.extend([stop_tokens] * self.r)
401
-
402
- # Concat the mel outputs into sequence
403
- mel_outputs = torch.cat(mel_outputs, dim=2)
404
-
405
- # Post-Process for Linear Spectrograms
406
- postnet_out = self.postnet(mel_outputs)
407
- linear = self.post_proj(postnet_out)
408
- linear = linear.transpose(1, 2)
409
-
410
- # For easy visualisation
411
- attn_scores = torch.cat(attn_scores, 1)
412
- # attn_scores = attn_scores.cpu().data.numpy()
413
- stop_outputs = torch.cat(stop_outputs, 1)
414
-
415
- return mel_outputs, linear, attn_scores, stop_outputs
416
-
417
- def generate(self, x, speaker_embedding=None, steps=2000):
418
- self.eval()
419
- device = next(self.parameters()).device # use same device as parameters
420
-
421
- batch_size, _ = x.size()
422
-
423
- # Need to initialise all hidden states and pack into tuple for tidyness
424
- attn_hidden = torch.zeros(batch_size, self.decoder_dims, device=device)
425
- rnn1_hidden = torch.zeros(batch_size, self.lstm_dims, device=device)
426
- rnn2_hidden = torch.zeros(batch_size, self.lstm_dims, device=device)
427
- hidden_states = (attn_hidden, rnn1_hidden, rnn2_hidden)
428
-
429
- # Need to initialise all lstm cell states and pack into tuple for tidyness
430
- rnn1_cell = torch.zeros(batch_size, self.lstm_dims, device=device)
431
- rnn2_cell = torch.zeros(batch_size, self.lstm_dims, device=device)
432
- cell_states = (rnn1_cell, rnn2_cell)
433
-
434
- # Need a <GO> Frame for start of decoder loop
435
- go_frame = torch.zeros(batch_size, self.n_mels, device=device)
436
-
437
- # Need an initial context vector
438
- context_vec = torch.zeros(batch_size, self.encoder_dims + self.speaker_embedding_size, device=device)
439
-
440
- # SV2TTS: Run the encoder with the speaker embedding
441
- # The projection avoids unnecessary matmuls in the decoder loop
442
- encoder_seq = self.encoder(x, speaker_embedding)
443
- encoder_seq_proj = self.encoder_proj(encoder_seq)
444
-
445
- # Need a couple of lists for outputs
446
- mel_outputs, attn_scores, stop_outputs = [], [], []
447
-
448
- # Run the decoder loop
449
- for t in range(0, steps, self.r):
450
- prenet_in = mel_outputs[-1][:, :, -1] if t > 0 else go_frame
451
- mel_frames, scores, hidden_states, cell_states, context_vec, stop_tokens = \
452
- self.decoder(encoder_seq, encoder_seq_proj, prenet_in,
453
- hidden_states, cell_states, context_vec, t, x)
454
- mel_outputs.append(mel_frames)
455
- attn_scores.append(scores)
456
- stop_outputs.extend([stop_tokens] * self.r)
457
- # Stop the loop when all stop tokens in batch exceed threshold
458
- if (stop_tokens > 0.5).all() and t > 10: break
459
-
460
- # Concat the mel outputs into sequence
461
- mel_outputs = torch.cat(mel_outputs, dim=2)
462
-
463
- # Post-Process for Linear Spectrograms
464
- postnet_out = self.postnet(mel_outputs)
465
- linear = self.post_proj(postnet_out)
466
-
467
-
468
- linear = linear.transpose(1, 2)
469
-
470
- # For easy visualisation
471
- attn_scores = torch.cat(attn_scores, 1)
472
- stop_outputs = torch.cat(stop_outputs, 1)
473
-
474
- self.train()
475
-
476
- return mel_outputs, linear, attn_scores
477
-
478
- def init_model(self):
479
- for p in self.parameters():
480
- if p.dim() > 1: nn.init.xavier_uniform_(p)
481
-
482
- def get_step(self):
483
- return self.step.data.item()
484
-
485
- def reset_step(self):
486
- # assignment to parameters or buffers is overloaded, updates internal dict entry
487
- self.step = self.step.data.new_tensor(1)
488
-
489
- def log(self, path, msg):
490
- with open(path, "a") as f:
491
- print(msg, file=f)
492
-
493
- def load(self, path, optimizer=None):
494
- # Use device of model params as location for loaded state
495
- device = next(self.parameters()).device
496
- checkpoint = torch.load(str(path), map_location=device)
497
- self.load_state_dict(checkpoint["model_state"])
498
-
499
- if "optimizer_state" in checkpoint and optimizer is not None:
500
- optimizer.load_state_dict(checkpoint["optimizer_state"])
501
-
502
- def save(self, path, optimizer=None):
503
- if optimizer is not None:
504
- torch.save({
505
- "model_state": self.state_dict(),
506
- "optimizer_state": optimizer.state_dict(),
507
- }, str(path))
508
- else:
509
- torch.save({
510
- "model_state": self.state_dict(),
511
- }, str(path))
512
-
513
-
514
- def num_params(self, print_out=True):
515
- parameters = filter(lambda p: p.requires_grad, self.parameters())
516
- parameters = sum([np.prod(p.size()) for p in parameters]) / 1_000_000
517
- if print_out:
518
- print("Trainable Parameters: %.3fM" % parameters)
519
- return parameters
 
1
+ import os
2
+ import numpy as np
3
+ import torch
4
+ import torch.nn as nn
5
+ import torch.nn.functional as F
6
+ from pathlib import Path
7
+ from typing import Union
8
+
9
+
10
+ class HighwayNetwork(nn.Module):
11
+ def __init__(self, size):
12
+ super().__init__()
13
+ self.W1 = nn.Linear(size, size)
14
+ self.W2 = nn.Linear(size, size)
15
+ self.W1.bias.data.fill_(0.)
16
+
17
+ def forward(self, x):
18
+ x1 = self.W1(x)
19
+ x2 = self.W2(x)
20
+ g = torch.sigmoid(x2)
21
+ y = g * F.relu(x1) + (1. - g) * x
22
+ return y
23
+
24
+
25
+ class Encoder(nn.Module):
26
+ def __init__(self, embed_dims, num_chars, encoder_dims, K, num_highways, dropout):
27
+ super().__init__()
28
+ prenet_dims = (encoder_dims, encoder_dims)
29
+ cbhg_channels = encoder_dims
30
+ self.embedding = nn.Embedding(num_chars, embed_dims)
31
+ self.pre_net = PreNet(embed_dims, fc1_dims=prenet_dims[0], fc2_dims=prenet_dims[1],
32
+ dropout=dropout)
33
+ self.cbhg = CBHG(K=K, in_channels=cbhg_channels, channels=cbhg_channels,
34
+ proj_channels=[cbhg_channels, cbhg_channels],
35
+ num_highways=num_highways)
36
+
37
+ def forward(self, x, speaker_embedding=None):
38
+ x = self.embedding(x)
39
+ x = self.pre_net(x)
40
+ x.transpose_(1, 2)
41
+ x = self.cbhg(x)
42
+ if speaker_embedding is not None:
43
+ x = self.add_speaker_embedding(x, speaker_embedding)
44
+ return x
45
+
46
+ def add_speaker_embedding(self, x, speaker_embedding):
47
+ # SV2TTS
48
+ # The input x is the encoder output and is a 3D tensor with size (batch_size, num_chars, tts_embed_dims)
49
+ # When training, speaker_embedding is also a 2D tensor with size (batch_size, speaker_embedding_size)
50
+ # (for inference, speaker_embedding is a 1D tensor with size (speaker_embedding_size))
51
+ # This concats the speaker embedding for each char in the encoder output
52
+
53
+ # Save the dimensions as human-readable names
54
+ batch_size = x.size()[0]
55
+ num_chars = x.size()[1]
56
+
57
+ if speaker_embedding.dim() == 1:
58
+ idx = 0
59
+ else:
60
+ idx = 1
61
+
62
+ # Start by making a copy of each speaker embedding to match the input text length
63
+ # The output of this has size (batch_size, num_chars * tts_embed_dims)
64
+ speaker_embedding_size = speaker_embedding.size()[idx]
65
+ e = speaker_embedding.repeat_interleave(num_chars, dim=idx)
66
+
67
+ # Reshape it and transpose
68
+ e = e.reshape(batch_size, speaker_embedding_size, num_chars)
69
+ e = e.transpose(1, 2)
70
+
71
+ # Concatenate the tiled speaker embedding with the encoder output
72
+ x = torch.cat((x, e), 2)
73
+ return x
74
+
75
+
76
+ class BatchNormConv(nn.Module):
77
+ def __init__(self, in_channels, out_channels, kernel, relu=True):
78
+ super().__init__()
79
+ self.conv = nn.Conv1d(in_channels, out_channels, kernel, stride=1, padding=kernel // 2, bias=False)
80
+ self.bnorm = nn.BatchNorm1d(out_channels)
81
+ self.relu = relu
82
+
83
+ def forward(self, x):
84
+ x = self.conv(x)
85
+ x = F.relu(x) if self.relu is True else x
86
+ return self.bnorm(x)
87
+
88
+
89
+ class CBHG(nn.Module):
90
+ def __init__(self, K, in_channels, channels, proj_channels, num_highways):
91
+ super().__init__()
92
+
93
+ # List of all rnns to call `flatten_parameters()` on
94
+ self._to_flatten = []
95
+
96
+ self.bank_kernels = [i for i in range(1, K + 1)]
97
+ self.conv1d_bank = nn.ModuleList()
98
+ for k in self.bank_kernels:
99
+ conv = BatchNormConv(in_channels, channels, k)
100
+ self.conv1d_bank.append(conv)
101
+
102
+ self.maxpool = nn.MaxPool1d(kernel_size=2, stride=1, padding=1)
103
+
104
+ self.conv_project1 = BatchNormConv(len(self.bank_kernels) * channels, proj_channels[0], 3)
105
+ self.conv_project2 = BatchNormConv(proj_channels[0], proj_channels[1], 3, relu=False)
106
+
107
+ # Fix the highway input if necessary
108
+ if proj_channels[-1] != channels:
109
+ self.highway_mismatch = True
110
+ self.pre_highway = nn.Linear(proj_channels[-1], channels, bias=False)
111
+ else:
112
+ self.highway_mismatch = False
113
+
114
+ self.highways = nn.ModuleList()
115
+ for i in range(num_highways):
116
+ hn = HighwayNetwork(channels)
117
+ self.highways.append(hn)
118
+
119
+ self.rnn = nn.GRU(channels, channels // 2, batch_first=True, bidirectional=True)
120
+ self._to_flatten.append(self.rnn)
121
+
122
+ # Avoid fragmentation of RNN parameters and associated warning
123
+ self._flatten_parameters()
124
+
125
+ def forward(self, x):
126
+ # Although we `_flatten_parameters()` on init, when using DataParallel
127
+ # the model gets replicated, making it no longer guaranteed that the
128
+ # weights are contiguous in GPU memory. Hence, we must call it again
129
+ self._flatten_parameters()
130
+
131
+ # Save these for later
132
+ residual = x
133
+ seq_len = x.size(-1)
134
+ conv_bank = []
135
+
136
+ # Convolution Bank
137
+ for conv in self.conv1d_bank:
138
+ c = conv(x) # Convolution
139
+ conv_bank.append(c[:, :, :seq_len])
140
+
141
+ # Stack along the channel axis
142
+ conv_bank = torch.cat(conv_bank, dim=1)
143
+
144
+ # dump the last padding to fit residual
145
+ x = self.maxpool(conv_bank)[:, :, :seq_len]
146
+
147
+ # Conv1d projections
148
+ x = self.conv_project1(x)
149
+ x = self.conv_project2(x)
150
+
151
+ # Residual Connect
152
+ x = x + residual
153
+
154
+ # Through the highways
155
+ x = x.transpose(1, 2)
156
+ if self.highway_mismatch is True:
157
+ x = self.pre_highway(x)
158
+ for h in self.highways: x = h(x)
159
+
160
+ # And then the RNN
161
+ x, _ = self.rnn(x)
162
+ return x
163
+
164
+ def _flatten_parameters(self):
165
+ """Calls `flatten_parameters` on all the rnns used by the WaveRNN. Used
166
+ to improve efficiency and avoid PyTorch yelling at us."""
167
+ [m.flatten_parameters() for m in self._to_flatten]
168
+
169
+ class PreNet(nn.Module):
170
+ def __init__(self, in_dims, fc1_dims=256, fc2_dims=128, dropout=0.5):
171
+ super().__init__()
172
+ self.fc1 = nn.Linear(in_dims, fc1_dims)
173
+ self.fc2 = nn.Linear(fc1_dims, fc2_dims)
174
+ self.p = dropout
175
+
176
+ def forward(self, x):
177
+ x = self.fc1(x)
178
+ x = F.relu(x)
179
+ x = F.dropout(x, self.p, training=True)
180
+ x = self.fc2(x)
181
+ x = F.relu(x)
182
+ x = F.dropout(x, self.p, training=True)
183
+ return x
184
+
185
+
186
+ class Attention(nn.Module):
187
+ def __init__(self, attn_dims):
188
+ super().__init__()
189
+ self.W = nn.Linear(attn_dims, attn_dims, bias=False)
190
+ self.v = nn.Linear(attn_dims, 1, bias=False)
191
+
192
+ def forward(self, encoder_seq_proj, query, t):
193
+
194
+ # print(encoder_seq_proj.shape)
195
+ # Transform the query vector
196
+ query_proj = self.W(query).unsqueeze(1)
197
+
198
+ # Compute the scores
199
+ u = self.v(torch.tanh(encoder_seq_proj + query_proj))
200
+ scores = F.softmax(u, dim=1)
201
+
202
+ return scores.transpose(1, 2)
203
+
204
+
205
+ class LSA(nn.Module):
206
+ def __init__(self, attn_dim, kernel_size=31, filters=32):
207
+ super().__init__()
208
+ self.conv = nn.Conv1d(1, filters, padding=(kernel_size - 1) // 2, kernel_size=kernel_size, bias=True)
209
+ self.L = nn.Linear(filters, attn_dim, bias=False)
210
+ self.W = nn.Linear(attn_dim, attn_dim, bias=True) # Include the attention bias in this term
211
+ self.v = nn.Linear(attn_dim, 1, bias=False)
212
+ self.cumulative = None
213
+ self.attention = None
214
+
215
+ def init_attention(self, encoder_seq_proj):
216
+ device = next(self.parameters()).device # use same device as parameters
217
+ b, t, c = encoder_seq_proj.size()
218
+ self.cumulative = torch.zeros(b, t, device=device)
219
+ self.attention = torch.zeros(b, t, device=device)
220
+
221
+ def forward(self, encoder_seq_proj, query, t, chars):
222
+
223
+ if t == 0: self.init_attention(encoder_seq_proj)
224
+
225
+ processed_query = self.W(query).unsqueeze(1)
226
+
227
+ location = self.cumulative.unsqueeze(1)
228
+ processed_loc = self.L(self.conv(location).transpose(1, 2))
229
+
230
+ u = self.v(torch.tanh(processed_query + encoder_seq_proj + processed_loc))
231
+ u = u.squeeze(-1)
232
+
233
+ # Mask zero padding chars
234
+ u = u * (chars != 0).float()
235
+
236
+ # Smooth Attention
237
+ # scores = torch.sigmoid(u) / torch.sigmoid(u).sum(dim=1, keepdim=True)
238
+ scores = F.softmax(u, dim=1)
239
+ self.attention = scores
240
+ self.cumulative = self.cumulative + self.attention
241
+
242
+ return scores.unsqueeze(-1).transpose(1, 2)
243
+
244
+
245
+ class Decoder(nn.Module):
246
+ # Class variable because its value doesn't change between classes
247
+ # yet ought to be scoped by class because its a property of a Decoder
248
+ max_r = 20
249
+ def __init__(self, n_mels, encoder_dims, decoder_dims, lstm_dims,
250
+ dropout, speaker_embedding_size):
251
+ super().__init__()
252
+ self.register_buffer("r", torch.tensor(1, dtype=torch.int))
253
+ self.n_mels = n_mels
254
+ prenet_dims = (decoder_dims * 2, decoder_dims * 2)
255
+ self.prenet = PreNet(n_mels, fc1_dims=prenet_dims[0], fc2_dims=prenet_dims[1],
256
+ dropout=dropout)
257
+ self.attn_net = LSA(decoder_dims)
258
+ self.attn_rnn = nn.GRUCell(encoder_dims + prenet_dims[1] + speaker_embedding_size, decoder_dims)
259
+ self.rnn_input = nn.Linear(encoder_dims + decoder_dims + speaker_embedding_size, lstm_dims)
260
+ self.res_rnn1 = nn.LSTMCell(lstm_dims, lstm_dims)
261
+ self.res_rnn2 = nn.LSTMCell(lstm_dims, lstm_dims)
262
+ self.mel_proj = nn.Linear(lstm_dims, n_mels * self.max_r, bias=False)
263
+ self.stop_proj = nn.Linear(encoder_dims + speaker_embedding_size + lstm_dims, 1)
264
+
265
+ def zoneout(self, prev, current, p=0.1):
266
+ device = next(self.parameters()).device # Use same device as parameters
267
+ mask = torch.zeros(prev.size(), device=device).bernoulli_(p)
268
+ return prev * mask + current * (1 - mask)
269
+
270
+ def forward(self, encoder_seq, encoder_seq_proj, prenet_in,
271
+ hidden_states, cell_states, context_vec, t, chars):
272
+
273
+ # Need this for reshaping mels
274
+ batch_size = encoder_seq.size(0)
275
+
276
+ # Unpack the hidden and cell states
277
+ attn_hidden, rnn1_hidden, rnn2_hidden = hidden_states
278
+ rnn1_cell, rnn2_cell = cell_states
279
+
280
+ # PreNet for the Attention RNN
281
+ prenet_out = self.prenet(prenet_in)
282
+
283
+ # Compute the Attention RNN hidden state
284
+ attn_rnn_in = torch.cat([context_vec, prenet_out], dim=-1)
285
+ attn_hidden = self.attn_rnn(attn_rnn_in.squeeze(1), attn_hidden)
286
+
287
+ # Compute the attention scores
288
+ scores = self.attn_net(encoder_seq_proj, attn_hidden, t, chars)
289
+
290
+ # Dot product to create the context vector
291
+ context_vec = scores @ encoder_seq
292
+ context_vec = context_vec.squeeze(1)
293
+
294
+ # Concat Attention RNN output w. Context Vector & project
295
+ x = torch.cat([context_vec, attn_hidden], dim=1)
296
+ x = self.rnn_input(x)
297
+
298
+ # Compute first Residual RNN
299
+ rnn1_hidden_next, rnn1_cell = self.res_rnn1(x, (rnn1_hidden, rnn1_cell))
300
+ if self.training:
301
+ rnn1_hidden = self.zoneout(rnn1_hidden, rnn1_hidden_next)
302
+ else:
303
+ rnn1_hidden = rnn1_hidden_next
304
+ x = x + rnn1_hidden
305
+
306
+ # Compute second Residual RNN
307
+ rnn2_hidden_next, rnn2_cell = self.res_rnn2(x, (rnn2_hidden, rnn2_cell))
308
+ if self.training:
309
+ rnn2_hidden = self.zoneout(rnn2_hidden, rnn2_hidden_next)
310
+ else:
311
+ rnn2_hidden = rnn2_hidden_next
312
+ x = x + rnn2_hidden
313
+
314
+ # Project Mels
315
+ mels = self.mel_proj(x)
316
+ mels = mels.view(batch_size, self.n_mels, self.max_r)[:, :, :self.r]
317
+ hidden_states = (attn_hidden, rnn1_hidden, rnn2_hidden)
318
+ cell_states = (rnn1_cell, rnn2_cell)
319
+
320
+ # Stop token prediction
321
+ s = torch.cat((x, context_vec), dim=1)
322
+ s = self.stop_proj(s)
323
+ stop_tokens = torch.sigmoid(s)
324
+
325
+ return mels, scores, hidden_states, cell_states, context_vec, stop_tokens
326
+
327
+
328
+ class Tacotron(nn.Module):
329
+ def __init__(self, embed_dims, num_chars, encoder_dims, decoder_dims, n_mels,
330
+ fft_bins, postnet_dims, encoder_K, lstm_dims, postnet_K, num_highways,
331
+ dropout, stop_threshold, speaker_embedding_size):
332
+ super().__init__()
333
+ self.n_mels = n_mels
334
+ self.lstm_dims = lstm_dims
335
+ self.encoder_dims = encoder_dims
336
+ self.decoder_dims = decoder_dims
337
+ self.speaker_embedding_size = speaker_embedding_size
338
+ self.encoder = Encoder(embed_dims, num_chars, encoder_dims,
339
+ encoder_K, num_highways, dropout)
340
+ self.encoder_proj = nn.Linear(encoder_dims + speaker_embedding_size, decoder_dims, bias=False)
341
+ self.decoder = Decoder(n_mels, encoder_dims, decoder_dims, lstm_dims,
342
+ dropout, speaker_embedding_size)
343
+ self.postnet = CBHG(postnet_K, n_mels, postnet_dims,
344
+ [postnet_dims, fft_bins], num_highways)
345
+ self.post_proj = nn.Linear(postnet_dims, fft_bins, bias=False)
346
+
347
+ self.init_model()
348
+ self.num_params()
349
+
350
+ self.register_buffer("step", torch.zeros(1, dtype=torch.long))
351
+ self.register_buffer("stop_threshold", torch.tensor(stop_threshold, dtype=torch.float32))
352
+
353
+ @property
354
+ def r(self):
355
+ return self.decoder.r.item()
356
+
357
+ @r.setter
358
+ def r(self, value):
359
+ self.decoder.r = self.decoder.r.new_tensor(value, requires_grad=False)
360
+
361
+ def forward(self, x, m, speaker_embedding):
362
+ device = next(self.parameters()).device # use same device as parameters
363
+
364
+ self.step += 1
365
+ batch_size, _, steps = m.size()
366
+
367
+ # Initialise all hidden states and pack into tuple
368
+ attn_hidden = torch.zeros(batch_size, self.decoder_dims, device=device)
369
+ rnn1_hidden = torch.zeros(batch_size, self.lstm_dims, device=device)
370
+ rnn2_hidden = torch.zeros(batch_size, self.lstm_dims, device=device)
371
+ hidden_states = (attn_hidden, rnn1_hidden, rnn2_hidden)
372
+
373
+ # Initialise all lstm cell states and pack into tuple
374
+ rnn1_cell = torch.zeros(batch_size, self.lstm_dims, device=device)
375
+ rnn2_cell = torch.zeros(batch_size, self.lstm_dims, device=device)
376
+ cell_states = (rnn1_cell, rnn2_cell)
377
+
378
+ # <GO> Frame for start of decoder loop
379
+ go_frame = torch.zeros(batch_size, self.n_mels, device=device)
380
+
381
+ # Need an initial context vector
382
+ context_vec = torch.zeros(batch_size, self.encoder_dims + self.speaker_embedding_size, device=device)
383
+
384
+ # SV2TTS: Run the encoder with the speaker embedding
385
+ # The projection avoids unnecessary matmuls in the decoder loop
386
+ encoder_seq = self.encoder(x, speaker_embedding)
387
+ encoder_seq_proj = self.encoder_proj(encoder_seq)
388
+
389
+ # Need a couple of lists for outputs
390
+ mel_outputs, attn_scores, stop_outputs = [], [], []
391
+
392
+ # Run the decoder loop
393
+ for t in range(0, steps, self.r):
394
+ prenet_in = m[:, :, t - 1] if t > 0 else go_frame
395
+ mel_frames, scores, hidden_states, cell_states, context_vec, stop_tokens = \
396
+ self.decoder(encoder_seq, encoder_seq_proj, prenet_in,
397
+ hidden_states, cell_states, context_vec, t, x)
398
+ mel_outputs.append(mel_frames)
399
+ attn_scores.append(scores)
400
+ stop_outputs.extend([stop_tokens] * self.r)
401
+
402
+ # Concat the mel outputs into sequence
403
+ mel_outputs = torch.cat(mel_outputs, dim=2)
404
+
405
+ # Post-Process for Linear Spectrograms
406
+ postnet_out = self.postnet(mel_outputs)
407
+ linear = self.post_proj(postnet_out)
408
+ linear = linear.transpose(1, 2)
409
+
410
+ # For easy visualisation
411
+ attn_scores = torch.cat(attn_scores, 1)
412
+ # attn_scores = attn_scores.cpu().data.numpy()
413
+ stop_outputs = torch.cat(stop_outputs, 1)
414
+
415
+ return mel_outputs, linear, attn_scores, stop_outputs
416
+
417
+ def generate(self, x, speaker_embedding=None, steps=2000):
418
+ self.eval()
419
+ device = next(self.parameters()).device # use same device as parameters
420
+
421
+ batch_size, _ = x.size()
422
+
423
+ # Need to initialise all hidden states and pack into tuple for tidyness
424
+ attn_hidden = torch.zeros(batch_size, self.decoder_dims, device=device)
425
+ rnn1_hidden = torch.zeros(batch_size, self.lstm_dims, device=device)
426
+ rnn2_hidden = torch.zeros(batch_size, self.lstm_dims, device=device)
427
+ hidden_states = (attn_hidden, rnn1_hidden, rnn2_hidden)
428
+
429
+ # Need to initialise all lstm cell states and pack into tuple for tidyness
430
+ rnn1_cell = torch.zeros(batch_size, self.lstm_dims, device=device)
431
+ rnn2_cell = torch.zeros(batch_size, self.lstm_dims, device=device)
432
+ cell_states = (rnn1_cell, rnn2_cell)
433
+
434
+ # Need a <GO> Frame for start of decoder loop
435
+ go_frame = torch.zeros(batch_size, self.n_mels, device=device)
436
+
437
+ # Need an initial context vector
438
+ context_vec = torch.zeros(batch_size, self.encoder_dims + self.speaker_embedding_size, device=device)
439
+
440
+ # SV2TTS: Run the encoder with the speaker embedding
441
+ # The projection avoids unnecessary matmuls in the decoder loop
442
+ encoder_seq = self.encoder(x, speaker_embedding)
443
+ encoder_seq_proj = self.encoder_proj(encoder_seq)
444
+
445
+ # Need a couple of lists for outputs
446
+ mel_outputs, attn_scores, stop_outputs = [], [], []
447
+
448
+ # Run the decoder loop
449
+ for t in range(0, steps, self.r):
450
+ prenet_in = mel_outputs[-1][:, :, -1] if t > 0 else go_frame
451
+ mel_frames, scores, hidden_states, cell_states, context_vec, stop_tokens = \
452
+ self.decoder(encoder_seq, encoder_seq_proj, prenet_in,
453
+ hidden_states, cell_states, context_vec, t, x)
454
+ mel_outputs.append(mel_frames)
455
+ attn_scores.append(scores)
456
+ stop_outputs.extend([stop_tokens] * self.r)
457
+ # Stop the loop when all stop tokens in batch exceed threshold
458
+ if (stop_tokens > 0.5).all() and t > 10: break
459
+
460
+ # Concat the mel outputs into sequence
461
+ mel_outputs = torch.cat(mel_outputs, dim=2)
462
+
463
+ # Post-Process for Linear Spectrograms
464
+ postnet_out = self.postnet(mel_outputs)
465
+ linear = self.post_proj(postnet_out)
466
+
467
+
468
+ linear = linear.transpose(1, 2)
469
+
470
+ # For easy visualisation
471
+ attn_scores = torch.cat(attn_scores, 1)
472
+ stop_outputs = torch.cat(stop_outputs, 1)
473
+
474
+ self.train()
475
+
476
+ return mel_outputs, linear, attn_scores
477
+
478
+ def init_model(self):
479
+ for p in self.parameters():
480
+ if p.dim() > 1: nn.init.xavier_uniform_(p)
481
+
482
+ def get_step(self):
483
+ return self.step.data.item()
484
+
485
+ def reset_step(self):
486
+ # assignment to parameters or buffers is overloaded, updates internal dict entry
487
+ self.step = self.step.data.new_tensor(1)
488
+
489
+ def log(self, path, msg):
490
+ with open(path, "a") as f:
491
+ print(msg, file=f)
492
+
493
+ def load(self, path, optimizer=None):
494
+ # Use device of model params as location for loaded state
495
+ device = next(self.parameters()).device
496
+ checkpoint = torch.load(str(path), map_location=device, weights_only=False)
497
+ self.load_state_dict(checkpoint["model_state"])
498
+
499
+ if "optimizer_state" in checkpoint and optimizer is not None:
500
+ optimizer.load_state_dict(checkpoint["optimizer_state"])
501
+
502
+ def save(self, path, optimizer=None):
503
+ if optimizer is not None:
504
+ torch.save({
505
+ "model_state": self.state_dict(),
506
+ "optimizer_state": optimizer.state_dict(),
507
+ }, str(path))
508
+ else:
509
+ torch.save({
510
+ "model_state": self.state_dict(),
511
+ }, str(path))
512
+
513
+
514
+ def num_params(self, print_out=True):
515
+ parameters = filter(lambda p: p.requires_grad, self.parameters())
516
+ parameters = sum([np.prod(p.size()) for p in parameters]) / 1_000_000
517
+ if print_out:
518
+ print("Trainable Parameters: %.3fM" % parameters)
519
+ return parameters
utils/default_models.py CHANGED
@@ -1,57 +1,25 @@
1
- import urllib.request
2
  from pathlib import Path
3
- from threading import Thread
4
- from urllib.error import HTTPError
5
-
6
- from tqdm import tqdm
7
 
 
8
 
 
 
 
9
  default_models = {
10
- "encoder": ("https://drive.google.com/uc?export=download&id=1q8mEGwCkFy23KZsinbuvdKAQLqNKbYf1", 17090379),
11
- #"synthesizer": ("https://drive.google.com/u/0/uc?id=1EqFMIbvxffxtjiVrtykroF6_mUh-5Z3s&export=download&confirm=t", 370554559),
12
- "synthesizer": ("https://sourceforge.net/projects/encoders/files/synthesizer.pt/download", 370554559),
13
- "vocoder": ("https://drive.google.com/uc?export=download&id=1cf2NO6FtI0jDuy8AV3Xgn6leO6dHjIgu", 53845290),
14
  }
15
 
16
 
17
- class DownloadProgressBar(tqdm):
18
- def update_to(self, b=1, bsize=1, tsize=None):
19
- if tsize is not None:
20
- self.total = tsize
21
- self.update(b * bsize - self.n)
22
-
23
-
24
- def download(url: str, target: Path, bar_pos=0):
25
- # Ensure the directory exists
26
- target.parent.mkdir(exist_ok=True, parents=True)
27
-
28
- desc = f"Downloading {target.name}"
29
- with DownloadProgressBar(unit="B", unit_scale=True, miniters=1, desc=desc, position=bar_pos, leave=False) as t:
30
- try:
31
- urllib.request.urlretrieve(url, filename=target, reporthook=t.update_to)
32
- except HTTPError:
33
- return
34
-
35
-
36
  def ensure_default_models(models_dir: Path):
37
- # Define download tasks
38
- jobs = []
39
- for model_name, (url, size) in default_models.items():
40
- target_path = models_dir / "default" / f"{model_name}.pt"
41
- if target_path.exists():
42
- if target_path.stat().st_size != size:
43
- print(f"File {target_path} is not of expected size, redownloading...")
44
- else:
45
- continue
46
-
47
- thread = Thread(target=download, args=(url, target_path, len(jobs)))
48
- thread.start()
49
- jobs.append((thread, target_path, size))
50
-
51
- # Run and join threads
52
- for thread, target_path, size in jobs:
53
- thread.join()
54
-
55
  assert target_path.exists() and target_path.stat().st_size == size, \
56
- f"Download for {target_path.name} failed. You may download models manually instead.\n" \
57
- f"https://drive.google.com/drive/folders/1fU6umc5uQAVR2udZdHX-lDgXYzTyqG_j"
 
 
1
  from pathlib import Path
 
 
 
 
2
 
3
+ from huggingface_hub import hf_hub_download
4
 
5
+ # The original Google Drive links are dead; the same pretrained SV2TTS checkpoints
6
+ # (identical byte sizes) are mirrored on the Hub.
7
+ MODEL_REPO = "CorentinJ/SV2TTS"
8
  default_models = {
9
+ "encoder": 17090379,
10
+ "synthesizer": 370554559,
11
+ "vocoder": 53845290,
 
12
  }
13
 
14
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
15
  def ensure_default_models(models_dir: Path):
16
+ target_dir = Path(models_dir) / "default"
17
+ target_dir.mkdir(exist_ok=True, parents=True)
18
+ for model_name, size in default_models.items():
19
+ target_path = target_dir / f"{model_name}.pt"
20
+ if target_path.exists() and target_path.stat().st_size == size:
21
+ continue
22
+ print(f"Downloading {model_name}.pt from {MODEL_REPO}...")
23
+ hf_hub_download(MODEL_REPO, f"{model_name}.pt", local_dir=str(target_dir))
 
 
 
 
 
 
 
 
 
 
24
  assert target_path.exists() and target_path.stat().st_size == size, \
25
+ f"Download for {target_path.name} failed."
 
vocoder/audio.py CHANGED
@@ -1,108 +1,108 @@
1
- import math
2
- import numpy as np
3
- import librosa
4
- import vocoder.hparams as hp
5
- from scipy.signal import lfilter
6
- import soundfile as sf
7
-
8
-
9
- def label_2_float(x, bits) :
10
- return 2 * x / (2**bits - 1.) - 1.
11
-
12
-
13
- def float_2_label(x, bits) :
14
- assert abs(x).max() <= 1.0
15
- x = (x + 1.) * (2**bits - 1) / 2
16
- return x.clip(0, 2**bits - 1)
17
-
18
-
19
- def load_wav(path) :
20
- return librosa.load(str(path), sr=hp.sample_rate)[0]
21
-
22
-
23
- def save_wav(x, path) :
24
- sf.write(path, x.astype(np.float32), hp.sample_rate)
25
-
26
-
27
- def split_signal(x) :
28
- unsigned = x + 2**15
29
- coarse = unsigned // 256
30
- fine = unsigned % 256
31
- return coarse, fine
32
-
33
-
34
- def combine_signal(coarse, fine) :
35
- return coarse * 256 + fine - 2**15
36
-
37
-
38
- def encode_16bits(x) :
39
- return np.clip(x * 2**15, -2**15, 2**15 - 1).astype(np.int16)
40
-
41
-
42
- mel_basis = None
43
-
44
-
45
- def linear_to_mel(spectrogram):
46
- global mel_basis
47
- if mel_basis is None:
48
- mel_basis = build_mel_basis()
49
- return np.dot(mel_basis, spectrogram)
50
-
51
-
52
- def build_mel_basis():
53
- return librosa.filters.mel(hp.sample_rate, hp.n_fft, n_mels=hp.num_mels, fmin=hp.fmin)
54
-
55
-
56
- def normalize(S):
57
- return np.clip((S - hp.min_level_db) / -hp.min_level_db, 0, 1)
58
-
59
-
60
- def denormalize(S):
61
- return (np.clip(S, 0, 1) * -hp.min_level_db) + hp.min_level_db
62
-
63
-
64
- def amp_to_db(x):
65
- return 20 * np.log10(np.maximum(1e-5, x))
66
-
67
-
68
- def db_to_amp(x):
69
- return np.power(10.0, x * 0.05)
70
-
71
-
72
- def spectrogram(y):
73
- D = stft(y)
74
- S = amp_to_db(np.abs(D)) - hp.ref_level_db
75
- return normalize(S)
76
-
77
-
78
- def melspectrogram(y):
79
- D = stft(y)
80
- S = amp_to_db(linear_to_mel(np.abs(D)))
81
- return normalize(S)
82
-
83
-
84
- def stft(y):
85
- return librosa.stft(y=y, n_fft=hp.n_fft, hop_length=hp.hop_length, win_length=hp.win_length)
86
-
87
-
88
- def pre_emphasis(x):
89
- return lfilter([1, -hp.preemphasis], [1], x)
90
-
91
-
92
- def de_emphasis(x):
93
- return lfilter([1], [1, -hp.preemphasis], x)
94
-
95
-
96
- def encode_mu_law(x, mu) :
97
- mu = mu - 1
98
- fx = np.sign(x) * np.log(1 + mu * np.abs(x)) / np.log(1 + mu)
99
- return np.floor((fx + 1) / 2 * mu + 0.5)
100
-
101
-
102
- def decode_mu_law(y, mu, from_labels=True) :
103
- if from_labels:
104
- y = label_2_float(y, math.log2(mu))
105
- mu = mu - 1
106
- x = np.sign(y) / mu * ((1 + mu) ** np.abs(y) - 1)
107
- return x
108
-
 
1
+ import math
2
+ import numpy as np
3
+ import librosa
4
+ import vocoder.hparams as hp
5
+ from scipy.signal import lfilter
6
+ import soundfile as sf
7
+
8
+
9
+ def label_2_float(x, bits) :
10
+ return 2 * x / (2**bits - 1.) - 1.
11
+
12
+
13
+ def float_2_label(x, bits) :
14
+ assert abs(x).max() <= 1.0
15
+ x = (x + 1.) * (2**bits - 1) / 2
16
+ return x.clip(0, 2**bits - 1)
17
+
18
+
19
+ def load_wav(path) :
20
+ return librosa.load(str(path), sr=hp.sample_rate)[0]
21
+
22
+
23
+ def save_wav(x, path) :
24
+ sf.write(path, x.astype(np.float32), hp.sample_rate)
25
+
26
+
27
+ def split_signal(x) :
28
+ unsigned = x + 2**15
29
+ coarse = unsigned // 256
30
+ fine = unsigned % 256
31
+ return coarse, fine
32
+
33
+
34
+ def combine_signal(coarse, fine) :
35
+ return coarse * 256 + fine - 2**15
36
+
37
+
38
+ def encode_16bits(x) :
39
+ return np.clip(x * 2**15, -2**15, 2**15 - 1).astype(np.int16)
40
+
41
+
42
+ mel_basis = None
43
+
44
+
45
+ def linear_to_mel(spectrogram):
46
+ global mel_basis
47
+ if mel_basis is None:
48
+ mel_basis = build_mel_basis()
49
+ return np.dot(mel_basis, spectrogram)
50
+
51
+
52
+ def build_mel_basis():
53
+ return librosa.filters.mel(sr=hp.sample_rate, n_fft=hp.n_fft, n_mels=hp.num_mels, fmin=hp.fmin)
54
+
55
+
56
+ def normalize(S):
57
+ return np.clip((S - hp.min_level_db) / -hp.min_level_db, 0, 1)
58
+
59
+
60
+ def denormalize(S):
61
+ return (np.clip(S, 0, 1) * -hp.min_level_db) + hp.min_level_db
62
+
63
+
64
+ def amp_to_db(x):
65
+ return 20 * np.log10(np.maximum(1e-5, x))
66
+
67
+
68
+ def db_to_amp(x):
69
+ return np.power(10.0, x * 0.05)
70
+
71
+
72
+ def spectrogram(y):
73
+ D = stft(y)
74
+ S = amp_to_db(np.abs(D)) - hp.ref_level_db
75
+ return normalize(S)
76
+
77
+
78
+ def melspectrogram(y):
79
+ D = stft(y)
80
+ S = amp_to_db(linear_to_mel(np.abs(D)))
81
+ return normalize(S)
82
+
83
+
84
+ def stft(y):
85
+ return librosa.stft(y=y, n_fft=hp.n_fft, hop_length=hp.hop_length, win_length=hp.win_length)
86
+
87
+
88
+ def pre_emphasis(x):
89
+ return lfilter([1, -hp.preemphasis], [1], x)
90
+
91
+
92
+ def de_emphasis(x):
93
+ return lfilter([1], [1, -hp.preemphasis], x)
94
+
95
+
96
+ def encode_mu_law(x, mu) :
97
+ mu = mu - 1
98
+ fx = np.sign(x) * np.log(1 + mu * np.abs(x)) / np.log(1 + mu)
99
+ return np.floor((fx + 1) / 2 * mu + 0.5)
100
+
101
+
102
+ def decode_mu_law(y, mu, from_labels=True) :
103
+ if from_labels:
104
+ y = label_2_float(y, math.log2(mu))
105
+ mu = mu - 1
106
+ x = np.sign(y) / mu * ((1 + mu) ** np.abs(y) - 1)
107
+ return x
108
+
vocoder/inference.py CHANGED
@@ -1,64 +1,64 @@
1
- from vocoder.models.fatchord_version import WaveRNN
2
- from vocoder import hparams as hp
3
- import torch
4
-
5
-
6
- _model = None # type: WaveRNN
7
-
8
- def load_model(weights_fpath, verbose=True):
9
- global _model, _device
10
-
11
- if verbose:
12
- print("Building Wave-RNN")
13
- _model = WaveRNN(
14
- rnn_dims=hp.voc_rnn_dims,
15
- fc_dims=hp.voc_fc_dims,
16
- bits=hp.bits,
17
- pad=hp.voc_pad,
18
- upsample_factors=hp.voc_upsample_factors,
19
- feat_dims=hp.num_mels,
20
- compute_dims=hp.voc_compute_dims,
21
- res_out_dims=hp.voc_res_out_dims,
22
- res_blocks=hp.voc_res_blocks,
23
- hop_length=hp.hop_length,
24
- sample_rate=hp.sample_rate,
25
- mode=hp.voc_mode
26
- )
27
-
28
- if torch.cuda.is_available():
29
- _model = _model.cuda()
30
- _device = torch.device('cuda')
31
- else:
32
- _device = torch.device('cpu')
33
-
34
- if verbose:
35
- print("Loading model weights at %s" % weights_fpath)
36
- checkpoint = torch.load(weights_fpath, _device)
37
- _model.load_state_dict(checkpoint['model_state'])
38
- _model.eval()
39
-
40
-
41
- def is_loaded():
42
- return _model is not None
43
-
44
-
45
- def infer_waveform(mel, normalize=True, batched=True, target=8000, overlap=800,
46
- progress_callback=None):
47
- """
48
- Infers the waveform of a mel spectrogram output by the synthesizer (the format must match
49
- that of the synthesizer!)
50
-
51
- :param normalize:
52
- :param batched:
53
- :param target:
54
- :param overlap:
55
- :return:
56
- """
57
- if _model is None:
58
- raise Exception("Please load Wave-RNN in memory before using it")
59
-
60
- if normalize:
61
- mel = mel / hp.mel_max_abs_value
62
- mel = torch.from_numpy(mel[None, ...])
63
- wav = _model.generate(mel, batched, target, overlap, hp.mu_law, progress_callback)
64
- return wav
 
1
+ from vocoder.models.fatchord_version import WaveRNN
2
+ from vocoder import hparams as hp
3
+ import torch
4
+
5
+
6
+ _model = None # type: WaveRNN
7
+
8
+ def load_model(weights_fpath, verbose=True):
9
+ global _model, _device
10
+
11
+ if verbose:
12
+ print("Building Wave-RNN")
13
+ _model = WaveRNN(
14
+ rnn_dims=hp.voc_rnn_dims,
15
+ fc_dims=hp.voc_fc_dims,
16
+ bits=hp.bits,
17
+ pad=hp.voc_pad,
18
+ upsample_factors=hp.voc_upsample_factors,
19
+ feat_dims=hp.num_mels,
20
+ compute_dims=hp.voc_compute_dims,
21
+ res_out_dims=hp.voc_res_out_dims,
22
+ res_blocks=hp.voc_res_blocks,
23
+ hop_length=hp.hop_length,
24
+ sample_rate=hp.sample_rate,
25
+ mode=hp.voc_mode
26
+ )
27
+
28
+ if torch.cuda.is_available():
29
+ _model = _model.cuda()
30
+ _device = torch.device('cuda')
31
+ else:
32
+ _device = torch.device('cpu')
33
+
34
+ if verbose:
35
+ print("Loading model weights at %s" % weights_fpath)
36
+ checkpoint = torch.load(weights_fpath, map_location=_device, weights_only=False)
37
+ _model.load_state_dict(checkpoint['model_state'])
38
+ _model.eval()
39
+
40
+
41
+ def is_loaded():
42
+ return _model is not None
43
+
44
+
45
+ def infer_waveform(mel, normalize=True, batched=True, target=8000, overlap=800,
46
+ progress_callback=None):
47
+ """
48
+ Infers the waveform of a mel spectrogram output by the synthesizer (the format must match
49
+ that of the synthesizer!)
50
+
51
+ :param normalize:
52
+ :param batched:
53
+ :param target:
54
+ :param overlap:
55
+ :return:
56
+ """
57
+ if _model is None:
58
+ raise Exception("Please load Wave-RNN in memory before using it")
59
+
60
+ if normalize:
61
+ mel = mel / hp.mel_max_abs_value
62
+ mel = torch.from_numpy(mel[None, ...])
63
+ wav = _model.generate(mel, batched, target, overlap, hp.mu_law, progress_callback)
64
+ return wav
vocoder/models/fatchord_version.py CHANGED
@@ -1,434 +1,434 @@
1
- import torch
2
- import torch.nn as nn
3
- import torch.nn.functional as F
4
- from vocoder.distribution import sample_from_discretized_mix_logistic
5
- from vocoder.display import *
6
- from vocoder.audio import *
7
-
8
-
9
- class ResBlock(nn.Module):
10
- def __init__(self, dims):
11
- super().__init__()
12
- self.conv1 = nn.Conv1d(dims, dims, kernel_size=1, bias=False)
13
- self.conv2 = nn.Conv1d(dims, dims, kernel_size=1, bias=False)
14
- self.batch_norm1 = nn.BatchNorm1d(dims)
15
- self.batch_norm2 = nn.BatchNorm1d(dims)
16
-
17
- def forward(self, x):
18
- residual = x
19
- x = self.conv1(x)
20
- x = self.batch_norm1(x)
21
- x = F.relu(x)
22
- x = self.conv2(x)
23
- x = self.batch_norm2(x)
24
- return x + residual
25
-
26
-
27
- class MelResNet(nn.Module):
28
- def __init__(self, res_blocks, in_dims, compute_dims, res_out_dims, pad):
29
- super().__init__()
30
- k_size = pad * 2 + 1
31
- self.conv_in = nn.Conv1d(in_dims, compute_dims, kernel_size=k_size, bias=False)
32
- self.batch_norm = nn.BatchNorm1d(compute_dims)
33
- self.layers = nn.ModuleList()
34
- for i in range(res_blocks):
35
- self.layers.append(ResBlock(compute_dims))
36
- self.conv_out = nn.Conv1d(compute_dims, res_out_dims, kernel_size=1)
37
-
38
- def forward(self, x):
39
- x = self.conv_in(x)
40
- x = self.batch_norm(x)
41
- x = F.relu(x)
42
- for f in self.layers: x = f(x)
43
- x = self.conv_out(x)
44
- return x
45
-
46
-
47
- class Stretch2d(nn.Module):
48
- def __init__(self, x_scale, y_scale):
49
- super().__init__()
50
- self.x_scale = x_scale
51
- self.y_scale = y_scale
52
-
53
- def forward(self, x):
54
- b, c, h, w = x.size()
55
- x = x.unsqueeze(-1).unsqueeze(3)
56
- x = x.repeat(1, 1, 1, self.y_scale, 1, self.x_scale)
57
- return x.view(b, c, h * self.y_scale, w * self.x_scale)
58
-
59
-
60
- class UpsampleNetwork(nn.Module):
61
- def __init__(self, feat_dims, upsample_scales, compute_dims,
62
- res_blocks, res_out_dims, pad):
63
- super().__init__()
64
- total_scale = np.cumproduct(upsample_scales)[-1]
65
- self.indent = pad * total_scale
66
- self.resnet = MelResNet(res_blocks, feat_dims, compute_dims, res_out_dims, pad)
67
- self.resnet_stretch = Stretch2d(total_scale, 1)
68
- self.up_layers = nn.ModuleList()
69
- for scale in upsample_scales:
70
- k_size = (1, scale * 2 + 1)
71
- padding = (0, scale)
72
- stretch = Stretch2d(scale, 1)
73
- conv = nn.Conv2d(1, 1, kernel_size=k_size, padding=padding, bias=False)
74
- conv.weight.data.fill_(1. / k_size[1])
75
- self.up_layers.append(stretch)
76
- self.up_layers.append(conv)
77
-
78
- def forward(self, m):
79
- aux = self.resnet(m).unsqueeze(1)
80
- aux = self.resnet_stretch(aux)
81
- aux = aux.squeeze(1)
82
- m = m.unsqueeze(1)
83
- for f in self.up_layers: m = f(m)
84
- m = m.squeeze(1)[:, :, self.indent:-self.indent]
85
- return m.transpose(1, 2), aux.transpose(1, 2)
86
-
87
-
88
- class WaveRNN(nn.Module):
89
- def __init__(self, rnn_dims, fc_dims, bits, pad, upsample_factors,
90
- feat_dims, compute_dims, res_out_dims, res_blocks,
91
- hop_length, sample_rate, mode='RAW'):
92
- super().__init__()
93
- self.mode = mode
94
- self.pad = pad
95
- if self.mode == 'RAW' :
96
- self.n_classes = 2 ** bits
97
- elif self.mode == 'MOL' :
98
- self.n_classes = 30
99
- else :
100
- RuntimeError("Unknown model mode value - ", self.mode)
101
-
102
- self.rnn_dims = rnn_dims
103
- self.aux_dims = res_out_dims // 4
104
- self.hop_length = hop_length
105
- self.sample_rate = sample_rate
106
-
107
- self.upsample = UpsampleNetwork(feat_dims, upsample_factors, compute_dims, res_blocks, res_out_dims, pad)
108
- self.I = nn.Linear(feat_dims + self.aux_dims + 1, rnn_dims)
109
- self.rnn1 = nn.GRU(rnn_dims, rnn_dims, batch_first=True)
110
- self.rnn2 = nn.GRU(rnn_dims + self.aux_dims, rnn_dims, batch_first=True)
111
- self.fc1 = nn.Linear(rnn_dims + self.aux_dims, fc_dims)
112
- self.fc2 = nn.Linear(fc_dims + self.aux_dims, fc_dims)
113
- self.fc3 = nn.Linear(fc_dims, self.n_classes)
114
-
115
- self.step = nn.Parameter(torch.zeros(1).long(), requires_grad=False)
116
- self.num_params()
117
-
118
- def forward(self, x, mels):
119
- self.step += 1
120
- bsize = x.size(0)
121
- if torch.cuda.is_available():
122
- h1 = torch.zeros(1, bsize, self.rnn_dims).cuda()
123
- h2 = torch.zeros(1, bsize, self.rnn_dims).cuda()
124
- else:
125
- h1 = torch.zeros(1, bsize, self.rnn_dims).cpu()
126
- h2 = torch.zeros(1, bsize, self.rnn_dims).cpu()
127
- mels, aux = self.upsample(mels)
128
-
129
- aux_idx = [self.aux_dims * i for i in range(5)]
130
- a1 = aux[:, :, aux_idx[0]:aux_idx[1]]
131
- a2 = aux[:, :, aux_idx[1]:aux_idx[2]]
132
- a3 = aux[:, :, aux_idx[2]:aux_idx[3]]
133
- a4 = aux[:, :, aux_idx[3]:aux_idx[4]]
134
-
135
- x = torch.cat([x.unsqueeze(-1), mels, a1], dim=2)
136
- x = self.I(x)
137
- res = x
138
- x, _ = self.rnn1(x, h1)
139
-
140
- x = x + res
141
- res = x
142
- x = torch.cat([x, a2], dim=2)
143
- x, _ = self.rnn2(x, h2)
144
-
145
- x = x + res
146
- x = torch.cat([x, a3], dim=2)
147
- x = F.relu(self.fc1(x))
148
-
149
- x = torch.cat([x, a4], dim=2)
150
- x = F.relu(self.fc2(x))
151
- return self.fc3(x)
152
-
153
- def generate(self, mels, batched, target, overlap, mu_law, progress_callback=None):
154
- mu_law = mu_law if self.mode == 'RAW' else False
155
- progress_callback = progress_callback or self.gen_display
156
-
157
- self.eval()
158
- output = []
159
- start = time.time()
160
- rnn1 = self.get_gru_cell(self.rnn1)
161
- rnn2 = self.get_gru_cell(self.rnn2)
162
-
163
- with torch.no_grad():
164
- if torch.cuda.is_available():
165
- mels = mels.cuda()
166
- else:
167
- mels = mels.cpu()
168
- wave_len = (mels.size(-1) - 1) * self.hop_length
169
- mels = self.pad_tensor(mels.transpose(1, 2), pad=self.pad, side='both')
170
- mels, aux = self.upsample(mels.transpose(1, 2))
171
-
172
- if batched:
173
- mels = self.fold_with_overlap(mels, target, overlap)
174
- aux = self.fold_with_overlap(aux, target, overlap)
175
-
176
- b_size, seq_len, _ = mels.size()
177
-
178
- if torch.cuda.is_available():
179
- h1 = torch.zeros(b_size, self.rnn_dims).cuda()
180
- h2 = torch.zeros(b_size, self.rnn_dims).cuda()
181
- x = torch.zeros(b_size, 1).cuda()
182
- else:
183
- h1 = torch.zeros(b_size, self.rnn_dims).cpu()
184
- h2 = torch.zeros(b_size, self.rnn_dims).cpu()
185
- x = torch.zeros(b_size, 1).cpu()
186
-
187
- d = self.aux_dims
188
- aux_split = [aux[:, :, d * i:d * (i + 1)] for i in range(4)]
189
-
190
- for i in range(seq_len):
191
-
192
- m_t = mels[:, i, :]
193
-
194
- a1_t, a2_t, a3_t, a4_t = (a[:, i, :] for a in aux_split)
195
-
196
- x = torch.cat([x, m_t, a1_t], dim=1)
197
- x = self.I(x)
198
- h1 = rnn1(x, h1)
199
-
200
- x = x + h1
201
- inp = torch.cat([x, a2_t], dim=1)
202
- h2 = rnn2(inp, h2)
203
-
204
- x = x + h2
205
- x = torch.cat([x, a3_t], dim=1)
206
- x = F.relu(self.fc1(x))
207
-
208
- x = torch.cat([x, a4_t], dim=1)
209
- x = F.relu(self.fc2(x))
210
-
211
- logits = self.fc3(x)
212
-
213
- if self.mode == 'MOL':
214
- sample = sample_from_discretized_mix_logistic(logits.unsqueeze(0).transpose(1, 2))
215
- output.append(sample.view(-1))
216
- if torch.cuda.is_available():
217
- # x = torch.FloatTensor([[sample]]).cuda()
218
- x = sample.transpose(0, 1).cuda()
219
- else:
220
- x = sample.transpose(0, 1)
221
-
222
- elif self.mode == 'RAW' :
223
- posterior = F.softmax(logits, dim=1)
224
- distrib = torch.distributions.Categorical(posterior)
225
-
226
- sample = 2 * distrib.sample().float() / (self.n_classes - 1.) - 1.
227
- output.append(sample)
228
- x = sample.unsqueeze(-1)
229
- else:
230
- raise RuntimeError("Unknown model mode value - ", self.mode)
231
-
232
- if i % 100 == 0:
233
- gen_rate = (i + 1) / (time.time() - start) * b_size / 1000
234
- progress_callback(i, seq_len, b_size, gen_rate)
235
-
236
- output = torch.stack(output).transpose(0, 1)
237
- output = output.cpu().numpy()
238
- output = output.astype(np.float64)
239
-
240
- if batched:
241
- output = self.xfade_and_unfold(output, target, overlap)
242
- else:
243
- output = output[0]
244
-
245
- if mu_law:
246
- output = decode_mu_law(output, self.n_classes, False)
247
- if hp.apply_preemphasis:
248
- output = de_emphasis(output)
249
-
250
- # Fade-out at the end to avoid signal cutting out suddenly
251
- fade_out = np.linspace(1, 0, 20 * self.hop_length)
252
- output = output[:wave_len]
253
- output[-20 * self.hop_length:] *= fade_out
254
-
255
- self.train()
256
-
257
- return output
258
-
259
-
260
- def gen_display(self, i, seq_len, b_size, gen_rate):
261
- pbar = progbar(i, seq_len)
262
- msg = f'| {pbar} {i*b_size}/{seq_len*b_size} | Batch Size: {b_size} | Gen Rate: {gen_rate:.1f}kHz | '
263
- stream(msg)
264
-
265
- def get_gru_cell(self, gru):
266
- gru_cell = nn.GRUCell(gru.input_size, gru.hidden_size)
267
- gru_cell.weight_hh.data = gru.weight_hh_l0.data
268
- gru_cell.weight_ih.data = gru.weight_ih_l0.data
269
- gru_cell.bias_hh.data = gru.bias_hh_l0.data
270
- gru_cell.bias_ih.data = gru.bias_ih_l0.data
271
- return gru_cell
272
-
273
- def pad_tensor(self, x, pad, side='both'):
274
- # NB - this is just a quick method i need right now
275
- # i.e., it won't generalise to other shapes/dims
276
- b, t, c = x.size()
277
- total = t + 2 * pad if side == 'both' else t + pad
278
- if torch.cuda.is_available():
279
- padded = torch.zeros(b, total, c).cuda()
280
- else:
281
- padded = torch.zeros(b, total, c).cpu()
282
- if side == 'before' or side == 'both':
283
- padded[:, pad:pad + t, :] = x
284
- elif side == 'after':
285
- padded[:, :t, :] = x
286
- return padded
287
-
288
- def fold_with_overlap(self, x, target, overlap):
289
-
290
- ''' Fold the tensor with overlap for quick batched inference.
291
- Overlap will be used for crossfading in xfade_and_unfold()
292
-
293
- Args:
294
- x (tensor) : Upsampled conditioning features.
295
- shape=(1, timesteps, features)
296
- target (int) : Target timesteps for each index of batch
297
- overlap (int) : Timesteps for both xfade and rnn warmup
298
-
299
- Return:
300
- (tensor) : shape=(num_folds, target + 2 * overlap, features)
301
-
302
- Details:
303
- x = [[h1, h2, ... hn]]
304
-
305
- Where each h is a vector of conditioning features
306
-
307
- Eg: target=2, overlap=1 with x.size(1)=10
308
-
309
- folded = [[h1, h2, h3, h4],
310
- [h4, h5, h6, h7],
311
- [h7, h8, h9, h10]]
312
- '''
313
-
314
- _, total_len, features = x.size()
315
-
316
- # Calculate variables needed
317
- num_folds = (total_len - overlap) // (target + overlap)
318
- extended_len = num_folds * (overlap + target) + overlap
319
- remaining = total_len - extended_len
320
-
321
- # Pad if some time steps poking out
322
- if remaining != 0:
323
- num_folds += 1
324
- padding = target + 2 * overlap - remaining
325
- x = self.pad_tensor(x, padding, side='after')
326
-
327
- if torch.cuda.is_available():
328
- folded = torch.zeros(num_folds, target + 2 * overlap, features).cuda()
329
- else:
330
- folded = torch.zeros(num_folds, target + 2 * overlap, features).cpu()
331
-
332
- # Get the values for the folded tensor
333
- for i in range(num_folds):
334
- start = i * (target + overlap)
335
- end = start + target + 2 * overlap
336
- folded[i] = x[:, start:end, :]
337
-
338
- return folded
339
-
340
- def xfade_and_unfold(self, y, target, overlap):
341
-
342
- ''' Applies a crossfade and unfolds into a 1d array.
343
-
344
- Args:
345
- y (ndarry) : Batched sequences of audio samples
346
- shape=(num_folds, target + 2 * overlap)
347
- dtype=np.float64
348
- overlap (int) : Timesteps for both xfade and rnn warmup
349
-
350
- Return:
351
- (ndarry) : audio samples in a 1d array
352
- shape=(total_len)
353
- dtype=np.float64
354
-
355
- Details:
356
- y = [[seq1],
357
- [seq2],
358
- [seq3]]
359
-
360
- Apply a gain envelope at both ends of the sequences
361
-
362
- y = [[seq1_in, seq1_target, seq1_out],
363
- [seq2_in, seq2_target, seq2_out],
364
- [seq3_in, seq3_target, seq3_out]]
365
-
366
- Stagger and add up the groups of samples:
367
-
368
- [seq1_in, seq1_target, (seq1_out + seq2_in), seq2_target, ...]
369
-
370
- '''
371
-
372
- num_folds, length = y.shape
373
- target = length - 2 * overlap
374
- total_len = num_folds * (target + overlap) + overlap
375
-
376
- # Need some silence for the rnn warmup
377
- silence_len = overlap // 2
378
- fade_len = overlap - silence_len
379
- silence = np.zeros((silence_len), dtype=np.float64)
380
-
381
- # Equal power crossfade
382
- t = np.linspace(-1, 1, fade_len, dtype=np.float64)
383
- fade_in = np.sqrt(0.5 * (1 + t))
384
- fade_out = np.sqrt(0.5 * (1 - t))
385
-
386
- # Concat the silence to the fades
387
- fade_in = np.concatenate([silence, fade_in])
388
- fade_out = np.concatenate([fade_out, silence])
389
-
390
- # Apply the gain to the overlap samples
391
- y[:, :overlap] *= fade_in
392
- y[:, -overlap:] *= fade_out
393
-
394
- unfolded = np.zeros((total_len), dtype=np.float64)
395
-
396
- # Loop to add up all the samples
397
- for i in range(num_folds):
398
- start = i * (target + overlap)
399
- end = start + target + 2 * overlap
400
- unfolded[start:end] += y[i]
401
-
402
- return unfolded
403
-
404
- def get_step(self) :
405
- return self.step.data.item()
406
-
407
- def checkpoint(self, model_dir, optimizer) :
408
- k_steps = self.get_step() // 1000
409
- self.save(model_dir.joinpath("checkpoint_%dk_steps.pt" % k_steps), optimizer)
410
-
411
- def log(self, path, msg) :
412
- with open(path, 'a') as f:
413
- print(msg, file=f)
414
-
415
- def load(self, path, optimizer) :
416
- checkpoint = torch.load(path)
417
- if "optimizer_state" in checkpoint:
418
- self.load_state_dict(checkpoint["model_state"])
419
- optimizer.load_state_dict(checkpoint["optimizer_state"])
420
- else:
421
- # Backwards compatibility
422
- self.load_state_dict(checkpoint)
423
-
424
- def save(self, path, optimizer) :
425
- torch.save({
426
- "model_state": self.state_dict(),
427
- "optimizer_state": optimizer.state_dict(),
428
- }, path)
429
-
430
- def num_params(self, print_out=True):
431
- parameters = filter(lambda p: p.requires_grad, self.parameters())
432
- parameters = sum([np.prod(p.size()) for p in parameters]) / 1_000_000
433
- if print_out :
434
- print('Trainable Parameters: %.3fM' % parameters)
 
1
+ import torch
2
+ import torch.nn as nn
3
+ import torch.nn.functional as F
4
+ from vocoder.distribution import sample_from_discretized_mix_logistic
5
+ from vocoder.display import *
6
+ from vocoder.audio import *
7
+
8
+
9
+ class ResBlock(nn.Module):
10
+ def __init__(self, dims):
11
+ super().__init__()
12
+ self.conv1 = nn.Conv1d(dims, dims, kernel_size=1, bias=False)
13
+ self.conv2 = nn.Conv1d(dims, dims, kernel_size=1, bias=False)
14
+ self.batch_norm1 = nn.BatchNorm1d(dims)
15
+ self.batch_norm2 = nn.BatchNorm1d(dims)
16
+
17
+ def forward(self, x):
18
+ residual = x
19
+ x = self.conv1(x)
20
+ x = self.batch_norm1(x)
21
+ x = F.relu(x)
22
+ x = self.conv2(x)
23
+ x = self.batch_norm2(x)
24
+ return x + residual
25
+
26
+
27
+ class MelResNet(nn.Module):
28
+ def __init__(self, res_blocks, in_dims, compute_dims, res_out_dims, pad):
29
+ super().__init__()
30
+ k_size = pad * 2 + 1
31
+ self.conv_in = nn.Conv1d(in_dims, compute_dims, kernel_size=k_size, bias=False)
32
+ self.batch_norm = nn.BatchNorm1d(compute_dims)
33
+ self.layers = nn.ModuleList()
34
+ for i in range(res_blocks):
35
+ self.layers.append(ResBlock(compute_dims))
36
+ self.conv_out = nn.Conv1d(compute_dims, res_out_dims, kernel_size=1)
37
+
38
+ def forward(self, x):
39
+ x = self.conv_in(x)
40
+ x = self.batch_norm(x)
41
+ x = F.relu(x)
42
+ for f in self.layers: x = f(x)
43
+ x = self.conv_out(x)
44
+ return x
45
+
46
+
47
+ class Stretch2d(nn.Module):
48
+ def __init__(self, x_scale, y_scale):
49
+ super().__init__()
50
+ self.x_scale = x_scale
51
+ self.y_scale = y_scale
52
+
53
+ def forward(self, x):
54
+ b, c, h, w = x.size()
55
+ x = x.unsqueeze(-1).unsqueeze(3)
56
+ x = x.repeat(1, 1, 1, self.y_scale, 1, self.x_scale)
57
+ return x.view(b, c, h * self.y_scale, w * self.x_scale)
58
+
59
+
60
+ class UpsampleNetwork(nn.Module):
61
+ def __init__(self, feat_dims, upsample_scales, compute_dims,
62
+ res_blocks, res_out_dims, pad):
63
+ super().__init__()
64
+ total_scale = np.cumproduct(upsample_scales)[-1]
65
+ self.indent = pad * total_scale
66
+ self.resnet = MelResNet(res_blocks, feat_dims, compute_dims, res_out_dims, pad)
67
+ self.resnet_stretch = Stretch2d(total_scale, 1)
68
+ self.up_layers = nn.ModuleList()
69
+ for scale in upsample_scales:
70
+ k_size = (1, scale * 2 + 1)
71
+ padding = (0, scale)
72
+ stretch = Stretch2d(scale, 1)
73
+ conv = nn.Conv2d(1, 1, kernel_size=k_size, padding=padding, bias=False)
74
+ conv.weight.data.fill_(1. / k_size[1])
75
+ self.up_layers.append(stretch)
76
+ self.up_layers.append(conv)
77
+
78
+ def forward(self, m):
79
+ aux = self.resnet(m).unsqueeze(1)
80
+ aux = self.resnet_stretch(aux)
81
+ aux = aux.squeeze(1)
82
+ m = m.unsqueeze(1)
83
+ for f in self.up_layers: m = f(m)
84
+ m = m.squeeze(1)[:, :, self.indent:-self.indent]
85
+ return m.transpose(1, 2), aux.transpose(1, 2)
86
+
87
+
88
+ class WaveRNN(nn.Module):
89
+ def __init__(self, rnn_dims, fc_dims, bits, pad, upsample_factors,
90
+ feat_dims, compute_dims, res_out_dims, res_blocks,
91
+ hop_length, sample_rate, mode='RAW'):
92
+ super().__init__()
93
+ self.mode = mode
94
+ self.pad = pad
95
+ if self.mode == 'RAW' :
96
+ self.n_classes = 2 ** bits
97
+ elif self.mode == 'MOL' :
98
+ self.n_classes = 30
99
+ else :
100
+ RuntimeError("Unknown model mode value - ", self.mode)
101
+
102
+ self.rnn_dims = rnn_dims
103
+ self.aux_dims = res_out_dims // 4
104
+ self.hop_length = hop_length
105
+ self.sample_rate = sample_rate
106
+
107
+ self.upsample = UpsampleNetwork(feat_dims, upsample_factors, compute_dims, res_blocks, res_out_dims, pad)
108
+ self.I = nn.Linear(feat_dims + self.aux_dims + 1, rnn_dims)
109
+ self.rnn1 = nn.GRU(rnn_dims, rnn_dims, batch_first=True)
110
+ self.rnn2 = nn.GRU(rnn_dims + self.aux_dims, rnn_dims, batch_first=True)
111
+ self.fc1 = nn.Linear(rnn_dims + self.aux_dims, fc_dims)
112
+ self.fc2 = nn.Linear(fc_dims + self.aux_dims, fc_dims)
113
+ self.fc3 = nn.Linear(fc_dims, self.n_classes)
114
+
115
+ self.step = nn.Parameter(torch.zeros(1).long(), requires_grad=False)
116
+ self.num_params()
117
+
118
+ def forward(self, x, mels):
119
+ self.step += 1
120
+ bsize = x.size(0)
121
+ if torch.cuda.is_available():
122
+ h1 = torch.zeros(1, bsize, self.rnn_dims).cuda()
123
+ h2 = torch.zeros(1, bsize, self.rnn_dims).cuda()
124
+ else:
125
+ h1 = torch.zeros(1, bsize, self.rnn_dims).cpu()
126
+ h2 = torch.zeros(1, bsize, self.rnn_dims).cpu()
127
+ mels, aux = self.upsample(mels)
128
+
129
+ aux_idx = [self.aux_dims * i for i in range(5)]
130
+ a1 = aux[:, :, aux_idx[0]:aux_idx[1]]
131
+ a2 = aux[:, :, aux_idx[1]:aux_idx[2]]
132
+ a3 = aux[:, :, aux_idx[2]:aux_idx[3]]
133
+ a4 = aux[:, :, aux_idx[3]:aux_idx[4]]
134
+
135
+ x = torch.cat([x.unsqueeze(-1), mels, a1], dim=2)
136
+ x = self.I(x)
137
+ res = x
138
+ x, _ = self.rnn1(x, h1)
139
+
140
+ x = x + res
141
+ res = x
142
+ x = torch.cat([x, a2], dim=2)
143
+ x, _ = self.rnn2(x, h2)
144
+
145
+ x = x + res
146
+ x = torch.cat([x, a3], dim=2)
147
+ x = F.relu(self.fc1(x))
148
+
149
+ x = torch.cat([x, a4], dim=2)
150
+ x = F.relu(self.fc2(x))
151
+ return self.fc3(x)
152
+
153
+ def generate(self, mels, batched, target, overlap, mu_law, progress_callback=None):
154
+ mu_law = mu_law if self.mode == 'RAW' else False
155
+ progress_callback = progress_callback or self.gen_display
156
+
157
+ self.eval()
158
+ output = []
159
+ start = time.time()
160
+ rnn1 = self.get_gru_cell(self.rnn1)
161
+ rnn2 = self.get_gru_cell(self.rnn2)
162
+
163
+ with torch.no_grad():
164
+ if torch.cuda.is_available():
165
+ mels = mels.cuda()
166
+ else:
167
+ mels = mels.cpu()
168
+ wave_len = (mels.size(-1) - 1) * self.hop_length
169
+ mels = self.pad_tensor(mels.transpose(1, 2), pad=self.pad, side='both')
170
+ mels, aux = self.upsample(mels.transpose(1, 2))
171
+
172
+ if batched:
173
+ mels = self.fold_with_overlap(mels, target, overlap)
174
+ aux = self.fold_with_overlap(aux, target, overlap)
175
+
176
+ b_size, seq_len, _ = mels.size()
177
+
178
+ if torch.cuda.is_available():
179
+ h1 = torch.zeros(b_size, self.rnn_dims).cuda()
180
+ h2 = torch.zeros(b_size, self.rnn_dims).cuda()
181
+ x = torch.zeros(b_size, 1).cuda()
182
+ else:
183
+ h1 = torch.zeros(b_size, self.rnn_dims).cpu()
184
+ h2 = torch.zeros(b_size, self.rnn_dims).cpu()
185
+ x = torch.zeros(b_size, 1).cpu()
186
+
187
+ d = self.aux_dims
188
+ aux_split = [aux[:, :, d * i:d * (i + 1)] for i in range(4)]
189
+
190
+ for i in range(seq_len):
191
+
192
+ m_t = mels[:, i, :]
193
+
194
+ a1_t, a2_t, a3_t, a4_t = (a[:, i, :] for a in aux_split)
195
+
196
+ x = torch.cat([x, m_t, a1_t], dim=1)
197
+ x = self.I(x)
198
+ h1 = rnn1(x, h1)
199
+
200
+ x = x + h1
201
+ inp = torch.cat([x, a2_t], dim=1)
202
+ h2 = rnn2(inp, h2)
203
+
204
+ x = x + h2
205
+ x = torch.cat([x, a3_t], dim=1)
206
+ x = F.relu(self.fc1(x))
207
+
208
+ x = torch.cat([x, a4_t], dim=1)
209
+ x = F.relu(self.fc2(x))
210
+
211
+ logits = self.fc3(x)
212
+
213
+ if self.mode == 'MOL':
214
+ sample = sample_from_discretized_mix_logistic(logits.unsqueeze(0).transpose(1, 2))
215
+ output.append(sample.view(-1))
216
+ if torch.cuda.is_available():
217
+ # x = torch.FloatTensor([[sample]]).cuda()
218
+ x = sample.transpose(0, 1).cuda()
219
+ else:
220
+ x = sample.transpose(0, 1)
221
+
222
+ elif self.mode == 'RAW' :
223
+ posterior = F.softmax(logits, dim=1)
224
+ distrib = torch.distributions.Categorical(posterior)
225
+
226
+ sample = 2 * distrib.sample().float() / (self.n_classes - 1.) - 1.
227
+ output.append(sample)
228
+ x = sample.unsqueeze(-1)
229
+ else:
230
+ raise RuntimeError("Unknown model mode value - ", self.mode)
231
+
232
+ if i % 100 == 0:
233
+ gen_rate = (i + 1) / (time.time() - start) * b_size / 1000
234
+ progress_callback(i, seq_len, b_size, gen_rate)
235
+
236
+ output = torch.stack(output).transpose(0, 1)
237
+ output = output.cpu().numpy()
238
+ output = output.astype(np.float64)
239
+
240
+ if batched:
241
+ output = self.xfade_and_unfold(output, target, overlap)
242
+ else:
243
+ output = output[0]
244
+
245
+ if mu_law:
246
+ output = decode_mu_law(output, self.n_classes, False)
247
+ if hp.apply_preemphasis:
248
+ output = de_emphasis(output)
249
+
250
+ # Fade-out at the end to avoid signal cutting out suddenly
251
+ fade_out = np.linspace(1, 0, 20 * self.hop_length)
252
+ output = output[:wave_len]
253
+ output[-20 * self.hop_length:] *= fade_out
254
+
255
+ self.train()
256
+
257
+ return output
258
+
259
+
260
+ def gen_display(self, i, seq_len, b_size, gen_rate):
261
+ pbar = progbar(i, seq_len)
262
+ msg = f'| {pbar} {i*b_size}/{seq_len*b_size} | Batch Size: {b_size} | Gen Rate: {gen_rate:.1f}kHz | '
263
+ stream(msg)
264
+
265
+ def get_gru_cell(self, gru):
266
+ gru_cell = nn.GRUCell(gru.input_size, gru.hidden_size)
267
+ gru_cell.weight_hh.data = gru.weight_hh_l0.data
268
+ gru_cell.weight_ih.data = gru.weight_ih_l0.data
269
+ gru_cell.bias_hh.data = gru.bias_hh_l0.data
270
+ gru_cell.bias_ih.data = gru.bias_ih_l0.data
271
+ return gru_cell
272
+
273
+ def pad_tensor(self, x, pad, side='both'):
274
+ # NB - this is just a quick method i need right now
275
+ # i.e., it won't generalise to other shapes/dims
276
+ b, t, c = x.size()
277
+ total = t + 2 * pad if side == 'both' else t + pad
278
+ if torch.cuda.is_available():
279
+ padded = torch.zeros(b, total, c).cuda()
280
+ else:
281
+ padded = torch.zeros(b, total, c).cpu()
282
+ if side == 'before' or side == 'both':
283
+ padded[:, pad:pad + t, :] = x
284
+ elif side == 'after':
285
+ padded[:, :t, :] = x
286
+ return padded
287
+
288
+ def fold_with_overlap(self, x, target, overlap):
289
+
290
+ ''' Fold the tensor with overlap for quick batched inference.
291
+ Overlap will be used for crossfading in xfade_and_unfold()
292
+
293
+ Args:
294
+ x (tensor) : Upsampled conditioning features.
295
+ shape=(1, timesteps, features)
296
+ target (int) : Target timesteps for each index of batch
297
+ overlap (int) : Timesteps for both xfade and rnn warmup
298
+
299
+ Return:
300
+ (tensor) : shape=(num_folds, target + 2 * overlap, features)
301
+
302
+ Details:
303
+ x = [[h1, h2, ... hn]]
304
+
305
+ Where each h is a vector of conditioning features
306
+
307
+ Eg: target=2, overlap=1 with x.size(1)=10
308
+
309
+ folded = [[h1, h2, h3, h4],
310
+ [h4, h5, h6, h7],
311
+ [h7, h8, h9, h10]]
312
+ '''
313
+
314
+ _, total_len, features = x.size()
315
+
316
+ # Calculate variables needed
317
+ num_folds = (total_len - overlap) // (target + overlap)
318
+ extended_len = num_folds * (overlap + target) + overlap
319
+ remaining = total_len - extended_len
320
+
321
+ # Pad if some time steps poking out
322
+ if remaining != 0:
323
+ num_folds += 1
324
+ padding = target + 2 * overlap - remaining
325
+ x = self.pad_tensor(x, padding, side='after')
326
+
327
+ if torch.cuda.is_available():
328
+ folded = torch.zeros(num_folds, target + 2 * overlap, features).cuda()
329
+ else:
330
+ folded = torch.zeros(num_folds, target + 2 * overlap, features).cpu()
331
+
332
+ # Get the values for the folded tensor
333
+ for i in range(num_folds):
334
+ start = i * (target + overlap)
335
+ end = start + target + 2 * overlap
336
+ folded[i] = x[:, start:end, :]
337
+
338
+ return folded
339
+
340
+ def xfade_and_unfold(self, y, target, overlap):
341
+
342
+ ''' Applies a crossfade and unfolds into a 1d array.
343
+
344
+ Args:
345
+ y (ndarry) : Batched sequences of audio samples
346
+ shape=(num_folds, target + 2 * overlap)
347
+ dtype=np.float64
348
+ overlap (int) : Timesteps for both xfade and rnn warmup
349
+
350
+ Return:
351
+ (ndarry) : audio samples in a 1d array
352
+ shape=(total_len)
353
+ dtype=np.float64
354
+
355
+ Details:
356
+ y = [[seq1],
357
+ [seq2],
358
+ [seq3]]
359
+
360
+ Apply a gain envelope at both ends of the sequences
361
+
362
+ y = [[seq1_in, seq1_target, seq1_out],
363
+ [seq2_in, seq2_target, seq2_out],
364
+ [seq3_in, seq3_target, seq3_out]]
365
+
366
+ Stagger and add up the groups of samples:
367
+
368
+ [seq1_in, seq1_target, (seq1_out + seq2_in), seq2_target, ...]
369
+
370
+ '''
371
+
372
+ num_folds, length = y.shape
373
+ target = length - 2 * overlap
374
+ total_len = num_folds * (target + overlap) + overlap
375
+
376
+ # Need some silence for the rnn warmup
377
+ silence_len = overlap // 2
378
+ fade_len = overlap - silence_len
379
+ silence = np.zeros((silence_len), dtype=np.float64)
380
+
381
+ # Equal power crossfade
382
+ t = np.linspace(-1, 1, fade_len, dtype=np.float64)
383
+ fade_in = np.sqrt(0.5 * (1 + t))
384
+ fade_out = np.sqrt(0.5 * (1 - t))
385
+
386
+ # Concat the silence to the fades
387
+ fade_in = np.concatenate([silence, fade_in])
388
+ fade_out = np.concatenate([fade_out, silence])
389
+
390
+ # Apply the gain to the overlap samples
391
+ y[:, :overlap] *= fade_in
392
+ y[:, -overlap:] *= fade_out
393
+
394
+ unfolded = np.zeros((total_len), dtype=np.float64)
395
+
396
+ # Loop to add up all the samples
397
+ for i in range(num_folds):
398
+ start = i * (target + overlap)
399
+ end = start + target + 2 * overlap
400
+ unfolded[start:end] += y[i]
401
+
402
+ return unfolded
403
+
404
+ def get_step(self) :
405
+ return self.step.data.item()
406
+
407
+ def checkpoint(self, model_dir, optimizer) :
408
+ k_steps = self.get_step() // 1000
409
+ self.save(model_dir.joinpath("checkpoint_%dk_steps.pt" % k_steps), optimizer)
410
+
411
+ def log(self, path, msg) :
412
+ with open(path, 'a') as f:
413
+ print(msg, file=f)
414
+
415
+ def load(self, path, optimizer) :
416
+ checkpoint = torch.load(path, weights_only=False)
417
+ if "optimizer_state" in checkpoint:
418
+ self.load_state_dict(checkpoint["model_state"])
419
+ optimizer.load_state_dict(checkpoint["optimizer_state"])
420
+ else:
421
+ # Backwards compatibility
422
+ self.load_state_dict(checkpoint)
423
+
424
+ def save(self, path, optimizer) :
425
+ torch.save({
426
+ "model_state": self.state_dict(),
427
+ "optimizer_state": optimizer.state_dict(),
428
+ }, path)
429
+
430
+ def num_params(self, print_out=True):
431
+ parameters = filter(lambda p: p.requires_grad, self.parameters())
432
+ parameters = sum([np.prod(p.size()) for p in parameters]) / 1_000_000
433
+ if print_out :
434
+ print('Trainable Parameters: %.3fM' % parameters)