Feat: support V220 train (maybe)
This commit is contained in:
@@ -6,7 +6,8 @@
|
|||||||
- Ver 2.1での学習をサポート(`train_ms_V210.py`)
|
- Ver 2.1での学習をサポート(`train_ms_V210.py`)
|
||||||
|
|
||||||
## TODO
|
## TODO
|
||||||
- [ ] Ver 2.2での学習をサポート
|
- [x] Ver 2.2での学習をサポート←たぶんやった、まだ確認してない
|
||||||
|
- [ ] Ver 2.1での感情のクラス数を10から少なくして実験
|
||||||
- [ ] Ver 2.1, 2.2での学習でのbf16対応
|
- [ ] Ver 2.1, 2.2での学習でのbf16対応
|
||||||
- [ ] 推論のWebUIでのバージョンに応じた感情指定のサポート
|
- [ ] 推論のWebUIでのバージョンに応じた感情指定のサポート
|
||||||
- [ ] より良い推論WebUI?
|
- [ ] より良い推論WebUI?
|
||||||
|
|||||||
64
clap_gen.py
Normal file
64
clap_gen.py
Normal file
@@ -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生成!")
|
||||||
63
configs/config-V210.json
Normal file
63
configs/config-V210.json
Normal file
@@ -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"
|
||||||
|
}
|
||||||
@@ -301,7 +301,7 @@ class DistributedBucketSampler(torch.utils.data.distributed.DistributedSampler):
|
|||||||
self.buckets, self.num_samples_per_bucket = self._create_buckets()
|
self.buckets, self.num_samples_per_bucket = self._create_buckets()
|
||||||
logger.info(f"Bucket info: {self.num_samples_per_bucket}")
|
logger.info(f"Bucket info: {self.num_samples_per_bucket}")
|
||||||
logger.info(
|
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.total_size = sum(self.num_samples_per_bucket)
|
||||||
self.num_samples = self.total_size // self.num_replicas
|
self.num_samples = self.total_size // self.num_replicas
|
||||||
|
|||||||
BIN
empty_emo.npy
Normal file
BIN
empty_emo.npy
Normal file
Binary file not shown.
@@ -307,7 +307,7 @@ class DistributedBucketSampler(torch.utils.data.distributed.DistributedSampler):
|
|||||||
self.buckets, self.num_samples_per_bucket = self._create_buckets()
|
self.buckets, self.num_samples_per_bucket = self._create_buckets()
|
||||||
logger.info(f"Bucket info: {self.num_samples_per_bucket}")
|
logger.info(f"Bucket info: {self.num_samples_per_bucket}")
|
||||||
logger.info(
|
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.total_size = sum(self.num_samples_per_bucket)
|
||||||
self.num_samples = self.total_size // self.num_replicas
|
self.num_samples = self.total_size // self.num_replicas
|
||||||
|
|||||||
425
oldVersion/V220/data_utils.py
Normal file
425
oldVersion/V220/data_utils.py
Normal file
@@ -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
|
||||||
@@ -740,8 +740,10 @@ def train_and_evaluate(
|
|||||||
n_ckpts_to_keep=keep_ckpts,
|
n_ckpts_to_keep=keep_ckpts,
|
||||||
sort_by_time=True,
|
sort_by_time=True,
|
||||||
)
|
)
|
||||||
save_compressed_models = hps.train.save_compressed_models
|
if (
|
||||||
if save_compressed_models:
|
"save_compressed_models" in hps.train.keys()
|
||||||
|
and hps.train.save_compressed_models is True
|
||||||
|
):
|
||||||
utils.save_compressed_models_checkpoint(
|
utils.save_compressed_models_checkpoint(
|
||||||
net_g,
|
net_g,
|
||||||
epoch,
|
epoch,
|
||||||
|
|||||||
@@ -329,6 +329,16 @@ def run():
|
|||||||
if net_dur_disc is not None:
|
if net_dur_disc is not None:
|
||||||
scheduler_dur_disc.step()
|
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(
|
def train_and_evaluate(
|
||||||
rank,
|
rank,
|
||||||
@@ -589,8 +599,10 @@ def train_and_evaluate(
|
|||||||
n_ckpts_to_keep=keep_ckpts,
|
n_ckpts_to_keep=keep_ckpts,
|
||||||
sort_by_time=True,
|
sort_by_time=True,
|
||||||
)
|
)
|
||||||
save_compressed_models = hps.train.save_compressed_models
|
if (
|
||||||
if save_compressed_models:
|
"save_compressed_models" in hps.train.keys()
|
||||||
|
and hps.train.save_compressed_models is True
|
||||||
|
):
|
||||||
utils.save_compressed_models_checkpoint(
|
utils.save_compressed_models_checkpoint(
|
||||||
net_g,
|
net_g,
|
||||||
epoch,
|
epoch,
|
||||||
|
|||||||
745
train_ms_V220.py
Normal file
745
train_ms_V220.py
Normal file
@@ -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()
|
||||||
Reference in New Issue
Block a user