diff --git a/README.md b/README.md index 3558195..b89b413 100644 --- a/README.md +++ b/README.md @@ -2,6 +2,13 @@ Generates samples using [Piper](https://github.com/rhasspy/piper/) for training a wake word system like [openWakeWord](https://github.com/dscripka/openWakeWord). +Available models: + +* [English](https://github.com/rhasspy/piper-sample-generator/releases/download/v2.0.0/en_US-libritts_r-medium.pt) +* [French](https://github.com/rhasspy/piper-sample-generator/releases/download/v2.0.0/fr_FR-mls-medium.pt) +* [German](https://github.com/rhasspy/piper-sample-generator/releases/download/v2.0.0/de_DE-mls-medium.pt) +* [Dutch](https://github.com/rhasspy/piper-sample-generator/releases/download/v2.0.0/nl_NL-mls-medium.pt) + ## Install @@ -23,6 +30,7 @@ Download the LibriTTS-R generator (exported from [checkpoint](https://huggingfac wget -O models/en-us-libritts-high.pt 'https://github.com/rhasspy/piper-sample-generator/releases/download/v2.0.0/en_US-libritts_r-medium.pt' ``` +See links above for models for other languages. ## Run @@ -72,3 +80,28 @@ This will do several things to each sample: * Change the acoustics of the sample to sound like the speaker was in a room with echo or using a poor quality microphone 3. Resample to 16Khz for training (e.g., [openWakeWord](https://github.com/dscripka/openWakeWord)) + +## Short Phrases + +Models that were trained on audio books tend to perform poorly when speaking short phrases or single words. +The French, German, and Dutch models trained from the [MLS](http://openslr.org/94/) have this problem. + +The problem can be mitigated by repeating the phrase over and over, and then clipping out a single sample. +To do this automatically, follow these steps: + +1. Ensure your short phrase ends with a comma (`,`) +2. Lower the noise settings with `--noise-scales 0.333` and `--noise-scale-ws 0.333` +3. Use `--min-phoneme-count 300` (the value 300 was determined empirically and may be less for some models) + +For example: + +``` sh +python3 generate_samples.py \ + 'framboise,' \ + --model models/fr_FR-mls-medium.pt \ + --noise-scales 0.333 \ + --noise-scale-ws 0.333 \ + --min-phoneme-count 300 + --max-samples 1 \ + --output-dir . +``` diff --git a/generate_samples.py b/generate_samples.py index 349ebde..d0a6e85 100755 --- a/generate_samples.py +++ b/generate_samples.py @@ -7,7 +7,7 @@ import logging import os import wave from pathlib import Path -from typing import List, Union +from typing import Any, Dict, List, Optional, Tuple, Union import numpy as np import torch @@ -24,21 +24,20 @@ logging.basicConfig(level=logging.DEBUG) # Main generation function def generate_samples( - text: Union[List, str], - output_dir: str, - max_samples: int = None, - file_names: List[str] = [], - model: str = os.path.join( - Path(__file__).parent, "models", "en_US-libritts_r-medium.pt" - ), + text: Union[List[str], str], + output_dir: Union[str, Path], + max_samples: Optional[int] = None, + file_names: Optional[List[str]] = None, + model: Union[str, Path] = _DIR / "models" / "en_US-libritts_r-medium.pt", batch_size: int = 1, - slerp_weights: List[float] = [0.5], - length_scales: List[float] = [0.75, 1, 1.25], - noise_scales: List[float] = [0.667], - noise_scale_ws: List[float] = [0.8], - max_speakers: float = None, + slerp_weights: Tuple[float, ...] = (0.5,), + length_scales: Tuple[float, ...] = (0.75, 1, 1.25), + noise_scales: Tuple[float, ...] = (0.667,), + noise_scale_ws: Tuple[float, ...] = (0.8,), + max_speakers: Optional[float] = None, verbose: bool = False, auto_reduce_batch_size: bool = False, + min_phoneme_count: Optional[int] = None, **kwargs, ) -> None: """ @@ -61,6 +60,8 @@ def generate_samples( verbose (bool): Enable or disable more detailed logging messages (default: False). auto_reduce_batch_size (bool): Automatically and temporarily reduce the batch size if CUDA OOM errors are detected, and try to resume generation. + min_phoneme_count (int): If set, ensure this number of phonemes is always sent to the model. + Clip audio to extract original phrase. Returns: None @@ -71,12 +72,13 @@ def generate_samples( _LOGGER.debug("Loading %s", model) model_path = Path(model) - model = torch.load(model_path) - model.eval() + + torch_model = torch.load(model_path) + torch_model.eval() _LOGGER.info("Successfully loaded the model") if torch.cuda.is_available(): - model.cuda() + torch_model.cuda() _LOGGER.debug("CUDA available, using GPU") output_dir = Path(output_dir) @@ -147,10 +149,14 @@ def generate_samples( speaker_1 = torch.LongTensor([s[0] for s in speakers_batch]) speaker_2 = torch.LongTensor([s[1] for s in speakers_batch]) - phoneme_ids = [ - get_phonemes(voice, config, next(texts), verbose) - for i in range(batch_size) - ] + phoneme_ids_by_batch = [] + clip_indexes_by_batch = [] + for i in range(batch_size): + phoneme_ids, clip_phoneme_index = get_phonemes( + voice, config, next(texts), verbose, min_phoneme_count + ) + phoneme_ids_by_batch.append(phoneme_ids) + clip_indexes_by_batch.append(clip_phoneme_index) def right_pad_lists(lists): max_length = max(len(lst) for lst in lists) @@ -162,18 +168,18 @@ def generate_samples( padded_lists.append(padded_l) return padded_lists - phoneme_ids = right_pad_lists(phoneme_ids) + phoneme_ids_by_batch = right_pad_lists(phoneme_ids_by_batch) if auto_reduce_batch_size: oom_error = True counter = 1 while oom_error is True: try: - audio = generate_audio( - model, + audio, phoneme_samples = generate_audio( + torch_model, speaker_1[0 : batch_size // counter], speaker_2[0 : batch_size // counter], - phoneme_ids[0 : batch_size // counter], + phoneme_ids_by_batch[0 : batch_size // counter], slerp_weight, noise_scale, noise_scale_w, @@ -186,11 +192,11 @@ def generate_samples( gc.collect() counter += 1 # reduce batch size to avoid OOM errors else: - audio = generate_audio( - model, + audio, phoneme_samples = generate_audio( + torch_model, speaker_1, speaker_2, - phoneme_ids, + phoneme_ids_by_batch, slerp_weight, noise_scale, noise_scale_w, @@ -198,19 +204,34 @@ def generate_samples( max_len, ) + # Clip audio when using min_phoneme_count + for i, clip_phoneme_index in enumerate(clip_indexes_by_batch): + if clip_phoneme_index is not None: + last_sample_idx = int( + phoneme_samples[i].flatten()[clip_phoneme_index:].sum().item() + ) + + # Fill remainder of audio with silence. + # It will be removed in the next stage. + audio[i, 0, :-last_sample_idx] = 0 + # Resample audio audio = resampler(audio.cpu()).numpy() audio_int16 = audio_float_to_int16(audio) for audio_idx in range(audio_int16.shape[0]): # Use webrtcvad to trip silence from the clips - audio_data = remove_silence(audio_int16[audio_idx].flatten())[None,] + audio_data = remove_silence(audio_int16[audio_idx].flatten())[ + None, + ] if isinstance(file_names, it.cycle): wav_path = output_dir / next(file_names) else: wav_path = output_dir / f"{sample_idx}.wav" - with wave.open(str(wav_path), "wb") as wav_file: + + wav_file: wave.Wave_write = wave.open(str(wav_path), "wb") + with wav_file: wav_file.setframerate(resample_rate) wav_file.setsampwidth(2) wav_file.setnchannels(1) @@ -224,17 +245,22 @@ def generate_samples( # print(f"Batch {batch_idx +1}/{max_samples//batch_size} complete", " "*200, end='\r') # Next batch - _LOGGER.debug(f"Batch {batch_idx +1}/{max_samples//batch_size} complete") + _LOGGER.debug("Batch %s/%s complete", batch_idx + 1, max_samples // batch_size) speakers_batch = list(it.islice(speakers_iter, 0, batch_size)) batch_idx += 1 _LOGGER.info("Done") -def remove_silence(x, frame_duration=0.030, sample_rate=16000, min_start=2000): +def remove_silence( + x: np.ndarray, + frame_duration: float = 0.030, + sample_rate: int = 16000, + min_start: int = 2000, +) -> np.ndarray: """Uses webrtc voice activity detection to remove silence from the clips""" vad = webrtcvad.Vad(0) - if x.dtype == np.float32 or x.dtype == np.float64: + if x.dtype in (np.float32, np.float64): x = (x * 32767).astype(np.int16) x_new = x[0:min_start].tolist() step_size = int(sample_rate * frame_duration) @@ -295,10 +321,18 @@ def generate_audio( o = model.dec((z * y_mask)[:, :, :max_len], g=g) audio = o - return audio + phoneme_samples = w_ceil * 256 # hop length + + return audio, phoneme_samples -def get_phonemes(voice, config, text, verbose): +def get_phonemes( + voice: str, + config: Dict[str, Any], + text: str, + verbose: bool = False, + min_phoneme_count: Optional[int] = None, +) -> Tuple[List[int], Optional[int]]: # Combine all sentences phonemes = [ p @@ -309,18 +343,41 @@ def get_phonemes(voice, config, text, verbose): _LOGGER.debug("Phonemes: %s", phonemes) id_map = config["phoneme_id_map"] + + # Beginning of utterance phoneme_ids = list(id_map["^"]) + + # Phoneme ids for just the text + text_phoneme_ids = [] + for phoneme in phonemes: p_ids = id_map.get(phoneme) if p_ids is not None: phoneme_ids.extend(p_ids) + text_phoneme_ids.extend(p_ids) phoneme_ids.extend(id_map["_"]) + text_phoneme_ids.extend(id_map["_"]) + # Index where audio should be clipped at. + # When None, all of the audio will be used. + clip_phoneme_index: Optional[int] = None + + if min_phoneme_count is not None: + # Repeat phrase until minimum phoneme count is met. + # NOTE: It is critical that the ^ and $ phonemes are not repeated here. + while (len(phoneme_ids) - 1) < min_phoneme_count: + # We will clip audio at the beginning of the last phrase + clip_phoneme_index = len(phoneme_ids) - 1 + + phoneme_ids.extend(text_phoneme_ids) + + # End of utterance phoneme_ids.extend(id_map["$"]) - return phoneme_ids + + return phoneme_ids, clip_phoneme_index -def slerp(v1, v2, t, DOT_THR=0.9995, zdim=-1): +def slerp(v1, v2, t: float, DOT_THR: float = 0.9995, zdim: int = -1): """SLERP for pytorch tensors interpolating `v1` to `v2` with scale of `t`. `DOT_THR` determines when the vectors are too close to parallel. @@ -375,7 +432,9 @@ def audio_float_to_int16( return audio_norm -if __name__ == "__main__": +def main() -> None: + """Main entry point.""" + # Get command line arguments parser = argparse.ArgumentParser() parser.add_argument("text") @@ -401,7 +460,13 @@ if __name__ == "__main__": type=int, help="Maximum number of speakers to use (default: all)", ) + parser.add_argument("--min-phoneme-count", type=int) + parser.add_argument("--verbose", action="store_true") args = parser.parse_args().__dict__ # Generate speech generate_samples(**args) + + +if __name__ == "__main__": + main() diff --git a/pylintrc b/pylintrc index 22a70d0..561b2f1 100644 --- a/pylintrc +++ b/pylintrc @@ -35,3 +35,10 @@ disable= [FORMAT] expected-line-ending-format=LF + +[TYPECHECK] + +# List of members which are set dynamically and missed by pylint inference +# system, and so shouldn't trigger E1101 when accessed. Python regular +# expressions are accepted. +generated-members=numpy.*,torch.* diff --git a/requirements.txt b/requirements.txt index c6ac75a..8443d84 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,6 +1,6 @@ audiomentations==0.33.0 piper-phonemize==1.1.0 numpy<2 -torch +torch<2 torchaudio webrtcvad