diff --git a/README.md b/README.md index 74504ea..cf69621 100644 --- a/README.md +++ b/README.md @@ -6,7 +6,8 @@ - Ver 2.1での学習をサポート(`train_ms_V210.py`) ## TODO -- [ ] Ver 2.2での学習をサポート +- [x] Ver 2.2での学習をサポート←たぶんやった、まだ確認してない +- [ ] Ver 2.1での感情のクラス数を10から少なくして実験 - [ ] Ver 2.1, 2.2での学習でのbf16対応 - [ ] 推論のWebUIでのバージョンに応じた感情指定のサポート - [ ] より良い推論WebUI? diff --git a/clap_gen.py b/clap_gen.py new file mode 100644 index 0000000..4895c70 --- /dev/null +++ b/clap_gen.py @@ -0,0 +1,64 @@ +import argparse +from multiprocessing import Pool, cpu_count + +import torch +import torch.multiprocessing as mp +from tqdm import tqdm + +import utils +from config import config +from oldVersion.V220.clap_wrapper import get_clap_audio_feature +import librosa +import os + +os.environ["OMP_NUM_THREADS"] = "1" +os.environ["MKL_NUM_THREADS"] = "1" + + +def process_line(line): + device = config.emo_gen_config.device + if config.emo_gen_config.use_multi_device: + rank = mp.current_process()._identity + rank = rank[0] if len(rank) > 0 else 0 + if torch.cuda.is_available(): + gpu_id = rank % torch.cuda.device_count() + device = torch.device(f"cuda:{gpu_id}") + else: + device = torch.device("cpu") + wav_path, _, language_str, text, phones, tone, word2ph = line.strip().split("|") + + clap_path = wav_path.replace(".WAV", ".wav").replace(".wav", ".emo.npy") + if os.path.isfile(clap_path): + return + + audio = librosa.load(wav_path, 48000)[0] + # audio = librosa.resample(audio, 44100, 48000) + + clap = get_clap_audio_feature(audio, device) + torch.save(clap, clap_path) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument( + "-c", "--config", type=str, default=config.emo_gen_config.config_path + ) + parser.add_argument( + "--num_processes", type=int, default=config.emo_gen_config.num_processes + ) + args, _ = parser.parse_known_args() + config_path = args.config + hps = utils.get_hparams_from_file(config_path) + lines = [] + with open(hps.data.training_files, encoding="utf-8") as f: + lines.extend(f.readlines()) + + with open(hps.data.validation_files, encoding="utf-8") as f: + lines.extend(f.readlines()) + if len(lines) != 0: + num_processes = min(args.num_processes, cpu_count()) + with Pool(processes=num_processes) as pool: + for _ in tqdm(pool.imap_unordered(process_line, lines), total=len(lines)): + pass + + print(f"clap生成完毕!, 共有{len(lines)}个emo.pt生成!") diff --git a/configs/config-V210.json b/configs/config-V210.json new file mode 100644 index 0000000..0134c8f --- /dev/null +++ b/configs/config-V210.json @@ -0,0 +1,63 @@ +{ + "train": { + "log_interval": 200, + "eval_interval": 1000, + "save_compressed_models": true, + "seed": 42, + "epochs": 100, + "learning_rate": 0.0002, + "betas": [0.8, 0.99], + "eps": 1e-9, + "batch_size": 4, + "fp16_run": false, + "lr_decay": 0.99995, + "segment_size": 16384, + "init_lr_ratio": 1, + "warmup_epochs": 0, + "c_mel": 45, + "c_kl": 1.0, + "skip_optimizer": true + }, + "data": { + "training_files": "Data/yksi/filelists/train.list", + "validation_files": "Data/yksi/filelists/val.list", + "max_wav_value": 32768.0, + "sampling_rate": 44100, + "filter_length": 2048, + "hop_length": 512, + "win_length": 2048, + "n_mel_channels": 128, + "mel_fmin": 0.0, + "mel_fmax": null, + "add_blank": true, + "n_speakers": 1, + "cleaned_text": true + }, + "model": { + "use_spk_conditioned_encoder": true, + "use_noise_scaled_mas": true, + "use_mel_posterior_encoder": false, + "use_duration_discriminator": true, + "inter_channels": 192, + "hidden_channels": 192, + "filter_channels": 768, + "n_heads": 2, + "n_layers": 6, + "kernel_size": 3, + "p_dropout": 0.1, + "resblock": "1", + "resblock_kernel_sizes": [3, 7, 11], + "resblock_dilation_sizes": [ + [1, 3, 5], + [1, 3, 5], + [1, 3, 5] + ], + "upsample_rates": [8, 8, 2, 2, 2], + "upsample_initial_channel": 512, + "upsample_kernel_sizes": [16, 16, 8, 2, 2], + "n_layers_q": 3, + "use_spectral_norm": false, + "gin_channels": 256 + }, + "version": "2.1" +} diff --git a/configs/config.json b/configs/config-V230.json similarity index 100% rename from configs/config.json rename to configs/config-V230.json diff --git a/data_utils.py b/data_utils.py index 04d9774..c692f26 100644 --- a/data_utils.py +++ b/data_utils.py @@ -301,7 +301,7 @@ class DistributedBucketSampler(torch.utils.data.distributed.DistributedSampler): self.buckets, self.num_samples_per_bucket = self._create_buckets() logger.info(f"Bucket info: {self.num_samples_per_bucket}") logger.info( - f"Unuseful samples: {len(self.lengths) - sum(self.num_samples_per_bucket)}" + f"Unused samples: {len(self.lengths) - sum(self.num_samples_per_bucket)}" ) self.total_size = sum(self.num_samples_per_bucket) self.num_samples = self.total_size // self.num_replicas diff --git a/empty_emo.npy b/empty_emo.npy new file mode 100644 index 0000000..6865293 Binary files /dev/null and b/empty_emo.npy differ diff --git a/oldVersion/V210/data_utils.py b/oldVersion/V210/data_utils.py index dc07f55..fe00f20 100644 --- a/oldVersion/V210/data_utils.py +++ b/oldVersion/V210/data_utils.py @@ -307,7 +307,7 @@ class DistributedBucketSampler(torch.utils.data.distributed.DistributedSampler): self.buckets, self.num_samples_per_bucket = self._create_buckets() logger.info(f"Bucket info: {self.num_samples_per_bucket}") logger.info( - f"Unuseful samples: {len(self.lengths) - sum(self.num_samples_per_bucket)}" + f"Unused samples: {len(self.lengths) - sum(self.num_samples_per_bucket)}" ) self.total_size = sum(self.num_samples_per_bucket) self.num_samples = self.total_size // self.num_replicas diff --git a/oldVersion/V220/data_utils.py b/oldVersion/V220/data_utils.py new file mode 100644 index 0000000..4bc0db4 --- /dev/null +++ b/oldVersion/V220/data_utils.py @@ -0,0 +1,425 @@ +import os +import random +import torch +import torch.utils.data +from tqdm import tqdm +import numpy as np +from tools.log import logger +import commons +from mel_processing import spectrogram_torch, mel_spectrogram_torch +from utils import load_wav_to_torch, load_filepaths_and_text +from text import cleaned_text_to_sequence +from config import config + +"""Multi speaker version""" + + +class TextAudioSpeakerLoader(torch.utils.data.Dataset): + """ + 1) loads audio, speaker_id, text pairs + 2) normalizes text and converts them to sequences of integers + 3) computes spectrograms from audio files. + """ + + def __init__(self, audiopaths_sid_text, hparams): + self.audiopaths_sid_text = load_filepaths_and_text(audiopaths_sid_text) + self.max_wav_value = hparams.max_wav_value + self.sampling_rate = hparams.sampling_rate + self.filter_length = hparams.filter_length + self.hop_length = hparams.hop_length + self.win_length = hparams.win_length + self.sampling_rate = hparams.sampling_rate + self.spk_map = hparams.spk2id + self.hparams = hparams + + self.use_mel_spec_posterior = getattr( + hparams, "use_mel_posterior_encoder", False + ) + if self.use_mel_spec_posterior: + self.n_mel_channels = getattr(hparams, "n_mel_channels", 80) + + self.cleaned_text = getattr(hparams, "cleaned_text", False) + + self.add_blank = hparams.add_blank + self.min_text_len = getattr(hparams, "min_text_len", 1) + self.max_text_len = getattr(hparams, "max_text_len", 384) + + self.empty_emo = torch.squeeze( + torch.load("empty_emo.npy", map_location="cpu"), dim=1 + ) + + random.seed(1234) + random.shuffle(self.audiopaths_sid_text) + self._filter() + + def _filter(self): + """ + Filter text & store spec lengths + """ + # Store spectrogram lengths for Bucketing + # wav_length ~= file_size / (wav_channels * Bytes per dim) = file_size / (1 * 2) + # spec_length = wav_length // hop_length + + audiopaths_sid_text_new = [] + lengths = [] + skipped = 0 + logger.info("Init dataset...") + for _id, spk, language, text, phones, tone, word2ph in tqdm( + self.audiopaths_sid_text + ): + audiopath = f"{_id}" + if self.min_text_len <= len(phones) and len(phones) <= self.max_text_len: + phones = phones.split(" ") + tone = [int(i) for i in tone.split(" ")] + word2ph = [int(i) for i in word2ph.split(" ")] + audiopaths_sid_text_new.append( + [audiopath, spk, language, text, phones, tone, word2ph] + ) + lengths.append(os.path.getsize(audiopath) // (2 * self.hop_length)) + else: + skipped += 1 + logger.info( + "skipped: " + + str(skipped) + + ", total: " + + str(len(self.audiopaths_sid_text)) + ) + self.audiopaths_sid_text = audiopaths_sid_text_new + self.lengths = lengths + + def get_audio_text_speaker_pair(self, audiopath_sid_text): + # separate filename, speaker_id and text + audiopath, sid, language, text, phones, tone, word2ph = audiopath_sid_text + + bert, ja_bert, en_bert, phones, tone, language = self.get_text( + text, word2ph, phones, tone, language, audiopath + ) + + spec, wav = self.get_audio(audiopath) + sid = torch.LongTensor([int(self.spk_map[sid])]) + + if np.random.rand() > 0.1: + emo = torch.squeeze( + torch.load(audiopath.replace(".wav", ".emo.npy"), map_location="cpu"), + dim=1, + ) + else: + emo = self.empty_emo + return (phones, spec, wav, sid, tone, language, bert, ja_bert, en_bert, emo) + + def get_audio(self, filename): + audio, sampling_rate = load_wav_to_torch(filename) + if sampling_rate != self.sampling_rate: + raise ValueError( + "{} {} SR doesn't match target {} SR".format( + filename, sampling_rate, self.sampling_rate + ) + ) + audio_norm = audio / self.max_wav_value + audio_norm = audio_norm.unsqueeze(0) + spec_filename = filename.replace(".wav", ".spec.pt") + if self.use_mel_spec_posterior: + spec_filename = spec_filename.replace(".spec.pt", ".mel.pt") + try: + spec = torch.load(spec_filename) + except: + if self.use_mel_spec_posterior: + spec = mel_spectrogram_torch( + audio_norm, + self.filter_length, + self.n_mel_channels, + self.sampling_rate, + self.hop_length, + self.win_length, + self.hparams.mel_fmin, + self.hparams.mel_fmax, + center=False, + ) + else: + spec = spectrogram_torch( + audio_norm, + self.filter_length, + self.sampling_rate, + self.hop_length, + self.win_length, + center=False, + ) + spec = torch.squeeze(spec, 0) + if config.train_ms_config.spec_cache: + torch.save(spec, spec_filename) + return spec, audio_norm + + def get_text(self, text, word2ph, phone, tone, language_str, wav_path): + phone, tone, language = cleaned_text_to_sequence(phone, tone, language_str) + if self.add_blank: + phone = commons.intersperse(phone, 0) + tone = commons.intersperse(tone, 0) + language = commons.intersperse(language, 0) + for i in range(len(word2ph)): + word2ph[i] = word2ph[i] * 2 + word2ph[0] += 1 + bert_path = wav_path.replace(".wav", ".bert.pt") + try: + bert_ori = torch.load(bert_path) + assert bert_ori.shape[-1] == len(phone) + except Exception as e: + logger.warning("Bert load Failed") + logger.warning(e) + + if language_str == "ZH": + bert = bert_ori + ja_bert = torch.rand(1024, len(phone)) + en_bert = torch.rand(1024, len(phone)) + elif language_str == "JP": + bert = torch.rand(1024, len(phone)) + ja_bert = bert_ori + en_bert = torch.rand(1024, len(phone)) + elif language_str == "EN": + bert = torch.rand(1024, len(phone)) + ja_bert = torch.rand(1024, len(phone)) + en_bert = bert_ori + phone = torch.LongTensor(phone) + tone = torch.LongTensor(tone) + language = torch.LongTensor(language) + return bert, ja_bert, en_bert, phone, tone, language + + def get_sid(self, sid): + sid = torch.LongTensor([int(sid)]) + return sid + + def __getitem__(self, index): + return self.get_audio_text_speaker_pair(self.audiopaths_sid_text[index]) + + def __len__(self): + return len(self.audiopaths_sid_text) + + +class TextAudioSpeakerCollate: + """Zero-pads model inputs and targets""" + + def __init__(self, return_ids=False): + self.return_ids = return_ids + + def __call__(self, batch): + """Collate's training batch from normalized text, audio and speaker identities + PARAMS + ------ + batch: [text_normalized, spec_normalized, wav_normalized, sid] + """ + # Right zero-pad all one-hot text sequences to max input length + _, ids_sorted_decreasing = torch.sort( + torch.LongTensor([x[1].size(1) for x in batch]), dim=0, descending=True + ) + + max_text_len = max([len(x[0]) for x in batch]) + max_spec_len = max([x[1].size(1) for x in batch]) + max_wav_len = max([x[2].size(1) for x in batch]) + + text_lengths = torch.LongTensor(len(batch)) + spec_lengths = torch.LongTensor(len(batch)) + wav_lengths = torch.LongTensor(len(batch)) + sid = torch.LongTensor(len(batch)) + + text_padded = torch.LongTensor(len(batch), max_text_len) + tone_padded = torch.LongTensor(len(batch), max_text_len) + language_padded = torch.LongTensor(len(batch), max_text_len) + bert_padded = torch.FloatTensor(len(batch), 1024, max_text_len) + ja_bert_padded = torch.FloatTensor(len(batch), 1024, max_text_len) + en_bert_padded = torch.FloatTensor(len(batch), 1024, max_text_len) + emo = torch.FloatTensor(len(batch), 512) + + spec_padded = torch.FloatTensor(len(batch), batch[0][1].size(0), max_spec_len) + wav_padded = torch.FloatTensor(len(batch), 1, max_wav_len) + text_padded.zero_() + tone_padded.zero_() + language_padded.zero_() + spec_padded.zero_() + wav_padded.zero_() + bert_padded.zero_() + ja_bert_padded.zero_() + en_bert_padded.zero_() + emo.zero_() + + for i in range(len(ids_sorted_decreasing)): + row = batch[ids_sorted_decreasing[i]] + + text = row[0] + text_padded[i, : text.size(0)] = text + text_lengths[i] = text.size(0) + + spec = row[1] + spec_padded[i, :, : spec.size(1)] = spec + spec_lengths[i] = spec.size(1) + + wav = row[2] + wav_padded[i, :, : wav.size(1)] = wav + wav_lengths[i] = wav.size(1) + + sid[i] = row[3] + + tone = row[4] + tone_padded[i, : tone.size(0)] = tone + + language = row[5] + language_padded[i, : language.size(0)] = language + + bert = row[6] + bert_padded[i, :, : bert.size(1)] = bert + + ja_bert = row[7] + ja_bert_padded[i, :, : ja_bert.size(1)] = ja_bert + + en_bert = row[8] + en_bert_padded[i, :, : en_bert.size(1)] = en_bert + + emo[i, :] = row[9] + + return ( + text_padded, + text_lengths, + spec_padded, + spec_lengths, + wav_padded, + wav_lengths, + sid, + tone_padded, + language_padded, + bert_padded, + ja_bert_padded, + en_bert_padded, + emo, + ) + + +class DistributedBucketSampler(torch.utils.data.distributed.DistributedSampler): + """ + Maintain similar input lengths in a batch. + Length groups are specified by boundaries. + Ex) boundaries = [b1, b2, b3] -> any batch is included either {x | b1 < length(x) <=b2} or {x | b2 < length(x) <= b3}. + + It removes samples which are not included in the boundaries. + Ex) boundaries = [b1, b2, b3] -> any x s.t. length(x) <= b1 or length(x) > b3 are discarded. + """ + + def __init__( + self, + dataset, + batch_size, + boundaries, + num_replicas=None, + rank=None, + shuffle=True, + ): + super().__init__(dataset, num_replicas=num_replicas, rank=rank, shuffle=shuffle) + self.lengths = dataset.lengths + self.batch_size = batch_size + self.boundaries = boundaries + + self.buckets, self.num_samples_per_bucket = self._create_buckets() + logger.info(f"Bucket info: {self.num_samples_per_bucket}") + logger.info( + f"Unused samples: {len(self.lengths) - sum(self.num_samples_per_bucket)}" + ) + self.total_size = sum(self.num_samples_per_bucket) + self.num_samples = self.total_size // self.num_replicas + + def _create_buckets(self): + buckets = [[] for _ in range(len(self.boundaries) - 1)] + for i in range(len(self.lengths)): + length = self.lengths[i] + idx_bucket = self._bisect(length) + if idx_bucket != -1: + buckets[idx_bucket].append(i) + + try: + for i in range(len(buckets) - 1, 0, -1): + if len(buckets[i]) == 0: + buckets.pop(i) + self.boundaries.pop(i + 1) + assert all(len(bucket) > 0 for bucket in buckets) + # When one bucket is not traversed + except Exception as e: + print("Bucket warning ", e) + for i in range(len(buckets) - 1, -1, -1): + if len(buckets[i]) == 0: + buckets.pop(i) + self.boundaries.pop(i + 1) + + num_samples_per_bucket = [] + for i in range(len(buckets)): + len_bucket = len(buckets[i]) + total_batch_size = self.num_replicas * self.batch_size + rem = ( + total_batch_size - (len_bucket % total_batch_size) + ) % total_batch_size + num_samples_per_bucket.append(len_bucket + rem) + return buckets, num_samples_per_bucket + + def __iter__(self): + # deterministically shuffle based on epoch + g = torch.Generator() + g.manual_seed(self.epoch) + + indices = [] + if self.shuffle: + for bucket in self.buckets: + indices.append(torch.randperm(len(bucket), generator=g).tolist()) + else: + for bucket in self.buckets: + indices.append(list(range(len(bucket)))) + + batches = [] + for i in range(len(self.buckets)): + bucket = self.buckets[i] + len_bucket = len(bucket) + if len_bucket == 0: + continue + ids_bucket = indices[i] + num_samples_bucket = self.num_samples_per_bucket[i] + + # add extra samples to make it evenly divisible + rem = num_samples_bucket - len_bucket + ids_bucket = ( + ids_bucket + + ids_bucket * (rem // len_bucket) + + ids_bucket[: (rem % len_bucket)] + ) + + # subsample + ids_bucket = ids_bucket[self.rank :: self.num_replicas] + + # batching + for j in range(len(ids_bucket) // self.batch_size): + batch = [ + bucket[idx] + for idx in ids_bucket[ + j * self.batch_size : (j + 1) * self.batch_size + ] + ] + batches.append(batch) + + if self.shuffle: + batch_ids = torch.randperm(len(batches), generator=g).tolist() + batches = [batches[i] for i in batch_ids] + self.batches = batches + + assert len(self.batches) * self.batch_size == self.num_samples + return iter(self.batches) + + def _bisect(self, x, lo=0, hi=None): + if hi is None: + hi = len(self.boundaries) - 1 + + if hi > lo: + mid = (hi + lo) // 2 + if self.boundaries[mid] < x and x <= self.boundaries[mid + 1]: + return mid + elif x <= self.boundaries[mid]: + return self._bisect(x, lo, mid) + else: + return self._bisect(x, mid + 1, hi) + else: + return -1 + + def __len__(self): + return self.num_samples // self.batch_size diff --git a/train_ms.py b/train_ms.py index da57c9a..16dfc80 100644 --- a/train_ms.py +++ b/train_ms.py @@ -740,8 +740,10 @@ def train_and_evaluate( n_ckpts_to_keep=keep_ckpts, sort_by_time=True, ) - save_compressed_models = hps.train.save_compressed_models - if save_compressed_models: + if ( + "save_compressed_models" in hps.train.keys() + and hps.train.save_compressed_models is True + ): utils.save_compressed_models_checkpoint( net_g, epoch, diff --git a/train_ms_V210.py b/train_ms_V210.py index a2d3171..66c35bd 100644 --- a/train_ms_V210.py +++ b/train_ms_V210.py @@ -329,6 +329,16 @@ def run(): if net_dur_disc is not None: scheduler_dur_disc.step() + if epoch == hps.train.epochs: + utils.save_compressed_models_checkpoint( + net_g, + epoch, + os.path.join( + hps.model_dir, + f"release_{global_step}.pth", + ), + ) + def train_and_evaluate( rank, @@ -589,8 +599,10 @@ def train_and_evaluate( n_ckpts_to_keep=keep_ckpts, sort_by_time=True, ) - save_compressed_models = hps.train.save_compressed_models - if save_compressed_models: + if ( + "save_compressed_models" in hps.train.keys() + and hps.train.save_compressed_models is True + ): utils.save_compressed_models_checkpoint( net_g, epoch, diff --git a/train_ms_V220.py b/train_ms_V220.py new file mode 100644 index 0000000..1fc579e --- /dev/null +++ b/train_ms_V220.py @@ -0,0 +1,745 @@ +# flake8: noqa: E402 +import platform +import os +import torch +from torch.nn import functional as F +from torch.utils.data import DataLoader +from torch.utils.tensorboard import SummaryWriter +import torch.distributed as dist +from torch.nn.parallel import DistributedDataParallel as DDP +from torch.cuda.amp import autocast, GradScaler +from tqdm import tqdm +import logging +from config import config +import argparse +import datetime +import gc + +logging.getLogger("numba").setLevel(logging.WARNING) +import commons +import utils +from oldVersion.V220.data_utils import ( + TextAudioSpeakerLoader, + TextAudioSpeakerCollate, + DistributedBucketSampler, +) +from oldVersion.V220.models import ( + SynthesizerTrn, + MultiPeriodDiscriminator, + DurationDiscriminator, +) +from losses import generator_loss, discriminator_loss, feature_loss, kl_loss +from mel_processing import mel_spectrogram_torch, spec_to_mel_torch +from oldVersion.V220.text.symbols import symbols + +torch.backends.cuda.matmul.allow_tf32 = True +torch.backends.cudnn.allow_tf32 = ( + True # If encontered training problem,please try to disable TF32. +) +torch.set_float32_matmul_precision("medium") +torch.backends.cuda.sdp_kernel("flash") +torch.backends.cuda.enable_flash_sdp(True) +torch.backends.cuda.enable_mem_efficient_sdp( + True +) # Not available if torch version is lower than 2.0 +torch.backends.cuda.enable_math_sdp(True) +global_step = 0 + + +def run(): + # 环境变量解析 + envs = config.train_ms_config.env + for env_name, env_value in envs.items(): + if env_name not in os.environ.keys(): + print("加载config中的配置{}".format(str(env_value))) + os.environ[env_name] = str(env_value) + print( + "加载环境变量 \nMASTER_ADDR: {},\nMASTER_PORT: {},\nWORLD_SIZE: {},\nRANK: {},\nLOCAL_RANK: {}".format( + os.environ["MASTER_ADDR"], + os.environ["MASTER_PORT"], + os.environ["WORLD_SIZE"], + os.environ["RANK"], + os.environ["LOCAL_RANK"], + ) + ) + + backend = "nccl" + if platform.system() == "Windows": + backend = "gloo" # If Windows,switch to gloo backend. + dist.init_process_group( + backend=backend, + init_method="env://", + timeout=datetime.timedelta(seconds=300), + ) # Use torchrun instead of mp.spawn + rank = dist.get_rank() + local_rank = int(os.environ["LOCAL_RANK"]) + n_gpus = dist.get_world_size() + + # 命令行/config.yml配置解析 + # hps = utils.get_hparams() + parser = argparse.ArgumentParser() + # 非必要不建议使用命令行配置,请使用config.yml文件 + parser.add_argument( + "-c", + "--config", + type=str, + default=config.train_ms_config.config_path, + help="JSON file for configuration", + ) + + parser.add_argument( + "-m", + "--model", + type=str, + help="数据集文件夹路径,请注意,数据不再默认放在/logs文件夹下。如果需要用命令行配置,请声明相对于根目录的路径", + default=config.dataset_path, + ) + args = parser.parse_args() + model_dir = os.path.join(args.model, config.train_ms_config.model) + if not os.path.exists(model_dir): + os.makedirs(model_dir) + hps = utils.get_hparams_from_file(args.config) + hps.model_dir = model_dir + # 比较路径是否相同 + if os.path.realpath(args.config) != os.path.realpath( + config.train_ms_config.config_path + ): + with open(args.config, "r", encoding="utf-8") as f: + data = f.read() + with open(config.train_ms_config.config_path, "w", encoding="utf-8") as f: + f.write(data) + + torch.manual_seed(hps.train.seed) + torch.cuda.set_device(local_rank) + + global global_step + if rank == 0: + logger = utils.get_logger(hps.model_dir) + logger.info(hps) + utils.check_git_hash(hps.model_dir) + writer = SummaryWriter(log_dir=hps.model_dir) + writer_eval = SummaryWriter(log_dir=os.path.join(hps.model_dir, "eval")) + train_dataset = TextAudioSpeakerLoader(hps.data.training_files, hps.data) + train_sampler = DistributedBucketSampler( + train_dataset, + hps.train.batch_size, + [32, 300, 400, 500, 600, 700, 800, 900, 1000], + num_replicas=n_gpus, + rank=rank, + shuffle=True, + ) + collate_fn = TextAudioSpeakerCollate() + train_loader = DataLoader( + train_dataset, + num_workers=min(config.train_ms_config.num_workers, os.cpu_count() - 1), + shuffle=False, + pin_memory=True, + collate_fn=collate_fn, + batch_sampler=train_sampler, + persistent_workers=True, + prefetch_factor=4, + ) # DataLoader config could be adjusted. + if rank == 0: + eval_dataset = TextAudioSpeakerLoader(hps.data.validation_files, hps.data) + eval_loader = DataLoader( + eval_dataset, + num_workers=0, + shuffle=False, + batch_size=1, + pin_memory=True, + drop_last=False, + collate_fn=collate_fn, + ) + if ( + "use_noise_scaled_mas" in hps.model.keys() + and hps.model.use_noise_scaled_mas is True + ): + print("Using noise scaled MAS for VITS2") + mas_noise_scale_initial = 0.01 + noise_scale_delta = 2e-6 + else: + print("Using normal MAS for VITS1") + mas_noise_scale_initial = 0.0 + noise_scale_delta = 0.0 + if ( + "use_duration_discriminator" in hps.model.keys() + and hps.model.use_duration_discriminator is True + ): + print("Using duration discriminator for VITS2") + net_dur_disc = DurationDiscriminator( + hps.model.hidden_channels, + hps.model.hidden_channels, + 3, + 0.1, + gin_channels=hps.model.gin_channels if hps.data.n_speakers != 0 else 0, + ).cuda(local_rank) + if ( + "use_spk_conditioned_encoder" in hps.model.keys() + and hps.model.use_spk_conditioned_encoder is True + ): + if hps.data.n_speakers == 0: + raise ValueError( + "n_speakers must be > 0 when using spk conditioned encoder to train multi-speaker model" + ) + else: + print("Using normal encoder for VITS1") + + net_g = SynthesizerTrn( + len(symbols), + hps.data.filter_length // 2 + 1, + hps.train.segment_size // hps.data.hop_length, + n_speakers=hps.data.n_speakers, + mas_noise_scale_initial=mas_noise_scale_initial, + noise_scale_delta=noise_scale_delta, + **hps.model, + ).cuda(local_rank) + + if getattr(hps.train, "freeze_ZH_bert", False): + print("Freezing ZH bert encoder !!!") + for param in net_g.enc_p.bert_proj.parameters(): + param.requires_grad = False + + if getattr(hps.train, "freeze_EN_bert", False): + print("Freezing EN bert encoder !!!") + for param in net_g.enc_p.en_bert_proj.parameters(): + param.requires_grad = False + + if getattr(hps.train, "freeze_JP_bert", False): + print("Freezing JP bert encoder !!!") + for param in net_g.enc_p.ja_bert_proj.parameters(): + param.requires_grad = False + + net_d = MultiPeriodDiscriminator(hps.model.use_spectral_norm).cuda(local_rank) + optim_g = torch.optim.AdamW( + filter(lambda p: p.requires_grad, net_g.parameters()), + hps.train.learning_rate, + betas=hps.train.betas, + eps=hps.train.eps, + ) + optim_d = torch.optim.AdamW( + net_d.parameters(), + hps.train.learning_rate, + betas=hps.train.betas, + eps=hps.train.eps, + ) + if net_dur_disc is not None: + optim_dur_disc = torch.optim.AdamW( + net_dur_disc.parameters(), + hps.train.learning_rate, + betas=hps.train.betas, + eps=hps.train.eps, + ) + else: + optim_dur_disc = None + net_g = DDP(net_g, device_ids=[local_rank], bucket_cap_mb=512) + net_d = DDP(net_d, device_ids=[local_rank], bucket_cap_mb=512) + dur_resume_lr = None + if net_dur_disc is not None: + net_dur_disc = DDP( + net_dur_disc, + device_ids=[local_rank], + find_unused_parameters=True, + bucket_cap_mb=512, + ) + + # 下载底模 + if config.train_ms_config.base["use_base_model"]: + utils.download_checkpoint( + hps.model_dir, + config.train_ms_config.base, + token=config.openi_token, + mirror=config.mirror, + ) + + try: + if net_dur_disc is not None: + _, _, dur_resume_lr, epoch_str = utils.load_checkpoint( + utils.latest_checkpoint_path(hps.model_dir, "DUR_*.pth"), + net_dur_disc, + optim_dur_disc, + skip_optimizer=hps.train.skip_optimizer + if "skip_optimizer" in hps.train + else True, + ) + _, optim_g, g_resume_lr, epoch_str = utils.load_checkpoint( + utils.latest_checkpoint_path(hps.model_dir, "G_*.pth"), + net_g, + optim_g, + skip_optimizer=hps.train.skip_optimizer + if "skip_optimizer" in hps.train + else True, + ) + _, optim_d, d_resume_lr, epoch_str = utils.load_checkpoint( + utils.latest_checkpoint_path(hps.model_dir, "D_*.pth"), + net_d, + optim_d, + skip_optimizer=hps.train.skip_optimizer + if "skip_optimizer" in hps.train + else True, + ) + if not optim_g.param_groups[0].get("initial_lr"): + optim_g.param_groups[0]["initial_lr"] = g_resume_lr + if not optim_d.param_groups[0].get("initial_lr"): + optim_d.param_groups[0]["initial_lr"] = d_resume_lr + if not optim_dur_disc.param_groups[0].get("initial_lr"): + optim_dur_disc.param_groups[0]["initial_lr"] = dur_resume_lr + + epoch_str = max(epoch_str, 1) + # global_step = (epoch_str - 1) * len(train_loader) + global_step = int( + utils.get_steps(utils.latest_checkpoint_path(hps.model_dir, "G_*.pth")) + ) + print( + f"******************检测到模型存在,epoch为 {epoch_str},gloabl step为 {global_step}*********************" + ) + except Exception as e: + print(e) + epoch_str = 1 + global_step = 0 + + scheduler_g = torch.optim.lr_scheduler.ExponentialLR( + optim_g, gamma=hps.train.lr_decay, last_epoch=epoch_str - 2 + ) + scheduler_d = torch.optim.lr_scheduler.ExponentialLR( + optim_d, gamma=hps.train.lr_decay, last_epoch=epoch_str - 2 + ) + if net_dur_disc is not None: + if not optim_dur_disc.param_groups[0].get("initial_lr"): + optim_dur_disc.param_groups[0]["initial_lr"] = dur_resume_lr + scheduler_dur_disc = torch.optim.lr_scheduler.ExponentialLR( + optim_dur_disc, gamma=hps.train.lr_decay, last_epoch=epoch_str - 2 + ) + else: + scheduler_dur_disc = None + scaler = GradScaler(enabled=hps.train.fp16_run) + + for epoch in range(epoch_str, hps.train.epochs + 1): + if rank == 0: + train_and_evaluate( + rank, + local_rank, + epoch, + hps, + [net_g, net_d, net_dur_disc], + [optim_g, optim_d, optim_dur_disc], + [scheduler_g, scheduler_d, scheduler_dur_disc], + scaler, + [train_loader, eval_loader], + logger, + [writer, writer_eval], + ) + else: + train_and_evaluate( + rank, + local_rank, + epoch, + hps, + [net_g, net_d, net_dur_disc], + [optim_g, optim_d, optim_dur_disc], + [scheduler_g, scheduler_d, scheduler_dur_disc], + scaler, + [train_loader, None], + None, + None, + ) + scheduler_g.step() + scheduler_d.step() + if net_dur_disc is not None: + scheduler_dur_disc.step() + if epoch == hps.train.epochs: + utils.save_compressed_models_checkpoint( + net_g, + epoch, + os.path.join( + hps.model_dir, + f"release_{global_step}.pth", + ), + ) + + +def train_and_evaluate( + rank, + local_rank, + epoch, + hps, + nets, + optims, + schedulers, + scaler, + loaders, + logger, + writers, +): + net_g, net_d, net_dur_disc = nets + optim_g, optim_d, optim_dur_disc = optims + scheduler_g, scheduler_d, scheduler_dur_disc = schedulers + train_loader, eval_loader = loaders + if writers is not None: + writer, writer_eval = writers + + train_loader.batch_sampler.set_epoch(epoch) + global global_step + + net_g.train() + net_d.train() + if net_dur_disc is not None: + net_dur_disc.train() + for batch_idx, ( + x, + x_lengths, + spec, + spec_lengths, + y, + y_lengths, + speakers, + tone, + language, + bert, + ja_bert, + en_bert, + emo, + ) in enumerate(tqdm(train_loader)): + if net_g.module.use_noise_scaled_mas: + current_mas_noise_scale = ( + net_g.module.mas_noise_scale_initial + - net_g.module.noise_scale_delta * global_step + ) + net_g.module.current_mas_noise_scale = max(current_mas_noise_scale, 0.0) + x, x_lengths = x.cuda(local_rank, non_blocking=True), x_lengths.cuda( + local_rank, non_blocking=True + ) + spec, spec_lengths = spec.cuda( + local_rank, non_blocking=True + ), spec_lengths.cuda(local_rank, non_blocking=True) + y, y_lengths = y.cuda(local_rank, non_blocking=True), y_lengths.cuda( + local_rank, non_blocking=True + ) + speakers = speakers.cuda(local_rank, non_blocking=True) + tone = tone.cuda(local_rank, non_blocking=True) + language = language.cuda(local_rank, non_blocking=True) + bert = bert.cuda(local_rank, non_blocking=True) + ja_bert = ja_bert.cuda(local_rank, non_blocking=True) + en_bert = en_bert.cuda(local_rank, non_blocking=True) + emo = emo.cuda(local_rank, non_blocking=True) + + with autocast(enabled=hps.train.fp16_run): + ( + y_hat, + l_length, + attn, + ids_slice, + x_mask, + z_mask, + (z, z_p, m_p, logs_p, m_q, logs_q), + (hidden_x, logw, logw_), + g, + loss_commit, + ) = net_g( + x, + x_lengths, + spec, + spec_lengths, + speakers, + tone, + language, + bert, + ja_bert, + en_bert, + emo, + ) + mel = spec_to_mel_torch( + spec, + hps.data.filter_length, + hps.data.n_mel_channels, + hps.data.sampling_rate, + hps.data.mel_fmin, + hps.data.mel_fmax, + ) + y_mel = commons.slice_segments( + mel, ids_slice, hps.train.segment_size // hps.data.hop_length + ) + y_hat_mel = mel_spectrogram_torch( + y_hat.squeeze(1), + hps.data.filter_length, + hps.data.n_mel_channels, + hps.data.sampling_rate, + hps.data.hop_length, + hps.data.win_length, + hps.data.mel_fmin, + hps.data.mel_fmax, + ) + + y = commons.slice_segments( + y, ids_slice * hps.data.hop_length, hps.train.segment_size + ) # slice + + # Discriminator + y_d_hat_r, y_d_hat_g, _, _ = net_d(y, y_hat.detach()) + with autocast(enabled=False): + loss_disc, losses_disc_r, losses_disc_g = discriminator_loss( + y_d_hat_r, y_d_hat_g + ) + loss_disc_all = loss_disc + if net_dur_disc is not None: + y_dur_hat_r, y_dur_hat_g = net_dur_disc( + hidden_x.detach(), + x_mask.detach(), + logw.detach(), + logw_.detach(), + g.detach(), + ) + with autocast(enabled=False): + # TODO: I think need to mean using the mask, but for now, just mean all + ( + loss_dur_disc, + losses_dur_disc_r, + losses_dur_disc_g, + ) = discriminator_loss(y_dur_hat_r, y_dur_hat_g) + loss_dur_disc_all = loss_dur_disc + optim_dur_disc.zero_grad() + scaler.scale(loss_dur_disc_all).backward() + scaler.unscale_(optim_dur_disc) + commons.clip_grad_value_(net_dur_disc.parameters(), None) + scaler.step(optim_dur_disc) + + optim_d.zero_grad() + scaler.scale(loss_disc_all).backward() + scaler.unscale_(optim_d) + grad_norm_d = commons.clip_grad_value_(net_d.parameters(), None) + scaler.step(optim_d) + + with autocast(enabled=hps.train.fp16_run): + # Generator + y_d_hat_r, y_d_hat_g, fmap_r, fmap_g = net_d(y, y_hat) + if net_dur_disc is not None: + y_dur_hat_r, y_dur_hat_g = net_dur_disc( + hidden_x, x_mask, logw, logw_, g + ) + with autocast(enabled=False): + loss_dur = torch.sum(l_length.float()) + loss_mel = F.l1_loss(y_mel, y_hat_mel) * hps.train.c_mel + loss_kl = kl_loss(z_p, logs_q, m_p, logs_p, z_mask) * hps.train.c_kl + + loss_fm = feature_loss(fmap_r, fmap_g) + loss_gen, losses_gen = generator_loss(y_d_hat_g) + loss_gen_all = ( + loss_gen + loss_fm + loss_mel + loss_dur + loss_kl + loss_commit + ) + if net_dur_disc is not None: + loss_dur_gen, losses_dur_gen = generator_loss(y_dur_hat_g) + loss_gen_all += loss_dur_gen + optim_g.zero_grad() + scaler.scale(loss_gen_all).backward() + scaler.unscale_(optim_g) + grad_norm_g = commons.clip_grad_value_(net_g.parameters(), None) + scaler.step(optim_g) + scaler.update() + + if rank == 0: + if global_step % hps.train.log_interval == 0: + lr = optim_g.param_groups[0]["lr"] + losses = [loss_disc, loss_gen, loss_fm, loss_mel, loss_dur, loss_kl] + logger.info( + "Train Epoch: {} [{:.0f}%]".format( + epoch, 100.0 * batch_idx / len(train_loader) + ) + ) + logger.info([x.item() for x in losses] + [global_step, lr]) + + scalar_dict = { + "loss/g/total": loss_gen_all, + "loss/d/total": loss_disc_all, + "learning_rate": lr, + "grad_norm_d": grad_norm_d, + "grad_norm_g": grad_norm_g, + } + scalar_dict.update( + { + "loss/g/fm": loss_fm, + "loss/g/mel": loss_mel, + "loss/g/dur": loss_dur, + "loss/g/kl": loss_kl, + } + ) + scalar_dict.update( + {"loss/g/{}".format(i): v for i, v in enumerate(losses_gen)} + ) + scalar_dict.update( + {"loss/d_r/{}".format(i): v for i, v in enumerate(losses_disc_r)} + ) + scalar_dict.update( + {"loss/d_g/{}".format(i): v for i, v in enumerate(losses_disc_g)} + ) + + image_dict = { + "slice/mel_org": utils.plot_spectrogram_to_numpy( + y_mel[0].data.cpu().numpy() + ), + "slice/mel_gen": utils.plot_spectrogram_to_numpy( + y_hat_mel[0].data.cpu().numpy() + ), + "all/mel": utils.plot_spectrogram_to_numpy( + mel[0].data.cpu().numpy() + ), + "all/attn": utils.plot_alignment_to_numpy( + attn[0, 0].data.cpu().numpy() + ), + } + utils.summarize( + writer=writer, + global_step=global_step, + images=image_dict, + scalars=scalar_dict, + ) + + if global_step % hps.train.eval_interval == 0: + evaluate(hps, net_g, eval_loader, writer_eval) + utils.save_checkpoint( + net_g, + optim_g, + hps.train.learning_rate, + epoch, + os.path.join(hps.model_dir, "G_{}.pth".format(global_step)), + ) + utils.save_checkpoint( + net_d, + optim_d, + hps.train.learning_rate, + epoch, + os.path.join(hps.model_dir, "D_{}.pth".format(global_step)), + ) + if net_dur_disc is not None: + utils.save_checkpoint( + net_dur_disc, + optim_dur_disc, + hps.train.learning_rate, + epoch, + os.path.join(hps.model_dir, "DUR_{}.pth".format(global_step)), + ) + keep_ckpts = config.train_ms_config.keep_ckpts + if keep_ckpts > 0: + utils.clean_checkpoints( + path_to_models=hps.model_dir, + n_ckpts_to_keep=keep_ckpts, + sort_by_time=True, + ) + save_compressed_models = hps.train.save_compressed_models + if save_compressed_models: + utils.save_compressed_models_checkpoint( + net_g, + epoch, + os.path.join( + hps.model_dir, + f"release_{global_step}.pth", + ), + ) + + global_step += 1 + + gc.collect() + torch.cuda.empty_cache() + if rank == 0: + logger.info("====> Epoch: {}".format(epoch)) + + +def evaluate(hps, generator, eval_loader, writer_eval): + generator.eval() + image_dict = {} + audio_dict = {} + print("Evaluating ...") + with torch.no_grad(): + for batch_idx, ( + x, + x_lengths, + spec, + spec_lengths, + y, + y_lengths, + speakers, + tone, + language, + bert, + ja_bert, + en_bert, + emo, + ) in enumerate(eval_loader): + x, x_lengths = x.cuda(), x_lengths.cuda() + spec, spec_lengths = spec.cuda(), spec_lengths.cuda() + y, y_lengths = y.cuda(), y_lengths.cuda() + speakers = speakers.cuda() + bert = bert.cuda() + ja_bert = ja_bert.cuda() + en_bert = en_bert.cuda() + tone = tone.cuda() + language = language.cuda() + emo = emo.cuda() + for use_sdp in [True, False]: + y_hat, attn, mask, *_ = generator.module.infer( + x, + x_lengths, + speakers, + tone, + language, + bert, + ja_bert, + en_bert, + emo, + y=spec, + max_len=1000, + sdp_ratio=0.0 if not use_sdp else 1.0, + ) + y_hat_lengths = mask.sum([1, 2]).long() * hps.data.hop_length + + mel = spec_to_mel_torch( + spec, + hps.data.filter_length, + hps.data.n_mel_channels, + hps.data.sampling_rate, + hps.data.mel_fmin, + hps.data.mel_fmax, + ) + y_hat_mel = mel_spectrogram_torch( + y_hat.squeeze(1).float(), + hps.data.filter_length, + hps.data.n_mel_channels, + hps.data.sampling_rate, + hps.data.hop_length, + hps.data.win_length, + hps.data.mel_fmin, + hps.data.mel_fmax, + ) + image_dict.update( + { + f"gen/mel_{batch_idx}": utils.plot_spectrogram_to_numpy( + y_hat_mel[0].cpu().numpy() + ) + } + ) + audio_dict.update( + { + f"gen/audio_{batch_idx}_{use_sdp}": y_hat[ + 0, :, : y_hat_lengths[0] + ] + } + ) + image_dict.update( + { + f"gt/mel_{batch_idx}": utils.plot_spectrogram_to_numpy( + mel[0].cpu().numpy() + ) + } + ) + audio_dict.update({f"gt/audio_{batch_idx}": y[0, :, : y_lengths[0]]}) + + utils.summarize( + writer=writer_eval, + global_step=global_step, + images=image_dict, + audios=audio_dict, + audio_sampling_rate=hps.data.sampling_rate, + ) + generator.train() + + +if __name__ == "__main__": + run()