From 208bb0f0ee75cec191e6e1f3875d4c9025dfa85a Mon Sep 17 00:00:00 2001 From: Edresson Date: Thu, 27 May 2021 20:01:00 -0300 Subject: [PATCH] add batched speaker encoder inference --- TTS/bin/compute_embeddings.py | 1 - TTS/speaker_encoder/models/resnet.py | 12 +++++++----- 2 files changed, 7 insertions(+), 6 deletions(-) diff --git a/TTS/bin/compute_embeddings.py b/TTS/bin/compute_embeddings.py index 045aa372..9affac64 100644 --- a/TTS/bin/compute_embeddings.py +++ b/TTS/bin/compute_embeddings.py @@ -2,7 +2,6 @@ import argparse import glob import os -import numpy as np import torch from tqdm import tqdm diff --git a/TTS/speaker_encoder/models/resnet.py b/TTS/speaker_encoder/models/resnet.py index fe89c5aa..aa2171ed 100644 --- a/TTS/speaker_encoder/models/resnet.py +++ b/TTS/speaker_encoder/models/resnet.py @@ -174,15 +174,17 @@ class ResNetSpeakerEncoder(nn.Module): offsets = np.linspace(0, max_len-num_frames, num=num_eval) - embeddings = [] + frames_batch = [] for offset in offsets: offset = int(offset) end_offset = int(offset+num_frames) frames = x[:, offset:end_offset] - embed = self.forward(frames, l2_norm=True) - embeddings.append(embed) + frames_batch.append(frames) + + frames_batch = torch.cat(frames_batch, dim=0) + embeddings = self.forward(frames_batch, l2_norm=True) - embeddings = torch.stack(embeddings) if return_mean: - embeddings = torch.mean(embeddings, dim=0) + embeddings = torch.mean(embeddings, dim=0, keepdim=True) + return embeddings