-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathcompute_secs.py
More file actions
76 lines (59 loc) · 2.44 KB
/
Copy pathcompute_secs.py
File metadata and controls
76 lines (59 loc) · 2.44 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
import torch
import os
import numpy as np
import librosa
import torch.nn.functional as F
from transformers import Wav2Vec2FeatureExtractor, WavLMForXVector
LENGTH = 63
TEST_DIR = "test"
SAMPLE_RATE = 16000
feature_extractor = Wav2Vec2FeatureExtractor.from_pretrained(
"microsoft/wavlm-base-plus-sv"
)
sv_model = WavLMForXVector.from_pretrained("microsoft/wavlm-base-plus-sv")
def SECS(ref, gen):
spk1_wav, _ = librosa.load(ref, sr=SAMPLE_RATE)
spk2_wav, _ = librosa.load(gen, sr=SAMPLE_RATE)
input1 = feature_extractor(
[spk1_wav], padding=True, return_tensors="pt", sampling_rate=SAMPLE_RATE
)
if torch.cuda.is_available():
for key in input1.keys():
input1[key] = input1[key].to(sv_model.device)
with torch.no_grad():
embds_1 = sv_model(**input1).embeddings
embds_1 = embds_1[0]
input2 = feature_extractor(
[spk2_wav], padding=True, return_tensors="pt", sampling_rate=SAMPLE_RATE
)
if torch.cuda.is_available():
for key in input2.keys():
input2[key] = input2[key].to(sv_model.device)
with torch.no_grad():
embds_2 = sv_model(**input2).embeddings
embds_2 = embds_2[0]
cos_sim = F.cosine_similarity(embds_1, embds_2, dim=-1).detach().cpu().numpy()
return cos_sim
if __name__ == "__main__":
# Список для хранения результатов SECS
secs_scores = []
# Пройдемся по всем файлам
for i in range(LENGTH):
original_audio_file = os.path.join(TEST_DIR, f"original_audio_{i}.wav")
generated_audio_file = os.path.join(TEST_DIR, f"generated_audio_{i}.wav")
# Проверка существования файлов
if not os.path.exists(original_audio_file) or not os.path.exists(generated_audio_file):
print(f"Files for index {i} not found. Skipping.")
continue
# Расчет SECS
secs_score = SECS(original_audio_file, generated_audio_file)
secs_scores.append(secs_score)
# Вывод результата для текущего файла
print(f"File: {i}")
print(f"Original: {original_audio_file}")
print(f"Generated: {generated_audio_file}")
print(f"SECS: {secs_score:.4f}")
print("-" * 50)
# Среднее значение SECS
average_secs = np.mean(secs_scores)
print(f"Average SECS: {average_secs:.4f}")