diff --git a/emo_gen.py b/emo_gen.py new file mode 100644 index 0000000..0856ab2 --- /dev/null +++ b/emo_gen.py @@ -0,0 +1,169 @@ +import torch +import torch.nn as nn +from torch.utils.data import Dataset +from torch.utils.data import DataLoader +from transformers import Wav2Vec2Processor +from transformers.models.wav2vec2.modeling_wav2vec2 import ( + Wav2Vec2Model, + Wav2Vec2PreTrainedModel, +) +import librosa +import numpy as np +import argparse +from config import config +import utils +import os +from tqdm import tqdm + + +class RegressionHead(nn.Module): + r"""Classification head.""" + + def __init__(self, config): + super().__init__() + + self.dense = nn.Linear(config.hidden_size, config.hidden_size) + self.dropout = nn.Dropout(config.final_dropout) + self.out_proj = nn.Linear(config.hidden_size, config.num_labels) + + def forward(self, features, **kwargs): + x = features + x = self.dropout(x) + x = self.dense(x) + x = torch.tanh(x) + x = self.dropout(x) + x = self.out_proj(x) + + return x + + +class EmotionModel(Wav2Vec2PreTrainedModel): + r"""Speech emotion classifier.""" + + def __init__(self, config): + super().__init__(config) + + self.config = config + self.wav2vec2 = Wav2Vec2Model(config) + self.classifier = RegressionHead(config) + self.init_weights() + + def forward( + self, + input_values, + ): + outputs = self.wav2vec2(input_values) + hidden_states = outputs[0] + hidden_states = torch.mean(hidden_states, dim=1) + logits = self.classifier(hidden_states) + + return hidden_states, logits + + +class AudioDataset(Dataset): + def __init__(self, list_of_wav_files, sr, processor): + self.list_of_wav_files = list_of_wav_files + self.processor = processor + self.sr = sr + + def __len__(self): + return len(self.list_of_wav_files) + + def __getitem__(self, idx): + wav_file = self.list_of_wav_files[idx] + audio_data, _ = librosa.load(wav_file, sr=self.sr) + processed_data = self.processor(audio_data, sampling_rate=self.sr)[ + "input_values" + ][0] + return torch.from_numpy(processed_data) + + +model_name = "./emotional/wav2vec2-large-robust-12-ft-emotion-msp-dim" +processor = Wav2Vec2Processor.from_pretrained(model_name) +model = EmotionModel.from_pretrained(model_name) + + +def process_func( + x: np.ndarray, + sampling_rate: int, + model: EmotionModel, + processor: Wav2Vec2Processor, + device: str, + embeddings: bool = False, +) -> np.ndarray: + r"""Predict emotions or extract embeddings from raw audio signal.""" + model = model.to(device) + y = processor(x, sampling_rate=sampling_rate) + y = y["input_values"][0] + y = torch.from_numpy(y).unsqueeze(0).to(device) + + # run through model + with torch.no_grad(): + y = model(y)[0 if embeddings else 1] + + # convert to numpy + y = y.detach().cpu().numpy() + + return y + + +def get_emo(path): + wav, sr = librosa.load(path, 16000) + device = config.bert_gen_config.device + return process_func( + np.expand_dims(wav, 0).astype(np.float), + sr, + model, + processor, + device, + embeddings=True, + ).squeeze(0) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument( + "-c", "--config", type=str, default=config.bert_gen_config.config_path + ) + parser.add_argument( + "--num_processes", type=int, default=config.bert_gen_config.num_processes + ) + args, _ = parser.parse_known_args() + config_path = args.config + hps = utils.get_hparams_from_file(config_path) + + device = config.bert_gen_config.device + + model_name = "./emotional/wav2vec2-large-robust-12-ft-emotion-msp-dim" + processor = ( + Wav2Vec2Processor.from_pretrained(model_name) + if processor is None + else processor + ) + model = ( + EmotionModel.from_pretrained(model_name).to(device) + if model is None + else model.to(device) + ) + + 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()) + + wavnames = [line.split("|")[0] for line in lines] + dataset = AudioDataset(wavnames, 16000, processor) + data_loader = DataLoader(dataset, batch_size=1, shuffle=False, num_workers=16) + + with torch.no_grad(): + for i, data in tqdm(enumerate(data_loader), total=len(data_loader)): + wavname = wavnames[i] + emo_path = wavname.replace(".wav", ".emo.npy") + if os.path.exists(emo_path): + continue + emb = model(data.to(device))[0].detach().cpu().numpy() + np.save(emo_path, emb) + + print("Emo vec 生成完毕!") diff --git a/models.py b/models.py index d53b430..d493878 100644 --- a/models.py +++ b/models.py @@ -10,6 +10,7 @@ import monotonic_align from torch.nn import Conv1d, ConvTranspose1d, Conv2d from torch.nn.utils import weight_norm, remove_weight_norm, spectral_norm +from vector_quantize_pytorch import VectorQuantize from commons import init_weights, get_padding from text import symbols, num_tones, num_languages @@ -321,6 +322,7 @@ class TextEncoder(nn.Module): n_layers, kernel_size, p_dropout, + n_speakers, gin_channels=0, ): super().__init__() @@ -342,6 +344,18 @@ class TextEncoder(nn.Module): self.bert_proj = nn.Conv1d(1024, hidden_channels, 1) self.ja_bert_proj = nn.Conv1d(1024, hidden_channels, 1) self.en_bert_proj = nn.Conv1d(1024, hidden_channels, 1) + self.emo_proj = nn.Linear(1024, 1024) + self.emo_quantizer = [ + VectorQuantize( + dim=1024, + codebook_size=5, + decay=0.8, + commitment_weight=1.0, + learnable_codebook=True, + ema_update=False, + ) + ] * n_speakers + self.emo_q_proj = nn.Linear(1024, hidden_channels) self.encoder = attentions.Encoder( hidden_channels, @@ -354,10 +368,33 @@ class TextEncoder(nn.Module): ) self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1) - def forward(self, x, x_lengths, tone, language, bert, ja_bert, en_bert, g=None): + def forward( + self, x, x_lengths, tone, language, bert, ja_bert, en_bert, emo, sid, g=None + ): + sid = sid.cpu() bert_emb = self.bert_proj(bert).transpose(1, 2) ja_bert_emb = self.ja_bert_proj(ja_bert).transpose(1, 2) en_bert_emb = self.en_bert_proj(en_bert).transpose(1, 2) + if emo.size(-1) == 1024: + emo_emb = self.emo_proj(emo.unsqueeze(1)) + emo_commit_loss = torch.zeros(1) + emo_emb_ = [] + for i in range(emo_emb.size(0)): + temp_emo_emb, _, temp_emo_commit_loss = self.emo_quantizer[sid[i]]( + emo_emb[i].unsqueeze(0).cpu() + ) + emo_commit_loss += temp_emo_commit_loss + emo_emb_.append(temp_emo_emb) + emo_emb = torch.cat(emo_emb_, dim=0).to(emo_emb.device) + emo_commit_loss = emo_commit_loss.to(emo_emb.device) + else: + emo_emb = ( + self.emo_quantizer[sid[0]] + .get_output_from_indices(emo.to(torch.int).cpu()) + .unsqueeze(0) + .to(emo.device) + ) + emo_commit_loss = torch.zeros(1) x = ( self.emb(x) + self.tone_emb(tone) @@ -365,6 +402,7 @@ class TextEncoder(nn.Module): + bert_emb + ja_bert_emb + en_bert_emb + + self.emo_q_proj(emo_emb) ) * math.sqrt( self.hidden_channels ) # [b, t, h] @@ -377,7 +415,7 @@ class TextEncoder(nn.Module): stats = self.proj(x) * x_mask m, logs = torch.split(stats, self.out_channels, dim=1) - return x, m, logs, x_mask + return x, m, logs, x_mask, emo_commit_loss class ResidualCouplingBlock(nn.Module): @@ -810,6 +848,7 @@ class SynthesizerTrn(nn.Module): n_layers, kernel_size, p_dropout, + self.n_speakers, gin_channels=self.enc_gin_channels, ) self.dec = Generator( @@ -860,7 +899,7 @@ class SynthesizerTrn(nn.Module): hidden_channels, 256, 3, 0.5, gin_channels=gin_channels ) - if n_speakers >= 1: + if n_speakers > =1: self.emb_g = nn.Embedding(n_speakers, gin_channels) else: self.ref_enc = ReferenceEncoder(spec_channels, gin_channels) @@ -877,13 +916,14 @@ class SynthesizerTrn(nn.Module): bert, ja_bert, en_bert, + emo=None, ): if self.n_speakers > 0: g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1] else: g = self.ref_enc(y.transpose(1, 2)).unsqueeze(-1) - x, m_p, logs_p, x_mask = self.enc_p( - x, x_lengths, tone, language, bert, ja_bert, en_bert, g=g + x, m_p, logs_p, x_mask, loss_commit = self.enc_p( + x, x_lengths, tone, language, bert, ja_bert, en_bert, emo, sid, g=g ) z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g) z_p = self.flow(z, y_mask, g=g) @@ -949,6 +989,7 @@ class SynthesizerTrn(nn.Module): y_mask, (z, z_p, m_p, logs_p, m_q, logs_q), (x, logw, logw_), + loss_commit, ) def infer( @@ -961,6 +1002,7 @@ class SynthesizerTrn(nn.Module): bert, ja_bert, en_bert, + emo=None, noise_scale=0.667, length_scale=1, noise_scale_w=0.8, @@ -974,8 +1016,8 @@ class SynthesizerTrn(nn.Module): g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1] else: g = self.ref_enc(y.transpose(1, 2)).unsqueeze(-1) - x, m_p, logs_p, x_mask = self.enc_p( - x, x_lengths, tone, language, bert, ja_bert, en_bert, g=g + x, m_p, logs_p, x_mask, _ = self.enc_p( + x, x_lengths, tone, language, bert, ja_bert, en_bert, emo, sid, g=g ) logw = self.sdp(x, x_mask, g=g, reverse=True, noise_scale=noise_scale_w) * ( sdp_ratio diff --git a/train_ms.py b/train_ms.py index 1bda2a5..37933d1 100644 --- a/train_ms.py +++ b/train_ms.py @@ -340,6 +340,7 @@ def train_and_evaluate( bert, ja_bert, en_bert, + emo, ) in tqdm(enumerate(train_loader)): if net_g.module.use_noise_scaled_mas: current_mas_noise_scale = ( @@ -362,6 +363,7 @@ def train_and_evaluate( bert = bert.cuda(rank, non_blocking=True) ja_bert = ja_bert.cuda(rank, non_blocking=True) en_bert = en_bert.cuda(rank, non_blocking=True) + emo = emo.cuda(rank, non_blocking=True) with autocast(enabled=hps.train.fp16_run): ( @@ -373,6 +375,7 @@ def train_and_evaluate( z_mask, (z, z_p, m_p, logs_p, m_q, logs_q), (hidden_x, logw, logw_), + loss_commit, ) = net_g( x, x_lengths, @@ -384,6 +387,7 @@ def train_and_evaluate( bert, ja_bert, en_bert, + emo, ) mel = spec_to_mel_torch( spec, @@ -454,7 +458,9 @@ def train_and_evaluate( 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_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 @@ -579,6 +585,7 @@ def evaluate(hps, generator, eval_loader, writer_eval): 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()