before changes
This commit is contained in:
@@ -0,0 +1,191 @@
|
||||
#!/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)
|
||||
Reference in New Issue
Block a user