deepsafe's picture
sync from GitHub (0154d02)
4b0b144 verified
Raw History Blame Contribute Delete
6.69 kB
"""Evaluate ShiftySpeech SSL-AASIST on the DeepSafe audio dataset.
Reports accuracy, precision, recall, F1, EER, and per-file results.
"""
import os
import sys
import time
import warnings
warnings.filterwarnings("ignore", category=DeprecationWarning)
# Monkey-patch omegaconf for fairseq compatibility
import omegaconf._utils as _omegaconf_utils
if not hasattr(_omegaconf_utils, "is_primitive_type"):
_omegaconf_utils.is_primitive_type = lambda t: t in (int, float, bool, str, bytes)
import argparse
import librosa
import numpy as np
import torch
# Add model code to path
SERVICE_DIR = os.path.dirname(os.path.abspath(__file__))
MODEL_CODE_PATH = os.path.join(
SERVICE_DIR, "synthetic_speech_detection", "SSL_Anti-spoofing"
)
sys.path.insert(0, MODEL_CODE_PATH)
from model import Model as SSLAASISTModel
SAMPLE_RATE = 16000
TARGET_SAMPLES = 64600
DATASET_DIR = os.path.join(
SERVICE_DIR, os.pardir, os.pardir, os.pardir, "dataset", "audio"
)
DATASET_DIR = os.path.normpath(DATASET_DIR)
def pad_audio(audio, target=TARGET_SAMPLES):
"""Pad/trim audio to target length using tiling."""
if len(audio) >= target:
return audio[:target]
num_repeats = target // len(audio) + 1
return np.tile(audio, num_repeats)[:target]
def compute_eer(target_scores, nontarget_scores):
"""Compute Equal Error Rate."""
n_scores = target_scores.size + nontarget_scores.size
all_scores = np.concatenate((target_scores, nontarget_scores))
labels = np.concatenate(
(np.ones(target_scores.size), np.zeros(nontarget_scores.size))
)
indices = np.argsort(all_scores, kind="mergesort")
labels = labels[indices]
tar_trial_sums = np.cumsum(labels)
nontarget_trial_sums = nontarget_scores.size - (
np.arange(1, n_scores + 1) - tar_trial_sums
)
frr = np.concatenate((np.atleast_1d(0), tar_trial_sums / target_scores.size))
far = np.concatenate(
(
np.atleast_1d(1),
nontarget_trial_sums / nontarget_scores.size,
)
)
abs_diffs = np.abs(frr - far)
min_index = np.argmin(abs_diffs)
eer = np.mean((frr[min_index], far[min_index]))
return eer
def main():
weights_path = os.path.join(SERVICE_DIR, "weights", "hfg_aug_1_2.pt")
os.makedirs(os.path.join(SERVICE_DIR, "models"), exist_ok=True)
print("=" * 70)
print("ShiftySpeech (SSL-AASIST) - DeepSafe Dataset Evaluation")
print("=" * 70)
print(f"Weights: {weights_path}")
print(f"Dataset: {DATASET_DIR}")
print(f"Device: cpu")
print()
# Load model
print("Loading model...")
start = time.time()
args_ns = argparse.Namespace()
ssl_model = SSLAASISTModel(args_ns, "cpu")
state_dict = torch.load(weights_path, map_location="cpu", weights_only=False)
ssl_model.load_state_dict(state_dict)
ssl_model.eval()
print(f"Model loaded in {time.time() - start:.1f}s")
print()
# Collect audio files
real_dir = os.path.join(DATASET_DIR, "real")
fake_dir = os.path.join(DATASET_DIR, "fake")
files = []
for fname in sorted(os.listdir(real_dir)):
if fname.endswith(".wav"):
files.append((os.path.join(real_dir, fname), 0, fname))
for fname in sorted(os.listdir(fake_dir)):
if fname.endswith(".wav"):
files.append((os.path.join(fake_dir, fname), 1, fname))
n_real = sum(1 for _, label, _ in files if label == 0)
n_fake = sum(1 for _, label, _ in files if label == 1)
print(f"Total files: {len(files)} (real: {n_real}, fake: {n_fake})")
print()
# Run inference
results = []
total_time = 0.0
print(
f"{'File':<20} {'True':>5} {'Pred':>5} {'P(fake)':>8} "
f"{'P(real)':>8} {'Time':>6}"
)
print("-" * 60)
for path, true_label, fname in files:
audio, sr = librosa.load(path, sr=SAMPLE_RATE, mono=True)
audio = pad_audio(audio)
x = torch.FloatTensor(audio).unsqueeze(0)
t0 = time.time()
with torch.no_grad():
out = ssl_model(x)
elapsed = time.time() - t0
total_time += elapsed
probs = torch.softmax(out, dim=1)
p_fake = probs[0, 0].item()
p_real = probs[0, 1].item()
pred = 1 if p_fake >= 0.5 else 0
results.append(
{
"file": fname,
"true_label": true_label,
"pred_label": pred,
"prob_fake": p_fake,
"prob_real": p_real,
}
)
true_str = "FAKE" if true_label == 1 else "REAL"
pred_str = "FAKE" if pred == 1 else "REAL"
correct = "ok" if pred == true_label else "XX"
print(
f"{fname:<20} {true_str:>5} {pred_str:>5} "
f"{p_fake:>8.4f} {p_real:>8.4f} {elapsed:>5.2f}s "
f"[{correct}]"
)
print()
print("=" * 70)
print("METRICS")
print("=" * 70)
# Compute metrics
true_labels = np.array([r["true_label"] for r in results])
pred_labels = np.array([r["pred_label"] for r in results])
tp = int(np.sum((pred_labels == 1) & (true_labels == 1)))
tn = int(np.sum((pred_labels == 0) & (true_labels == 0)))
fp = int(np.sum((pred_labels == 1) & (true_labels == 0)))
fn = int(np.sum((pred_labels == 0) & (true_labels == 1)))
accuracy = (tp + tn) / len(results) if len(results) > 0 else 0
precision = tp / (tp + fp) if (tp + fp) > 0 else 0
recall = tp / (tp + fn) if (tp + fn) > 0 else 0
f1 = (
2 * precision * recall / (precision + recall) if (precision + recall) > 0 else 0
)
specificity = tn / (tn + fp) if (tn + fp) > 0 else 0
# EER using bonafide scores (prob_real: higher = more real)
bonafide_scores = np.array(
[r["prob_real"] for r in results if r["true_label"] == 0]
)
spoof_scores = np.array([r["prob_real"] for r in results if r["true_label"] == 1])
if len(bonafide_scores) > 0 and len(spoof_scores) > 0:
eer = compute_eer(bonafide_scores, spoof_scores)
else:
eer = float("nan")
print(f"Accuracy: {accuracy:.4f} ({accuracy * 100:.1f}%)")
print(f"Precision: {precision:.4f}")
print(f"Recall: {recall:.4f}")
print(f"F1 Score: {f1:.4f}")
print(f"Specificity: {specificity:.4f}")
print(f"EER: {eer:.4f} ({eer * 100:.1f}%)")
print()
print(f"Confusion Matrix:")
print(f" TP={tp:>3d} FP={fp:>3d}")
print(f" FN={fn:>3d} TN={tn:>3d}")
print()
print(f"Total inference time: {total_time:.1f}s")
print(f"Avg per file: {total_time / len(results):.3f}s")
print(f"Total files: {len(results)}")
if __name__ == "__main__":
main()