"""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()