192 lines
8.3 KiB
Python
192 lines
8.3 KiB
Python
#!/usr/bin/env python3
|
|
# -*- coding: utf-8 -*-
|
|
|
|
import sys
|
|
import os
|
|
import argparse
|
|
|
|
# load coqui-ai/TTS libraries
|
|
from TTS.bin.resample import resample_files
|
|
from TTS.bin.compute_embeddings import compute_embeddings
|
|
|
|
import train_config as tc
|
|
|
|
|
|
#
|
|
# parse command line arguments
|
|
#
|
|
def parse_cmdline_args():
|
|
parser = argparse.ArgumentParser(
|
|
description = "Code to prepare voice dataset for multi-speaker baseline model training (i.e., adjust sampling rate and compute speaker embeddings).")
|
|
parser.add_argument("--dataset_preset", type = str, choices = ("VCTK", "LibriTTS_tc360", "DAPS", "POTION_Salut", "potion_voice_cloning"), required = True,
|
|
help = "Path the voice dataset archive (.zip, .tar.gz, .tgz, .tar.bz2, and .tbz are supported)")
|
|
parser.add_argument("--dataset_archive_path", type = str, required = True,
|
|
help = "Path the voice dataset archive (.zip, .tar.gz, .tgz, .tar.bz2, and .tbz are supported)")
|
|
parser.add_argument("--output_path", type = str, default = "results/datasets",
|
|
help = "Path to store augmented dataset")
|
|
parser.add_argument("--sampling_rate", type = int, default = 22050, choices = (16000, 22050, 32000, 48000), # 32k & 48k are untested
|
|
help = "Sampling rate for training run")
|
|
|
|
return parser.parse_args()
|
|
|
|
|
|
#
|
|
# utility functuion to extract archives (zip, tar, tgz, ...)
|
|
# - returns first entry in archive (typically the main directory name contained in the archive)
|
|
#
|
|
def extract_archive(archive_path, dest_path):
|
|
|
|
from zipfile import ZipFile
|
|
import tarfile
|
|
|
|
if archive_path.endswith('.zip'):
|
|
opener, getnames, mode = ZipFile, ZipFile.namelist, 'r'
|
|
|
|
elif (archive_path.endswith('.tar.gz')) or (archive_path.endswith('.tgz')):
|
|
opener, getnames, mode = tarfile.open, tarfile.TarFile.getnames, 'r:gz'
|
|
|
|
elif (archive_path.endswith('.tar.bz2')) or (archive_path.endswith('.tbz')):
|
|
opener, getnames, mode = tarfile.open, tarfile.TarFile.getnames, 'r:bz2'
|
|
|
|
else:
|
|
print("Extracting archive " + archive_path + " is not supported.")
|
|
return
|
|
|
|
# extract archive
|
|
with opener(archive_path, mode) as archive:
|
|
archive_dir = archive.getnames()[0]
|
|
archive.extractall(path = dest_path)
|
|
|
|
return archive_dir
|
|
|
|
|
|
#
|
|
# main training method (VITS multi-speaker model)
|
|
#
|
|
def main(args):
|
|
print("Commencing preparation of dataset for multi-speaker baseline model training:")
|
|
print("")
|
|
print(" + Dataset preset: {}" . format(args.dataset_preset))
|
|
print(" + Dataset : {}" . format(args.dataset_archive_path))
|
|
print(" + Output path : {}" . format(args.output_path))
|
|
print(" + Sampling rate : {}" . format(args.sampling_rate))
|
|
print("")
|
|
|
|
# set parameters according to dataset preset
|
|
if args.dataset_preset == "VCTK":
|
|
DATASET_NAME = tc.VCTK_DATASET_NAME
|
|
DATASET_FORMATTER = tc.VCTK_DATASET_FORMATTER
|
|
DATASET_FILE_FORMAT = tc.VCTK_DATASET_FILE_FORMAT
|
|
NO_EVAL = False
|
|
elif args.dataset_preset == "LibriTTS_tc360":
|
|
DATASET_NAME = tc.LIBRITTS_TC360_DATASET_NAME
|
|
DATASET_FORMATTER = tc.LIBRITTS_TC360_DATASET_FORMATTER
|
|
DATASET_FILE_FORMAT = tc.LIBRITTS_TC360_DATASET_FILE_FORMAT
|
|
NO_EVAL = False
|
|
elif args.dataset_preset == "POTION_Salut":
|
|
DATASET_NAME = tc.POTION_SALUT_DATASET_NAME
|
|
DATASET_FORMATTER = tc.POTION_SALUT_DATASET_FORMATTER
|
|
DATASET_FILE_FORMAT = tc.POTION_SALUT_DATASET_FILE_FORMAT
|
|
NO_EVAL = False
|
|
elif args.dataset_preset == "potion_voice_cloning":
|
|
DATASET_NAME = tc.POTION_SALUT_DATASET_NAME
|
|
DATASET_FORMATTER = tc.POTION_SALUT_DATASET_FORMATTER
|
|
DATASET_FILE_FORMAT = tc.POTION_SALUT_DATASET_FILE_FORMAT
|
|
NO_EVAL = True
|
|
|
|
# define sampling rate for computing speaker embeddings
|
|
SPK_EMB_SAMPLING_RATE = 16000
|
|
|
|
# define the number of threads used during audio resampling
|
|
NUM_RESAMPLE_THREADS = 10
|
|
|
|
# extract dataset archive
|
|
print(f">>> Extracting archive ...")
|
|
dataset_root = extract_archive(args.dataset_archive_path, os.path.join(args.output_path, "sr" + str(args.sampling_rate)))
|
|
|
|
# set dataset path (there should only be ONE directory in the extracted archive location)
|
|
dataset_path = os.path.join(args.output_path, "sr" + str(args.sampling_rate), dataset_root)
|
|
|
|
# ensure the dataset_path exists
|
|
os.makedirs(dataset_path, exist_ok = True)
|
|
|
|
# resample dataset for speaker embeddings computation
|
|
print(f">>> Resampling audio files to 16000Hz ...")
|
|
resample_files(dataset_path, 16000, file_ext = DATASET_FILE_FORMAT, n_jobs = NUM_RESAMPLE_THREADS)
|
|
|
|
# compute speaker embeddings
|
|
SPEAKER_ENCODER_CHECKPOINT_PATH = "assets/speaker_encoder_model/model_se.pth.tar"
|
|
SPEAKER_ENCODER_CONFIG_PATH = "assets/speaker_encoder_model/config_se.json"
|
|
|
|
# init list speaker embeddings/d-vectors to be used during the training
|
|
d_vector_files = []
|
|
|
|
# check if the speakers embeddings are already computated, if not compute them
|
|
embeddings_file = os.path.join(dataset_path, "speakers.pth")
|
|
|
|
if not os.path.isfile(embeddings_file):
|
|
print(f">>> Computing speaker embeddings ...")
|
|
compute_embeddings(
|
|
SPEAKER_ENCODER_CHECKPOINT_PATH,
|
|
SPEAKER_ENCODER_CONFIG_PATH,
|
|
embeddings_file,
|
|
old_spakers_file = None,
|
|
config_dataset_path = None,
|
|
formatter_name = DATASET_FORMATTER,
|
|
dataset_name = DATASET_NAME,
|
|
dataset_path = dataset_path,
|
|
meta_file_train = "",
|
|
meta_file_val = "",
|
|
disable_cuda = False,
|
|
no_eval = NO_EVAL
|
|
)
|
|
|
|
d_vector_files.append(embeddings_file)
|
|
|
|
# if targetted sampling rate is not the same as that used for computing speaker embeddings, replace and resample audio files
|
|
if not args.sampling_rate == SPK_EMB_SAMPLING_RATE:
|
|
print(f">>> Extracting original archive again (overwritting previously resampled files)...")
|
|
extract_archive(args.dataset_archive_path, os.path.join(args.output_path, "sr" + str(args.sampling_rate)))
|
|
print(f">>> Resampling audio files to {args.sampling_rate}Hz ...")
|
|
resample_files(dataset_path, args.sampling_rate, file_ext = DATASET_FILE_FORMAT, n_jobs = NUM_RESAMPLE_THREADS)
|
|
|
|
# exit gracefully
|
|
print("")
|
|
print("Completed preparing voice dataset for multi-speaker baseline model training; generated asset locations are as follows:")
|
|
print(" --> {}" . format(dataset_path))
|
|
print(" --> {}" . format(embeddings_file))
|
|
print("")
|
|
print("Done; bye.")
|
|
print("")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
# parse command line arguments
|
|
args = parse_cmdline_args()
|
|
|
|
# clear command line arguments to avoid triggering argparse features part of Trainer / coqpit imports
|
|
# Traceback (most recent call last):
|
|
# File "train_multispeaker_baseline_model.py", line 208, in <module>
|
|
# main(args)
|
|
# File "train_multispeaker_baseline_model.py", line 177, in main
|
|
# trainer = Trainer(
|
|
# File "/home/ubuntu/dev/potion-voice_venv/lib/python3.8/site-packages/trainer/trainer.py", line 360, in __init__
|
|
# config, new_fields = self.init_training(args, coqpit_overrides, config)
|
|
# File "/home/ubuntu/dev/potion-voice_venv/lib/python3.8/site-packages/trainer/trainer.py", line 594, in init_training
|
|
# config.parse_known_args(coqpit_overrides, relaxed_parser=True)
|
|
# File "/home/ubuntu/dev/potion-voice_venv/lib/python3.8/site-packages/coqpit/coqpit.py", line 843, in parse_known_args
|
|
# parser = self.init_argparse(arg_prefix=arg_prefix, relaxed_parser=relaxed_parser)
|
|
# File "/home/ubuntu/dev/potion-voice_venv/lib/python3.8/site-packages/coqpit/coqpit.py", line 881, in init_argparse
|
|
# _init_argparse(
|
|
# File "/home/ubuntu/dev/potion-voice_venv/lib/python3.8/site-packages/coqpit/coqpit.py", line 529, in _init_argparse
|
|
# parser = _init_argparse(
|
|
# File "/home/ubuntu/dev/potion-voice_venv/lib/python3.8/site-packages/coqpit/coqpit.py", line 550, in _init_argparse
|
|
# return default.init_argparse(
|
|
# AttributeError: 'str' object has no attribute 'init_argparse'
|
|
sys.argv = [sys.argv[0]]
|
|
|
|
# ensure the output path exists
|
|
os.makedirs(args.output_path, exist_ok = True)
|
|
|
|
main(args)
|