Files
2026-10-05 16:14:53 -04:00

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)