diff --git a/.gitignore b/.gitignore index 9425a0aa..7ce70e40 100644 --- a/.gitignore +++ b/.gitignore @@ -117,6 +117,7 @@ venv.bak/ # pytorch models *.pth *.pth.tar +!dummy_speakers.pth result/ # setup.py diff --git a/TTS/bin/compute_embeddings.py b/TTS/bin/compute_embeddings.py index c58123df..d7fe3c4b 100644 --- a/TTS/bin/compute_embeddings.py +++ b/TTS/bin/compute_embeddings.py @@ -2,41 +2,34 @@ import argparse import os from argparse import RawTextHelpFormatter +import torch from tqdm import tqdm from TTS.config import load_config from TTS.tts.datasets import load_tts_samples +from TTS.tts.utils.managers import save_file from TTS.tts.utils.speakers import SpeakerManager parser = argparse.ArgumentParser( description="""Compute embedding vectors for each wav file in a dataset.\n\n""" """ Example runs: - python TTS/bin/compute_embeddings.py speaker_encoder_model.pth speaker_encoder_config.json dataset_config.json embeddings_output_path/ + python TTS/bin/compute_embeddings.py speaker_encoder_model.pth speaker_encoder_config.json dataset_config.json """, formatter_class=RawTextHelpFormatter, ) parser.add_argument("model_path", type=str, help="Path to model checkpoint file.") -parser.add_argument( - "config_path", - type=str, - help="Path to model config file.", -) - -parser.add_argument( - "config_dataset_path", - type=str, - help="Path to dataset config file.", -) -parser.add_argument("output_path", type=str, help="path for output speakers.json and/or speakers.npy.") -parser.add_argument( - "--old_file", type=str, help="Previous speakers.json file, only compute for new audios.", default=None -) -parser.add_argument("--use_cuda", type=bool, help="flag to set cuda. Default False", default=False) +parser.add_argument("config_path", type=str, help="Path to model config file.") +parser.add_argument("config_dataset_path", type=str, help="Path to dataset config file.") +parser.add_argument("--output_path", type=str, help="Path for output `pth` or `json` file.", default="speakers.pth") +parser.add_argument("--old_file", type=str, help="Previous embedding file to only compute new audios.", default=None) +parser.add_argument("--disable_cuda", type=bool, help="Flag to disable cuda.", default=False) parser.add_argument("--no_eval", type=bool, help="Do not compute eval?. Default False", default=False) args = parser.parse_args() +use_cuda = torch.cuda.is_available() and not args.disable_cuda + c_dataset = load_config(args.config_dataset_path) meta_data_train, meta_data_eval = load_tts_samples(c_dataset.datasets, eval_split=not args.no_eval) @@ -50,7 +43,7 @@ encoder_manager = SpeakerManager( encoder_model_path=args.model_path, encoder_config_path=args.config_path, d_vectors_file_path=args.old_file, - use_cuda=args.use_cuda, + use_cuda=use_cuda, ) class_name_key = encoder_manager.encoder_config.class_name_key @@ -79,13 +72,13 @@ for idx, wav_file in enumerate(tqdm(wav_files)): if speaker_mapping: # save speaker_mapping if target dataset is defined - if ".json" not in args.output_path: - mapping_file_path = os.path.join(args.output_path, "speakers.json") + if os.path.isdir(args.output_path): + mapping_file_path = os.path.join(args.output_path, "speakers.pth") else: mapping_file_path = args.output_path - os.makedirs(os.path.dirname(mapping_file_path), exist_ok=True) + if os.path.dirname(mapping_file_path) != "": + os.makedirs(os.path.dirname(mapping_file_path), exist_ok=True) - # pylint: disable=W0212 - encoder_manager._save_json(mapping_file_path, speaker_mapping) + save_file(speaker_mapping, mapping_file_path) print("Speaker embeddings saved at:", mapping_file_path) diff --git a/TTS/tts/models/base_tts.py b/TTS/tts/models/base_tts.py index c71872d3..c86bd391 100644 --- a/TTS/tts/models/base_tts.py +++ b/TTS/tts/models/base_tts.py @@ -407,16 +407,16 @@ class BaseTTS(BaseTrainerModel): return test_figures, test_audios def on_init_start(self, trainer): - """Save the speaker.json and language_ids.json at the beginning of the training. Also update both paths.""" + """Save the speaker.pth and language_ids.json at the beginning of the training. Also update both paths.""" if self.speaker_manager is not None: - output_path = os.path.join(trainer.output_path, "speakers.json") + output_path = os.path.join(trainer.output_path, "speakers.pth") self.speaker_manager.save_ids_to_file(output_path) trainer.config.speakers_file = output_path # some models don't have `model_args` set if hasattr(trainer.config, "model_args"): trainer.config.model_args.speakers_file = output_path trainer.config.save_json(os.path.join(trainer.output_path, "config.json")) - print(f" > `speakers.json` is saved to {output_path}.") + print(f" > `speakers.pth` is saved to {output_path}.") print(" > `speakers_file` is updated in the config.json.") if hasattr(self, "language_manager") and self.language_manager is not None: diff --git a/TTS/tts/utils/managers.py b/TTS/tts/utils/managers.py index 5415d52c..0243d3b4 100644 --- a/TTS/tts/utils/managers.py +++ b/TTS/tts/utils/managers.py @@ -11,6 +11,28 @@ from TTS.encoder.utils.generic_utils import setup_encoder_model from TTS.utils.audio import AudioProcessor +def load_file(path: str): + if path.endswith(".json"): + with fsspec.open(path, "r") as f: + return json.load(f) + elif path.endswith(".pth"): + with fsspec.open(path, "rb") as f: + return torch.load(f, map_location="cpu") + else: + raise ValueError("Unsupported file type") + + +def save_file(obj: Any, path: str): + if path.endswith(".json"): + with fsspec.open(path, "w") as f: + json.dump(obj, f, indent=4) + elif path.endswith(".pth"): + with fsspec.open(path, "wb") as f: + torch.save(obj, f) + else: + raise ValueError("Unsupported file type") + + class BaseIDManager: """Base `ID` Manager class. Every new `ID` manager must inherit this. It defines common `ID` manager specific functions. @@ -46,7 +68,7 @@ class BaseIDManager: Args: file_path (str): Path to the file. """ - self.ids = self._load_json(file_path) + self.ids = load_file(file_path) def save_ids_to_file(self, file_path: str) -> None: """Save IDs to a json file. @@ -54,7 +76,7 @@ class BaseIDManager: Args: file_path (str): Path to the output file. """ - self._save_json(file_path, self.ids) + save_file(self.ids, file_path) def get_random_id(self) -> Any: """Get a random embedding. @@ -125,7 +147,7 @@ class EmbeddingManager(BaseIDManager): Args: file_path (str): Path to the output file. """ - self._save_json(file_path, self.embeddings) + save_file(self.embeddings, file_path) def load_embeddings_from_file(self, file_path: str) -> None: """Load embeddings from a json file. @@ -133,7 +155,7 @@ class EmbeddingManager(BaseIDManager): Args: file_path (str): Path to the target json file. """ - self.embeddings = self._load_json(file_path) + self.embeddings = load_file(file_path) speakers = sorted({x["name"] for x in self.embeddings.values()}) self.ids = {name: i for i, name in enumerate(speakers)} diff --git a/tests/aux_tests/test_speaker_manager.py b/tests/aux_tests/test_speaker_manager.py index 7552e0a5..890cc023 100644 --- a/tests/aux_tests/test_speaker_manager.py +++ b/tests/aux_tests/test_speaker_manager.py @@ -16,6 +16,7 @@ encoder_model_path = os.path.join(get_tests_input_path(), "checkpoint_0.pth") sample_wav_path = os.path.join(get_tests_input_path(), "../data/ljspeech/wavs/LJ001-0001.wav") sample_wav_path2 = os.path.join(get_tests_input_path(), "../data/ljspeech/wavs/LJ001-0002.wav") d_vectors_file_path = os.path.join(get_tests_input_path(), "../data/dummy_speakers.json") +d_vectors_file_pth_path = os.path.join(get_tests_input_path(), "../data/dummy_speakers.pth") class SpeakerManagerTest(unittest.TestCase): @@ -58,12 +59,13 @@ class SpeakerManagerTest(unittest.TestCase): # remove dummy model os.remove(encoder_model_path) - @staticmethod - def test_speakers_file_processing(): + def test_speakers_file_processing(self): manager = SpeakerManager(d_vectors_file_path=d_vectors_file_path) - print(manager.num_speakers) - print(manager.embedding_dim) - print(manager.clip_ids) + self.assertEqual(manager.num_speakers, 1) + self.assertEqual(manager.embedding_dim, 256) + manager = SpeakerManager(d_vectors_file_path=d_vectors_file_pth_path) + self.assertEqual(manager.num_speakers, 1) + self.assertEqual(manager.embedding_dim, 256) d_vector = manager.get_embedding_by_clip(manager.clip_ids[0]) assert len(d_vector) == 256 d_vectors = manager.get_embeddings_by_name(manager.speaker_names[0]) diff --git a/tests/data/dummy_speakers.pth b/tests/data/dummy_speakers.pth new file mode 100644 index 00000000..4ba7f7bc Binary files /dev/null and b/tests/data/dummy_speakers.pth differ